Power Tower Modint
(modint/power-tower-modint.hpp)
- View this file on GitHub
- Last update: 2026-07-28 20:47:27+09:00
- Include:
#include "modint/power-tower-modint.hpp"
非負整数の加算,乗算,累乗からなる式を法 $m$ で計算する.
通常の modint と異なり PowerTowerModInt<m> 自身を指数に指定できる.法と底が互いに素でない場合や,0 の累乗も扱える.$0^0=1$ とする.
$1\leq m\lt 2^{31}$ である必要がある.通常の減算,除算,負数は扱わない.
使い方
-
PowerTowerModInt<m>(x):非負の 64 bit 整数 $x$ から構築する. -
get_mod():法 $m$ を返す. -
val():式の値を $m$ で割った余りを返す. -
large_val():式の真の値を $x$ として,$x\lt m$ ならば $x$,そうでなければ $m+(x\bmod m)$ を返す.返り値は $[0,2m)$ に属する. -
operator+,operator*:式を加算または乗算する. -
unsafe_subtract(rhs):自身とrhsが表す真の値をそれぞれ $x,y$ として,$x\gets x-y$ とする.$x\lt m$ ならば $0\leq y\leq x$,そうでなければ $x-y\geq m$ を仮定する. -
unsafe_subtract(int64_t y):自身が表す真の値を $x$ として,$x\gets x-y$ とする.$x\lt m$ ならば $0\leq y\leq x$,そうでなければ $x-y\geq m$ を仮定する. -
pow(e):自身をe乗する. -
operator==,operator!=:val()が等しいかを判定する.
unsafe_subtract で結果が $m$ 以上となる場合の前提条件は内部表現から検査できない.val() の大小ではなく,加算,乗算,累乗からなる元の式の真の値について保証する必要がある.
仕組み
$m,\varphi(m),\varphi(\varphi(m)),\dots,1$ のそれぞれを法とする値を保持する.各法 $p$ では,非負整数 $x$ を
\[\begin{cases} x & (x\lt p),\\ p+(x\bmod p) & (x\geq p) \end{cases}\]で表す.これにより,指数が周期へ入る前の小さい値と周期へ入った値を区別する.
計算量
オブジェクトの空間計算量は $O(\log m)$.構築,加算,乗算,unsafe_subtract は $O(\log m)$ 時間,pow は $O((\log m)^2)$ 時間,val() と large_val() は $O(1)$ 時間.
資料
Verified with
Code
#pragma once
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 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
*/