m1une's library

This documentation is automatically generated by online-judge-tools/verification-helper

View on GitHub

:heavy_check_mark: Modular Kth Root
(math/modular_kth_root.hpp)

Overview

modular_kth_root solves

\[x^k \equiv a \pmod p\]

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
Back to top page