Download data.py from HCho/bitstream-modmul-model: direct link, hf CLI and curl.
- Browser
- Download file 5.39 kB
-
https://huggingface.co/HCho/bitstream-modmul-model/resolve/main/data.py
- Command line
-
hf download hf://HCho/bitstream-modmul-model/data.py
-
curl -L -o data.py https://huggingface.co/HCho/bitstream-modmul-model/resolve/main/data.py
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 | |