src/math/natural.bend source
src/math/natural.bend on the hub · documented module
import Baseimport ./pow2.bend as P2# Exact integer functions of a Python-style math library on the natural# numbers (Python's non-negative int), after the design reference's# section 4.2 (math module integer functions), 3.1 (divmod, three-argument# pow) and 10 (ilog, iroot, clamp, mod_inverse).## gcd(a, b), gcd_all(xs) greatest common divisor; gcd_all([]) == 0# lcm(a, b), lcm_all(xs) least common multiple; lcm_all([]) == 1# isqrt(n) floor(sqrt(n)), exact for every n# iroot(n, k) floor of the k-th root (k >= 1)# ilog(n, b) floor(log_b(n)), exact (n >= 1, b >= 2)# factorial(n), perm(n, k) n!, n! / (n - k)! (0 when k > n)# comb(n, k) n choose k (0 when k > n)# prod(xs), sum(xs) product and sum of a list; prod([]) == 1# pow_mod(b, e, m) b^e mod m by binary exponentiation# mod_inverse(a, m) x with a*x == 1 (mod m), x < m# divmod(a, b) (a // b, a % b)# bit_length(n) bits of n without leading zeros# clamp(x, lo, hi) x limited to [lo, hi]## Errors follow the reference's error model (section 2.3) as values, not# exceptions: ZeroDivision where Python raises ZeroDivisionError or the# "cannot be 0" ValueError of pow, Domain for the other ValueErrors, and# NotInvertible for pow(a, -1, m) without an inverse.## Nat is exact; Bend's native backend stops the program on a Nat above# 2^48 - 1, so results (and the intermediate products of comb, lcm and# pow_mod) must stay below that at run time. The proofs are over the# mathematical naturals.type MathError is Data: ZeroDivision{} Domain{} NotInvertible{}# the last nonzero remainder of the extended Euclid loop and its# coefficient: a * coef == gcd (mod m)type Bezout is Data: BZ{gcd: Nat, coef: Nat}# a // b and a % btype QuotRem is Data: QR{quot: Nat, rem: Nat}# ---- greatest common divisor, least common multiple ----# Euclid's algorithm; b <= fuel throughout (every step lowers b).def gcd_go(fuel: Nat, +a: Nat, +b: Nat) -> Nat: match fuel b: case 0n _: a case 1n+f 0n: a case 1n+f 1n+ +bp: gcd_go(f, 1n+bp, Nat.mod(a, 1n+bp))def gcd(+a: Nat, +b: Nat) -> Nat: gcd_go(b, a, b)def lcm(+a: Nat, +b: Nat) -> Nat: match a b: case 0n _: 0n case 1n+ap 0n: 0n case 1n+ +ap 1n+ +bp: Nat.mul(Nat.div(1n+ap, gcd(1n+ap, 1n+bp)), 1n+bp)def gcd_all_go(xs: List<&2, Nat>, +acc: Nat) -> Nat: match xs: case Nil{}: acc case Con{+x, rest}: gcd_all_go(rest, gcd(acc, x))def gcd_all(xs: List<&2, Nat>) -> Nat: gcd_all_go(xs, 0n)def lcm_all_go(xs: List<&2, Nat>, +acc: Nat) -> Nat: match xs: case Nil{}: acc case Con{+x, rest}: lcm_all_go(rest, lcm(acc, x))def lcm_all(xs: List<&2, Nat>) -> Nat: lcm_all_go(xs, 1n)# ---- sums and products ----def sum_go(xs: List<&2, Nat>, +acc: Nat) -> Nat: match xs: case Nil{}: acc case Con{+x, rest}: sum_go(rest, Nat.add(acc, x))def sum(xs: List<&2, Nat>) -> Nat: sum_go(xs, 0n)def prod_go(xs: List<&2, Nat>, +acc: Nat) -> Nat: match xs: case Nil{}: acc case Con{+x, rest}: prod_go(rest, Nat.mul(acc, x))def prod(xs: List<&2, Nat>) -> Nat: prod_go(xs, 1n)# ---- factorial, permutations, combinations ----def factorial_go(n: Nat, +acc: Nat) -> Nat: match n: case 0n: acc case 1n+ +p: factorial_go(p, Nat.mul(1n+p, acc))def factorial(+n: Nat) -> Nat: factorial_go(n, 1n)# n * (n - 1) * ... * (n - k + 1)def perm_go(k: Nat, +m: Nat, +acc: Nat) -> Nat: match k: case 0n: acc case 1n+j: perm_go(j, Nat.sub(m, 1n), Nat.mul(acc, m))def perm(+n: Nat, +k: Nat) -> Nat: perm_go(k, n, 1n)# r_{i+1} = r_i * (n - i) / (i + 1): every r_i is C(n, i), so the division# is exact (reference section 8.6).def comb_go(j: Nat, +n: Nat, +i: Nat, +r: Nat) -> Nat: match j: case 0n: r case 1n+ +jp: comb_go(jp, n, 1n+i, Nat.div(Nat.mul(r, Nat.sub(n, i)), 1n+i))def comb_small(+n: Nat, +k: Nat, +big: Bool) -> Nat: match big: case True{}: 0n case False{}: comb_go(Nat.min(k, Nat.sub(n, k)), n, 0n, 1n)def comb(+n: Nat, +k: Nat) -> Nat: comb_small(n, k, Nat.is_lt(n, k))# ---- bits ----def bit_length_go(fuel: Nat, +n: Nat, +k: Nat) -> Nat: match fuel n: case 0n _: k case 1n+f 0n: k case 1n+f 1n+ +np: bit_length_go(f, Nat.div(1n+np, 2n), 1n+k)def bit_length(+n: Nat) -> Nat: bit_length_go(n, n, 0n)# ---- roots and logarithms ----# Binary search for the last r in [lo, hi) where the test holds, given it# holds at lo and fails at hi (reference section 8.5 notes the float-sqrt# shortcut is wrong above 2^52; this search is exact). `more` is lo + 1 < hi# and `ok` is the test at the midpoint; the fuel bounds the halvings.def mid(+lo: Nat, +hi: Nat) -> Nat: Nat.div(Nat.add(lo, hi), 2n)# r^k <= n without forming a product above n: `ok` is acc * r <= n, asked# as acc <= n / r, and acc * r is only formed once it is known to fit.def pow_le_go(k: Nat, +r: Nat, +acc: Nat, +n: Nat, +ok: Bool) -> Bool: match k ok: case 0n _: True{} case 1n+j False{}: False{} case 1n+j True{}: pow_le_go(j, r, Nat.mul(acc, r), n, Nat.is_le(Nat.mul(acc, r), Nat.div(n, r)))def root_le(+n: Nat, +k: Nat, +r: Nat) -> Bool: match k r: case 0n _: Nat.is_le(1n, n) case 1n+kp 0n: True{} case 1n+ +kp 1n+ +rp: pow_le_go(1n+kp, 1n+rp, 1n, n, Nat.is_le(1n, Nat.div(n, 1n+rp)))# the test is r^k <= ndef search_go(fuel: Nat, +n: Nat, +k: Nat, +lo: Nat, +hi: Nat, +more: Bool, +ok: Bool) -> Nat: match fuel more ok: case 0n _ _: lo case 1n+f False{} _: lo case 1n+f True{} True{}: search_go(f, n, k, mid(lo, hi), hi, Nat.is_lt(1n+mid(lo, hi), hi), root_le(n, k, mid(mid(lo, hi), hi))) case 1n+f True{} False{}: search_go(f, n, k, lo, mid(lo, hi), Nat.is_lt(1n+lo, mid(lo, hi)), root_le(n, k, mid(lo, mid(lo, hi))))def search(+n: Nat, +k: Nat, +lo: Nat, +hi: Nat) -> Nat: search_go(Nat.sub(hi, lo), n, k, lo, hi, Nat.is_lt(1n+lo, hi), root_le(n, k, mid(lo, hi)))# the largest r with r * r <= n, by Heron's (Newton's) iteration as in Lean 4# Mathlib's Nat.sqrt: g := (g + n / g) / 2 while that decreases g. Started# from 2^(b/2 + 1) >= sqrt(n) (n < 2^b, b = bit_length(n)) it converges in# O(log b) steps; the fuel g bounds the (strictly decreasing) steps.def sqrt_next(+n: Nat, +g: Nat) -> Nat: Nat.div(Nat.add(g, Nat.div(n, g)), 2n)def sqrt_iter(fuel: Nat, +n: Nat, +g: Nat, +m: Nat, +down: Bool) -> Nat: match fuel down: case 0n _: g case 1n+f False{}: g case 1n+f True{}: sqrt_iter(f, n, m, sqrt_next(n, m), Nat.is_lt(sqrt_next(n, m), m))def isqrt_from(+n: Nat, +g: Nat) -> Nat: sqrt_iter(g, n, g, sqrt_next(n, g), Nat.is_lt(sqrt_next(n, g), g))def isqrt(+n: Nat) -> Nat: isqrt_from(n, P2.pow2t(1n+Nat.div(bit_length(n), 2n)))# the largest r with r^k <= n (k >= 1)def iroot_k(+n: Nat, +k: Nat) -> Nat: match k: case 0n: 0n case 1n: n case 2n+ +kp: search(n, 2n+kp, 0n, P2.pow2t(1n+Nat.div(bit_length(n), 2n+kp)))def iroot(+n: Nat, +k: Nat) -> Result<&2, &2, MathError, Nat>: match k: case 0n: Fail{Domain{}} case 1n+ +kp: Done{iroot_k(n, 1n+kp)}# the largest k with b^k <= n: p is b^k, `up` is p * b <= n, asked as# p <= n / b so no product above n is formeddef ilog_go(fuel: Nat, +n: Nat, +b: Nat, +k: Nat, +p: Nat, +up: Bool) -> Nat: match fuel up: case 0n _: k case 1n+f False{}: k case 1n+f True{}: ilog_go(f, n, b, 1n+k, Nat.mul(p, b), Nat.is_le(Nat.mul(p, b), Nat.div(n, b)))def ilog_ok(+n: Nat, +b: Nat, +bad: Bool) -> Result<&2, &2, MathError, Nat>: match bad: case True{}: Fail{Domain{}} case False{}: Done{ilog_go(n, n, b, 0n, 1n, Nat.is_le(1n, Nat.div(n, b)))}def ilog(+n: Nat, +b: Nat) -> Result<&2, &2, MathError, Nat>: ilog_ok(n, b, Bool.or(Nat.is_eq(n, 0n), Nat.is_lt(b, 2n)))# ---- modular arithmetic ----# right-to-left binary exponentiation, reducing mod m after every product:# acc * base^e == b^e0 (mod m) throughoutdef pow_mod_odd(+m: Nat, +bit: Nat, +base: Nat, +acc: Nat) -> Nat: match bit: case 0n: acc case 1n+z: Nat.mod(Nat.mul(acc, base), m)def pow_mod_go(fuel: Nat, +m: Nat, +e: Nat, +base: Nat, +acc: Nat) -> Nat: match fuel e: case 0n _: acc case 1n+f 0n: acc case 1n+f 1n+ +ep: pow_mod_go(f, m, Nat.div(1n+ep, 2n), Nat.mod(Nat.mul(base, base), m), pow_mod_odd(m, Nat.mod(1n+ep, 2n), base, acc))def pow_mod(+b: Nat, +e: Nat, +m: Nat) -> Result<&2, &2, MathError, Nat>: match m: case 0n: Fail{ZeroDivision{}} case 1n+ +mp: Done{pow_mod_go(e, 1n+mp, e, Nat.mod(b, 1n+mp), Nat.mod(1n, 1n+mp))}# Extended Euclid on (m, a mod m) keeping only the coefficient of a, mod m:# a * s0 == r0 and a * s1 == r1 (mod m). The remainders are gcd's.def inv_step(+m: Nat, +q: Nat, +s0: Nat, +s1: Nat) -> Nat: Nat.mod(Nat.add(s0, Nat.sub(m, Nat.mod(Nat.mul(q, s1), m))), m)def inv_go(fuel: Nat, +m: Nat, +r0: Nat, +s0: Nat, +r1: Nat, +s1: Nat) -> Bezout: match fuel r1: case 0n _: BZ{r0, s0} case 1n+f 0n: BZ{r0, s0} case 1n+f 1n+ +rp: inv_go(f, m, 1n+rp, s1, Nat.mod(r0, 1n+rp), inv_step(m, Nat.div(r0, 1n+rp), s0, s1))def inv_fin(+m: Nat, bz: Bezout) -> Result<&2, &2, MathError, Nat>: match bz: case BZ{1n, s}: Done{Nat.mod(s, m)} case BZ{0n, s}: Fail{NotInvertible{}} case BZ{2n+g, s}: Fail{NotInvertible{}}# pow(a, -1, m): the x < m with a * x == 1 (mod m)def mod_inverse(+a: Nat, +m: Nat) -> Result<&2, &2, MathError, Nat>: match m: case 0n: Fail{ZeroDivision{}} case 1n+ +mp: inv_fin(1n+mp, inv_go(1n+mp, 1n+mp, 1n+mp, 0n, Nat.mod(a, 1n+mp), 1n))# ---- division and clamping ----def divmod(+a: Nat, +b: Nat) -> Result<&2, &2, MathError, QuotRem>: match b: case 0n: Fail{ZeroDivision{}} case 1n+ +bp: Done{QR{Nat.div(a, 1n+bp), Nat.mod(a, 1n+bp)}}def clamp_ok(+x: Nat, +lo: Nat, +hi: Nat, +bad: Bool) -> Result<&2, &2, MathError, Nat>: match bad: case True{}: Fail{Domain{}} case False{}: Done{Nat.min(Nat.max(x, lo), hi)}def clamp(+x: Nat, +lo: Nat, +hi: Nat) -> Result<&2, &2, MathError, Nat>: clamp_ok(x, lo, hi, Nat.is_lt(hi, lo))