Files
abap-llm/train/accept.py
2026-10-05 12:03:54 +02:00

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()