Download code/tiny_agent/rollout.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 16.8 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tiny_agent/rollout.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/tiny_agent/rollout.py
-
curl -L -o rollout.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/tiny_agent/rollout.py
16.8 kB
| """Batched multi-turn agent rollouts. | |
| Two schedulers over the same Episode records: | |
| * Engine (default, continuous batching): every cache row ("slot") runs its own episode and is | |
| refilled from a queue as soon as its episode ends, so nobody waits for the slowest rollout. | |
| Each tick: start queued episodes in free slots, collect finished tool calls, prefill pending | |
| prompts/tool results for a sub-batch of rows (batched: only when enough rows wait, or they waited | |
| long enough, or nothing is decoding), then decode one token for every generating row. | |
| Tool calls and scoring run in a thread pool (bash is a subprocess in bubblewrap). | |
| * Roller (fallback, lockstep): all rows generate a turn together, then run tools together. | |
| Every generated token records the log-probability it was sampled with (`logp`), so the learner can | |
| form a per-token importance ratio even when an episode spans policy updates. Assistant turns record | |
| their token spans and a flag for bad tool-call behavior (malformed call, unknown tool, bad | |
| arguments, repeated call), for segment-level penalties. | |
| """ | |
| from __future__ import annotations | |
| from collections import deque | |
| from concurrent.futures import ThreadPoolExecutor | |
| from dataclasses import dataclass, field | |
| import torch | |
| import torch.nn.functional as F | |
| from tokenizers import Tokenizer | |
| from tiny_agent.chat import parse_assistant, render, repeated_calls | |
| from tiny_agent.generate import KVCache, append, sample | |
| from tiny_agent.tasks import Task, check, grounded, invented | |
| from tiny_agent.tools import Workspace | |
| _FORMAT_ERRORS = ("Error: unknown tool", "Error: bad arguments") | |
| class Episode: | |
| task: Task | |
| ws: Workspace | None | |
| messages: list | |
| tokens: list = field(default_factory=list) # full token sequence | |
| gen_mask: list = field(default_factory=list) # 1 where the model generated the token | |
| logp: list = field(default_factory=list) # sampling log-prob of generated tokens (0 elsewhere) | |
| turn_spans: list = field(default_factory=list) # [start, end) token span of each assistant turn | |
| flags: list = field(default_factory=list) # per assistant turn: bad tool-call behavior | |
| turns: int = 0 | |
| done: bool = False | |
| truncated: bool = False | |
| correct: bool = False | |
| grounded: bool = False | |
| invented: bool = False # wrong answer seen in no tool result | |
| gen_tokens: int = 0 | |
| tool_tokens: int = 0 | |
| parse_errors: int = 0 | |
| repeats: int = 0 | |
| reward: float = 0.0 | |
| group: int = -1 | |
| idx: int = -1 | |
| version: int = 0 # policy version when the episode started | |
| def signals(self) -> dict: | |
| """Length signals for the in-group length penalty.""" | |
| return {"turns": max(1, self.turns), "input_tokens": self.tool_tokens + 1, | |
| "output_tokens": max(1, self.gen_tokens)} | |
| def finalize(e: Episode) -> Episode: | |
| """Score an ended episode and flag its bad tool-call turns; closes the workspace.""" | |
| try: | |
| e.correct = bool(check(e.task, e.ws.submitted, e.ws)) | |
| e.grounded = bool(grounded(e.task, e.messages)) | |
| e.invented = not e.correct and invented(e.task, e.ws.submitted, e.messages) | |
| except Exception: # e.g. the model wrote a config the checker cannot parse | |
| e.correct = e.grounded = e.invented = False | |
| finally: | |
| e.ws.close() | |
| reps = repeated_calls(e.messages) | |
| e.repeats = sum(reps) | |
| flags, ai, msgs = [], 0, e.messages | |
| for j, m in enumerate(msgs): | |
| if m["role"] != "assistant": | |
| continue | |
| bad = any("error" in c for c in m.get("tool_calls") or []) | |
| if j + 1 < len(msgs) and msgs[j + 1]["role"] == "tool": | |
| bad |= any(r.startswith(_FORMAT_ERRORS) for r in msgs[j + 1]["results"]) | |
| flags.append(bool(bad or reps[ai])) | |
| ai += 1 | |
| e.flags = flags | |
| e.done = True | |
| return e | |
| def run_calls(e: Episode, calls) -> list[str]: | |
| out = [] | |
| for c in calls: | |
| if "error" in c: | |
| e.parse_errors += 1 | |
| out.append(c["error"]) | |
| else: | |
| try: | |
| out.append(e.ws.call(c["name"], c["arguments"])) | |
| except Exception as ex: # never let one episode's tool call take down the run | |
| out.append(f"Error: tool failed: {type(ex).__name__}: {ex}") | |
| if e.ws.submitted is not None: | |
| break | |
| return out | |
| def _sample_with_logp(logits, temperature): | |
| tok = sample(logits, temperature) | |
| lp = F.log_softmax(logits / max(temperature, 1e-6), dim=-1).gather(1, tok[:, None]).squeeze(1) | |
| both = torch.stack([tok.float(), lp]).tolist() # one host sync; token ids < 2**24 are exact | |
| return tok, [int(t) for t in both[0]], both[1] | |
| class Engine: | |
| def __init__(self, model, tok: Tokenizer, device="xpu", slots=256, max_len=4096, max_turns=8, | |
| max_turn_tokens=768, temperature=1.0, prefill_rows=16, prefill_wait=8, workers=24): | |
| self.model, self.tok, self.device = model, tok, device | |
| self.B, self.max_len, self.max_turns, self.max_turn_tokens = slots, max_len, max_turns, max_turn_tokens | |
| self.temperature, self.prefill_rows, self.prefill_wait = temperature, prefill_rows, prefill_wait | |
| self.im_end = tok.token_to_id("<|im_end|>") | |
| self.pool = ThreadPoolExecutor(workers) | |
| self.cache = KVCache(model, slots, max_len, device) | |
| self.logits = torch.zeros(slots, model.cfg.vocab_size, device=device) | |
| self.ep: list[Episode | None] = [None] * slots | |
| self.state = ["free"] * slots # free | prefill | gen | tools | |
| self.pending = [None] * slots # tokens waiting to be prefilled | |
| self.waited = [0] * slots | |
| self.turn = [None] * slots # tokens generated in the current turn | |
| self.turn_start = [0] * slots | |
| self.fut = [None] * slots | |
| self.queue: deque[Episode] = deque() | |
| self.scoring = [] | |
| self.version = 0 | |
| self.stats = {"ticks": 0, "prefill_calls": 0, "decode_calls": 0} | |
| def enc(self, text: str) -> list[int]: | |
| return self.tok.encode(text, add_special_tokens=False).ids | |
| def submit(self, tasks: list[Task], group: int = -1, start_idx: int = 0) -> list[Episode]: | |
| eps = [] | |
| for i, t in enumerate(tasks): | |
| e = Episode(t, None, t.messages(), group=group, idx=start_idx + i) | |
| e._ws = self.pool.submit(Workspace, t.files) # build the workspace off the tick loop | |
| self.queue.append(e) | |
| eps.append(e) | |
| return eps | |
| def busy(self) -> bool: | |
| return bool(self.queue) or any(s != "free" for s in self.state) or bool(self.scoring) | |
| def in_flight(self) -> int: | |
| return len(self.queue) + sum(s != "free" for s in self.state) | |
| def tick(self) -> list[Episode]: | |
| """Advance every slot a little; returns episodes that finished (scored) since the last tick.""" | |
| self.stats["ticks"] += 1 | |
| self._fill() | |
| self._poll_tools() | |
| self._prefill() | |
| self._decode() | |
| return self._collect() | |
| def run(self, tasks: list[Task]) -> list[Episode]: | |
| """Run tasks to completion; returns episodes in input order.""" | |
| eps = self.submit(tasks) | |
| while any(not e.done for e in eps): | |
| self.tick() | |
| return eps | |
| # -- internals | |
| def _fill(self): | |
| free = [s for s in range(self.B) if self.state[s] == "free"] | |
| started = [] | |
| for s in free: | |
| if not self.queue: | |
| break | |
| e = self.queue.popleft() | |
| e.ws = e._ws.result() | |
| del e._ws | |
| e.version = self.version | |
| self.ep[s], self.state[s], self.waited[s] = e, "prefill", 0 | |
| self.pending[s] = self.enc(render(e.messages, add_generation_prompt=True)) | |
| started.append(s) | |
| if started: | |
| self.cache.reset(started) | |
| def _poll_tools(self): | |
| for s in range(self.B): | |
| if self.state[s] != "tools" or not self.fut[s].done(): | |
| continue | |
| e, results = self.ep[s], self.fut[s].result() | |
| self.fut[s] = None | |
| msg = {"role": "tool", "results": results} | |
| e.messages.append(msg) | |
| if e.ws.submitted is not None or e.turns >= self.max_turns: | |
| self._finish(s) | |
| continue | |
| ids = self.enc("\n" + render([msg], add_generation_prompt=True)) | |
| if len(e.tokens) + len(ids) + 16 > self.max_len: | |
| e.truncated = True | |
| self._finish(s) | |
| continue | |
| e.tool_tokens += len(ids) | |
| self.pending[s], self.state[s], self.waited[s] = ids, "prefill", 0 | |
| def _prefill(self): | |
| rows = [s for s in range(self.B) if self.state[s] == "prefill"] | |
| if not rows: | |
| return | |
| generating = any(st == "gen" for st in self.state) | |
| if generating and len(rows) < self.prefill_rows and max(self.waited[s] for s in rows) < self.prefill_wait: | |
| for s in rows: | |
| self.waited[s] += 1 | |
| return | |
| ns = [len(self.pending[s]) for s in rows] | |
| idx = torch.zeros(len(rows), max(ns), dtype=torch.long) | |
| for j, s in enumerate(rows): | |
| idx[j, : ns[j]] = torch.tensor(self.pending[s]) | |
| out = append(self.model, self.cache, idx.to(self.device), ns, rows=rows) | |
| self.logits[torch.tensor(rows, device=self.device)] = out | |
| self.stats["prefill_calls"] += 1 | |
| for s in rows: | |
| e, p = self.ep[s], self.pending[s] | |
| e.tokens += p | |
| e.gen_mask += [0] * len(p) | |
| e.logp += [0.0] * len(p) | |
| self.pending[s], self.state[s], self.turn[s], self.turn_start[s] = None, "gen", [], len(e.tokens) | |
| def _decode(self): | |
| rows = [s for s in range(self.B) if self.state[s] == "gen"] | |
| if not rows: | |
| return | |
| tok, tl, lp = self._sample(self.logits, self.ep) | |
| n = torch.zeros(self.B, dtype=torch.long) | |
| n[rows] = 1 | |
| ended = [] | |
| for s in rows: | |
| e, t = self.ep[s], tl[s] | |
| e.tokens.append(t) | |
| e.gen_mask.append(1) | |
| e.logp.append(lp[s]) | |
| e.gen_tokens += 1 | |
| self.turn[s].append(t) | |
| if t == self.im_end: | |
| ended.append((s, False)) | |
| elif len(self.turn[s]) >= self.max_turn_tokens or len(e.tokens) + 2 >= self.max_len: | |
| ended.append((s, True)) | |
| out = append(self.model, self.cache, tok[:, None], n) | |
| self.logits = torch.where(n.to(self.device, non_blocking=True)[:, None] > 0, out, self.logits) | |
| self.stats["decode_calls"] += 1 | |
| for s, trunc in ended: | |
| self._end_turn(s, trunc) | |
| def _sample(self, logits, row_eps): | |
| """(tokens on device, tokens as list, log-probs as list) for every row; tests override this.""" | |
| return _sample_with_logp(logits, self.temperature) | |
| def _end_turn(self, s, truncated): | |
| e = self.ep[s] | |
| msg = parse_assistant(self.tok.decode(self.turn[s], skip_special_tokens=False)) | |
| e.messages.append(msg) | |
| e.turns += 1 | |
| e.turn_spans.append((self.turn_start[s], len(e.tokens))) | |
| self.turn[s] = None | |
| if truncated: | |
| e.truncated = True | |
| calls = msg["tool_calls"] | |
| if not calls or truncated: | |
| self._finish(s) # a turn without tool calls ends the episode (no submit = no answer) | |
| return | |
| self.fut[s], self.state[s] = self.pool.submit(run_calls, e, calls), "tools" | |
| def _finish(self, s): | |
| e = self.ep[s] | |
| self.ep[s], self.state[s] = None, "free" | |
| self.scoring.append(self.pool.submit(finalize, e)) | |
| def _collect(self) -> list[Episode]: | |
| done = [f for f in self.scoring if f.done()] | |
| if done: | |
| self.scoring = [f for f in self.scoring if not f.done()] | |
| return [f.result() for f in done] | |
| class Roller: | |
| """Lockstep fallback: all rows generate one turn together, then run tools together.""" | |
| def __init__(self, model, tok: Tokenizer, device="xpu", max_len=4096, max_turns=8, max_turn_tokens=768, | |
| temperature=1.0, workers=16): | |
| self.model, self.tok, self.device = model, tok, device | |
| self.max_len, self.max_turns, self.max_turn_tokens, self.temperature = max_len, max_turns, max_turn_tokens, temperature | |
| self.im_end = tok.token_to_id("<|im_end|>") | |
| self.pool = ThreadPoolExecutor(workers) | |
| self.version = 0 | |
| def enc(self, text: str) -> list[int]: | |
| return self.tok.encode(text, add_special_tokens=False).ids | |
| def run(self, tasks: list[Task]) -> list[Episode]: | |
| eps = [Episode(t, Workspace(t.files), t.messages(), idx=i, version=self.version) for i, t in enumerate(tasks)] | |
| B = len(eps) | |
| cache = KVCache(self.model, B, self.max_len, self.device) | |
| prompts = [self.enc(render(e.messages, add_generation_prompt=True)) for e in eps] | |
| logits = self._feed(cache, prompts, eps) | |
| live = set(range(B)) | |
| while live: | |
| active = sorted(live) | |
| starts = {i: len(eps[i].tokens) for i in active} | |
| turn_toks = self._generate_turn(cache, logits, eps, active) | |
| chunks = [[] for _ in range(B)] | |
| jobs = {} | |
| for i in active: | |
| e = eps[i] | |
| e.turns += 1 | |
| e.turn_spans.append((starts[i], len(e.tokens))) | |
| msg = parse_assistant(self.tok.decode(turn_toks[i], skip_special_tokens=False)) | |
| e.messages.append(msg) | |
| calls = msg["tool_calls"] | |
| if not calls or e.truncated: | |
| live.discard(i) | |
| continue | |
| jobs[i] = self.pool.submit(run_calls, e, calls) | |
| for i, fut in jobs.items(): | |
| e = eps[i] | |
| results = fut.result() | |
| e.messages.append({"role": "tool", "results": results}) | |
| if e.ws.submitted is not None or e.turns >= self.max_turns: | |
| live.discard(i) | |
| continue | |
| ids = self.enc("\n" + render([{"role": "tool", "results": results}], add_generation_prompt=True)) | |
| if len(e.tokens) + len(ids) + 16 > self.max_len: | |
| e.truncated = True | |
| live.discard(i) | |
| continue | |
| e.tool_tokens += len(ids) | |
| chunks[i] = ids | |
| if any(chunks): | |
| logits = self._feed(cache, chunks, eps) | |
| return list(self.pool.map(finalize, eps)) | |
| def _feed(self, cache, chunks, eps): | |
| n = torch.tensor([len(c) for c in chunks]) # CPU: append() mirrors lengths on the host | |
| x = torch.zeros(len(chunks), int(n.max()), dtype=torch.long) | |
| for b, c in enumerate(chunks): | |
| if c: | |
| x[b, : len(c)] = torch.tensor(c) | |
| eps[b].tokens.extend(c) | |
| eps[b].gen_mask.extend([0] * len(c)) | |
| eps[b].logp.extend([0.0] * len(c)) | |
| return append(self.model, cache, x.to(self.device), n) | |
| def _sample(self, logits, row_eps): | |
| return _sample_with_logp(logits, self.temperature) | |
| def _generate_turn(self, cache, logits, eps, active): | |
| B = len(eps) | |
| out = {i: [] for i in active} | |
| live = torch.zeros(B, dtype=torch.bool) | |
| live[active] = True | |
| for _ in range(self.max_turn_tokens): | |
| nxt, nl, lp = self._sample(logits, eps) | |
| lv = live.tolist() | |
| for i in range(B): | |
| if lv[i]: | |
| out[i].append(nl[i]) | |
| eps[i].tokens.append(nl[i]) | |
| eps[i].gen_mask.append(1) | |
| eps[i].logp.append(lp[i]) | |
| eps[i].gen_tokens += 1 | |
| logits = append(self.model, cache, nxt[:, None], live.long()) | |
| for i in range(B): | |
| if lv[i] and (nl[i] == self.im_end or len(eps[i].tokens) + 2 >= self.max_len): | |
| live[i] = False | |
| if nl[i] != self.im_end: | |
| eps[i].truncated = True | |
| if not live.any(): | |
| break | |
| for i in active: | |
| if out[i] and out[i][-1] != self.im_end: | |
| eps[i].truncated = True | |
| return out | |
| def make_roller(kind: str, model, tok, **kw): | |
| if kind == "lockstep": | |
| kw.pop("slots", None) | |
| kw.pop("prefill_rows", None) | |
| kw.pop("prefill_wait", None) | |
| return Roller(model, tok, **kw) | |
| return Engine(model, tok, **kw) | |