darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
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