Modular Kth Root
(math/modular_kth_root.hpp)
- View this file on GitHub
- Last update: 2026-07-16 21:30:39+09:00
- Include:
#include "math/modular_kth_root.hpp"
Overview
modular_kth_root solves
for a prime modulus $p$. It returns any valid root, or std::nullopt when no
root exists. The implementation reduces the exponent by
$\gcd(k,p-1)$ and extracts its prime-power factors with a generalized
Tonelli–Shanks algorithm.
Interface
std::optional<uint64_t> modular_kth_root(
uint64_t value,
uint64_t degree,
uint64_t prime
);
template <class Mint>
std::optional<Mint> modular_kth_root(
Mint value,
uint64_t degree
);
| Function | Description |
|---|---|
modular_kth_root(value, degree, prime) |
Reduces value modulo prime and returns any kth root. |
modular_kth_root(value, degree) |
Uses Mint::mod() and returns the root as Mint. |
prime must be at least 2 and prime; primality is not tested. The modular
integer overload requires val(), static mod(), and construction from an
integer.
For degree == 0, exponentiation follows the convention $x^0=1$, including
$0^0=1$. The function therefore returns 0 when value is congruent to 1
and no result otherwise. For positive degree, zero has root zero.
The answer is not canonical when several roots exist.
Complexity
Let $g=\gcd(k,p-1)$. Factoring $g$ by trial division takes $O(\sqrt g)$ integer operations. For a prime factor $q^e$ of $g$, write $p-1=mq^s$. Its generalized Tonelli–Shanks correction uses $O(\sqrt{(s-e)q})$ stored group elements and $s-e$ correction rounds, plus a search for a q-th non-residue. Every modular exponentiation costs $O(\log p)$ modular multiplications.
This is designed for the common competitive-programming regime of many prime moduli up to about $10^9$ and avoids a full discrete logarithm per query.
Example
#include "math/modular_kth_root.hpp"
#include <cassert>
int main() {
auto root = m1une::math::modular_kth_root(8, 3, 13);
assert(root.has_value());
assert(*root * *root % 13 * *root % 13 == 8);
assert(!m1une::math::modular_kth_root(3, 2, 7).has_value());
}
Depends on
Required by
Verified with
Code
#ifndef M1UNE_MATH_MODULAR_KTH_ROOT_HPP
#define M1UNE_MATH_MODULAR_KTH_ROOT_HPP 1
#include <algorithm>
#include <cassert>
#include <cstdint>
#include <numeric>
#include <optional>
#include <utility>
#include <vector>
#include "integer_arithmetic.hpp"
namespace m1une {
namespace math {
namespace modular_kth_root_detail {
inline uint64_t multiply(uint64_t first, uint64_t second, uint64_t mod) {
return static_cast<uint64_t>(static_cast<__uint128_t>(first) * second % mod);
}
inline uint64_t power(uint64_t base, uint64_t exponent, uint64_t mod) {
uint64_t result = 1 % mod;
while (exponent != 0) {
if (exponent & 1) result = multiply(result, base, mod);
base = multiply(base, base, mod);
exponent >>= 1;
}
return result;
}
inline uint64_t integer_power(uint64_t base, int exponent) {
uint64_t result = 1;
for (int i = 0; i < exponent; i++) result *= base;
return result;
}
inline uint64_t inverse(uint64_t value, uint64_t mod) {
if (mod == 1) return 0;
value %= mod;
uint64_t old_remainder = mod;
uint64_t remainder = value;
__int128_t old_coefficient = 0;
__int128_t coefficient = 1;
while (remainder != 0) {
const uint64_t quotient = old_remainder / remainder;
const uint64_t next_remainder =
old_remainder - quotient * remainder;
old_remainder = remainder;
remainder = next_remainder;
const __int128_t next_coefficient =
old_coefficient - static_cast<__int128_t>(quotient) * coefficient;
old_coefficient = coefficient;
coefficient = next_coefficient;
}
assert(old_remainder == 1);
old_coefficient %= static_cast<__int128_t>(mod);
if (old_coefficient < 0) old_coefficient += mod;
return static_cast<uint64_t>(old_coefficient);
}
inline uint64_t extract_prime_power_root(
uint64_t value,
uint64_t root_prime,
int exponent,
uint64_t prime
) {
uint64_t coprime_part = prime - 1;
int available_exponent = 0;
while (coprime_part % root_prime == 0) {
coprime_part /= root_prime;
available_exponent++;
}
assert(exponent <= available_exponent);
const uint64_t root_prime_power = integer_power(root_prime, exponent);
const uint64_t inverse_coprime_part = inverse(
coprime_part, root_prime_power
);
const uint64_t residue = static_cast<uint64_t>(
static_cast<__uint128_t>(root_prime_power - 1) *
inverse_coprime_part % root_prime_power
);
const uint64_t root_exponent = static_cast<uint64_t>(
(static_cast<__uint128_t>(residue) * coprime_part + 1) /
root_prime_power
);
uint64_t root = power(value, root_exponent, prime);
if (exponent == available_exponent) return root;
uint64_t non_residue = 2;
while (power(non_residue, (prime - 1) / root_prime, prime) == 1) {
non_residue++;
}
const uint64_t generator = power(non_residue, coprime_part, prime);
const uint64_t digit_generator = power(
generator,
integer_power(root_prime, available_exponent - 1),
prime
);
const uint64_t step = isqrt(
static_cast<uint64_t>(available_exponent - exponent) * root_prime
) + 1;
const uint64_t giant_factor = power(digit_generator, step, prime);
std::vector<std::pair<uint64_t, uint64_t>> baby_steps;
baby_steps.reserve(step + 1);
uint64_t current = 1;
for (uint64_t index = 0; index <= step; index++) {
baby_steps.emplace_back(current, index);
current = multiply(current, giant_factor, prime);
}
std::sort(baby_steps.begin(), baby_steps.end());
const uint64_t inverse_digit_generator = power(
digit_generator, prime - 2, prime
);
for (int level = exponent; level < available_exponent; level++) {
const uint64_t root_power = power(root, root_prime_power, prime);
const uint64_t error = multiply(
power(root_power, prime - 2, prime), value, prime
);
uint64_t target = power(
error,
integer_power(root_prime, available_exponent - 1 - level),
prime
);
bool found = false;
uint64_t logarithm = 0;
for (uint64_t remainder = 0; remainder <= step; remainder++) {
auto iterator = std::upper_bound(
baby_steps.begin(),
baby_steps.end(),
target,
[](uint64_t key, const std::pair<uint64_t, uint64_t>& entry) {
return key < entry.first;
}
);
if (iterator != baby_steps.begin()) {
--iterator;
if (iterator->first == target) {
logarithm = remainder + step * iterator->second;
found = true;
break;
}
}
target = multiply(target, inverse_digit_generator, prime);
}
assert(found);
if (!found) return 0;
const uint64_t correction_exponent =
logarithm * integer_power(root_prime, level - exponent);
root = multiply(
root,
power(generator, correction_exponent, prime),
prime
);
}
return root;
}
} // namespace modular_kth_root_detail
// Returns x such that x^degree = value (mod prime), or nullopt when no root
// exists. The modulus must be prime.
inline std::optional<uint64_t> modular_kth_root(
uint64_t value,
uint64_t degree,
uint64_t prime
) {
assert(prime >= 2);
value %= prime;
if (degree == 0) {
if (value == 1) return uint64_t(0);
return std::nullopt;
}
if (value == 0) return uint64_t(0);
if (prime == 2) return uint64_t(1);
const uint64_t group_order = prime - 1;
degree %= group_order;
const uint64_t common_divisor = std::gcd(degree, group_order);
if (
modular_kth_root_detail::power(
value, group_order / common_divisor, prime
) != 1
) {
return std::nullopt;
}
const uint64_t reduced_order = group_order / common_divisor;
uint64_t transformed = 1;
if (reduced_order != 1) {
const uint64_t inverse_degree = modular_kth_root_detail::inverse(
degree / common_divisor, reduced_order
);
transformed = modular_kth_root_detail::power(
value, inverse_degree, prime
);
}
uint64_t remaining = common_divisor;
int exponent = 0;
while ((remaining & 1) == 0) {
remaining >>= 1;
exponent++;
}
if (exponent != 0) {
transformed = modular_kth_root_detail::extract_prime_power_root(
transformed, 2, exponent, prime
);
}
for (uint64_t divisor = 3; divisor <= remaining / divisor; divisor += 2) {
exponent = 0;
while (remaining % divisor == 0) {
remaining /= divisor;
exponent++;
}
if (exponent != 0) {
transformed = modular_kth_root_detail::extract_prime_power_root(
transformed, divisor, exponent, prime
);
}
}
if (remaining != 1) {
transformed = modular_kth_root_detail::extract_prime_power_root(
transformed, remaining, 1, prime
);
}
return transformed;
}
template <class Mint>
std::optional<Mint> modular_kth_root(Mint value, uint64_t degree) {
auto root = modular_kth_root(
static_cast<uint64_t>(value.val()),
degree,
static_cast<uint64_t>(Mint::mod())
);
if (!root.has_value()) return std::nullopt;
return Mint(*root);
}
} // namespace math
} // namespace m1une
#endif // M1UNE_MATH_MODULAR_KTH_ROOT_HPP#line 1 "math/modular_kth_root.hpp"
#include <algorithm>
#include <cassert>
#include <cstdint>
#include <numeric>
#include <optional>
#include <utility>
#include <vector>
#line 1 "math/integer_arithmetic.hpp"
#line 5 "math/integer_arithmetic.hpp"
#include <concepts>
#include <limits>
#line 8 "math/integer_arithmetic.hpp"
#include <type_traits>
namespace m1une {
namespace math {
namespace integer_arithmetic_detail {
template <std::integral T>
requires(!std::same_as<std::remove_cv_t<T>, bool>)
constexpr std::optional<T> checked_multiply(T first, T second) {
constexpr T minimum = std::numeric_limits<T>::min();
constexpr T maximum = std::numeric_limits<T>::max();
if constexpr (std::unsigned_integral<T>) {
if (second != 0 && maximum / second < first) return std::nullopt;
} else {
if (0 < first) {
if (0 < second) {
if (maximum / second < first) return std::nullopt;
} else if (second < minimum / first) {
return std::nullopt;
}
} else if (first < 0) {
if (0 < second) {
if (first < minimum / second) return std::nullopt;
} else if (second < maximum / first) {
return std::nullopt;
}
}
}
return T(first * second);
}
template <std::unsigned_integral T>
constexpr bool kth_power_leq(T base, unsigned exponent, T limit) {
assert(exponent > 0);
if (base <= 1) return base <= limit;
const T multiplication_limit = limit / base;
T product = 1;
for (unsigned i = 0; i < exponent; i++) {
if (product > multiplication_limit) return false;
product *= base;
}
return true;
}
} // namespace integer_arithmetic_detail
// Returns floor(sqrt(value)) exactly, without floating-point arithmetic.
template <std::integral T>
requires(!std::same_as<std::remove_cv_t<T>, bool>)
constexpr T isqrt(T value) {
if constexpr (std::signed_integral<T>) assert(0 <= value);
if (value <= 1) return value;
T low = 1;
T high = value / 2 + 1;
while (low < high) {
T middle = low + (high - low + 1) / 2;
if (middle <= value / middle) {
low = middle;
} else {
high = middle - 1;
}
}
return low;
}
template <std::integral T>
requires(!std::same_as<std::remove_cv_t<T>, bool>)
constexpr T floor_sqrt(T value) {
return isqrt(value);
}
// Returns ceil(sqrt(value)) exactly, without floating-point arithmetic.
template <std::integral T>
requires(!std::same_as<std::remove_cv_t<T>, bool>)
constexpr T ceil_sqrt(T value) {
T result = isqrt(value);
if (result == 0) return 0;
if (result != 0 && value / result == result && value % result == 0) {
return result;
}
return result + 1;
}
// Returns floor(value^(1 / degree)) exactly, without floating-point arithmetic.
template <std::integral T, std::integral Degree>
requires(
!std::same_as<std::remove_cv_t<T>, bool>
&& !std::same_as<std::remove_cv_t<Degree>, bool>
)
constexpr T floor_kth_root(T value, Degree degree) {
if constexpr (std::signed_integral<T>) {
assert(0 <= value);
if (value < 0) return T();
}
assert(0 < degree);
if (degree <= 0) return T();
if (value <= 1 || degree == 1) return value;
if (degree == 2) return isqrt(value);
using U = std::make_unsigned_t<T>;
using UDegree = std::make_unsigned_t<Degree>;
constexpr int digits = std::numeric_limits<U>::digits;
const UDegree unsigned_degree = static_cast<UDegree>(degree);
if (unsigned_degree >= static_cast<UDegree>(digits)) return T(1);
const unsigned exponent = static_cast<unsigned>(unsigned_degree);
const U unsigned_value = static_cast<U>(value);
int bit_width = 0;
for (U remaining = unsigned_value; remaining != 0; remaining >>= 1) {
bit_width++;
}
const int root_bits =
(bit_width + static_cast<int>(exponent) - 1) /
static_cast<int>(exponent);
U low = 1;
U high = U(1) << root_bits;
while (high - low > 1) {
const U middle = low + (high - low) / 2;
if (
integer_arithmetic_detail::kth_power_leq(
middle, exponent, unsigned_value
)
) {
low = middle;
} else {
high = middle;
}
}
return static_cast<T>(low);
}
// Returns base^exponent, or nullopt when the result does not fit in T.
template <std::integral T, std::unsigned_integral Exponent>
requires(
!std::same_as<std::remove_cv_t<T>, bool>
&& !std::same_as<std::remove_cv_t<Exponent>, bool>
)
constexpr std::optional<T> checked_ipow(T base, Exponent exponent) {
T result = 1;
while (exponent != 0) {
if (exponent & 1) {
auto product =
integer_arithmetic_detail::checked_multiply(result, base);
if (!product.has_value()) return std::nullopt;
result = *product;
}
exponent >>= 1;
if (exponent != 0) {
auto square =
integer_arithmetic_detail::checked_multiply(base, base);
if (!square.has_value()) return std::nullopt;
base = *square;
}
}
return result;
}
template <std::integral T, std::unsigned_integral Exponent>
requires(
!std::same_as<std::remove_cv_t<T>, bool>
&& !std::same_as<std::remove_cv_t<Exponent>, bool>
)
constexpr std::optional<T> checked_integer_pow(T base, Exponent exponent) {
return checked_ipow(base, exponent);
}
// Returns base^exponent. The result must be representable by T.
template <std::integral T, std::unsigned_integral Exponent>
requires(
!std::same_as<std::remove_cv_t<T>, bool>
&& !std::same_as<std::remove_cv_t<Exponent>, bool>
)
constexpr T ipow(T base, Exponent exponent) {
std::optional<T> result = checked_ipow(base, exponent);
assert(result.has_value());
return result.value_or(T());
}
template <std::integral T, std::unsigned_integral Exponent>
requires(
!std::same_as<std::remove_cv_t<T>, bool>
&& !std::same_as<std::remove_cv_t<Exponent>, bool>
)
constexpr T integer_pow(T base, Exponent exponent) {
return ipow(base, exponent);
}
} // namespace math
} // namespace m1une
#line 13 "math/modular_kth_root.hpp"
namespace m1une {
namespace math {
namespace modular_kth_root_detail {
inline uint64_t multiply(uint64_t first, uint64_t second, uint64_t mod) {
return static_cast<uint64_t>(static_cast<__uint128_t>(first) * second % mod);
}
inline uint64_t power(uint64_t base, uint64_t exponent, uint64_t mod) {
uint64_t result = 1 % mod;
while (exponent != 0) {
if (exponent & 1) result = multiply(result, base, mod);
base = multiply(base, base, mod);
exponent >>= 1;
}
return result;
}
inline uint64_t integer_power(uint64_t base, int exponent) {
uint64_t result = 1;
for (int i = 0; i < exponent; i++) result *= base;
return result;
}
inline uint64_t inverse(uint64_t value, uint64_t mod) {
if (mod == 1) return 0;
value %= mod;
uint64_t old_remainder = mod;
uint64_t remainder = value;
__int128_t old_coefficient = 0;
__int128_t coefficient = 1;
while (remainder != 0) {
const uint64_t quotient = old_remainder / remainder;
const uint64_t next_remainder =
old_remainder - quotient * remainder;
old_remainder = remainder;
remainder = next_remainder;
const __int128_t next_coefficient =
old_coefficient - static_cast<__int128_t>(quotient) * coefficient;
old_coefficient = coefficient;
coefficient = next_coefficient;
}
assert(old_remainder == 1);
old_coefficient %= static_cast<__int128_t>(mod);
if (old_coefficient < 0) old_coefficient += mod;
return static_cast<uint64_t>(old_coefficient);
}
inline uint64_t extract_prime_power_root(
uint64_t value,
uint64_t root_prime,
int exponent,
uint64_t prime
) {
uint64_t coprime_part = prime - 1;
int available_exponent = 0;
while (coprime_part % root_prime == 0) {
coprime_part /= root_prime;
available_exponent++;
}
assert(exponent <= available_exponent);
const uint64_t root_prime_power = integer_power(root_prime, exponent);
const uint64_t inverse_coprime_part = inverse(
coprime_part, root_prime_power
);
const uint64_t residue = static_cast<uint64_t>(
static_cast<__uint128_t>(root_prime_power - 1) *
inverse_coprime_part % root_prime_power
);
const uint64_t root_exponent = static_cast<uint64_t>(
(static_cast<__uint128_t>(residue) * coprime_part + 1) /
root_prime_power
);
uint64_t root = power(value, root_exponent, prime);
if (exponent == available_exponent) return root;
uint64_t non_residue = 2;
while (power(non_residue, (prime - 1) / root_prime, prime) == 1) {
non_residue++;
}
const uint64_t generator = power(non_residue, coprime_part, prime);
const uint64_t digit_generator = power(
generator,
integer_power(root_prime, available_exponent - 1),
prime
);
const uint64_t step = isqrt(
static_cast<uint64_t>(available_exponent - exponent) * root_prime
) + 1;
const uint64_t giant_factor = power(digit_generator, step, prime);
std::vector<std::pair<uint64_t, uint64_t>> baby_steps;
baby_steps.reserve(step + 1);
uint64_t current = 1;
for (uint64_t index = 0; index <= step; index++) {
baby_steps.emplace_back(current, index);
current = multiply(current, giant_factor, prime);
}
std::sort(baby_steps.begin(), baby_steps.end());
const uint64_t inverse_digit_generator = power(
digit_generator, prime - 2, prime
);
for (int level = exponent; level < available_exponent; level++) {
const uint64_t root_power = power(root, root_prime_power, prime);
const uint64_t error = multiply(
power(root_power, prime - 2, prime), value, prime
);
uint64_t target = power(
error,
integer_power(root_prime, available_exponent - 1 - level),
prime
);
bool found = false;
uint64_t logarithm = 0;
for (uint64_t remainder = 0; remainder <= step; remainder++) {
auto iterator = std::upper_bound(
baby_steps.begin(),
baby_steps.end(),
target,
[](uint64_t key, const std::pair<uint64_t, uint64_t>& entry) {
return key < entry.first;
}
);
if (iterator != baby_steps.begin()) {
--iterator;
if (iterator->first == target) {
logarithm = remainder + step * iterator->second;
found = true;
break;
}
}
target = multiply(target, inverse_digit_generator, prime);
}
assert(found);
if (!found) return 0;
const uint64_t correction_exponent =
logarithm * integer_power(root_prime, level - exponent);
root = multiply(
root,
power(generator, correction_exponent, prime),
prime
);
}
return root;
}
} // namespace modular_kth_root_detail
// Returns x such that x^degree = value (mod prime), or nullopt when no root
// exists. The modulus must be prime.
inline std::optional<uint64_t> modular_kth_root(
uint64_t value,
uint64_t degree,
uint64_t prime
) {
assert(prime >= 2);
value %= prime;
if (degree == 0) {
if (value == 1) return uint64_t(0);
return std::nullopt;
}
if (value == 0) return uint64_t(0);
if (prime == 2) return uint64_t(1);
const uint64_t group_order = prime - 1;
degree %= group_order;
const uint64_t common_divisor = std::gcd(degree, group_order);
if (
modular_kth_root_detail::power(
value, group_order / common_divisor, prime
) != 1
) {
return std::nullopt;
}
const uint64_t reduced_order = group_order / common_divisor;
uint64_t transformed = 1;
if (reduced_order != 1) {
const uint64_t inverse_degree = modular_kth_root_detail::inverse(
degree / common_divisor, reduced_order
);
transformed = modular_kth_root_detail::power(
value, inverse_degree, prime
);
}
uint64_t remaining = common_divisor;
int exponent = 0;
while ((remaining & 1) == 0) {
remaining >>= 1;
exponent++;
}
if (exponent != 0) {
transformed = modular_kth_root_detail::extract_prime_power_root(
transformed, 2, exponent, prime
);
}
for (uint64_t divisor = 3; divisor <= remaining / divisor; divisor += 2) {
exponent = 0;
while (remaining % divisor == 0) {
remaining /= divisor;
exponent++;
}
if (exponent != 0) {
transformed = modular_kth_root_detail::extract_prime_power_root(
transformed, divisor, exponent, prime
);
}
}
if (remaining != 1) {
transformed = modular_kth_root_detail::extract_prime_power_root(
transformed, remaining, 1, prime
);
}
return transformed;
}
template <class Mint>
std::optional<Mint> modular_kth_root(Mint value, uint64_t degree) {
auto root = modular_kth_root(
static_cast<uint64_t>(value.val()),
degree,
static_cast<uint64_t>(Mint::mod())
);
if (!root.has_value()) return std::nullopt;
return Mint(*root);
}
} // namespace math
} // namespace m1une