import Base import ./logic.bend as Lg # Bits.bend — U32 bit tricks and branch-free integer helpers. # # Efficiency notes (Bend2): # - Everything here runs on U32 (32-bit words). Nat is unary and slow: # never use Nat for counting in hot code; Nat appears only as small # structural fuel (<= 64n) so the termination checker is happy. # - Loops always run a FIXED number of steps (fuel-bounded) and freeze # their state with Bool.pick instead of early-exiting. Early exit would # need a match on a computed Bool plus a recursive helper, which is a # use-before-definition cycle that user code cannot write (the Base # library itself is exempt from that rule; user modules are not). # Fixed-step loops cost O(32) worst case, with zero branching surprise. # - U32.div / U32.mod cost ~32 internal steps each. Prefer shifts, masks, # and the helpers below inside inner loops. # popcount: number of set bits. Fixed 32 steps. def pop_go(fuel: Nat, +x: U32, acc: U32) -> U32: match fuel: case 0n: acc case 1n+p: pop_go(p, U32.shr(x), (acc + U32.and(x, 1) : U32)) def pop(+x: U32) -> U32: pop_go(32n, x, 0) # ctz: count trailing zeros. ctz(0) == 32. Fixed 32 steps, state frozen # with picks once the first set bit is seen. def ctz_go(fuel: Nat, +x: U32, +acc: U32) -> U32: match fuel: case 0n: acc case 1n+p: +bit = U32.and(x, 1) +zero = U32.is_zero(bit) ctz_go(p, Bool.pick(U32, zero, U32.shr(x), x), Bool.pick(U32, zero, (acc + 1 : U32), acc)) def ctz(+x: U32) -> U32: ctz_go(32n, x, 0) # clz: count leading zeros. clz(0) == 32. Via Base U32.log2 (no loop). def clz(+x: U32) -> U32: Bool.pick(U32, U32.is_zero(x), 32, U32.from_nat(Nat.sub(31n, U32.log2(x)))) # rotl / rotr by k bits (k: Nat). Straight-line, no loop. def rotl(+x: U32, +k: Nat) -> U32: U32.or(U32.shln(x, k), U32.shrn(x, Nat.sub(32n, k))) def rotr(+x: U32, +k: Nat) -> U32: U32.or(U32.shrn(x, k), U32.shln(x, Nat.sub(32n, k))) # is_pow2: True for 1, 2, 4, ... (and False for 0). def is_pow2(+x: U32) -> Bool: Bool.and(Bool.not(U32.is_zero(x)), U32.is_zero(U32.and(x, U32.sub(x, 1)))) # next_pow2: smallest power of two >= x (with next_pow2(0) == 1). # Wraps to 0 for x > 0x80000000 (no 33-bit values exist). def next_pow2_big(p: U32) -> U32: U32.shln(1, U32.to_nat(U32.sub(32, clz(p)))) def next_pow2(+x: U32) -> U32: Bool.pick(U32, U32.is_le(x, 1), 1, next_pow2_big(U32.sub(x, 1))) # ceil_div: (a + b - 1) / b. Returns 0 when b == 0. def ceil_div(a: U32, +b: U32) -> U32: U32.div((a + U32.sub(b, 1) : U32), b) # abs_diff: |a - b| without signed types (sub wraps, pick keeps the good side). def abs_diff(+a: U32, +b: U32) -> U32: Bool.pick(U32, U32.is_lt(a, b), U32.sub(b, a), U32.sub(a, b)) # avg: (a + b) / 2 with no overflow: (a & b) + ((a ^ b) >> 1). def avg(+a: U32, +b: U32) -> U32: (U32.and(a, b) + U32.shr(U32.xor(a, b)) : U32) # parity: popcount mod 2 (0 or 1). is_odd: lowest bit set. def parity(+x: U32) -> U32: U32.and(pop(x), 1) def is_odd(+x: U32) -> Bool: Bool.not(U32.is_zero(U32.and(x, 1))) # mix32: Murmur3 fmix32 avalanche. Use to hash U32 keys. def mix32(+x: U32) -> U32: +a = U32.xor(x, U32.shrn(x, 16n)) +b = U32.mul(a, 2246822507) +c = U32.xor(b, U32.shrn(b, 13n)) +d = U32.mul(c, 3266489909) U32.xor(d, U32.shrn(d, 16n)) # word_cmp_refl: cmp(n, w, w) == EQ (induction on n + w). law word_cmp_refl: for n: Nat for w: Word(n) {Word.cmp(n, w, w) == EQ{} : Cmp} def word_cmp_refl(n, w): match n w: case 0n WNil{}: {==} case 1n+p WCon{+b, t}: Equal.trans(Cmp, Word.cmp.fin(b, b, Word.cmp(p, t, t)), Word.cmp.fin(b, b, EQ{}), EQ{}, Equal.cong(Cmp, Cmp, w => Word.cmp.fin(b, b, w), Word.cmp(p, t, t), EQ{}, word_cmp_refl(p, t)), Lg.bool_cmp_refl(b)) # u32_eq_refl: (x == x) == True. law u32_eq_refl: for a: U32 {U32.is_eq(a, a) == True{} : Bool} def u32_eq_refl(a): match a: case U32{+w}: Equal.cong(Cmp, Bool, Cmp.is_eq, Word.cmp(32n, w, w), EQ{}, word_cmp_refl(32n, w)) # fnv_step: one FNV-1a byte fold. h = (h ^ byte) * 16777619. def fnv_step(h: U32, b: U32) -> U32: U32.mul(U32.xor(h, b), 16777619) # sat_add / sat_sub: saturating arithmetic (clamp instead of wrap). def sat_add(+a: U32, b: U32) -> U32: +s = (a + b : U32) Bool.pick(U32, U32.is_lt(s, a), 4294967295, s) def sat_sub(+a: U32, +b: U32) -> U32: Bool.pick(U32, U32.is_lt(a, b), 0, U32.sub(a, b)) # lowbit: lowest set bit as a mask (lowbit(12) == 4). def lowbit(+x: U32) -> U32: U32.and(x, U32.sub(0, x)) # bswap32: reverse byte order. def bswap32(+x: U32) -> U32: +a = U32.shln(U32.and(x, 255), 24n) +b = U32.shln(U32.and(x, 65280), 8n) +c = U32.shrn(U32.and(x, 16711680), 8n) +d = U32.shrn(x, 24n) U32.or(U32.or(a, b), U32.or(c, d)) # mul_hi: upper 32 bits of the 64-bit product (16-bit splitting, exact). def mul_hi(+a: U32, +b: U32) -> U32: +ahi = U32.shrn(a, 16n) +alo = U32.and(a, 65535) +bhi = U32.shrn(b, 16n) +blo = U32.and(b, 65535) +m1 = U32.mul(ahi, blo) m2 = U32.mul(alo, bhi) +mw = (m1 + m2 : U32) +cm = Bool.pick(U32, U32.is_lt(mw, m1), 1, 0) mhi = (U32.shrn(mw, 16n) + U32.shln(cm, 16n) : U32) mlo = U32.and(mw, 65535) l = U32.mul(alo, blo) lhi = U32.shrn(l, 16n) s = (mlo + lhi : U32) c = U32.shrn(s, 16n) ((U32.mul(ahi, bhi) + mhi : U32) + c : U32)