openenv/echo_env
6
1"""Graph -> the JSON a trainer consumes.2 3One document per rollout. Every field a trainer needs is precomputed and validated; nothing4downstream has to re-derive, re-tokenize, or guess.5 6 {7 "session_id": ...,8 "stats": {...}, graph shape: turns, roots, forks, discards9 "sequences": [ one per root-to-leaf path10 {"input_ids", "loss_mask", "logprobs", "prompt_len", "n_turns",11 "turn_lengths", sampled tokens per turn: the join key against a harness trace12 "role", "agent" | "auxiliary" | "discarded"13 "validation": [...]}14 ],15 "validation": [...], rollout-level findings16 "trainable": bool the single gate: did anything survive17 }18 19Sequences are labelled rather than filtered. A caller that silently drops rows cannot be20distinguished from one that had none to drop, and "the group quietly shrank" is far harder to21diagnose than "three rows were labelled auxiliary". The trainer picks by `role`.22 23ROLE ASSIGNMENT is structural first, heuristic only as a tiebreak. A rollout's real work is the24longest path with tool access; a title generator or a summariser is a short toolless root. The old25approach (matching known system-prompt strings) needed a new entry per harness and failed silently26on the harnesses nobody had profiled yet.27"""28 29from __future__ import annotations30 31from typing import Any32 33from .validate import check_rollout, check_sequence34 35AGENT, AUXILIARY, DISCARDED = "agent", "auxiliary", "discarded"36 37 38def _assign_roles(graph, sequences, *, trainable_capture: bool = True) -> list[str]:39 """Label each flattened path. Purely structural.40 41 - a path ending in a discarded node (a sibling that never continued) is a retry42 - a path whose turns carry a TOOL MANIFEST is the agent working43 - anything else is auxiliary: title generators, summarisers, classifiers44 45 **Multiple agent paths are normal and all of them are trainable.** A harness that rewrites its46 system prompt mid-run breaks token-prefix continuity and starts a new root, even though the47 conversation continued: claude-code does exactly this, swapping a 12118-char system prompt for a48 12541-char one at call 6 while its message list grows 2 -> 26 unbroken. Both roots are the agent49 doing real work on the same task and earn the same reward.50 51 An earlier version kept only the single longest tool-using path. On opencode that was52 indistinguishable from correct (its second root really is a title generator), but on claude-code53 it silently discarded 6 genuine agent turns. Tool access alone is the honest signal: aux calls54 essentially never pass a tool manifest, and coding agents essentially always do.55 """56 discarded_ids = {n.node_id for n in graph.discarded_nodes()}57 live = [58 (i, s) for i, s in enumerate(sequences) if s.node_ids[-1] not in discarded_ids59 ]60 61 # Tools only DISCRIMINATE when some paths have them and others do not. That is the opencode62 # shape: an agent chain with a manifest plus a toolless title generator.63 #64 # Some harnesses never send a manifest at all. terminus-2 parses tool calls out of raw model65 # text, so every one of its paths has n_tools == 0. Applying the tool rule there labels the whole66 # rollout auxiliary, and an earlier "keep the longest" fallback then kept exactly ONE of its 1367 # turns -- a harness-trace cross-check caught it as `captured [263] vs trace [167, ..., 136]`.68 #69 # So: if nothing in the rollout uses tools, tools carry no signal and every live path is agent70 # work. If something does, the toolless paths really are auxiliary.71 any_tools = any(graph.get(nid).n_tools > 0 for _, s in live for nid in s.node_ids)72 73 roles = []74 for i, seq in enumerate(sequences):75 if seq.node_ids[-1] in discarded_ids:76 roles.append(DISCARDED)77 continue78 if not any_tools:79 # `n_trainable` is the tiebreak only where it can mean something. On an eval endpoint it80 # is 0 for every sequence by construction, so using it there labelled EVERY path auxiliary81 # for a harness that sends no tool manifest — terminus-2 parses tool calls out of raw text82 # — which emptied `result.turns` and mistagged the conversations, on a rollout that had83 # captured perfectly well. With neither tools nor token counts to discriminate on, a live84 # path is the agent working: that is the same conclusion the tool rule reaches, and the85 # cost of being wrong is a mislabelled trace rather than a mistrained token, since nothing86 # here is trainable anyway.87 usable = seq.n_trainable if trainable_capture else bool(seq.node_ids)88 roles.append(AGENT if usable else AUXILIARY)89 continue90 has_tools = any(graph.get(nid).n_tools > 0 for nid in seq.node_ids)91 roles.append(AGENT if has_tools else AUXILIARY)92 return roles93 94 95def export_session(96 session,97 *,98 include_discarded: bool = False,99 include_messages: bool = False,100 capture_level: str = "tokens",101) -> dict[str, Any]:102 """Build the document for one rollout: training rows when there are any, the trace always.103 104 Args:105 session:106 The live capture session.107 include_discarded (`bool`, *optional*, defaults to `False`):108 Keep paths that led nowhere (retries, resamples) as labelled rows.109 include_messages (`bool`, *optional*, defaults to `False`):110 Add each turn's request messages, tools and response message. Off by default because it111 multiplies payload size by the full conversation text, on when you need to feed TRL's112 `TraceEntry` contract or measure re-tokenization skew.113 capture_level (`str`, *optional*, defaults to `"tokens"`):114 What the upstream could return. Below `tokens` this is an eval rollout: the trace and the115 graph structure are complete, no row is trainable, and the token arrays are empty — by116 construction rather than by accident.117 118 Returns:119 `dict[str, Any]`: The rollout document.120 """121 graph = session.graph122 rollout_report = check_rollout(123 graph,124 capture_level=capture_level,125 budget_stop_count=getattr(session, "budget_stop_count", 0),126 )127 trainable_capture = (128 capture_level == "tokens" and getattr(session, "purpose", "auto") != "eval"129 )130 131 # Sequences are built at every level, because they are the rollout's STRUCTURE — which calls132 # belong to which conversation, which path is the agent working, which branches died — and that133 # structure is real whether or not token ids came back. Everything that reads a rollout as a134 # trace (`conversations_from_document`, `turns_from_document`, the UI transcript) walks these135 # rows, so dropping them on an eval rollout deletes exactly the payload an eval rollout is for.136 #137 # What is withheld at a lower level is the *training* claim: `check_sequence` is skipped, since138 # every one of its findings is about token arrays that are empty by design, and no row is ever139 # marked trainable. The token fields stay as the empty lists the graph produced.140 sequences = graph.sequences()141 roles = _assign_roles(graph, sequences, trainable_capture=trainable_capture)142 143 rows: list[dict[str, Any]] = []144 for seq, role in zip(sequences, roles):145 if role == DISCARDED and not include_discarded:146 continue147 report = check_sequence(seq) if trainable_capture else None148 rows.append(149 {150 "role": role,151 "root_id": seq.root_id,152 "node_ids": seq.node_ids,153 "n_turns": seq.n_turns,154 "prompt_len": seq.prompt_len,155 "n_trainable": seq.n_trainable if trainable_capture else 0,156 "turn_lengths": seq.turn_lengths(),157 "input_ids": seq.input_ids,158 "loss_mask": seq.loss_mask159 if trainable_capture160 else [0] * len(seq.input_ids),161 "logprobs": seq.logprobs,162 "trainable": bool(report and report.ok and role == AGENT),163 "validation": [str(f) for f in report.findings] if report else [],164 }165 )166 167 # Every call in arrival order, including the ones excluded from training. This is what an168 # external trace can be reconciled against: the harness logged every LLM call it made,169 # so comparing only the surviving path would report a mismatch on any rollout that retried.170 discarded_ids = {n.node_id for n in graph.discarded_nodes()}171 turns = [172 {173 "node_id": node.node_id,174 "index": node.index,175 "root_id": graph.root_of(node.node_id),176 "n_sampled": len(node.sampled_ids),177 "n_prompt": len(node.prompt_ids),178 "n_tools": node.n_tools,179 "finish_reason": node.finish_reason,180 "harness_session_id": node.harness_session_id,181 # Travels with the turn because it decides whether a trainer's recompute is comparable to182 # the captured logprob at all. See `TurnNode.sampling_params`.183 "sampling_params": node.sampling_params,184 "requested_sampling_params": node.requested_sampling_params,185 "sampled_logprobs": node.sampled_logprobs,186 "discarded": node.node_id in discarded_ids,187 **(188 {189 "request_messages": node.request_messages,190 "request_tools": node.request_tools,191 "response_message": node.response_message,192 }193 if include_messages194 else {}195 ),196 }197 for node in graph.nodes()198 ]199 200 trainable_rows = [r for r in rows if r["trainable"]]201 return {202 "session_id": session.session_id,203 "metadata": session.metadata,204 "budget_stop_count": getattr(session, "budget_stop_count", 0),205 "turns": turns,206 "stats": {207 **graph.stats(),208 "n_sequences": len(rows),209 "n_trainable_sequences": len(trainable_rows),210 "n_trainable_tokens": sum(r["n_trainable"] for r in trainable_rows),211 },212 "sequences": rows,213 "validation": [str(f) for f in rollout_report.findings] + session.findings,214 "trainable": bool(trainable_rows) and rollout_report.ok,215 # Why this rollout is or is not trainable, travelling with the data rather than living only216 # in the server's log. A consumer that reads `trainable` alone is safe; one that wants to217 # explain an empty `sequences` list to a human needs these two.218 "capture_level": capture_level,219 "rollout_type": "train" if trainable_capture else "eval",220 }221 222 223def summarise(document: dict[str, Any]) -> str:224 """One screen of text. What you actually read after a rollout."""225 stats = document["stats"]226 lines = [227 f"session {document['session_id']} "228 f"{document.get('rollout_type', 'train')} "229 f"trainable={document['trainable']}",230 f" graph: {stats['n_turns']} turns, {stats['n_roots']} roots, "231 f"{stats['n_forks']} forks, {stats['n_discarded']} discarded",232 f" training: {stats['n_trainable_sequences']} sequence(s), "233 f"{stats['n_trainable_tokens']} trainable tokens",234 ]235 for row in document["sequences"]:236 lines.append(237 f" [{row['role']:<10}] turns={row['n_turns']:<3} prompt={row['prompt_len']:<6} "238 f"len={len(row['input_ids']):<6} trainable={row['n_trainable']:<5} "239 f"turn_lengths={row['turn_lengths']}"240 )241 for finding in row["validation"]:242 if not finding.startswith("[INFO]"):243 lines.append(f" {finding}")244 for finding in document["validation"]:245 lines.append(f" {finding}")246 return "\n".join(lines)247 