~/bend-docscommunity

proofs/math/random/proof_pcg.bend source

proofs/math/random/proof_pcg.bend on the hub · documented module

import Baseimport ../../../spec/lib/common.bend as Cimport ../../../spec/math/w64.bend as SWimport ../../../spec/math/random/pcg.bend as SPimport ../../../spec/math/random.bend as SRMimport ../../../src/math/u64.bend as WUimport ../../../src/math/w64.bend as Ximport ../../../src/math/random/pcg.bend as Pimport ../../lib/lemmas/proofs/nat_algebra.bend as NAimport ../typed/width.bend as WWimport ../typed/w64add.bend as WAimport ../typed/w64sh.bend as SHimport ./pcg/xor.bend as XOimport ../../lib/lemmas/spec/numeric.bend as Simport ../typed/w64m128.bend as M128import ./pcg/step.bend as STimport ./pcg/impl.bend as IM# Entry point: `bend proofs/math/random/proof_pcg.bend` checks the PCG clauses# of spec/math/random.bend under the clause's name, for every state and every# multiplier and increment; no holes, no axioms. One step of# src/math/random/pcg.bend is the 128-bit LCG (the width-generic algebra is in# pcg/step.bend, instantiated at 64 in pcg/impl.bend) and its DXSM output is# spec/math/random/pcg.bend's dxsm: xor with a right shift, products mod 2^64# and "or 1" on the low word, each by its value lemma. The factors of each# product stay variables (mulg), so no product is ever normalized, and the# clause types are stated once (a clause type is expensive to normalize).# The other clauses are checked by proofs/math/random/proof.bend,# proof_draws.bend and proof_float.bend.def PCG.step(+mh: WU.U64, +ml: WU.U64, +ih: WU.U64, +il: WU.U64, +p: P.PCG) -> SRM.PCG.step(mh, ml, ih, il, p):  match p:    case P.P{+hi, +lo}:      IM.step_pair(mh, ml, ih, il, hi, lo, X.mul128(lo, ml), IM.m128(lo, ml))def PCG.constants(+p: P.PCG) -> SRM.PCG.constants(p):  {==}def v(+x: U32) -> Nat:  U32.to_nat(x)# x ^ (x >> k)def xs(+a: WU.U64, +k: Nat) -> {SW.value(P.xor64(a, X.shr(a, k))) == SP.xor_bits(64n, SW.value(a), C.high(k, SW.value(a))) : Nat}:  Equal.trans(Nat, SW.value(P.xor64(a, X.shr(a, k))), SP.xor_bits(64n, SW.value(a), SW.value(X.shr(a, k))), SP.xor_bits(64n, SW.value(a), C.high(k, SW.value(a))),    XO.xor64_value(a, X.shr(a, k)),    Equal.cong(Nat, Nat, z => SP.xor_bits(64n, SW.value(a), z), SW.value(X.shr(a, k)), C.high(k, SW.value(a)), SH.shr_value(a, k)))def plus1(+x: Nat) -> {Nat.add(x, 1n) == 1n+x : Nat}:  Equal.trans(Nat, Nat.add(x, 1n), 1n+Nat.add(x, 0n), 1n+x, NA.add_succ(x, 0n), Equal.cong(Nat, Nat, z => 1n+z, Nat.add(x, 0n), x, NA.add_zero(x)))# lo | 1def or_value(+lo: WU.U64) -> {SW.value(WU.U64{U32.or(X.lo(lo), 1), X.hi(lo)}) == Nat.add(Nat.double(C.half(SW.value(lo))), 1n) : Nat}:  match lo:    case WU.U64{+l, +h}:      +vl = v(l)      +vh = v(h)      +D = Nat.add(Nat.double(C.half(vl)), C.shift(32n, vh))      +e1 = Equal.cong(Nat, Nat, z => Nat.add(z, C.shift(32n, vh)), v(U32.or(l, 1)), 1n+Nat.double(C.half(vl)), SH.or1(l))      +e2 = Equal.cong(Nat, Nat, z => 1n+z, D, Nat.double(Nat.add(C.half(vl), C.shift(31n, vh))), Equal.sym(Nat, Nat.double(Nat.add(C.half(vl), C.shift(31n, vh))), D, NA.double_add(C.half(vl), C.shift(31n, vh))))      +e3 = Equal.cong(Nat, Nat, z => 1n+Nat.double(z), Nat.add(C.half(vl), C.shift(31n, vh)), C.half(Nat.add(vl, C.shift(32n, vh))), Equal.sym(Nat, C.half(Nat.add(vl, C.shift(32n, vh))), Nat.add(C.half(vl), C.shift(31n, vh)), WW.half_dbl(vl, C.shift(31n, vh))))      +e4 = Equal.sym(Nat, Nat.add(Nat.double(C.half(Nat.add(vl, C.shift(32n, vh)))), 1n), 1n+Nat.double(C.half(Nat.add(vl, C.shift(32n, vh)))), plus1(Nat.double(C.half(Nat.add(vl, C.shift(32n, vh))))))      Equal.trans(Nat, Nat.add(v(U32.or(l, 1)), C.shift(32n, vh)), 1n+D, Nat.add(Nat.double(C.half(Nat.add(vl, C.shift(32n, vh)))), 1n), e1,        Equal.trans(Nat, 1n+D, 1n+Nat.double(Nat.add(C.half(vl), C.shift(31n, vh))), Nat.add(Nat.double(C.half(Nat.add(vl, C.shift(32n, vh)))), 1n), e2,          Equal.trans(Nat, 1n+Nat.double(Nat.add(C.half(vl), C.shift(31n, vh))), 1n+Nat.double(C.half(Nat.add(vl, C.shift(32n, vh)))), Nat.add(Nat.double(C.half(Nat.add(vl, C.shift(32n, vh)))), 1n), e3, e4)))# the value of a product mod 2^64 from the values of its factors; the# factors stay variables here, so no product is ever normalizeddef mulg(+x: WU.U64, +o: WU.U64, +xv: Nat, +hx: {SW.value(x) == xv : Nat}, +ov: Nat, +ho: {SW.value(o) == ov : Nat}) -> {SW.value(X.mul(x, o)) == C.low(64n, Nat.mul(xv, ov)) : Nat}:  Equal.trans(Nat, SW.value(X.mul(x, o)), C.low(64n, Nat.mul(SW.value(x), SW.value(o))), C.low(64n, Nat.mul(xv, ov)), WA.mul_value(x, o),    Equal.trans(Nat, C.low(64n, Nat.mul(SW.value(x), SW.value(o))), C.low(64n, Nat.mul(xv, SW.value(o))), C.low(64n, Nat.mul(xv, ov)),      Equal.cong(Nat, Nat, z => C.low(64n, Nat.mul(z, SW.value(o))), SW.value(x), xv, hx),      Equal.cong(Nat, Nat, z => C.low(64n, Nat.mul(xv, z)), SW.value(o), ov, ho)))def xs32(+h: WU.U64) -> {SW.value(P.xor64(h, X.shr(h, 32n))) == SP.xor_bits(64n, SW.value(h), C.high(32n, SW.value(h))) : Nat}:  xs(h, 32n)def xs48(+h: WU.U64) -> {SW.value(P.xor64(h, X.shr(h, 48n))) == SP.xor_bits(64n, SW.value(h), C.high(48n, SW.value(h))) : Nat}:  xs(h, 48n)# x ^ (x >> 48) for a word of value tdef xt(+H: WU.U64, +t: Nat, +ht: {SW.value(H) == t : Nat}) -> {SW.value(P.xor64(H, X.shr(H, 48n))) == SP.xor_bits(64n, t, C.high(48n, t)) : Nat}:  Equal.trans(Nat, SW.value(P.xor64(H, X.shr(H, 48n))), SP.xor_bits(64n, SW.value(H), C.high(48n, SW.value(H))), SP.xor_bits(64n, t, C.high(48n, t)),    xs48(H), Equal.cong(Nat, Nat, z => SP.xor_bits(64n, z, C.high(48n, z)), SW.value(H), t, ht))# the first half's value is tdef half1(+cm: WU.U64, +hi: WU.U64, +t: Nat, +ht: {t == SP.dxsm1(SW.value(cm), SW.value(hi)) : Nat}) -> {SW.value(X.mul(P.xor64(hi, X.shr(hi, 32n)), cm)) == t : Nat}:  Equal.trans(Nat, SW.value(X.mul(P.xor64(hi, X.shr(hi, 32n)), cm)), SP.dxsm1(SW.value(cm), SW.value(hi)), t,    mulg(P.xor64(hi, X.shr(hi, 32n)), cm, SP.xor_bits(64n, SW.value(hi), C.high(32n, SW.value(hi))), xs32(hi), SW.value(cm), {==}),    Equal.sym(Nat, t, SP.dxsm1(SW.value(cm), SW.value(hi)), ht))# THEOREM (PCG.output)# values are known (x ^ x >> 48 with x of value t, and lo | 1)def PCG.output(+cm: WU.U64, +hi: WU.U64, +lo: WU.U64, +t: Nat, +ht: {t == SP.dxsm1(SW.value(cm), SW.value(hi)) : Nat}) -> SRM.PCG.output(cm, hi, lo, t, ht):  +H = X.mul(P.xor64(hi, X.shr(hi, 32n)), cm)  mulg(P.xor64(H, X.shr(H, 48n)), WU.U64{U32.or(X.lo(lo), 1), X.hi(lo)}, SP.xor_bits(64n, t, C.high(48n, t)), xt(H, t, half1(cm, hi, t, ht)), Nat.add(Nat.double(C.half(SW.value(lo))), 1n), or_value(lo))