src/math/random/rand.bend source
src/math/random/rand.bend on the hub · documented module
import Baseimport ../u64.bend as Wimport ../w64.bend as Ximport ../f64.bend as F# Random values from any source, written once: Go's math/rand/v2 Rand# methods, bit for bit (the 64-bit code paths).## The Source interface. A source is any state type S with a step## ~next: S -> W.U64 & S the next 64-bit output and the new state## passed as templates (~S, ~next), as src/math/num.bend passes a numeric# type's ~op and ~test: each source compiles to its own copy with next# inlined, and the state is threaded explicitly (Bend is pure: every# function returns its value next to the advanced state). The sources here# are chacha8.bend's ChaCha8 and pcg.bend's PCG:## import ./src/math/random.bend as R# import ./src/math/random/chacha8.bend as C8# (k, g) = R.uint64n(~C8.ChaCha8, ~R.chacha8_next, g, n)## A later module (math/statistics: normal, exponential, ...) is written the# same way, once for every source, on top of these functions:## def normal(~S: Data, ~next: S -> W.U64 & S, s: S) -> F.F64 & S:# ... R.float64(~S, ~next, s) ...## and its laws can be stated and proved for an arbitrary ~next (a template# proof is checked once, for every source; proofs/math/random/rand.bend).## uint64(s) next(s) Go Uint64# uint32(s) the top 32 bits Go Uint32# int64(s) the low 63 bits Go Int64# int32(s) the top 31 bits Go Int32# uint64n(s, n) uniform in [0, n) for n > 0 Go Uint64N# (uint_below is the same function; n == 0 is read# as 2^64, Go's internal uint64n(0): a full uint64)# uint32n(s, n) uniform in [0, n), n: U32 > 0 Go Uint32N# intn(s, n) uniform in [0, n), n: Nat, Go IntN# 0 < n < 2^48 (a Nat's run-time bound); 0 for n == 0# int_range(s, lo, hi) uniform in [lo, hi), lo < hi lo + IntN(hi - lo)# float64(s) uniform multiple of 2^-53 in [0, 1) Go Float64# shuffle(s, xs) a Fisher-Yates permutation of xs Go Shuffle# perm(s, n) a random permutation of 0..n-1 Go Perm## uint64n is Lemire's nearly divisionless method ("Fast Random Integer# Generation in an Interval", ACM TOMACS 2019), exactly as Go's uint64n:# a power of two n masks the low bits; otherwise the 128-bit product x * n# is split into hi:lo, and x is rejected while lo < 2^64 mod n (computed,# with a division, only when lo < n). The accepted x map onto every k < n# equally often (proved: proofs/math/random/lemire.bend), so the output is# exactly uniform for a uniform source. Go retries forever; here at most# 128 draws are made (Bend requires termination): a uniform source rejects# with probability below 1/2 per draw, so the 128th rejection has# probability below 2^-128, and even then the result is below n.def fst(-A: Data, -S: Data, p: A & S) -> A: (a, s) = p adef snd(-A: Data, -S: Data, p: A & S) -> S: (a, s) = p sdef uint64(~S: Data, ~next: S -> W.U64 & S, s: S) -> W.U64 & S: next(s)def top32(-S: Data, p: W.U64 & S) -> U32 & S: (x, s) = p (X.hi(x), s)def uint32(~S: Data, ~next: S -> W.U64 & S, s: S) -> U32 & S: top32(S, next(s))def low63(-S: Data, p: W.U64 & S) -> W.U64 & S: (+x, s) = p (W.U64{X.lo(x), U32.and(X.hi(x), 2147483647)}, s)def int64(~S: Data, ~next: S -> W.U64 & S, s: S) -> W.U64 & S: low63(S, next(s))def top31(-S: Data, p: W.U64 & S) -> U32 & S: (x, s) = p (U32.shr(X.hi(x)), s)def int32(~S: Data, ~next: S -> W.U64 & S, s: S) -> U32 & S: top31(S, next(s))# ---- bounded integers: Lemire ----def and64(+a: W.U64, +b: W.U64) -> W.U64: W.U64{U32.and(X.lo(a), X.lo(b)), U32.and(X.hi(a), X.hi(b))}# n & (n - 1) == 0: n is a power of two (or zero)def is_pow2(+n: W.U64) -> Bool: X.is_zero(and64(n, X.sub(n, W.U64{1, 0})))def mask(-S: Data, +n: W.U64, p: W.U64 & S) -> W.U64 & S: (+x, s) = p (and64(x, X.sub(n, W.U64{1, 0})), s)# 2^64 mod n = (2^64 - n) mod n, for n > 0 (Go: -n % n)def thresh(+n: W.U64) -> W.U64: X.rem(X.sub(W.U64{0, 0}, n), n)# the draw x as the product x * n = hi:lo, kept with the statetype Draw<-S: Data> is Data: D{hi: W.U64, lo: W.U64, s: S}def draw_fin(-S: Data, s: S, p: W.U64 & W.U64) -> Draw<S>: (lo, hi) = p D{hi, lo, s}def draw(-S: Data, +n: W.U64, p: W.U64 & S) -> Draw<S>: (+x, s) = p draw_fin(S, s, X.mul128(x, n))# one rejection step: while lo < t, draw again; an accepted draw is keptdef again(~S: Data, ~next: S -> W.U64 & S, +n: W.U64, +hi: W.U64, +lo: W.U64, s: S, reject: Bool) -> Draw<S>: match reject: case False{}: D{hi, lo, s} case True{}: draw(S, n, next(s))def step(~S: Data, ~next: S -> W.U64 & S, +n: W.U64, +t: W.U64, d: Draw<S>) -> Draw<S>: match d: case D{+hi, +lo, s}: again(~S, ~next, n, hi, lo, s, X.lt(lo, t))# fuel rejection steps (an accepted draw is a fixed point of step)def retry(~S: Data, ~next: S -> W.U64 & S, fuel: Nat, +n: W.U64, +t: W.U64, d: Draw<S>) -> Draw<S>: match fuel: case 0n: d case 1n+f: retry(~S, ~next, f, n, t, step(~S, ~next, n, t, d))def result(-S: Data, d: Draw<S>) -> W.U64 & S: match d: case D{hi, lo, s}: (hi, s)# lo >= n >= t accepts at once; otherwise t = 2^64 mod n is computed and at# most 127 more draws are madedef lemire_small(~S: Data, ~next: S -> W.U64 & S, +n: W.U64, +hi: W.U64, +lo: W.U64, s: S, small: Bool) -> W.U64 & S: match small: case False{}: (hi, s) case True{}: result(S, retry(~S, ~next, 127n, n, thresh(n), D{hi, lo, s}))def lemire_first(~S: Data, ~next: S -> W.U64 & S, +n: W.U64, d: Draw<S>) -> W.U64 & S: match d: case D{+hi, +lo, s}: lemire_small(~S, ~next, n, hi, lo, s, X.lt(lo, n))def uint64n_pick(~S: Data, ~next: S -> W.U64 & S, s: S, +n: W.U64, pow2: Bool) -> W.U64 & S: match pow2: case True{}: mask(S, n, next(s)) case False{}: lemire_first(~S, ~next, n, draw(S, n, next(s)))# uniform in [0, n) for n > 0 (Go's Uint64N)def uint64n(~S: Data, ~next: S -> W.U64 & S, s: S, +n: W.U64) -> W.U64 & S: uint64n_pick(~S, ~next, s, n, is_pow2(n))def uint_below(~S: Data, ~next: S -> W.U64 & S, s: S, +n: W.U64) -> W.U64 & S: uint64n(~S, ~next, s, n)def lo32(-S: Data, p: W.U64 & S) -> U32 & S: (x, s) = p (X.lo(x), s)# uniform in [0, n) for n > 0 (Go's Uint32N)def uint32n(~S: Data, ~next: S -> W.U64 & S, s: S, +n: U32) -> U32 & S: lo32(S, uint64n(~S, ~next, s, W.U64{n, 0}))# the 32-bit word of the low 8k bits of n, a byte at a timedef w_go(k: Nat, +n: Nat) -> U32: match k: case 0n: 0 case 1n+j: U32.add(U32.from_nat(Nat.mod(n, 256n)), U32.mul(w_go(j, Nat.div(n, 256n)), X.pow2(8n)))# n div 2^(8k)def skip(k: Nat, +n: Nat) -> Nat: match k: case 0n: n case 1n+j: skip(j, Nat.div(n, 256n))# the 64-bit word of n < 2^64def nat64(+n: Nat) -> W.U64: W.U64{w_go(4n, n), w_go(4n, skip(4n, n))}def to_nat(-S: Data, p: W.U64 & S) -> Nat & S: (x, s) = p (X.n48(x), s)def intn_z(~S: Data, ~next: S -> W.U64 & S, s: S, +n: Nat, z: Bool) -> Nat & S: match z: case True{}: (0n, s) case False{}: to_nat(S, uint64n(~S, ~next, s, nat64(n)))# uniform in [0, n) for 0 < n < 2^48 (Go's IntN); (0, s) for n == 0def intn(~S: Data, ~next: S -> W.U64 & S, s: S, +n: Nat) -> Nat & S: intn_z(~S, ~next, s, n, Nat.is_eq(n, 0n))def plus(-S: Data, +lo: Nat, p: Nat & S) -> Nat & S: (k, s) = p (Nat.add(lo, k), s)# uniform in [lo, hi) for lo < hi, hi - lo < 2^48; (lo, s) when hi <= lodef int_range(~S: Data, ~next: S -> W.U64 & S, s: S, +lo: Nat, +hi: Nat) -> Nat & S: plus(S, lo, intn(~S, ~next, s, Nat.sub(hi, lo)))# ---- floats ----# x >> 11 as a double divided by 2^53: for m = x >> 11 > 0 with c leading# zeros (11 <= c <= 63), m << (c - 11) has its top bit at 52, and the value# m 2^-53 is that significand with biased exponent 1033 - c (below 1023)def f53(+m: W.U64, +c: Nat) -> F.F64: F.Bits{X.lo(X.shl(m, Nat.sub(c, 11n))), U32.add(U32.mul(U32.from_nat(Nat.sub(1033n, c)), 1048576), U32.and(X.hi(X.shl(m, Nat.sub(c, 11n))), 1048575))}def f53_z(+m: W.U64, z: Bool) -> F.F64: match z: case True{}: F.zero(False{}) case False{}: f53(m, X.clz(m))# Go's Float64: float64(x << 11 >> 11) / (1 << 53), the low 53 bitsdef to_float(+x: W.U64) -> F.F64: f53_z(W.U64{X.lo(x), U32.and(X.hi(x), 2097151)}, X.is_zero(W.U64{X.lo(x), U32.and(X.hi(x), 2097151)}))def float_of(-S: Data, p: W.U64 & S) -> F.F64 & S: (+x, s) = p (to_float(x), s)def float64(~S: Data, ~next: S -> W.U64 & S, s: S) -> F.F64 & S: float_of(S, next(s))# ---- permutations ----def nth(-A: Data, xs: List<&2, A>, i: Nat) -> Maybe<&2, A>: match xs i: case Nil{} _: None{} case Con{h, t} 0n: Some{h} case Con{h, t} 1n+p: nth(A, t, p)def put(-A: Data, xs: List<&2, A>, i: Nat, x: A) -> List<&2, A>: match xs i: case Nil{} _: Nil{} case Con{h, t} 0n: Con{x, t} case Con{h, t} 1n+p: Con{h, put(A, t, p, x)}def swap_m(-A: Data, xs: List<&2, A>, +i: Nat, +j: Nat, a: Maybe<&2, A>, b: Maybe<&2, A>) -> List<&2, A>: match a b: case Some{x} Some{y}: put(A, put(A, xs, i, y), j, x) case _ _: xs# exchange elements i and j (unchanged when either is out of range)def swap(-A: Data, +xs: List<&2, A>, +i: Nat, +j: Nat) -> List<&2, A>: swap_m(A, xs, i, j, nth(A, xs, i), nth(A, xs, j))def swap_at(-A: Data, -S: Data, xs: List<&2, A>, +i: Nat, p: W.U64 & S) -> List<&2, A> & S: (+j, s) = p (swap(A, xs, i, X.n48(j)), s)# the length as a 64-bit worddef len64(-A: Data, xs: List<&2, A>) -> W.U64: match xs: case Nil{}: W.U64{0, 0} case Con{x, rest}: X.add(len64(A, rest), W.U64{1, 0})# swap element i with a uniform j < n = i + 1def shuffle_step(~A: Data, ~S: Data, ~next: S -> W.U64 & S, +i: Nat, +n: W.U64, st: List<&2, A> & S) -> List<&2, A> & S: (xs, s) = st swap_at(A, S, xs, i, uint64n(~S, ~next, s, n))# i = k, k - 1, ..., 1 with the bound n = i + 1 as a 64-bit worddef shuffle_go(~A: Data, ~S: Data, ~next: S -> W.U64 & S, k: Nat, +n: W.U64, st: List<&2, A> & S) -> List<&2, A> & S: match k: case 0n: st case 1n+ +p: shuffle_go(~A, ~S, ~next, p, X.sub(n, W.U64{1, 0}), shuffle_step(~A, ~S, ~next, 1n+p, n, st))# Go's Shuffle (Fisher-Yates, from the last index down), for fewer than 2^48# elementsdef shuffle(~A: Data, ~S: Data, ~next: S -> W.U64 & S, s: S, +xs: List<&2, A>) -> List<&2, A> & S: shuffle_go(~A, ~S, ~next, Nat.sub(List.length(&2, A, xs), 1n), len64(A, xs), (xs, s))def range_go(n: Nat, acc: List<&2, Nat>) -> List<&2, Nat>: match n: case 0n: acc case 1n+ +p: range_go(p, Con{p, acc})# [0, 1, ..., n - 1]def range(n: Nat) -> List<&2, Nat>: range_go(n, [])# Go's Perm: a shuffle of 0..n-1def perm(~S: Data, ~next: S -> W.U64 & S, s: S, +n: Nat) -> List<&2, Nat> & S: shuffle(~Nat, ~S, ~next, s, range(n))