153 lines
5.5 KiB
Python
153 lines
5.5 KiB
Python
#!/usr/bin/env python3
|
|
"""Compare original and whole-layer-curve Frontier against the BC-8 real pilot."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import statistics
|
|
from collections import Counter
|
|
from pathlib import Path
|
|
|
|
|
|
ROOT = Path(__file__).resolve().parent
|
|
REPO = ROOT.parents[1]
|
|
CONFIGS = ("tp1_mns16", "tp2_mns16", "tp4_mns16", "tp4_mns32")
|
|
REPLAY = ROOT / "replay/bc8"
|
|
ORIGINAL = REPO / "runs/frontier-knee-sweep-v0/raw/fixed/rho0p02"
|
|
REAL = REPO / "runs/frontier-pilot-v0/results/pilot-surface.json"
|
|
|
|
|
|
def find_one(root: Path, name: str) -> Path:
|
|
matches = list(root.glob(f"**/{name}"))
|
|
if len(matches) != 1:
|
|
raise ValueError(f"expected one {name} under {root}, got {matches}")
|
|
return matches[0]
|
|
|
|
|
|
def summarize_sim(root: Path) -> dict:
|
|
metrics = json.loads(find_one(root, "system_metrics.json").read_text())
|
|
ledger = find_one(root, "frontier_stage_batch_ledger.jsonl")
|
|
histogram: Counter[int] = Counter()
|
|
service_ms: dict[int, set[float]] = {}
|
|
for line in ledger.read_text().splitlines():
|
|
row = json.loads(line)
|
|
tokens = row["request_num_tokens"]
|
|
if tokens and all(int(value) == 1 for value in tokens):
|
|
batch = len(tokens)
|
|
histogram[batch] += 1
|
|
service_ms.setdefault(batch, set()).add(
|
|
float(row["execution_time"]["model_time_ms"])
|
|
)
|
|
return {
|
|
"tpot_mean_ms": metrics["tpot_statistics"]["mean"],
|
|
"ttft_mean_ms": metrics["ttft_statistics"]["mean"],
|
|
"e2e_mean_ms": metrics["request_e2e_time_statistics"]["mean"],
|
|
"decode_batch_histogram": dict(sorted(histogram.items())),
|
|
"decode_service_ms": {
|
|
str(batch): sorted(values) for batch, values in sorted(service_ms.items())
|
|
},
|
|
"decode_b_gt_1_fraction": sum(
|
|
count for batch, count in histogram.items() if batch > 1
|
|
)
|
|
/ sum(histogram.values()),
|
|
}
|
|
|
|
|
|
def main() -> None:
|
|
real_rows = [
|
|
row
|
|
for row in json.loads(REAL.read_text())
|
|
if row["load"] == "p4" and row["canonical"]
|
|
]
|
|
points = []
|
|
for config in CONFIGS:
|
|
real = [row for row in real_rows if row["config"] == config]
|
|
original = summarize_sim(ORIGINAL / config)
|
|
corrected = summarize_sim(REPLAY / "raw" / config)
|
|
usage = json.loads((REPLAY / "raw" / config / "usage.json").read_text())
|
|
points.append(
|
|
{
|
|
"config": config,
|
|
"tp": int(config[2]),
|
|
"mns": int(config.split("mns")[1]),
|
|
"real_tpot_mean_ms": statistics.fmean(
|
|
row["tpot_mean_ms"] for row in real
|
|
),
|
|
"original": original,
|
|
"whole_curve": corrected,
|
|
"original_tpot_residual_ms": (
|
|
original["tpot_mean_ms"]
|
|
- statistics.fmean(row["tpot_mean_ms"] for row in real)
|
|
),
|
|
"whole_curve_tpot_residual_ms": (
|
|
corrected["tpot_mean_ms"]
|
|
- statistics.fmean(row["tpot_mean_ms"] for row in real)
|
|
),
|
|
"usage": usage,
|
|
}
|
|
)
|
|
|
|
def ranking(key):
|
|
return [
|
|
point["config"]
|
|
for point in sorted(points, key=lambda point: (key(point), point["config"]))
|
|
if point["mns"] == 16
|
|
]
|
|
|
|
real_ranking = ranking(lambda point: point["real_tpot_mean_ms"])
|
|
original_ranking = ranking(lambda point: point["original"]["tpot_mean_ms"])
|
|
corrected_ranking = ranking(lambda point: point["whole_curve"]["tpot_mean_ms"])
|
|
payload = {
|
|
"schema": "frontier-decode-grid-bc8-verdict.v1",
|
|
"points": points,
|
|
"mns16_rankings_fast_to_slow": {
|
|
"real": real_ranking,
|
|
"original": original_ranking,
|
|
"whole_curve": corrected_ranking,
|
|
},
|
|
"whole_curve_restores_real_mns16_ranking": corrected_ranking == real_ranking,
|
|
"all_corrected_batches_within_measured_support": all(
|
|
max(map(int, point["whole_curve"]["decode_batch_histogram"])) <= 8
|
|
for point in points
|
|
),
|
|
"decision": (
|
|
"STATIC_WHOLE_CURVE_FIXES_BC8"
|
|
if corrected_ranking == real_ranking
|
|
else "ESCALATE_EVENT_LEVEL_STATE"
|
|
),
|
|
}
|
|
output = ROOT / "results/bc8-replay-verdict.json"
|
|
output.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")
|
|
|
|
lines = [
|
|
"# BC-8 whole-layer decode-curve replay",
|
|
"",
|
|
"| Config | Real TPOT | Original sim | Whole-curve sim | Residual before -> after | b>1 |",
|
|
"|---|---:|---:|---:|---:|---:|",
|
|
]
|
|
for point in points:
|
|
lines.append(
|
|
f"| {point['config']} | {point['real_tpot_mean_ms']:.3f} | "
|
|
f"{point['original']['tpot_mean_ms']:.3f} | "
|
|
f"{point['whole_curve']['tpot_mean_ms']:.3f} | "
|
|
f"{point['original_tpot_residual_ms']:+.3f} -> "
|
|
f"{point['whole_curve_tpot_residual_ms']:+.3f} | "
|
|
f"{point['whole_curve']['decode_b_gt_1_fraction']:.2%} |"
|
|
)
|
|
lines.extend(
|
|
[
|
|
"",
|
|
f"- Real MNS16: `{' < '.join(real_ranking)}`",
|
|
f"- Original: `{' < '.join(original_ranking)}`",
|
|
f"- Whole curve: `{' < '.join(corrected_ranking)}`",
|
|
f"- Decision: **{payload['decision']}**.",
|
|
"",
|
|
]
|
|
)
|
|
(ROOT / "results/bc8-replay-verdict.md").write_text("\n".join(lines))
|
|
print(json.dumps({"output": str(output), "decision": payload["decision"]}))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|