Aleph-w 3.0
A C++ Library for Data Structures and Algorithms
Loading...
Searching...
No Matches
modular_arithmetic.H
Go to the documentation of this file.
1/*
2 Aleph_w
3
4 Data structures & Algorithms
5 version 2.0.0b
6 https://github.com/lrleon/Aleph-w
7
8 This file is part of Aleph-w library
9
10 Copyright (c) 2002-2026 Leandro Rabindranath Leon
11
12 Permission is hereby granted, free of charge, to any person obtaining a copy
13 of this software and associated documentation files (the "Software"), to deal
14 in the Software without restriction, including without limitation the rights
15 to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
16 copies of the Software, and to permit persons to whom the Software is
17 furnished to do so, subject to the following conditions:
18
19 The above copyright notice and this permission notice shall be included in all
20 copies or substantial portions of the Software.
21
22 THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
23 IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
24 FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
25 AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
26 LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
27 OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
28 SOFTWARE.
29*/
30
42# ifndef MODULAR_ARITHMETIC_H
43# define MODULAR_ARITHMETIC_H
44
45# include <cstdint>
46# include <type_traits>
47# ifdef _WIN32
48# include <intrin.h>
49# endif
50
51# include <ah-errors.H>
52# include <tpl_array.H>
53
54namespace Aleph
55{
68 {
69 ah_invalid_argument_if(m == 0) << "mod_mul: modulus must be > 0";
70
71# if defined(__SIZEOF_INT128__) && !defined(_WIN32)
72 return static_cast<uint64_t>((static_cast<__uint128_t>(a) * b) % m);
73# elif defined(_MSC_VER) && !defined(__clang__) && (defined(_M_X64) || defined(_M_ARM64))
74 // MSVC cl (64-bit): no __int128, but _umul128 + _udiv128 are available
75 unsigned long long hi;
76 unsigned long long lo = _umul128(a, b, &hi);
77 unsigned long long rem;
78 _udiv128(hi, lo, m, &rem);
79 return static_cast<uint64_t>(rem);
80# else
81 // clang-cl: has __uint128_t but no __umodti3 runtime and no _udiv128.
82 // MSVC 32-bit: no _umul128/_udiv128. Fall through to binary method.
83 uint64_t res = 0;
84 a %= m;
85 while (b)
86 {
87 if (b & 1)
88 {
89 if (m - res <= a) res = a - (m - res);
90 else res += a;
91 }
92 if (m - a <= a) a = a - (m - a);
93 else a <<= 1;
94 b >>= 1;
95 }
96 return res;
97# endif
98 }
99
111 {
112 ah_invalid_argument_if(m == 0) << "mod_exp: modulus must be > 0";
113 if (m == 1) return 0;
114 uint64_t res = 1;
115 base %= m;
116 while (exp > 0)
117 {
118 if (exp & 1)
119 res = mod_mul(res, base, m);
120 base = mod_mul(base, base, m);
121 exp >>= 1;
122 }
123 return res;
124 }
125
136 template <typename T>
137 requires (std::is_integral_v<T> and std::is_signed_v<T>)
138 [[nodiscard]] T ext_gcd(T a, T b, T & x, T & y) noexcept
139 {
140 if (b == 0)
141 {
142 x = 1;
143 y = 0;
144 return a;
145 }
146 T x1, y1;
147 T d = ext_gcd(b, a % b, x1, y1);
148 x = y1;
149 y = x1 - static_cast<T>(a / b) * y1;
150 return d;
151 }
152
169 [[nodiscard]] inline uint64_t mod_inv(const uint64_t a, const uint64_t m)
170 {
171 ah_domain_error_if(m == 0) << "mod_inv: modulus cannot be 0";
172
173 if (m == 1)
174 return 0;
175
176 const uint64_t a_mod = a % m;
178 << "Modular inverse does not exist (" << a << " is 0 mod " << m << ")";
179
180 // Iterative extended Euclidean algorithm using unsigned arithmetic.
181 // We track t0 such that a * t0 ≡ r0 (mod m) at each step,
182 // keeping t0 in [0, m) to avoid signed overflow.
183 uint64_t r0 = a_mod, r1 = m;
184 uint64_t t0 = 1, t1 = 0;
185
186 while (r1 != 0)
187 {
188 const uint64_t q = r0 / r1;
189
190 const uint64_t tmp_r = r0 - q * r1;
191 r0 = r1;
192 r1 = tmp_r;
193
194 // t_new = t0 - q * t1 (mod m), computed in unsigned
195 const uint64_t qt1 = mod_mul(q, t1, m);
196 const uint64_t tmp_t = (t0 >= qt1) ? (t0 - qt1) : (m - (qt1 - t0));
197 t0 = t1;
198 t1 = tmp_t;
199 }
200
202 << "Modular inverse does not exist (numbers " << a
203 << " and " << m << " are not coprime)";
204
205 return t0;
206 }
207
208# if defined(__SIZEOF_INT128__) && !defined(_WIN32)
210 struct MontgomeryCtx;
211
212 namespace detail
213 {
215 [[nodiscard]] constexpr MontgomeryCtx
216 montgomery_ctx_unchecked(const uint64_t mod) noexcept;
217 }
218
234 struct MontgomeryCtx
235 {
236 public:
237 MontgomeryCtx() = delete;
238
242 [[nodiscard]] constexpr uint64_t mod() const noexcept { return mod_; }
243
247 [[nodiscard]] constexpr uint64_t mod2() const noexcept { return mod2_; }
248
252 [[nodiscard]] constexpr uint64_t r() const noexcept { return r_; }
253
257 [[nodiscard]] constexpr uint64_t r2() const noexcept { return r2_; }
258
263 {
264 return mod_inv_neg_;
265 }
266
267 private:
268 friend constexpr MontgomeryCtx
269 detail::montgomery_ctx_unchecked(const uint64_t mod) noexcept;
270
272 constexpr MontgomeryCtx(const uint64_t mod,
273 const uint64_t mod2,
274 const uint64_t r,
275 const uint64_t r2,
277 : mod_(mod), mod2_(mod2), r_(r), r2_(r2), mod_inv_neg_(mod_inv_neg)
278 {
279 /* empty */
280 }
281
282 uint64_t mod_;
284 uint64_t r_;
285 uint64_t r2_;
287 };
288
289 namespace detail
290 {
291 [[nodiscard]] constexpr uint64_t
292 montgomery_neg_inverse(const uint64_t mod) noexcept
293 {
294 uint64_t inv = 1;
295 for (size_t i = 0; i < 6; ++i)
296 inv *= 2 - mod * inv;
297 return ~inv + 1;
298 }
299
300 [[nodiscard]] constexpr MontgomeryCtx
301 montgomery_ctx_unchecked(const uint64_t mod) noexcept
302 {
303 const __uint128_t r128 = static_cast<__uint128_t>(1) << 64;
304 const auto r = static_cast<uint64_t>(r128 % mod);
305 const auto r2 = static_cast<uint64_t>((static_cast<__uint128_t>(r) * r) % mod);
306
307 return MontgomeryCtx(mod,
308 mod <= UINT64_MAX - mod ? mod + mod : 0,
309 r,
310 r2,
312 }
313 }
314
321 [[nodiscard]] inline MontgomeryCtx
322 montgomery_ctx(const uint64_t mod)
323 {
324 ah_invalid_argument_if(mod <= 1)
325 << "montgomery_ctx: modulus must be > 1";
326 ah_invalid_argument_if((mod & 1ULL) == 0)
327 << "montgomery_ctx: modulus " << mod << " must be odd";
328 return detail::montgomery_ctx_unchecked(mod);
329 }
330
336 template <uint64_t Mod>
337 [[nodiscard]] consteval MontgomeryCtx
339 {
340 static_assert(Mod > 1, "montgomery_ctx_for_mod: modulus must be > 1");
341 static_assert((Mod & 1ULL) == 1ULL,
342 "montgomery_ctx_for_mod: modulus must be odd");
343 return detail::montgomery_ctx_unchecked(Mod);
344 }
345
369 [[nodiscard]] constexpr uint64_t
370 mont_reduce(const __uint128_t x,
371 const MontgomeryCtx & ctx) noexcept
372 {
373 const uint64_t q = static_cast<uint64_t>(x) * ctx.mod_inv_neg();
374 const auto x_lo = static_cast<uint64_t>(x);
375 const auto x_hi = static_cast<uint64_t>(x >> 64);
376 const __uint128_t qmod = static_cast<__uint128_t>(q) * ctx.mod();
377 const auto qmod_lo = static_cast<uint64_t>(qmod);
378 const auto qmod_hi = static_cast<uint64_t>(qmod >> 64);
379 const uint64_t carry = qmod_lo > UINT64_MAX - x_lo ? 1ULL : 0ULL;
380
381 // t = (x + q·p) / R; REDC guarantees 0 ≤ t < 2p.
382 const __uint128_t t = static_cast<__uint128_t>(x_hi) + qmod_hi + carry;
383
384 // Fast path: 2p < 2^64 so t fits in uint64_t; one conditional subtraction
385 // reduces to [0, p) using only additions and comparisons.
386 if (ctx.mod2() != 0)
387 {
388 const uint64_t t64 = static_cast<uint64_t>(t);
389 return t64 >= ctx.mod() ? t64 - ctx.mod() : t64;
390 }
391
392 // Fallback for large primes (p ≥ 2^63): t may exceed UINT64_MAX, so keep
393 // it in __uint128_t and use a single 128-bit division to finish.
394 return static_cast<uint64_t>(t % ctx.mod());
395 }
396
406 [[nodiscard]] constexpr uint64_t
407 mont_mul(const uint64_t a,
408 const uint64_t b,
409 const MontgomeryCtx & ctx) noexcept
410 {
411 return mont_reduce(static_cast<__uint128_t>(a) * b, ctx);
412 }
413
420 [[nodiscard]] constexpr uint64_t
421 to_mont(const uint64_t a,
422 const MontgomeryCtx & ctx) noexcept
423 {
424 return mont_mul(a % ctx.mod(), ctx.r2(), ctx);
425 }
426
433 [[nodiscard]] constexpr uint64_t
434 from_mont(const uint64_t a,
435 const MontgomeryCtx & ctx) noexcept
436 {
437 return mont_reduce(a, ctx);
438 }
439
449 [[nodiscard]] constexpr uint64_t
450 mont_exp(uint64_t base,
452 const MontgomeryCtx & ctx) noexcept
453 {
454 uint64_t result = to_mont(1, ctx);
455 while (exp > 0)
456 {
457 if (exp & 1ULL)
458 result = mont_mul(result, base, ctx);
459 base = mont_mul(base, base, ctx);
460 exp >>= 1;
461 }
462 return result;
463 }
464# endif
465
479 const Array<uint64_t> & mod)
480 {
481 ah_invalid_argument_if(rem.size() != mod.size())
482 << "crt: arrays must have the same size (got " << rem.size()
483 << " vs " << mod.size() << ")";
484
485 const size_t n = rem.size();
486 if (n == 0)
487 return 0;
488
489 // Compute product of all moduli with overflow detection
490 uint64_t prod = 1;
491 for (size_t i = 0; i < n; ++i)
492 {
493 ah_invalid_argument_if(mod[i] <= 1)
494 << "crt: all moduli must be > 1 (got " << mod[i] << " at index " << i << ")";
495
497 << "crt: product of moduli overflows uint64_t at index " << i;
498 prod *= mod[i];
499 }
500
501 uint64_t result = 0;
502 for (size_t i = 0; i < n; ++i)
503 {
504 const uint64_t p = prod / mod[i];
505 const uint64_t inv = mod_inv(p, mod[i]);
506 uint64_t term = mod_mul(rem[i], p, prod);
507 term = mod_mul(term, inv, prod);
508 if (result >= prod - term)
509 result -= (prod - term);
510 else
511 result += term;
512 }
513
514 return result;
515 }
516} // namespace Aleph
517
518# endif // MODULAR_ARITHMETIC_H
Exception handling system with formatted messages for Aleph-w.
#define ah_overflow_error_if(C)
Throws std::overflow_error if condition holds.
Definition ah-errors.H:468
#define ah_domain_error_if(C)
Throws std::domain_error if condition holds.
Definition ah-errors.H:527
#define ah_invalid_argument_if(C)
Throws std::invalid_argument if condition holds.
Definition ah-errors.H:644
Simple dynamic array with automatic resizing and functional operations.
Definition tpl_array.H:138
__gmp_expr< T, __gmp_unary_expr< __gmp_expr< T, U >, __gmp_y1_function > > y1(const __gmp_expr< T, U > &expr)
Definition gmpfrxx.h:4114
__gmp_expr< T, __gmp_unary_expr< __gmp_expr< T, U >, __gmp_exp_function > > exp(const __gmp_expr< T, U > &expr)
Definition gmpfrxx.h:4077
size_t blossom_maximum_cardinality_matching(const GT &g, DynDlist< typename GT::Arc * > &matching, SA sa=SA())
Alias of compute_maximum_cardinality_general_matching().
Definition Blossom.H:466
static mpfr_t y
Definition mpfr_mul_d.c:3
Main namespace for Aleph-w library functions.
Definition ah-arena.H:89
uint64_t mod_inv(const uint64_t a, const uint64_t m)
Modular Inverse.
T ext_gcd(T a, T b, T &x, T &y) noexcept
Extended Euclidean Algorithm.
and
Check uniqueness with explicit hash + equality functors.
std::decay_t< typename HeadC::Item_Type > T
Definition ah-zip.H:105
uint64_t mod_exp(uint64_t base, uint64_t exp, const uint64_t m)
Modular exponentiation.
uint64_t mod_mul(uint64_t a, uint64_t b, uint64_t m)
Safe 64-bit modular multiplication.
uint64_t crt(const Array< uint64_t > &rem, const Array< uint64_t > &mod)
Chinese Remainder Theorem (CRT).
double mod(double a, double b)
FooMap m(5, fst_unit_pair_hash, snd_unit_pair_hash)
gsl_rng * r
Dynamic array container with automatic resizing.