Acceptance filter: score, end reason, harness errors, loop trimming, repair marker
Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
This commit is contained in:
141
train/accept.py
Normal file
141
train/accept.py
Normal file
@@ -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()
|
||||
Reference in New Issue
Block a user