142 lines
5.6 KiB
Python
142 lines
5.6 KiB
Python
"""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()
|