from __future__ import annotations import random import numpy as np import torch PAD_HEAD = 3 def _rand_bits_int(rng: random.Random, l: int) -> int: if l == 1: return 1 return (1 << (l - 1)) | rng.getrandbits(l - 1) def sample_modulus(rng: random.Random, n: int) -> int: lmax = n - PAD_HEAD r = rng.random() if r < 0.50: l = lmax elif r < 0.80: l = rng.randint(2, lmax) else: l = min(lmax, 1 + int(2 ** (rng.random() * 4))) l = max(2, l) if l <= 3 and rng.random() < 0.7: return rng.choice([2, 3, 5, 7][: 2 if l == 2 else 4]) if l >= 5 and rng.random() < 0.15: c = rng.choice([1, 3, 5, 7, 9, 15, 17, 31, 33, 63, rng.randint(1, 99)]) if rng.random() < 0.7: m = (1 << l) - c else: m = (1 << (l - 1)) + c if rng.random() < 0.9: m |= 1 if 2 <= m and m.bit_length() <= l: return m m = _rand_bits_int(rng, l) if rng.random() < 0.75: m |= 1 return max(2, m) def _carry_stress(rng: random.Random, hi: int) -> int: nbits = max(2, hi.bit_length()) j = rng.randint(1, nbits) i = rng.randint(0, j - 1) v = (1 << j) - (1 << i) if rng.random() < 0.5: v |= rng.getrandbits(max(1, i)) return v % hi def _sample_x(rng: random.Random, m: int, hi_mult: int, w: int) -> int: hi = hi_mult * m r = rng.random() if r < 0.40: for _ in range(8): q = rng.randint(0, hi_mult) lo_s = max(q * m, w) hi_s = min((q + 1) * m, hi + w) if lo_s < hi_s: x = rng.randrange(lo_s, hi_s) - w if 0 <= x < hi: return x return rng.randrange(hi) if r < 0.50: return rng.randrange(hi) if r < 0.58: u = rng.randrange(m) d = rng.randrange(hi_mult) return min(hi - 1, hi_mult * u + d) if r < 0.64: return rng.randrange(m) if r < 0.78: k = rng.randint(1, hi_mult) delta = rng.choice([0, 1, 2, 3, rng.randint(0, 8)]) s_t = k * m + (delta if rng.random() < 0.5 else -delta) x = s_t - w return x if 0 <= x < hi else rng.randrange(hi) if r < 0.96: if rng.random() < 0.5: x = _carry_stress(rng, hi) else: k = rng.randint(1, hi_mult) x = k * m - w + (1 << rng.randint(0, max(1, hi.bit_length() - 2))) \ - rng.randint(0, 3) if not (0 <= x < hi): x = _carry_stress(rng, hi) return x return rng.choice([0, 1, 2, 3]) def sample_reduce(rng: random.Random, n: int) -> tuple[int, int]: m = sample_modulus(rng, n) return m, _sample_x(rng, m, 4, 0) def _pack_bits(vals: list[int], n: int) -> np.ndarray: out = np.empty((len(vals), n), dtype=np.uint8) for i, v in enumerate(vals): s = np.frombuffer(format(v, f"0{n}b").encode(), dtype=np.uint8) out[i] = s - 48 return out def _T(vals, n): return torch.from_numpy(_pack_bits(vals, n)).float() def make_reduce_batch(rng, n, bsz, instances=None): mask = (1 << n) - 1 ms, xs, zs, qs, p3s = [], [], [], [], [] borrows = [[], [], []] for j in range(bsz): if instances is not None: m, x = instances[j % len(instances)] else: m, x = sample_reduce(rng, n) q = x // m zs.append(x - q * m) qs.append(q) for k in (1, 2, 3): km = k * m diff = (x - km) & mask borrows[k - 1].append((x ^ km ^ diff) & mask) ms.append(m); xs.append(x); p3s.append(3 * m) batch = { "x": _T(xs, n), "p": _T(ms, n), "p3": _T(p3s, n), "z": _T(zs, n), "borrow": torch.stack([_T(borrows[k], n) for k in range(3)], dim=-1), "q": torch.tensor(qs, dtype=torch.long), "raw": list(zip(ms, xs)), } return batch def sample_add(rng: random.Random, n: int) -> tuple[int, int, int]: r = rng.random() if r < 0.45: x = rng.getrandbits(rng.randint(1, n - 2)) if rng.random() < 0.5 \ else rng.randrange(1 << (n - 2)) elif r < 0.85: x = _carry_stress(rng, 1 << (n - 2)) elif r < 0.95: x = rng.choice([0, 1, 2, 3]) else: x = (1 << (n - 2)) - rng.randint(1, 4) if rng.random() < 0.7: x &= ~1 r = rng.random() if r < 0.5: y = rng.getrandbits(rng.randint(1, n - 3)) if rng.random() < 0.5 \ else rng.randrange(1 << (n - 3)) elif r < 0.9: y = _carry_stress(rng, 1 << (n - 3)) else: y = rng.choice([0, 1, (1 << (n - 3)) - 1]) g = rng.randint(0, 1) return x, y, g def make_add_batch(rng, n, bsz, instances=None): mask = (1 << n) - 1 xs, ys, gs, ss, cs = [], [], [], [], [] for j in range(bsz): if instances is not None: x, y, g = instances[j % len(instances)] else: x, y, g = sample_add(rng, n) w = g * y s = x + w ss.append(s & mask) cs.append((x ^ w ^ s) & mask) xs.append(x); ys.append(y); gs.append(g) batch = { "x": _T(xs, n), "y": _T(ys, n), "g": torch.tensor(gs, dtype=torch.float32), "z": _T(ss, n), "carry": _T(cs, n), "raw": list(zip(xs, ys, gs)), } return batch