File size: 7,133 Bytes
94c46f3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 | """Deterministic streaming loader over packed uint16 shards.
- One PackedStream per language: contiguous (seq_len+1)-token windows,
epoch-level window permutation (seeded), transparent epoch rollover with
epoch counting (matters for AR, whose pool may be smaller than the budget).
- MixedStream: bilingual token-level mixing. For every global sample slot a
stateless hash of (seed, slot) picks the language, so the mixture is exactly
reproducible, independent of world size, and resumable from a slot counter.
Mixing mode (thesis decision, configurable in the run config):
token-level p=0.5 matches ATLAS's compute-matched framing: each language
contributes half the training *tokens*; under a starved tokenizer the
cross-script partner therefore sees less *content* -- that is exactly the
effect under audit. `content` mode instead sets p from the byte premiums so
content is matched and token counts differ.
"""
import json
from pathlib import Path
import numpy as np
from ..paths import shard_dir
def _splitmix64(x: int) -> int:
x = (x + 0x9E3779B97F4A7C15) & 0xFFFFFFFFFFFFFFFF
z = x
z = ((z ^ (z >> 30)) * 0xBF58476D1CE4E5B9) & 0xFFFFFFFFFFFFFFFF
z = ((z ^ (z >> 27)) * 0x94D049BB133111EB) & 0xFFFFFFFFFFFFFFFF
return z ^ (z >> 31)
class PackedStream:
def __init__(self, lang: str, tok_name: str, seq_len: int, seed: int):
self.lang, self.seq_len, self.seed = lang, seq_len, seed
d = shard_dir(lang, tok_name)
index = json.loads((d / "index.json").read_text())
# Pool checkpoints can create complete but empty zstd shards. Packing
# those historically left zero-token .bin entries in index.json; mmap
# cannot open a zero-byte file when a window crosses that boundary.
# They carry no data, so filter them losslessly at load time.
names = [name for name in sorted(index) if index[name] > 0]
self.paths = [d / name for name in names]
self.lens = np.array([index[name] for name in names], dtype=np.int64)
self.cum = np.concatenate([[0], np.cumsum(self.lens)])
self.total_tokens = int(self.cum[-1])
self.n_windows = self.total_tokens // (seq_len + 1)
if self.n_windows == 0:
raise RuntimeError(f"shards for {lang}/{tok_name} too small")
self._mm = [None] * len(self.paths)
self._epoch = -1
self._perm = None
self.consumed = 0 # windows consumed (resume state)
def _memmap(self, i: int):
if self._mm[i] is None:
self._mm[i] = np.memmap(self.paths[i], dtype=np.uint16, mode="r")
return self._mm[i]
def _window(self, w: int) -> np.ndarray:
start = w * (self.seq_len + 1)
i = int(np.searchsorted(self.cum, start, side="right") - 1)
out = np.empty(self.seq_len + 1, dtype=np.uint16)
filled = 0
while filled < self.seq_len + 1:
off = start + filled - self.cum[i]
take = int(min(self.lens[i] - off, self.seq_len + 1 - filled))
out[filled:filled + take] = self._memmap(i)[off:off + take]
filled += take
i += 1
return out
def get(self, k: int) -> np.ndarray:
"""k-th window in the global consumption order (deterministic)."""
epoch, pos = divmod(k, self.n_windows)
if epoch != self._epoch:
rng = np.random.default_rng([self.seed, epoch, 0xDA7A])
self._perm = rng.permutation(self.n_windows)
self._epoch = epoch
return self._window(int(self._perm[pos]))
def next(self) -> np.ndarray:
x = self.get(self.consumed)
self.consumed += 1
return x
@property
def epochs(self) -> float:
return self.consumed / self.n_windows
def state_dict(self):
return {"consumed": self.consumed}
def load_state_dict(self, s):
self.consumed = s["consumed"]
class MixedStream:
"""Token-level language mixing across one or two PackedStreams."""
def __init__(self, langs: list[str], tok_name: str, seq_len: int, seed: int,
probs: list[float] | None = None):
self.langs = langs
self.streams = {l: PackedStream(l, tok_name, seq_len, seed + j)
for j, l in enumerate(langs)}
self.probs = probs or [1.0 / len(langs)] * len(langs)
assert abs(sum(self.probs) - 1.0) < 1e-6
self.seed = seed
self.slot = 0 # global sample counter (resume state)
def _choose(self, slot: int) -> str:
if len(self.langs) == 1:
return self.langs[0]
u = _splitmix64(self.seed * 0x100000001 + slot) / 2**64
acc = 0.0
for l, p in zip(self.langs, self.probs):
acc += p
if u < acc:
return l
return self.langs[-1]
def next_batch(self, n: int) -> tuple[np.ndarray, dict[str, int]]:
"""n windows -> array (n, seq_len+1); also per-language window counts."""
rows, counts = [], {l: 0 for l in self.langs}
for _ in range(n):
l = self._choose(self.slot)
rows.append(self.streams[l].next())
counts[l] += 1
self.slot += 1
return np.stack(rows), counts
def draw(self) -> tuple[str, int]:
"""Advance one global slot; return (lang, occurrence-index) WITHOUT reading.
Occurrence-index is the k-th window of that language in global order, so
stream.get(occ) is a pure deterministic lookup. All DDP ranks call draw()
for every slot (keeping counters identical) but only materialise their own
rows via streams[lang].get(occ) -- giving each rank a disjoint, reproducible
1/world shard of the global batch.
"""
l = self._choose(self.slot)
occ = self.streams[l].consumed
self.streams[l].consumed += 1
self.slot += 1
return l, occ
def rank_batch(self, n_global: int, rank: int, world: int
) -> tuple[np.ndarray, dict[str, int]]:
"""n_global windows drawn globally; rows n_global/world owned by `rank`."""
rows, counts = [], {l: 0 for l in self.langs}
for j in range(n_global):
l, occ = self.draw()
if j % world == rank:
rows.append(self.streams[l].get(occ))
counts[l] += 1
return np.stack(rows), counts
def skip_to(self, slot: int) -> None:
"""Fast-forward mixture state without reading data (for resume)."""
assert slot >= self.slot
for s in range(self.slot, slot):
self.streams[self._choose(s)].consumed += 1
self.slot = slot
def state_dict(self):
return {"slot": self.slot,
"streams": {l: s.state_dict() for l, s in self.streams.items()}}
def load_state_dict(self, s):
self.slot = s["slot"]
for l, st in s["streams"].items():
self.streams[l].load_state_dict(st)
def stats(self) -> dict:
return {l: {"windows": s.consumed, "epochs": round(s.epochs, 4)}
for l, s in self.streams.items()}
|