From 61b903f7b3106aafc911e95179e8745721ffeff7 Mon Sep 17 00:00:00 2001 From: Kral Date: Mon, 5 Oct 2026 12:03:54 +0200 Subject: [PATCH] Acceptance filter: score, end reason, harness errors, loop trimming, repair marker Co-Authored-By: Claude Sonnet 5.5 --- train/accept.py | 141 ++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 141 insertions(+) create mode 100644 train/accept.py diff --git a/train/accept.py b/train/accept.py new file mode 100644 index 0000000..d341454 --- /dev/null +++ b/train/accept.py @@ -0,0 +1,141 @@ +"""Acceptance filter for trajectory records (stage 2). + + python3 train/accept.py [--traj runs/traj] [--min-score 80] [--out runs/traj/accepted.jsonl] + +A record is accepted when: it exists and the run had no harness error and no setup failure; end reason is +"report" (the model ended itself); the score reaches the threshold; no raw tool result shows a harness +problem (RFC lock, session or HTTP error). Loop parts are trimmed: when a call repeats with the identical +arguments and the identical result right after itself, only the first pair stays (the world did not change, +so the later context is the same). Repair trajectories (a failed write followed by a good write) stay and +are marked. Output: one JSON record per line (messages trimmed, metadata, reasons). +""" +import argparse +import glob +import json +import os +import re + +HARNESS_ERRORS = re.compile(r"Concurrent call detected|\[LOCK\]|HTTP (?:4|5)\d\d|Connection refused|timed out|" + r"session (?:dropped|expired)|Traceback|RFC_(?:COMMUNICATION|SYSTEM)_FAILURE", re.I) +ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + + +def tool_pairs(messages): + """[(assistant index, [tool message indices])] for assistant messages with tool calls.""" + out = [] + for i, m in enumerate(messages): + if m["role"] == "assistant" and m.get("tool_calls"): + j = i + 1 + while j < len(messages) and messages[j]["role"] == "tool": + j += 1 + out.append((i, list(range(i + 1, j)))) + return out + + +def call_key(messages, i, tools): + calls = messages[i]["tool_calls"] + sig = [(c["function"]["name"], json.dumps(_args(c), sort_keys=True)) for c in calls] + return sig, [messages[t]["content"] for t in tools] + + +def _args(c): + a = c["function"].get("arguments") or "{}" + try: + return json.loads(a) if isinstance(a, str) else a + except ValueError: + return {"_raw": a} + + +def trim_loops(messages): + """Remove an assistant tool call + its results when it repeats the previous call and result exactly.""" + drop, last_key, last_end, removed = set(), None, -1, 0 + for i, tools in tool_pairs(messages): + key = call_key(messages, i, tools) + end = (tools[-1] if tools else i) + 1 + if key == last_key and i == last_end: # directly after the same call with the same result + drop.update([i] + tools) + removed += 1 + else: + last_key = key + last_end = end + kept = [m for k, m in enumerate(messages) if k not in drop] + return kept, removed + + +def is_failed_write(content): + c = content.replace(" ", "") + return c.startswith("ERROR:") or '"success":false' in c + + +def has_repair(messages): + """A write tool result that failed, followed later by a write result that succeeded.""" + failed = False + names = {} + for m in messages: + if m["role"] == "assistant": + for c in m.get("tool_calls") or []: + names[c.get("id")] = c["function"]["name"] + elif m["role"] == "tool" and re.match(r"(sap_push|sap_activate|write_source|activate)", + names.get(m.get("tool_call_id"), "")): + if is_failed_write(m["content"]): + failed = True + elif failed and '"success":true' in m["content"].replace(" ", ""): + return True + return False + + +def judge(rec, summary, min_score): + """Return (accepted, reasons).""" + md = rec["metadata"] + why = [] + if summary.get("harness_error") or md.get("harness_error"): + why.append("harness_error") + if md.get("setup_failed"): + why.append("setup_failed") + if md.get("end_reason") != "report": + why.append(f"end_reason={md.get('end_reason')}") + if md.get("score") is None or md["score"] < min_score: + why.append(f"score={md.get('score')}") + bad = [r for r in rec["tool_results_raw"] if HARNESS_ERRORS.search(r.get("result") or "") + and not r.get("tool", "").startswith("sap_run_unit")] + if bad: + why.append("harness_text_in_result") + if not (rec["messages"][-1].get("content") or "").strip(): + why.append("no_final_report") + return not why, why + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--traj", default=os.path.join(ROOT, "runs", "traj")) + ap.add_argument("--min-score", type=float, default=80) + ap.add_argument("--out") + a = ap.parse_args() + summary = {} + sp = os.path.join(a.traj, "summary.jsonl") + if os.path.exists(sp): + for r in map(json.loads, open(sp)): + summary[r.get("run_dir")] = r + out = a.out or os.path.join(a.traj, "accepted.jsonl") + n_all = n_ok = 0 + rejected = {} + with open(out, "w") as f: + for p in sorted(glob.glob(os.path.join(a.traj, "*", "record.json"))): + rec = json.load(open(p)) + n_all += 1 + ok, why = judge(rec, summary.get(os.path.basename(os.path.dirname(p)), {}), a.min_score) + if not ok: + for w in why: + rejected[w.split("=")[0]] = rejected.get(w.split("=")[0], 0) + 1 + continue + msgs, removed = trim_loops(rec["messages"]) + n_ok += 1 + f.write(json.dumps({"id": f"{rec['task']['id']}_r{rec['run']}", "run_dir": os.path.dirname(p), + "task": rec["task"], "teacher": rec["teacher"], "metadata": rec["metadata"], + "messages": msgs, "tools": rec["tools"], "loop_pairs_removed": removed, + "repair": has_repair(msgs)}) + "\n") + print(f"{n_ok} accepted of {n_all} records -> {out}; rejected by reason: {rejected}") + + +if __name__ == "__main__": + main()