| 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 |
|
|