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