Modular Square Root
(math/modular_square_root.hpp)
- View this file on GitHub
- Last update: 2026-07-11 19:26:27+09:00
- Include:
#include "math/modular_square_root.hpp"
Overview
modular_square_root solves
for a prime modulus p. It uses the Tonelli-Shanks algorithm and returns no
value when a is a quadratic non-residue.
The runtime-modulus overload is useful when each query has a different prime.
The one-argument overload works with modular integer types that provide
val(), static mod(), and construction from an integer, including
m1une::math::ModInt.
API
std::optional<uint64_t> modular_square_root(
uint64_t value,
uint64_t prime
);
template <class Mint>
std::optional<Mint> modular_square_root(Mint value);
| Function | Description | Complexity |
|---|---|---|
modular_square_root(value, prime) |
Returns a square root modulo the prime, or std::nullopt if none exists. |
$O(\log^2 p)$ modular multiplications |
modular_square_root(value) |
Uses Mint::mod() and returns the result as Mint. |
$O(\log^2 p)$ modular multiplications |
prime must be at least 2 and prime. The function does not perform a
primality test. value is reduced modulo prime, and either square root may be
returned. In particular, the result is not guaranteed to be the smaller root.
Multiplication uses an unsigned 128-bit intermediate, so the runtime overload
supports the full uint64_t range of prime moduli.
Example
#include "math/modint.hpp"
#include "math/modular_square_root.hpp"
#include <iostream>
int main() {
auto root = m1une::math::modular_square_root(10, 13);
if (root.has_value()) {
std::cout << *root << "\n"; // 6 or 7
}
using Mint = m1une::math::modint998244353;
auto mint_root = m1une::math::modular_square_root(Mint(4));
std::cout << mint_root->val() << "\n"; // 2 or 998244351
}
Required by
Graph All
(graph/all.hpp)
Graph Counting
(graph/counting.hpp)
Math All
(math/all.hpp)
Math All
(math/all.hpp)
Bernoulli Numbers and Power Sums
(math/bernoulli.hpp)
Combinatorial Sequences
(math/combinatorial_sequences.hpp)
Formal Power Series All
(math/fps/all.hpp)
Formal Power Series Composition
(math/fps/composition.hpp)
Compositional Inverse of Formal Power Series
(math/fps/compositional_inverse.hpp)
Formal Power Series
(math/fps/formal_power_series.hpp)
Geometric-Sequence Polynomial Evaluation and Interpolation
(math/fps/geometric_sequence_evaluation.hpp)
Polynomial Half-GCD
(math/fps/half_gcd.hpp)
Lagrange Inversion Formula
(math/fps/lagrange_inversion.hpp)
Linear Recurrences and Bostan-Mori
(math/fps/linear_recurrence.hpp)
Multipoint Evaluation and Interpolation
(math/fps/multipoint_evaluation.hpp)
Polynomial Factorization
(math/fps/polynomial_factorization.hpp)
Polynomial Roots over a Finite Field
(math/fps/polynomial_roots.hpp)
Solve Formal Power Series Equation
(math/fps/solve_fps_equation.hpp)
Sparse Formal Power Series
(math/fps/sparse_formal_power_series.hpp)
Newton Method
(math/newton_method.hpp)
Partition Function
(math/partition_function.hpp)
Verified with
verify/graph/cow_game.test.cpp
verify/graph/graph_algorithms.test.cpp
verify/graph/graph_counting.test.cpp
verify/graph/range_edge_graph.test.cpp
verify/math/bell_number.test.cpp
verify/math/bernoulli_number.test.cpp
verify/math/bernoulli_utilities.test.cpp
verify/math/fps/composition.test.cpp
verify/math/fps/compositional_inverse.test.cpp
verify/math/fps/exp_of_formal_power_series.test.cpp
verify/math/fps/exp_of_formal_power_series_sparse.test.cpp
verify/math/fps/find_linear_recurrence.test.cpp
verify/math/fps/fps_algorithms.test.cpp
verify/math/fps/half_gcd.test.cpp
verify/math/fps/inv_of_formal_power_series.test.cpp
verify/math/fps/inv_of_formal_power_series_sparse.test.cpp
verify/math/fps/kth_term_of_linearly_recurrent_sequence.test.cpp
verify/math/fps/lagrange_inversion.test.cpp
verify/math/fps/log_of_formal_power_series.test.cpp
verify/math/fps/log_of_formal_power_series_sparse.test.cpp
verify/math/fps/multipoint_evaluation.test.cpp
verify/math/fps/multipoint_evaluation_geometric.test.cpp
verify/math/fps/polynomial_factorization.test.cpp
verify/math/fps/polynomial_interpolation.test.cpp
verify/math/fps/polynomial_interpolation_geometric.test.cpp
verify/math/fps/polynomial_roots.test.cpp
verify/math/fps/polynomial_taylor_shift.test.cpp
verify/math/fps/pow_of_formal_power_series.test.cpp
verify/math/fps/pow_of_formal_power_series_sparse.test.cpp
verify/math/fps/sqrt_of_formal_power_series.test.cpp
verify/math/fps/sqrt_of_formal_power_series_sparse.test.cpp
verify/math/math_algorithms.test.cpp
verify/math/math_algorithms.test.cpp
verify/math/modular_square_root.test.cpp
verify/math/newton_method.test.cpp
verify/math/partition_function.test.cpp
verify/math/stirling_number_of_the_second_kind.test.cpp
Code
#ifndef M1UNE_MATH_MODULAR_SQUARE_ROOT_HPP
#define M1UNE_MATH_MODULAR_SQUARE_ROOT_HPP 1
#include <cassert>
#include <cstdint>
#include <optional>
namespace m1une {
namespace math {
namespace internal {
inline uint64_t modular_square_root_multiply(uint64_t lhs, uint64_t rhs, uint64_t mod) {
return static_cast<uint64_t>(static_cast<unsigned __int128>(lhs) * rhs % mod);
}
inline uint64_t modular_square_root_power(uint64_t base, uint64_t exponent, uint64_t mod) {
uint64_t result = 1 % mod;
while (exponent > 0) {
if (exponent & 1) result = modular_square_root_multiply(result, base, mod);
base = modular_square_root_multiply(base, base, mod);
exponent >>= 1;
}
return result;
}
} // namespace internal
// Returns x such that x * x = value (mod prime), or nullopt when no such x exists.
// The modulus must be prime.
inline std::optional<uint64_t> modular_square_root(uint64_t value, uint64_t prime) {
assert(prime >= 2);
value %= prime;
if (value == 0 || prime == 2) return value;
if (internal::modular_square_root_power(value, (prime - 1) / 2, prime) != 1) {
return std::nullopt;
}
if (prime % 4 == 3) {
return internal::modular_square_root_power(value, prime / 4 + 1, prime);
}
uint64_t odd_part = prime - 1;
int power_of_two = 0;
while ((odd_part & 1) == 0) {
odd_part >>= 1;
power_of_two++;
}
uint64_t non_residue = 2;
while (internal::modular_square_root_power(non_residue, (prime - 1) / 2, prime) == 1) {
non_residue++;
}
uint64_t c = internal::modular_square_root_power(non_residue, odd_part, prime);
uint64_t root = internal::modular_square_root_power(value, odd_part / 2 + 1, prime);
uint64_t remainder = internal::modular_square_root_power(value, odd_part, prime);
int remaining_power = power_of_two;
while (remainder != 1) {
int exponent = 1;
uint64_t squared = internal::modular_square_root_multiply(remainder, remainder, prime);
while (squared != 1) {
squared = internal::modular_square_root_multiply(squared, squared, prime);
exponent++;
}
uint64_t correction = c;
for (int i = 0; i < remaining_power - exponent - 1; i++) {
correction = internal::modular_square_root_multiply(correction, correction, prime);
}
root = internal::modular_square_root_multiply(root, correction, prime);
c = internal::modular_square_root_multiply(correction, correction, prime);
remainder = internal::modular_square_root_multiply(remainder, c, prime);
remaining_power = exponent;
}
return root;
}
template <class Mint>
std::optional<Mint> modular_square_root(Mint value) {
auto root = modular_square_root(static_cast<uint64_t>(value.val()),
static_cast<uint64_t>(Mint::mod()));
if (!root.has_value()) return std::nullopt;
return Mint(*root);
}
} // namespace math
} // namespace m1une
#endif // M1UNE_MATH_MODULAR_SQUARE_ROOT_HPP#line 1 "math/modular_square_root.hpp"
#include <cassert>
#include <cstdint>
#include <optional>
namespace m1une {
namespace math {
namespace internal {
inline uint64_t modular_square_root_multiply(uint64_t lhs, uint64_t rhs, uint64_t mod) {
return static_cast<uint64_t>(static_cast<unsigned __int128>(lhs) * rhs % mod);
}
inline uint64_t modular_square_root_power(uint64_t base, uint64_t exponent, uint64_t mod) {
uint64_t result = 1 % mod;
while (exponent > 0) {
if (exponent & 1) result = modular_square_root_multiply(result, base, mod);
base = modular_square_root_multiply(base, base, mod);
exponent >>= 1;
}
return result;
}
} // namespace internal
// Returns x such that x * x = value (mod prime), or nullopt when no such x exists.
// The modulus must be prime.
inline std::optional<uint64_t> modular_square_root(uint64_t value, uint64_t prime) {
assert(prime >= 2);
value %= prime;
if (value == 0 || prime == 2) return value;
if (internal::modular_square_root_power(value, (prime - 1) / 2, prime) != 1) {
return std::nullopt;
}
if (prime % 4 == 3) {
return internal::modular_square_root_power(value, prime / 4 + 1, prime);
}
uint64_t odd_part = prime - 1;
int power_of_two = 0;
while ((odd_part & 1) == 0) {
odd_part >>= 1;
power_of_two++;
}
uint64_t non_residue = 2;
while (internal::modular_square_root_power(non_residue, (prime - 1) / 2, prime) == 1) {
non_residue++;
}
uint64_t c = internal::modular_square_root_power(non_residue, odd_part, prime);
uint64_t root = internal::modular_square_root_power(value, odd_part / 2 + 1, prime);
uint64_t remainder = internal::modular_square_root_power(value, odd_part, prime);
int remaining_power = power_of_two;
while (remainder != 1) {
int exponent = 1;
uint64_t squared = internal::modular_square_root_multiply(remainder, remainder, prime);
while (squared != 1) {
squared = internal::modular_square_root_multiply(squared, squared, prime);
exponent++;
}
uint64_t correction = c;
for (int i = 0; i < remaining_power - exponent - 1; i++) {
correction = internal::modular_square_root_multiply(correction, correction, prime);
}
root = internal::modular_square_root_multiply(root, correction, prime);
c = internal::modular_square_root_multiply(correction, correction, prime);
remainder = internal::modular_square_root_multiply(remainder, c, prime);
remaining_power = exponent;
}
return root;
}
template <class Mint>
std::optional<Mint> modular_square_root(Mint value) {
auto root = modular_square_root(static_cast<uint64_t>(value.val()),
static_cast<uint64_t>(Mint::mod()));
if (!root.has_value()) return std::nullopt;
return Mint(*root);
}
} // namespace math
} // namespace m1une