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()