Download code/tiny_agent/data.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 4.5 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tiny_agent/data.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/tiny_agent/data.py
-
curl -L -o data.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tiny_agent/data.py
4.5 kB
| """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 | |