"""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") @dataclass 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) @torch.no_grad() 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() @torch.no_grad() 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 @torch.no_grad() 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)