HCho's picture
BitStream Modular Machine - submission 1
4fc906a verified
Raw
History Blame Contribute Delete
5.39 kB
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