lalmodel-code / scripts /probe_report.py
Z User
feat(probe): 训练质量探针 v1 (同步自 main e6bfd12, 仅源码+节点配置)
bb9690c
Raw
History Blame Contribute Delete
8.84 kB
#!/usr/bin/env python3
"""lalmodel 质量探针 — 离线报告生成器
数据源 (二选一):
1. HF ckpts 仓: gasschina/lal-ckpts 的 probe_samples.jsonl (Space 每 200 步上传)
2. 本地文件: --jsonl path/to/probe_samples.jsonl (冒烟/调试)
输出 (默认 /home/z/my-project/download/):
quality_report.md 质量时间线: eval_loss 趋势 + 各 prompt 生成样本演进 + 指标表
quality_curve.png eval_loss / distinct-1,2 / loop / E[step] 曲线
用法:
python3 scripts/probe_report.py # 拉 HF 数据生成报告
python3 scripts/probe_report.py --jsonl local.jsonl # 本地数据
"""
import os
import sys
import json
import argparse
import datetime
def _load_token() -> str:
"""HF token: 优先 env, 兜底读仓内 hfkey.txt (secret 扫描不允许硬编码 token)."""
t = os.environ.get("HF_TOKEN", "")
if t:
return t
for p in [os.path.join(os.path.dirname(__file__), "..", "hfkey.txt"),
"/home/z/my-project/lalmodel-code/hfkey.txt"]:
if os.path.exists(p):
return open(p).read().strip()
return ""
HF_TOKEN = _load_token()
CKPT_REPO = "gasschina/lal-ckpts"
OUT_DIR = "/home/z/my-project/download"
Q_LABELS = ["你好", "什么是火?", "中国的首都是", "从前有一座山"]
def fetch_jsonl_from_hf() -> str:
"""从 HF ckpts 仓下载 probe_samples.jsonl, 返回本地路径."""
from huggingface_hub import hf_hub_download
p = hf_hub_download(repo_id=CKPT_REPO, filename="probe_samples.jsonl",
repo_type="model", token=HF_TOKEN,
local_dir="/home/z/my-project/run-smoke/probe_cache")
return p
def load_events(path: str):
"""解析 jsonl → [event dict], 跳过坏行."""
events = []
with open(path, encoding="utf-8") as f:
for line in f:
line = line.strip()
if not line:
continue
try:
events.append(json.loads(line))
except json.JSONDecodeError:
continue
events.sort(key=lambda e: e.get("global", e.get("step", 0)))
return events
def render_markdown(events, out_path):
now = datetime.datetime.now().strftime("%Y-%m-%d %H:%M UTC")
lines = []
lines.append("# lalmodel 质量探针报告\n")
lines.append(f"*生成时间: {now} · 数据点: {len(events)} 个探针事件 · "
f"每事件 = held-out eval + 4 固定 prompt 生成*\n")
# ---- eval_loss 趋势 ----
lines.append("## 1. 客观指标趋势 (held-out eval loss)\n")
lines.append("| global step | eval_loss | Δ vs 上一点 | train avg |")
lines.append("|---:|---:|---:|---:|")
prev = None
for e in events:
ev = e.get("eval_loss", -1)
d = f"{ev - prev:+.3f}" if prev is not None else "—"
lines.append(f"| {e.get('global', e.get('step'))} | {ev:.4f} | {d} | "
f"{e.get('train_avg', float('nan')):.3f} |" if "train_avg" in e else
f"| {e.get('global', e.get('step'))} | {ev:.4f} | {d} | — |")
if ev > 0:
prev = ev
lines.append("")
# ---- 生成样本时间线 (每个 prompt 一个小节) ----
lines.append("## 2. 生成样本时间线 (temp=0.2, top_k=8, 无人工润色)\n")
for qi in range(4):
lines.append(f"### Prompt {qi+1}: 「{Q_LABELS[qi]}」\n")
for e in events:
gens = e.get("gens", [])
if qi >= len(gens):
continue
g = gens[qi]
text = (g.get("text") or "").strip()
if len(text) > 120:
text = text[:117] + "..."
lines.append(f"- **step {e.get('global', e.get('step'))}** "
f"`n={g.get('n')} d1={g.get('d1', 0):.2f} d2={g.get('d2', 0):.2f} "
f"loop={g.get('loop', 0)} E={g.get('e_step', 0):.2f}`")
lines.append(f" > {text if text else '*(空输出)*'}")
lines.append("")
# ---- 指标解读 ----
lines.append("## 3. 指标说明\n")
lines.append("- **eval_loss**: 数据集末尾 4 样本 (训练随机抽样几乎不触碰) 的下一 token "
"CE, 与训练同款前向语义 → 跨步可比的客观曲线, 持续下降 = 模型在学习")
lines.append("- **d1 / d2**: 生成 token 的 distinct-1/2 (0~1, 越高越多样; "
"早期模型通常低 = 高重复)")
lines.append("- **loop**: 末端循环周期长度 (0 = 无循环; >0 表示末尾陷入 k-token 循环)")
lines.append("- **E**: PonderNet 思考深度 E[step] (生成路径); "
"**关键观察**: 是否随训练分化 — 难 token 想得深、易 token 想得浅")
lines.append("- **exit**: 早退 token 占比 (生成路径提前停机的比例)")
with open(out_path, "w", encoding="utf-8") as f:
f.write("\n".join(lines) + "\n")
print(f"[report] markdown -> {out_path}")
def render_chart(events, out_path):
import matplotlib.font_manager as fm
for fp in ["/usr/share/fonts/truetype/chinese/SarasaMonoSC-Regular.ttf",
"/usr/share/fonts/truetype/dejavu/DejaVuSans.ttf"]:
if os.path.exists(fp):
fm.fontManager.addfont(fp)
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
plt.rcParams["font.sans-serif"] = ["Sarasa Mono SC", "DejaVu Sans"]
plt.rcParams["axes.unicode_minus"] = False
steps = [e.get("global", e.get("step", 0)) for e in events]
ev = [e.get("eval_loss", float("nan")) for e in events]
d1s = [[g.get("d1", 0) for g in e.get("gens", [])] for e in events]
d2s = [[g.get("d2", 0) for g in e.get("gens", [])] for e in events]
lps = [[g.get("loop", 0) for g in e.get("gens", [])] for e in events]
es = [[g.get("e_step", 0) for g in e.get("gens", [])] for e in events]
avg1 = [sum(x) / len(x) if x else 0 for x in d1s]
avg2 = [sum(x) / len(x) if x else 0 for x in d2s]
avgl = [sum(x) / len(x) if x else 0 for x in lps]
avge = [sum(x) / len(x) if x else 0 for x in es]
fig, axes = plt.subplots(2, 2, figsize=(11, 7), constrained_layout=True)
ax = axes[0][0]
ax.plot(steps, ev, "o-", color="#c0392b", lw=1.8, ms=4)
ax.set_title("held-out eval loss (越低越好)")
ax.set_xlabel("global step"); ax.set_ylabel("CE loss"); ax.grid(alpha=0.3)
ax = axes[0][1]
ax.plot(steps, avg1, "o-", color="#2980b9", lw=1.8, ms=4, label="distinct-1")
ax.plot(steps, avg2, "s-", color="#8e44ad", lw=1.8, ms=4, label="distinct-2")
ax.set_title("生成多样性 (越高越不重复)")
ax.set_xlabel("global step"); ax.legend(); ax.grid(alpha=0.3)
ax = axes[1][0]
ax.plot(steps, avgl, "o-", color="#d35400", lw=1.8, ms=4)
ax.set_title("末端循环周期 (0=无循环, 越低越好)")
ax.set_xlabel("global step"); ax.set_ylabel("loop k"); ax.grid(alpha=0.3)
ax = axes[1][1]
ax.plot(steps, avge, "o-", color="#16a085", lw=1.8, ms=4, label="E[step]")
ax2 = ax.twinx()
exitr = []
for e in events:
gs = e.get("gens", [])
exitr.append(sum(g.get("exit", 0) for g in gs) / len(gs) if gs else 0)
ax2.plot(steps, exitr, "^--", color="#7f8c8d", lw=1.2, ms=4, alpha=0.7)
ax.set_title("PonderNet 思考深度 (分化 = 循环思考生效)")
ax.set_xlabel("global step"); ax.set_ylabel("E[step]", color="#16a085")
ax2.set_ylabel("早退占比", color="#7f8c8d"); ax.grid(alpha=0.3)
fig.suptitle("lalmodel · CCI2 训练质量探针", fontsize=13)
fig.savefig(out_path, dpi=140)
plt.close(fig)
print(f"[report] chart -> {out_path}")
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--jsonl", help="本地 probe_samples.jsonl (默认拉 HF)")
ap.add_argument("--out", default=OUT_DIR)
args = ap.parse_args()
if args.jsonl:
path = args.jsonl
else:
try:
path = fetch_jsonl_from_hf()
except Exception as e:
print(f"[!] HF 拉取失败: {e}\n (probe_samples.jsonl 每 200 步才上传一次, "
f"刚部署时可能还没有)")
sys.exit(1)
events = load_events(path)
if not events:
print("[!] 0 个探针事件 — Space 还没跑满第一个 probe 周期 (100 步)")
sys.exit(1)
os.makedirs(args.out, exist_ok=True)
render_markdown(events, os.path.join(args.out, "quality_report.md"))
try:
render_chart(events, os.path.join(args.out, "quality_curve.png"))
except Exception as e:
print(f"[!] 图表生成失败 (报告不受影响): {e}")
print(f"[report] {len(events)} 个探针事件, 步数范围 "
f"{events[0].get('global', 0)}..{events[-1].get('global', 0)}")
if __name__ == "__main__":
main()