"""Weighted mixture loader over tokenized .bin shards (uint16, docs separated by EOS). Each row picks a source by weight, a file by length, and a uniform random window. Document ids for attention masking come from EOS positions (the EOS token stays with the doc it ends). """ from __future__ import annotations import glob import os import queue import threading import numpy as np import torch from tiny_agent.text import DATA, EOS_ID # Stable phase: mostly code + math, ~25% English for reading ability. Weights are renormalized # over sources that exist on disk, so synthetic sources can be added later without edits here. MIXTURES = { "stable": { "code_python": 0.16, "code_shell": 0.06, "code_javascript": 0.04, "code_typescript": 0.04, "code_markdown": 0.04, "code_go": 0.02, "code_rust": 0.02, "code_c": 0.02, "math": 0.25, "english": 0.25, "synth_grounded": 0.10, }, "decay": { "code_python": 0.12, "code_shell": 0.06, "code_javascript": 0.02, "code_typescript": 0.02, "code_markdown": 0.03, "math": 0.12, "english": 0.15, "synth_grounded": 0.15, "synth_tools": 0.22, "synth_teacher": 0.06, "synth_reasoning": 0.07, }, # short warm start before RL r2 (from final.pt, low LR): demos of the new training kinds and # reworded questions, with enough of the decay mix to not forget it "sft_agent": { "synth_agent_v2": 0.30, "synth_teacher_v2": 0.15, "synth_tools": 0.10, "synth_grounded": 0.10, "synth_reasoning": 0.05, "code_python": 0.08, "code_shell": 0.05, "math": 0.07, "english": 0.10, }, } class MixtureLoader: def __init__(self, split: str, mixture: str | dict, T: int, B: int, seed: int = 0, prefetch: int = 4, data_dir: str = DATA, stream: bool = True): weights = MIXTURES[mixture] if isinstance(mixture, str) else mixture self.T, self.B = T, B self.sources, self.files, self.file_p = [], [], [] for name, w in weights.items(): paths = sorted(glob.glob(f"{data_dir}/tok/{split}/{name}/*.bin")) arrs = [np.memmap(p, dtype=np.uint16, mode="r") for p in paths if os.path.getsize(p) > 2 * (T + 2)] if not arrs or w <= 0: continue lens = np.array([a.size for a in arrs], dtype=np.float64) self.sources.append((name, w)) self.files.append(arrs) self.file_p.append(lens / lens.sum()) if not self.sources: raise RuntimeError(f"no data for split={split} mixture={mixture} in {data_dir}/tok") w = np.array([w for _, w in self.sources]) self.src_p = w / w.sum() self.rng = np.random.default_rng(seed) self.q: queue.Queue = queue.Queue(maxsize=prefetch) self._stop = False if stream: self.thread = threading.Thread(target=self._fill, daemon=True) self.thread.start() def describe(self) -> dict: return {n: round(float(p), 4) for (n, _), p in zip(self.sources, self.src_p)} def _one(self, rng, src: int | None = None): T = self.T s = rng.choice(len(self.sources), p=self.src_p) if src is None else src f = rng.choice(len(self.files[s]), p=self.file_p[s]) arr = self.files[s][f] o = int(rng.integers(0, arr.size - T - 1)) return np.asarray(arr[o:o + T + 1], dtype=np.int64) def _make(self, rng, src=None): rows = np.stack([self._one(rng, src) for _ in range(self.B)]) x = torch.from_numpy(rows) inp, tgt = x[:, :-1], x[:, 1:] is_eos = (inp == EOS_ID).long() doc = is_eos.cumsum(1) - is_eos return inp.pin_memory(), tgt.pin_memory(), doc.pin_memory() def _fill(self): while not self._stop: self.q.put(self._make(self.rng)) def next(self, device="xpu"): inp, tgt, doc = self.q.get() return (inp.to(device, non_blocking=True), tgt.to(device, non_blocking=True), doc.to(device, non_blocking=True)) def fixed_batches(self, n: int, seed: int = 1234) -> dict[str, list]: """Deterministic per-source eval batches. Uses its own generator: the prefetch thread keeps drawing from self.rng concurrently, so sharing it would make val sets differ between runs.""" rng = np.random.default_rng(seed) return {name: [self._make(rng, i) for _ in range(n)] for i, (name, _) in enumerate(self.sources)} def close(self): self._stop = True