Files
abap-llm/train/analyze.py

237 lines
11 KiB
Python

"""Analysis of the accepted DeepSeek trajectories (Opus item B, 2026-10-06): repair taxonomy, near duplicates,
empty_response, and the behaviors the teacher shows and Qwen lacks.
python3 train/analyze.py writes runs/analysis/analysis.json and prints the numbers
"""
import glob
import json
import os
import re
import statistics
import sys
from collections import Counter
ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, ROOT)
sys.path.insert(0, os.path.join(ROOT, "train"))
import accept as acc # noqa: E402
from harness import mix # noqa: E402
from harness.generator import RESERVED # noqa: E402
WRITE = ("sap_push_source", "sap_push_element", "sap_push_message", "sap_activate", "sap_create_object")
SEARCH_LIKE = ("sap_search_object", "sap_pull_source", "sap_object_members", "sap_object_structure", "sap_element_info",
"sap_usage_references", "sap_sql_query", "sap_inactive_objects")
SIG = re.compile(r"(?:IMPORTING|EXPORTING|RETURNING|CHANGING)[^.]*?\bTYPE\s+(?:c|n|p|x)\s+LENGTH\b|\bVALUE\([^)]*\)\s+TYPE\s+(?:c|n|p|x)\s+LENGTH\b", re.I | re.S)
def pct(v, p):
v = sorted(v)
return v[min(len(v) - 1, int(p * len(v)))] if v else None
def calls_of(messages):
"""[(index, tool name, args dict, result text)] in order."""
res = {m.get("tool_call_id"): m["content"] for m in messages if m["role"] == "tool"}
out = []
for i, m in enumerate(messages):
if m["role"] != "assistant":
continue
for c in m.get("tool_calls") or []:
a = c["function"].get("arguments") or "{}"
try:
a = json.loads(a) if isinstance(a, str) else a
except ValueError:
a = {}
out.append((i, c["function"]["name"], a, res.get(c.get("id"), "")))
return out
def failed(tool, text):
t = (text or "").replace(" ", "")
return tool in WRITE and (t.startswith("ERROR:") or '"success":false' in t)
def error_text(text):
try:
o = json.loads(text.replace("ERROR: ", "", 1) if text.startswith("ERROR:") else text)
except ValueError:
return text[:200]
if not isinstance(o, dict):
return text[:200]
msgs = [m.get("message", "") for m in (o.get("activation") or {}).get("messages", []) if m.get("severity") == "E"]
out = " | ".join(msgs) or str(o.get("error") or o.get("message") or "")
sc = o.get("syntaxCheck")
if sc:
out += " | syntaxCheck: " + "; ".join(m.get("message", "") for m in sc.get("messages", []))
return out[:400]
def classify(src, err):
"""Labels of a failed write (a write can match more than one)."""
labels = []
if src and SIG.search(src):
labels.append("type_c_length_in_signature")
low = (err or "").lower()
if "longer than the allowed 30 characters" in low or "30 characters" in low:
labels.append("name_over_30")
words = {w for w in re.findall(r"[A-Za-z_]+", err or "")}
if "is not valid" in low or "reserved" in low or any(w.upper() in RESERVED and w.upper() in (err or "").upper() for w in words if len(w) > 3 and w.isupper()):
if re.search(r"\b(?:" + "|".join(sorted(RESERVED, key=len, reverse=True)) + r")\b", err or "", re.I) or "reserved" in low:
labels.append("reserved_word")
return labels
def norm(err):
e = re.sub(r"\bZ[0-9A-Z]{7}_\w+", "<OBJ>", err or "")
e = re.sub(r"\b[A-Z][A-Z0-9_]{5,}\b", "<NAME>", e)
e = re.sub(r"\d+", "N", e)
return e[:110]
def behaviors(messages):
cs = calls_of(messages)
out = {"first_write": None, "searches_before_first_write": 0, "max_consecutive_search_object": 0, "calls": len(cs),
"failed_writes": 0, "repaired": False, "calls_to_repair": None, "same_push_repeats": 0}
run_so = 0
first_fail = None
last_src = {}
for n, (i, tool, a, res) in enumerate(cs):
if tool == "sap_search_object":
run_so += 1
out["max_consecutive_search_object"] = max(out["max_consecutive_search_object"], run_so)
else:
run_so = 0
if tool in ("sap_push_source", "sap_push_element", "sap_push_message") and out["first_write"] is None:
out["first_write"] = n
if out["first_write"] is None and tool in SEARCH_LIKE:
out["searches_before_first_write"] += 1
if tool == "sap_push_source":
key = (a.get("objectName"), a.get("includeType"))
if last_src.get(key) == a.get("source") and a.get("source"):
out["same_push_repeats"] += 1
last_src[key] = a.get("source")
if failed(tool, res):
out["failed_writes"] += 1
if first_fail is None:
first_fail = (n, a.get("objectName"), a.get("source"))
elif first_fail and tool in ("sap_push_source", "sap_push_element") and '"success":true' in (res or "").replace(" ", "") \
and a.get("objectName") == first_fail[1] and a.get("source") != first_fail[2] and not out["repaired"]:
out["repaired"] = True
out["calls_to_repair"] = n - first_fail[0]
return out
def shingles(text, k=5):
t = re.sub(r"\bz[0-9a-z]{7}_", "", text.lower())
w = re.findall(r"\w+", t)
return {" ".join(w[i:i + k]) for i in range(max(len(w) - k + 1, 0))}
def final_sources(messages):
"""Last pushed source per (object, include): the trajectory's own code."""
last = {}
for i, tool, a, res in calls_of(messages):
if tool == "sap_push_source" and a.get("source") and '"success":true' in (res or "").replace(" ", ""):
last[(a.get("objectName"), a.get("includeType"))] = a["source"]
return "\n".join(last.values())
def main():
rows = [json.loads(l) for l in open(os.path.join(ROOT, "runs", "traj", "summary.jsonl"))]
acc_recs, empty_runs, all_turn_out = [], [], []
for r in rows:
p = os.path.join(ROOT, "runs", "traj", r.get("run_dir") or "-", "record.json")
if not os.path.exists(p):
continue
rec = json.load(open(p))
ok, why = acc.judge(rec, r, 80)
md = rec["metadata"]
if ok:
acc_recs.append((r, rec))
all_turn_out += [u.get("completion_tokens", 0) for u in md.get("turn_usage", [])]
if md.get("end_reason") == "empty_response":
empty_runs.append((r, rec))
out = {"accepted": len(acc_recs), "runs": len(rows)}
# ---- repair taxonomy
fail_ct, labels_ct, traj_with, repaired_with, top = 0, Counter(), Counter(), Counter(), Counter()
per_label_examples = {}
kinds = Counter()
for r, rec in acc_recs:
seen = {}
for i, tool, a, res in calls_of(rec["messages"]):
if failed(tool, res):
fail_ct += 1
err = error_text(res)
top[norm(err)] += 1
for l in classify(a.get("source"), err):
labels_ct[l] += 1
seen[l] = True
per_label_examples.setdefault(l, (r["task"], err[:160]))
b = behaviors(rec["messages"])
for l in seen:
traj_with[l] += 1
if b["repaired"]:
repaired_with[l] += 1
out["repair_taxonomy"] = {"failed_writes_in_accepted": fail_ct, "step3_errors": {
l: {"failed_writes": labels_ct[l], "trajectories": traj_with[l], "trajectories_with_repair": repaired_with[l],
"example": per_label_examples.get(l)} for l in ("type_c_length_in_signature", "reserved_word", "name_over_30")},
"top_error_messages": top.most_common(15)}
# ---- teacher behaviors vs Qwen
def stats(recs):
b = [behaviors(x["messages"]) for x in recs]
had_fail = [x for x in b if x["failed_writes"]]
sb = [x["searches_before_first_write"] for x in b if x["first_write"] is not None]
return {"n": len(b), "with_a_failed_write": len(had_fail),
"repair_after_first_error": sum(x["repaired"] for x in had_fail), "median_calls_to_repair":
statistics.median([x["calls_to_repair"] for x in had_fail if x["calls_to_repair"] is not None] or [None]),
"median_searches_before_first_write": statistics.median(sb) if sb else None, "p90_searches_before_first_write": pct(sb, .9),
"max_consecutive_search_object": max([x["max_consecutive_search_object"] for x in b] or [0]),
"runs_with_no_write": sum(1 for x in b if x["first_write"] is None),
"runs_with_same_push_repeat": sum(1 for x in b if x["same_push_repeats"] > 0)}
out["teacher"] = stats([rec for _, rec in acc_recs])
qrecs = []
for p in glob.glob(os.path.join(ROOT, "runs", "local_qwen", "runs", "*", "record.json")):
qrecs.append(json.load(open(p)))
out["qwen_series_A"] = stats(qrecs)
out["qwen_series_A"]["note"] = "all 20 runs of series A (accepted and failed)"
# ---- near duplicates
sh = [(r["task"], r["run"], shingles(final_sources(rec["messages"])), shingles(rec["messages"][-1].get("content") or "")) for r, rec in acc_recs]
dup = []
for i in range(len(sh)):
for j in range(i + 1, len(sh)):
a, b = sh[i], sh[j]
for k, nm in ((2, "code"), (3, "report")):
if a[k] and b[k]:
jac = len(a[k] & b[k]) / len(a[k] | b[k])
if jac >= (0.8 if nm == "code" else 0.85):
dup.append({"a": f"{a[0]}_r{a[1]}", "b": f"{b[0]}_r{b[1]}", "what": nm, "jaccard": round(jac, 2), "same_task": a[0] == b[0]})
out["near_duplicates"] = {"pairs_flagged": len(dup), "same_task_pairs": sum(1 for d in dup if d["same_task"]), "examples": dup[:12]}
# ---- empty_response
er = []
for r, rec in empty_runs:
md = rec["metadata"]
tu = md.get("turn_usage", [])
last = [u.get("completion_tokens", 0) for u in tu[-3:]]
cs = calls_of(rec["messages"])
lastcall = cs[-1] if cs else None
er.append({"task": r["task"], "kind": mix.kind_of_task_dir(r["task"]), "category": rec["task"]["category"], "calls": md.get("tool_calls"),
"turns": len(tu), "last3_output_tokens": last, "last_tool": lastcall[1] if lastcall else None,
"last_tool_result_chars": len(lastcall[3]) if lastcall else None,
"context_tokens_last": (tu[-1].get("prompt_tokens") if tu else None)})
out["empty_response"] = {"runs": len(er), "of": len(rows), "by_kind": Counter(e["kind"] for e in er),
"by_last_tool": Counter(e["last_tool"] for e in er), "detail": er,
"output_tokens_of_normal_turns": {"p50": pct(all_turn_out, .5), "p90": pct(all_turn_out, .9),
"p99": pct(all_turn_out, .99), "p999": pct(all_turn_out, .999), "max": max(all_turn_out or [0]),
"turns": len(all_turn_out), "turns_over_16k": sum(1 for x in all_turn_out if x > 16000),
"turns_over_20k": sum(1 for x in all_turn_out if x > 20000)}}
json.dump(out, open(os.path.join(ROOT, "runs", "analysis", "analysis.json"), "w"), indent=1, default=dict)
print(json.dumps(out, indent=1, default=dict)[:6500])
if __name__ == "__main__":
main()