237 lines
11 KiB
Python
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()
|