File size: 7,243 Bytes
880dff9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 | #!/usr/bin/env python3
"""四个 baseline 臂串行训练 + 评测的队列:linear → xattn(旧报告口径 98.6M)→ prompt → adaln。
每臂:scripts/train_baseline.sh <arm> --max_steps 5000 --warmup_steps 200(GA2,与 arope_5k_ga2 完全同配方)
→ bg50 四科(1×)+ 训练场景 8 条演示(复用 overnight.evaluate)。最后写 outputs/baselines/summary.md:
五臂并排(arope_5k_ga2 / arope_10k_ga2 + 四个 baseline)。只做三件事:等 GPU 空、subprocess 跑本仓库脚本、写汇总。
nohup .venv/bin/python scripts/baselines_queue.py > outputs/baselines/queue.log 2>&1 &
(--resume 跳过 status.json 里已完成的阶段)
"""
from __future__ import annotations
import json
import os
import subprocess
import sys
import time
from datetime import datetime
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import overnight as ON # noqa: E402 复用 log / stage / run / wait_train_done / evaluate / pick
ROOT = ON.ROOT
OUT = f"{ROOT}/outputs/baselines"
ON.OUT = OUT # overnight 的 status.json / eval 日志都落到这里
ARMS = [
# (臂, run 名, 环境变量, 说明)
("linear", "linear_5k_ga2", {}, "ReactiveGWM 逐块 Linear,0.18M"),
("xattn", "xattn_5k_ga2", {"ARM_KWARGS": '{"enable_mouse": false, "window_frames": 1}'}, "Matrix-Game 3 键盘 cross-attn,旧报告口径 98.6M"),
("prompt", "prompt_5k_ga2", {}, "Incantation 逐 cell 文本,0 参数"),
("adaln", "adaln_5k_ga2", {}, "AlayaWorld adaLN,66.5M"),
]
AROPE_RUNS = [("arope_5k_ga2", "ARoPE 5k(同配方)"), ("arope_10k_ga2", "ARoPE 10k")]
def gpus_free() -> bool:
out = subprocess.run(["nvidia-smi", "--query-gpu=memory.used", "--format=csv,noheader,nounits"],
capture_output=True, text=True).stdout.split()
return all(int(x) < 1024 for x in out)
def wait_gpus_free(hold_sec: int = 90, max_hours: float = 3.0):
"""等 8 卡都空闲并保持 hold_sec 秒(前面的评测可能还在收尾)。"""
t0 = time.time()
quiet = None
while True:
if gpus_free():
quiet = quiet or time.time()
if time.time() - quiet >= hold_sec:
return {"waited_sec": round(time.time() - t0)}
else:
quiet = None
if time.time() - t0 > max_hours * 3600:
raise RuntimeError("等 GPU 空闲超时")
time.sleep(15)
def train(arm, run_name, env_extra):
if os.path.exists(f"{ROOT}/outputs/{run_name}/step-5000.safetensors"):
return {"skipped": "step-5000.safetensors 已存在"}
wait_gpus_free()
launch_log = f"{ROOT}/outputs/{run_name}_launch.log"
pid_file = f"{ROOT}/outputs/{run_name}.pid"
env = dict(ON.ENV, RUN=run_name, GA="2", **env_extra)
with open(launch_log, "a", encoding="utf-8") as lf:
p = subprocess.Popen(["bash", f"{ROOT}/scripts/train_baseline.sh", arm, "--max_steps", "5000", "--warmup_steps", "200"],
cwd=ROOT, env=env, stdout=lf, stderr=subprocess.STDOUT)
open(pid_file, "w").write(str(p.pid))
ON.log(f"{run_name} 已启动 pid={p.pid} env={env_extra}")
r = ON.wait_train_done(run_name, pid_file)
p.wait(timeout=600)
r["rc"] = p.returncode
return r
def evaluate(run_name):
wait_gpus_free(hold_sec=30)
return ON.evaluate(run_name)
def write_summary():
runs = [(n, d) for n, d in AROPE_RUNS] + [(r, d) for _, r, _, d in ARMS]
def summ(run):
p = f"{ROOT}/outputs/eval_bg50/{run}/summary.json"
return json.load(open(p, encoding="utf-8")) if os.path.exists(p) else {}
S = {run: summ(run) for run, _ in runs}
metrics = [
("增益线性 斜率", ("gain", "slope")), ("增益线性 R²", ("gain", "r2")),
("折返 PSNR 中位", ("back", "psnr_med")), ("折返 SSIM 中位", ("back", "ssim_med")), ("折返 NCC 中位", ("back", "ncc_med")),
("折返 回归残差中位 px", ("back", "resid_px_med")), ("折返 补偿后 NCC", ("back", "ncc_comp_med")),
("倒退帧总数", ("smooth", "backward_frames_total")), ("速度波动 std/mean 中位", ("smooth", "speed_cv_med")),
("平均速度/承诺", ("smooth", "speed_ratio_med")),
("角度误差中位 °", ("direction", "angle_abs_med")), ("角度误差 p90 °", ("direction", "angle_abs_p90")),
("幅度比中位", ("direction", "ratio_med")), ("失控数", ("direction", "n_fail")),
]
lines = ["# 五臂 baseline 对比(bg50 四科,全部 1×;同数据、同配方:8 卡 × GA2、lr 1e-5、5k 步)", "",
f"生成于 {datetime.now():%Y-%m-%d %H:%M}。旧报告世界 RoPE@5k / 普通 RoPE@5k 的数字见 README。", "",
"| 指标 | " + " | ".join(f"{r}<br>{d}" for r, d in runs) + " |", "|---|" + "---|" * len(runs)]
for label, keys in metrics:
vals = []
for run, _ in runs:
v = ON.pick(S[run], *keys)
vals.append("—" if v is None else (f"{v:.3f}" if isinstance(v, float) else str(v)))
lines.append(f"| {label} | " + " | ".join(vals) + " |")
lines += ["", "## 训练场景首帧演示:指令 vs 实测背景位移(帧 0→80,px)", ""]
for run, d in runs:
st = ON.STATUS["stages"].get(f"eval:{run}", {}).get("result") or {}
demo = st.get("demo")
if not demo: # arope 的两轮评测不在本队列里,直接读 samples 目录
demo = {}
dd = f"{ROOT}/outputs/samples/{run}"
for name in ("replay", "right_x1", "left_x1", "up_x1", "down_x1", "right_x0.5", "right_x1.5", "there_back"):
p = f"{dd}/{name}.json"
if os.path.exists(p):
j = json.load(open(p, encoding="utf-8"))
demo[name] = {"cmd_bg_shift_80": [-v for v in j["frame_offset_px_80"]],
"measured": (j.get("measured") or {}).get("sift")}
lines.append(f"### {run}({d})")
lines.append("| 指令 | 指令位移 | 实测 |")
lines.append("|---|---|---|")
for name, x in demo.items():
c = x["cmd_bg_shift_80"]; m = x["measured"]
lines.append(f"| {name} | ({c[0]:+.0f},{c[1]:+.0f}) | " + (f"({m[0]:+.0f},{m[1]:+.0f})" if m else "测不出") + " |")
lines.append("")
lines += ["## 阶段状态", "", "```", json.dumps(ON.STATUS["stages"], ensure_ascii=False, indent=1, default=str)[:8000], "```"]
os.makedirs(OUT, exist_ok=True)
with open(f"{OUT}/summary.md", "w", encoding="utf-8") as f:
f.write("\n".join(lines))
return f"{OUT}/summary.md"
def main():
os.makedirs(OUT, exist_ok=True)
if os.path.exists(f"{OUT}/status.json") and "--resume" in sys.argv:
ON.STATUS.update(json.load(open(f"{OUT}/status.json", encoding="utf-8")))
ON.save_status()
for arm, run_name, env_extra, _ in ARMS:
ON.stage(f"train:{run_name}", lambda: train(arm, run_name, env_extra))
ON.stage(f"eval:{run_name}", lambda: evaluate(run_name))
ON.stage("summary", write_summary)
ON.STATUS["finished"] = datetime.now().isoformat(timespec="seconds")
ON.save_status()
ON.log("全部结束")
if __name__ == "__main__":
main()
|