ounce100m-code / kernels /p6_cli_probe.py
Cion-lab's picture
p6_cli_probe v3: build metric keys from the YAML's own metric/aggregation fields and rank candidate config files -- the first pass guessed metric_uri and picked a mmlu subject file
5f5b188 verified
Raw History Blame Contribute Delete
8.66 kB
"""Phase 6 contract check on free CPU: the CLI flags we pass, and the metric keys we promise.
E-047 and E-048 both had the same shape -- a documented lm-eval behaviour that turned out not to be true of
0.4.13 (`--no-deps` leaves the package unimportable; `--tasks list` is not a listing command). Two
assumptions of exactly that kind are still unverified and both would only be discovered after the model is
trained and published:
1. `run_task()` builds a command line of eight flags. If 0.4.13's CLI (which was restructured into
`lm_eval/_cli/harness.py` subcommands -- visible in E-048's traceback) renames or drops one of them, the
failure lands in Phase 6, not now. So the help text of the harness's own `run` command is read and every
flag `run_task()` passes is checked against it.
2. `PRIMARY_METRIC` promises `acc,none` and `exact_match,flexible-extract` will be the keys of
`results.json`, and docs/05 §7 promises which split each task is scored on and whether it declares
`num_fewshot`. E-048 could not read that from `TaskManager().task_index` (an `Entry` whose config is not
populated), so it is read from the source of truth instead: the YAMLs shipped inside the wheel, with
`include:` chains resolved.
**No eval data is touched, and none can be from this code**: it opens `.yaml` files inside site-packages and
prints `--help`. No task object is constructed, no `datasets.load_dataset` call is reachable here, no
benchmark row is read -- §3.3 and Phase 6's ordering both stay intact.
"""
import glob
import json
import os
import subprocess
import sys
PIN = "0.4.13"
TASKS = ["arc_challenge", "arc_easy", "hellaswag", "mmlu", "piqa", "truthfulqa_mc1",
"truthfulqa_mc2", "winogrande", "gsm8k"]
PRIMARY_METRIC = {"arc_challenge": "acc,none", "arc_easy": "acc,none", "hellaswag": "acc,none",
"mmlu": "acc,none", "piqa": "acc,none", "truthfulqa_mc1": "acc,none",
"truthfulqa_mc2": "acc,none", "winogrande": "acc,none",
"gsm8k": "exact_match,flexible-extract"}
# Every flag run_benchmarks.py's run_task() passes, in its own words, so the two cannot drift apart.
OURS = ["--model", "--model_args", "--tasks", "--batch_size", "--seed", "--output_path",
"--log_samples", "--num_fewshot", "--limit"]
rc = subprocess.run([sys.executable, "-m", "pip", "install", "--user", "--quiet",
"lm-eval==" + PIN]).returncode
print("pip rc", rc, flush=True)
def sh(argv, timeout=600):
p = subprocess.run(argv, capture_output=True, text=True, timeout=timeout)
return p.returncode, (p.stdout or "") + (p.stderr or "")
rc, top = sh([sys.executable, "-m", "lm_eval", "--help"])
print("TOP_HELP rc", rc, "chars", len(top), flush=True)
# Find the subcommand that evaluates (0.4.13's harness registered `run`; fall back to scanning --help).
subs = []
for line in top.splitlines():
s = line.strip()
if s.startswith("run") or (s and s.split()[0] in {"run", "eval", "simple"}):
subs.append(s.split()[0])
subs = list(dict.fromkeys(subs)) or ["run"]
helps = {}
for s in subs:
rc, h = sh([sys.executable, "-m", "lm_eval", s, "--help"])
helps[s] = (rc, h)
print("SUBCOMMAND", s, "help rc", rc, "chars", len(h), flush=True)
best = max(helps.items(), key=lambda kv: len(kv[1][1]))[1][1]
missing = [f for f in OURS if f not in best]
print("FLAGS_CHECKED", json.dumps({f: (f in best) for f in OURS}), flush=True)
print("FLAGS_MISSING", missing, flush=True)
import yaml # shipped with the image and a dependency of lm_eval
# The user-site directory pip created a minute ago is not on *this* process's sys.path: `site` adds it only
# if it exists at interpreter start, which it did not in a fresh container (the same reason the earlier probe
# only ever touched lm_eval through subprocesses). Ask a child where the package lives and read files from
# the parent -- reading a YAML needs no import.
rc, out = sh([sys.executable, "-c", "import lm_eval, os; print(os.path.dirname(lm_eval.__file__))"], 300)
loc = next((l.strip() for l in (out or "").splitlines() if l.strip().endswith("lm_eval")), "")
if not loc:
print("LM_EVAL_PATH_FAILED", repr((out or "")[-400:]), flush=True)
raise SystemExit(8)
tasks_dir = os.path.join(loc, "tasks")
def load_yaml(path):
try:
with open(path, encoding="utf-8") as fh:
d = yaml.safe_load(fh)
return d if isinstance(d, dict) else {}
except Exception:
return {}
def resolve(path, depth=0, seen=None):
"""Merge a task YAML with whatever it `include`s, the way lm-eval does."""
seen = seen or set()
if path in seen or depth > 6 or not os.path.exists(path):
return {}
seen.add(path)
d = load_yaml(path)
inc = d.pop("include", None)
if inc:
base = os.path.normpath(os.path.join(os.path.dirname(path), inc))
parent = resolve(base + ".yaml", depth + 1, seen) or resolve(base, depth + 1, seen)
merged = dict(parent)
merged.update(d)
return merged
return d
def candidates(name):
"""Every YAML that could be the config `--tasks <name>` means, best-guess first.
The first pass picked `mmlu/flan_cot_fewshot/mmlu_formal_logic.yaml` for `mmlu` -- a subject file of a
different variant -- so ranking is explicit now, and the winner is printed alongside the raw metric
block so a wrong pick is visible rather than silent.
"""
hits = glob.glob(os.path.join(tasks_dir, "**", name + ".yaml"), recursive=True)
exact = [h for h in hits if os.path.basename(h) == name + ".yaml"]
tmpl = [h for h in hits if "_default_template" in h or h.endswith("default.yaml")]
named = []
for p in glob.glob(os.path.join(tasks_dir, "**", "*.yaml"), recursive=True):
if os.path.basename(p).startswith("_"):
continue
d = load_yaml(p)
if d.get("task") == name or d.get("group") == name:
named.append(p)
ordered = []
for h in exact + named + tmpl + hits:
if h not in ordered:
ordered.append(h)
return ordered
def raw_metrics(path):
"""The metric block exactly as the YAML spells it, with the key names lm-eval actually uses."""
try:
with open(path, encoding="utf-8") as fh:
txt = fh.read()
except Exception:
return "UNREADABLE"
lines, out = txt.splitlines(), []
for i, l in enumerate(lines):
if "metric" in l.lower() or "include" in l.lower() or "num_fewshot" in l:
out.append(l.strip()[:90])
if len(out) > 7:
break
return " | ".join(out) or "no metric/include lines"
print("TASKS_DIR", tasks_dir, os.path.isdir(tasks_dir), flush=True)
report = {}
for t in TASKS:
cands = candidates(t)
if not cands:
report[t] = {"file": None}
print("YAML %-16s NO CANDIDATE FILES" % t, flush=True)
continue
y = cands[0]
c = resolve(y)
ml = c.get("metric_list") or []
names = [(m.get("metric") if isinstance(m, dict) else m) for m in ml]
aggs = [(m.get("aggregation") if isinstance(m, dict) else None) for m in ml]
keys = ["%s,%s" % (n, a) for n, a in zip(names, aggs)]
report[t] = {"file": os.path.relpath(y, tasks_dir), "n_candidates": len(cands),
"others": [os.path.relpath(x, tasks_dir) for x in cands[1:4]],
"num_fewshot": c.get("num_fewshot"), "test_split": c.get("test_split"),
"validation_split": c.get("validation_split"), "fewshot_split": c.get("fewshot_split"),
"metric_names": names, "metric_keys": keys, "dataset_path": c.get("dataset_path"),
"dataset_name": c.get("dataset_name"),
"metric_declared": PRIMARY_METRIC[t] in keys}
print("YAML %-16s %-38s fewshot=%-6s test=%-9s val=%-10s fs=%-7s keys=%s declared=%s" % (
t, report[t]["file"], c.get("num_fewshot"), c.get("test_split"), c.get("validation_split"),
c.get("fewshot_split"), keys, report[t]["metric_declared"]), flush=True)
print(" RAW %s" % raw_metrics(y)[:220], flush=True)
ok = not missing and all(report[t].get("metric_declared") for t in TASKS) and \
all(report[t].get("file") for t in TASKS)
print("CLI_PROBE_JSON", json.dumps({"subcommands": {k: v[0] for k, v in helps.items()},
"flags_missing": missing, "tasks": report}, default=str)[:5000],
flush=True)
print("VERDICT P6_CLI_PROBE", "OK" if ok else "PROBLEM", "| missing_flags", missing,
"| metrics_not_declared", [t for t in TASKS if not report[t].get("metric_declared")], flush=True)
raise SystemExit(0 if ok else 7)