#define PROBLEM "https://judge.yosupo.jp/problem/aplusb"
#include "template/template.hpp"
#include "modint/power-tower-modint.hpp"
uint64_t mod_pow ( uint64_t a , uint64_t n , uint64_t mod ) {
uint64_t res = 1 % mod ;
while ( n ) {
if ( n & 1 ) res = __uint128_t ( res ) * a % mod ;
a = __uint128_t ( a ) * a % mod ;
n >>= 1 ;
}
return res ;
}
uint64_t pow_u64 ( uint64_t a , uint64_t n ) {
uint64_t res = 1 ;
while ( n ) {
if ( n & 1 ) res *= a ;
a *= a ;
n >>= 1 ;
}
return res ;
}
template < uint32_t mod >
void test () {
using mint = PowerTowerModInt < mod > ;
static_assert ( mint :: get_mod () == mod );
const vector < uint64_t > values = {
0 , 1 , mod - 1 , mod , uint64_t ( mod ) + 1 , uint64_t ( mod ) * 2 ,
numeric_limits < uint32_t >:: max (), numeric_limits < uint64_t >:: max ()};
for ( uint64_t a : values )
for ( uint64_t b : values ) {
mint x = a , y = b ;
assert ( x . val () == a % mod );
assert ( x . large_val () == ( a < mod ? a : mod + a % mod ));
assert (( x + y ). val () == uint64_t (( __uint128_t ( a ) + b ) % mod ));
assert (( x * y ). val () == uint64_t ( __uint128_t ( a ) * b % mod ));
assert ( x . pow ( y ). val () == mod_pow ( a % mod , b , mod ));
if ( __uint128_t ( a ) >= __uint128_t ( b ) + mod ) {
mint z = x ;
z . unsafe_subtract ( y );
assert ( z . val () == ( a - b ) % mod );
assert ( mint ( 7 ). pow ( z ). val () == mod_pow ( 7 , a - b , mod ));
}
if ( a >= mod ) {
mint z = x + y ;
z . unsafe_subtract ( y );
assert ( z . val () == a % mod );
assert ( mint ( 7 ). pow ( z ). val () == mod_pow ( 7 , a , mod ));
}
}
for ( uint64_t a = 0 ; a <= 5 ; a ++ )
for ( uint64_t b = 0 ; b <= 5 ; b ++ )
for ( uint64_t c = 0 ; c <= 8 ; c ++ ) {
uint64_t exponent = pow_u64 ( b , c );
assert ( mint ( a ). pow ( mint ( b ). pow ( mint ( c ))). val () == mod_pow ( a , exponent , mod ));
}
for ( uint64_t x = 0 ; x <= min < uint64_t > ( mod - 1 , 100 ); x ++ )
for ( uint64_t y = 0 ; y <= x ; y ++ ) {
mint difference = x ;
difference . unsafe_subtract ( int64_t ( y ));
assert ( difference . val () == x - y );
assert ( difference . large_val () == x - y );
assert ( mint ( 7 ). pow ( difference ). val () == mod_pow ( 7 , x - y , mod ));
difference = mint ( x );
difference . unsafe_subtract ( mint ( y ));
assert ( difference . val () == x - y );
assert ( difference . large_val () == x - y );
assert ( mint ( 7 ). pow ( difference ). val () == mod_pow ( 7 , x - y , mod ));
}
mint large_difference = mint ( uint64_t ( mod ) + 123 );
large_difference . unsafe_subtract ( int64_t ( 123 ));
assert ( large_difference . val () == 0 );
assert ( large_difference . large_val () == mod );
assert ( mint ( 7 ). pow ( large_difference ). val () == mod_pow ( 7 , mod , mod ));
assert ( mint ( 0 ). pow ( mint ( 0 )). val () == 1 % mod );
assert ( mint ( 0 ). pow ( mint ( 0 ). pow ( mint ( 0 ))). val () == 0 );
assert ( mint ( 0 ). pow ( mint ( 0 ). pow ( mint ( 1 ))). val () == 1 % mod );
assert ( mint ( 0 ) == mint ( mod ));
mint huge = mint ( 2 ). pow ( mint ( 100 ));
mint difference = huge + mint ( mod );
difference . unsafe_subtract ( huge );
assert ( difference . val () == 0 );
assert ( difference . large_val () == mod );
assert ( mint ( 7 ). pow ( difference ). val () == mod_pow ( 7 , mod , mod ));
stringstream ss ;
ss << mint ( numeric_limits < uint64_t >:: max ());
assert ( ss . str () == to_string ( numeric_limits < uint64_t >:: max () % mod ));
}
int main () {
test < 1 > ();
test < 2 > ();
test < 3 > ();
test < 4 > ();
test < 6 > ();
test < 10 > ();
test < 998244353 > ();
test < 2147483647 > ();
using mint = PowerTowerModInt < 10 > ;
assert ( mint ( 2 ). pow ( mint ( 1 )). val () == 2 );
assert ( mint ( 2 ). pow ( mint ( 4 )). val () == 6 );
assert ( mint ( 2 ). pow ( mint ( 8 )). val () == 6 );
int a , b ;
in ( a , b );
out ( a + b );
}
#line 1 "verify/modint/UNIT_power_tower_modint.test.cpp"
#define PROBLEM "https://judge.yosupo.jp/problem/aplusb"
#line 2 "template/template.hpp"
#include <bits/stdc++.h>
using namespace std ;
#line 2 "template/macro.hpp"
#define rep(i, a, b) for (int i = (a); i < (int)(b); i++)
#define rrep(i, a, b) for (int i = (int)(b) - 1; i >= (a); i--)
#define ALL(v) (v).begin(), (v).end()
#define UNIQUE(v) sort(ALL(v)), (v).erase(unique(ALL(v)), (v).end())
#define SZ(v) (int)v.size()
#define MIN(v) *min_element(ALL(v))
#define MAX(v) *max_element(ALL(v))
#define LB(v, x) int(lower_bound(ALL(v), (x)) - (v).begin())
#define UB(v, x) int(upper_bound(ALL(v), (x)) - (v).begin())
#define YN(b) cout << ((b) ? "YES" : "NO") << "\n";
#define Yn(b) cout << ((b) ? "Yes" : "No") << "\n";
#define yn(b) cout << ((b) ? "yes" : "no") << "\n";
#line 6 "template/template.hpp"
#line 2 "template/util.hpp"
using uint = unsigned int ;
using ll = long long int ;
using ull = unsigned long long ;
using i128 = __int128_t ;
using u128 = __uint128_t ;
template < class T >
using priority_queue_asc = priority_queue < T , vector < T > , greater < T >> ;
template < class T , class S = T >
S SUM ( const vector < T >& a ) {
return accumulate ( ALL ( a ), S ( 0 ));
}
template < class T1 , class T2 >
inline bool chmin ( T1 & a , T2 b ) {
if ( a > b ) {
a = b ;
return true ;
}
return false ;
}
template < class T1 , class T2 >
inline bool chmax ( T1 & a , T2 b ) {
if ( a < b ) {
a = b ;
return true ;
}
return false ;
}
template < class T1 , class T2 >
inline bool chmin_opt ( optional < T1 >& a , T2 b ) {
if ( ! a || a > b ) {
a = b ;
return true ;
}
return false ;
}
template < class T1 , class T2 >
inline bool chmax_opt ( optional < T1 >& a , T2 b ) {
if ( ! a || a < b ) {
a = b ;
return true ;
}
return false ;
}
template < class T >
int popcnt ( T x ) {
return __builtin_popcountll ( x );
}
template < class T >
int topbit ( T x ) {
return ( x == 0 ? - 1 : 63 - __builtin_clzll ( x ));
}
template < class T >
int lowbit ( T x ) {
return ( x == 0 ? - 1 : __builtin_ctzll ( x ));
}
#line 8 "template/template.hpp"
#line 2 "template/inout.hpp"
struct Fast {
Fast () {
cin . tie ( nullptr );
ios_base :: sync_with_stdio ( false );
cout << fixed << setprecision ( 15 );
}
} fast ;
ostream & operator << ( ostream & os , __uint128_t x ) {
char buf [ 40 ];
size_t k = 0 ;
while ( x > 0 ) buf [ k ++ ] = ( char )( x % 10 + '0' ), x /= 10 ;
if ( k == 0 ) buf [ k ++ ] = '0' ;
while ( k ) os << buf [ -- k ];
return os ;
}
ostream & operator << ( ostream & os , __int128_t x ) {
return x < 0 ? ( os << '-' << ( __uint128_t )( - x )) : ( os << ( __uint128_t ) x );
}
template < class T , size_t N >
ostream & operator << ( ostream & os , const array < T , N >& a );
template < class T1 , class T2 >
istream & operator >> ( istream & is , pair < T1 , T2 >& p ) {
return is >> p . first >> p . second ;
}
template < class T1 , class T2 >
ostream & operator << ( ostream & os , const pair < T1 , T2 >& p ) {
return os << p . first << " " << p . second ;
}
template < class T >
istream & operator >> ( istream & is , vector < T >& a ) {
for ( auto & v : a ) is >> v ;
return is ;
}
template < class T >
ostream & operator << ( ostream & os , const vector < T >& a ) {
for ( auto it = a . begin (); it != a . end ();) {
os << * it ;
if ( ++ it != a . end ()) os << " " ;
}
return os ;
}
template < class T , size_t N >
ostream & operator << ( ostream & os , const array < T , N >& a ) {
for ( auto it = a . begin (); it != a . end ();) {
os << * it ;
if ( ++ it != a . end ()) os << " " ;
}
return os ;
}
template < class T >
ostream & operator << ( ostream & os , const set < T >& st ) {
os << "{" ;
for ( auto it = st . begin (); it != st . end ();) {
os << * it ;
if ( ++ it != st . end ()) os << "," ;
}
os << "}" ;
return os ;
}
template < class T1 , class T2 >
ostream & operator << ( ostream & os , const map < T1 , T2 >& mp ) {
os << "{" ;
for ( auto it = mp . begin (); it != mp . end ();) {
os << it -> first << ":" << it -> second ;
if ( ++ it != mp . end ()) os << "," ;
}
os << "}" ;
return os ;
}
void in () {}
template < typename T , class ... U >
void in ( T & t , U & ... u ) {
cin >> t ;
in ( u ...);
}
template < class ... T >
void in_zip ( int n , T & ... t ) {
assert ( n >= 0 && (( size ( t ) >= static_cast < size_t > ( n )) && ...));
for ( int i = 0 ; i < n ; i ++ ) in ( t [ i ]...);
}
void out () { cout << " \n " ; }
template < typename T , class ... U , char sep = ' ' >
void out ( const T & t , const U & ... u ) {
cout << t ;
if ( sizeof ...( u )) cout << sep ;
out ( u ...);
}
template < class T , class U >
void out_opt ( const optional < T >& opt , const U & fallback , ostream & os = cout ) {
if ( opt . has_value ())
os << opt . value ();
else
os << fallback ;
os << " \n " ;
}
template < class T , class U >
void out_opt ( const vector < optional < T >>& vec , const U & fallback , ostream & os = cout ) {
for ( auto it = vec . begin (); it != vec . end ();) {
if (( * it ). has_value ())
os << ( * it ). value ();
else
os << fallback ;
if ( ++ it != vec . end ()) os << " " ;
}
os << " \n " ;
}
namespace IO {
template < class T , class ... U >
T read ( U && ... u ) {
T t = T ( forward < U > ( u )...);
in ( t );
return t ;
}
namespace Graph {
vector < vector < int >> unweighted ( int n , int m , bool directed = false , int offset = 1 ) {
vector < vector < int >> g ( n );
for ( int i = 0 ; i < m ; i ++ ) {
int u , v ;
cin >> u >> v ;
u -= offset , v -= offset ;
g [ u ]. push_back ( v );
if ( ! directed ) g [ v ]. push_back ( u );
}
return g ;
}
template < class T >
vector < vector < pair < int , T >>> weighted ( int n , int m , bool directed = false , int offset = 1 ) {
vector < vector < pair < int , T >>> g ( n );
for ( int i = 0 ; i < m ; i ++ ) {
int u , v ;
T w ;
cin >> u >> v >> w ;
u -= offset , v -= offset ;
g [ u ]. push_back ({ v , w });
if ( ! directed ) g [ v ]. push_back ({ u , w });
}
return g ;
}
} // namespace Graph
namespace Tree {
vector < vector < int >> unweighted ( int n , bool directed = false , int offset = 1 ) {
return Graph :: unweighted ( n , n - 1 , directed , offset );
}
template < class T >
vector < vector < pair < int , T >>> weighted ( int n , bool directed = false , int offset = 1 ) {
return Graph :: weighted < T > ( n , n - 1 , directed , offset );
}
vector < vector < int >> rooted ( int n , bool to_root = true , bool to_leaf = true , int offset = 1 ) {
vector < vector < int >> g ( n );
for ( int i = 1 ; i < n ; i ++ ) {
int p ;
cin >> p ;
p -= offset ;
if ( to_root ) g [ i ]. push_back ( p );
if ( to_leaf ) g [ p ]. push_back ( i );
}
return g ;
}
} // namespace Tree
} // namespace IO
#line 10 "template/template.hpp"
#line 2 "template/debug.hpp"
#ifdef LOCAL
#define debug 1
#define show(...) _show(0, #__VA_ARGS__, __VA_ARGS__)
#else
#define debug 0
#define show(...) true
#endif
template < class T >
void _show ( int , T ) {
cerr << '\n' ;
}
template < class T1 , class T2 , class ... T3 >
void _show ( int i , const T1 & a , const T2 & b , const T3 & ... c ) {
for (; a [ i ] != ',' && a [ i ] != '\0' ; i ++ ) cerr << a [ i ];
cerr << ":" << b << " " ;
_show ( i + 1 , a , c ...);
}
#line 2 "modint/power-tower-modint.hpp"
namespace PowerTowerModIntInternal {
constexpr uint32_t totient ( uint32_t n ) {
uint32_t res = n ;
for ( uint32_t p = 2 ; p <= n / p ; p ++ )
if ( n % p == 0 ) {
res = res / p * ( p - 1 );
do n /= p ;
while ( n % p == 0 );
}
if ( n > 1 ) res = res / n * ( n - 1 );
return res ;
}
} // namespace PowerTowerModIntInternal
template < uint32_t m >
struct PowerTowerModInt {
static_assert ( 1 <= m && m < 0x80000000u );
using mint = PowerTowerModInt ;
private:
static constexpr uint32_t phi = PowerTowerModIntInternal :: totient ( m );
using lower_mint = PowerTowerModInt < phi > ;
public:
static constexpr uint32_t get_mod () { return m ; }
PowerTowerModInt () : _v ( 0 ), _lower ( 0 ) {}
PowerTowerModInt ( uint64_t v ) : _v ( normalize ( v )), _lower ( v ) {}
uint32_t val () const { return _v < m ? _v : _v - m ; }
uint32_t large_val () const { return _v ; }
mint & operator += ( const mint & rhs ) {
_v = normalize ( uint64_t ( _v ) + rhs . _v );
_lower += rhs . _lower ;
return * this ;
}
mint & operator *= ( const mint & rhs ) {
_v = normalize ( uint64_t ( _v ) * rhs . _v );
_lower *= rhs . _lower ;
return * this ;
}
mint & unsafe_subtract ( const mint & rhs ) {
if ( _v < m ) {
assert ( rhs . _v < m && rhs . _v <= _v );
return * this = mint ( _v - rhs . _v );
}
_v = m + uint32_t (( uint64_t ( val ()) + m - rhs . val ()) % m );
_lower . unsafe_subtract ( rhs . _lower );
return * this ;
}
mint & unsafe_subtract ( int64_t rhs ) {
assert ( rhs >= 0 );
if ( _v < m ) {
assert ( uint64_t ( rhs ) <= _v );
return * this = mint ( _v - uint64_t ( rhs ));
}
return unsafe_subtract ( mint ( uint64_t ( rhs )));
}
mint pow ( const mint & exponent ) const {
return raw ( pow_mod ( _v , exponent . _lower . _v ), _lower . pow ( exponent . _lower ));
}
friend mint operator + ( const mint & lhs , const mint & rhs ) { return mint ( lhs ) += rhs ; }
friend mint operator * ( const mint & lhs , const mint & rhs ) { return mint ( lhs ) *= rhs ; }
friend bool operator == ( const mint & lhs , const mint & rhs ) { return lhs . val () == rhs . val (); }
friend bool operator != ( const mint & lhs , const mint & rhs ) { return ! ( lhs == rhs ); }
friend ostream & operator << ( ostream & os , const mint & x ) { return os << x . val (); }
private:
uint32_t _v ;
lower_mint _lower ;
template < uint32_t >
friend struct PowerTowerModInt ;
static uint32_t normalize ( uint64_t v ) {
if ( v < uint64_t ( m ) * 2 ) return uint32_t ( v );
return uint32_t ( v % m ) + m ;
}
static uint32_t pow_mod ( uint32_t a , uint32_t n ) {
uint32_t res = 1 ;
while ( n ) {
if ( n & 1 ) res = normalize ( uint64_t ( res ) * a );
a = normalize ( uint64_t ( a ) * a );
n >>= 1 ;
}
return res ;
}
static mint raw ( uint32_t v , const lower_mint & lower ) {
mint res ;
res . _v = v ;
res . _lower = lower ;
return res ;
}
};
template < >
struct PowerTowerModInt < 1 > {
using mint = PowerTowerModInt ;
static constexpr uint32_t get_mod () { return 1 ; }
PowerTowerModInt () : _v ( false ) {}
PowerTowerModInt ( uint64_t v ) : _v ( v != 0 ) {}
uint32_t val () const { return 0 ; }
uint32_t large_val () const { return _v ; }
mint & operator += ( const mint & rhs ) {
_v = _v || rhs . _v ;
return * this ;
}
mint & operator *= ( const mint & rhs ) {
_v = _v && rhs . _v ;
return * this ;
}
mint & unsafe_subtract ( const mint & rhs ) {
if ( ! _v ) assert ( ! rhs . _v );
return * this ;
}
mint & unsafe_subtract ( int64_t rhs ) {
assert ( rhs >= 0 );
if ( ! _v ) {
assert ( rhs == 0 );
return * this ;
}
return unsafe_subtract ( mint ( uint64_t ( rhs )));
}
mint pow ( const mint & exponent ) const { return raw ( _v || ! exponent . _v ); }
friend mint operator + ( const mint & lhs , const mint & rhs ) { return mint ( lhs ) += rhs ; }
friend mint operator * ( const mint & lhs , const mint & rhs ) { return mint ( lhs ) *= rhs ; }
friend bool operator == ( const mint & , const mint & ) { return true ; }
friend bool operator != ( const mint & , const mint & ) { return false ; }
friend ostream & operator << ( ostream & os , const mint & ) { return os << 0 ; }
private:
bool _v ;
template < uint32_t >
friend struct PowerTowerModInt ;
static mint raw ( bool positive ) {
mint res ;
res . _v = positive ;
return res ;
}
};
/**
* @brief Power Tower Modint
* @docs docs/modint/power-tower-modint.md
*/
#line 5 "verify/modint/UNIT_power_tower_modint.test.cpp"
uint64_t mod_pow ( uint64_t a , uint64_t n , uint64_t mod ) {
uint64_t res = 1 % mod ;
while ( n ) {
if ( n & 1 ) res = __uint128_t ( res ) * a % mod ;
a = __uint128_t ( a ) * a % mod ;
n >>= 1 ;
}
return res ;
}
uint64_t pow_u64 ( uint64_t a , uint64_t n ) {
uint64_t res = 1 ;
while ( n ) {
if ( n & 1 ) res *= a ;
a *= a ;
n >>= 1 ;
}
return res ;
}
template < uint32_t mod >
void test () {
using mint = PowerTowerModInt < mod > ;
static_assert ( mint :: get_mod () == mod );
const vector < uint64_t > values = {
0 , 1 , mod - 1 , mod , uint64_t ( mod ) + 1 , uint64_t ( mod ) * 2 ,
numeric_limits < uint32_t >:: max (), numeric_limits < uint64_t >:: max ()};
for ( uint64_t a : values )
for ( uint64_t b : values ) {
mint x = a , y = b ;
assert ( x . val () == a % mod );
assert ( x . large_val () == ( a < mod ? a : mod + a % mod ));
assert (( x + y ). val () == uint64_t (( __uint128_t ( a ) + b ) % mod ));
assert (( x * y ). val () == uint64_t ( __uint128_t ( a ) * b % mod ));
assert ( x . pow ( y ). val () == mod_pow ( a % mod , b , mod ));
if ( __uint128_t ( a ) >= __uint128_t ( b ) + mod ) {
mint z = x ;
z . unsafe_subtract ( y );
assert ( z . val () == ( a - b ) % mod );
assert ( mint ( 7 ). pow ( z ). val () == mod_pow ( 7 , a - b , mod ));
}
if ( a >= mod ) {
mint z = x + y ;
z . unsafe_subtract ( y );
assert ( z . val () == a % mod );
assert ( mint ( 7 ). pow ( z ). val () == mod_pow ( 7 , a , mod ));
}
}
for ( uint64_t a = 0 ; a <= 5 ; a ++ )
for ( uint64_t b = 0 ; b <= 5 ; b ++ )
for ( uint64_t c = 0 ; c <= 8 ; c ++ ) {
uint64_t exponent = pow_u64 ( b , c );
assert ( mint ( a ). pow ( mint ( b ). pow ( mint ( c ))). val () == mod_pow ( a , exponent , mod ));
}
for ( uint64_t x = 0 ; x <= min < uint64_t > ( mod - 1 , 100 ); x ++ )
for ( uint64_t y = 0 ; y <= x ; y ++ ) {
mint difference = x ;
difference . unsafe_subtract ( int64_t ( y ));
assert ( difference . val () == x - y );
assert ( difference . large_val () == x - y );
assert ( mint ( 7 ). pow ( difference ). val () == mod_pow ( 7 , x - y , mod ));
difference = mint ( x );
difference . unsafe_subtract ( mint ( y ));
assert ( difference . val () == x - y );
assert ( difference . large_val () == x - y );
assert ( mint ( 7 ). pow ( difference ). val () == mod_pow ( 7 , x - y , mod ));
}
mint large_difference = mint ( uint64_t ( mod ) + 123 );
large_difference . unsafe_subtract ( int64_t ( 123 ));
assert ( large_difference . val () == 0 );
assert ( large_difference . large_val () == mod );
assert ( mint ( 7 ). pow ( large_difference ). val () == mod_pow ( 7 , mod , mod ));
assert ( mint ( 0 ). pow ( mint ( 0 )). val () == 1 % mod );
assert ( mint ( 0 ). pow ( mint ( 0 ). pow ( mint ( 0 ))). val () == 0 );
assert ( mint ( 0 ). pow ( mint ( 0 ). pow ( mint ( 1 ))). val () == 1 % mod );
assert ( mint ( 0 ) == mint ( mod ));
mint huge = mint ( 2 ). pow ( mint ( 100 ));
mint difference = huge + mint ( mod );
difference . unsafe_subtract ( huge );
assert ( difference . val () == 0 );
assert ( difference . large_val () == mod );
assert ( mint ( 7 ). pow ( difference ). val () == mod_pow ( 7 , mod , mod ));
stringstream ss ;
ss << mint ( numeric_limits < uint64_t >:: max ());
assert ( ss . str () == to_string ( numeric_limits < uint64_t >:: max () % mod ));
}
int main () {
test < 1 > ();
test < 2 > ();
test < 3 > ();
test < 4 > ();
test < 6 > ();
test < 10 > ();
test < 998244353 > ();
test < 2147483647 > ();
using mint = PowerTowerModInt < 10 > ;
assert ( mint ( 2 ). pow ( mint ( 1 )). val () == 2 );
assert ( mint ( 2 ). pow ( mint ( 4 )). val () == 6 );
assert ( mint ( 2 ). pow ( mint ( 8 )). val () == 6 );
int a , b ;
in ( a , b );
out ( a + b );
}