File size: 4,562 Bytes
4fd79a1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Fair multi-seed + long-horizon benchmark for Spectral World Models."""
from __future__ import annotations
import argparse, csv, json
from pathlib import Path
from statistics import mean, stdev
import torch
from spectral_world_models.data import DatasetConfig, generate_benchmark_npz
from spectral_world_models.train import train_one_model

DEFAULT_MODELS=["dreamer","transformer","koopman","neural_operator","swm","no_spectral_state","no_spectral_transition","no_text_branch","no_stability_penalty"]
BASE_METRICS=[("test_psnr",lambda r:r["test"]["psnr"]),("test_ssim",lambda r:r["test"]["ssim"]),("test_nll",lambda r:r["test"]["nll"]),("test_mse",lambda r:r["test"]["mse"]),("test_text_acc",lambda r:r["test"]["text_acc"])]
def stats(xs): return (mean(xs), stdev(xs) if len(xs)>1 else 0.0)

def main():
    ap=argparse.ArgumentParser(description="Run fair multi-seed SWM benchmark with H=5/10/20/30 rollouts.")
    ap.add_argument("--data",default=None); ap.add_argument("--rollout-data",default=None); ap.add_argument("--out",default=None)
    ap.add_argument("--epochs",type=int,default=3); ap.add_argument("--batch-size",type=int,default=64)
    ap.add_argument("--seeds",type=int,nargs="+",default=[0,1,2,3,4]); ap.add_argument("--horizons",type=int,nargs="+",default=[5,10,20,30])
    ap.add_argument("--device",default="auto"); ap.add_argument("--models",nargs="+",default=DEFAULT_MODELS)
    args=ap.parse_args(); root=Path(__file__).resolve().parents[1]
    data=Path(args.data) if args.data else root/"data"/"mini_moving_shapes.npz"
    out=Path(args.out) if args.out else root/"results"/"multiseed_horizons"
    rollout_data=Path(args.rollout_data) if args.rollout_data else root/"data"/"mini_moving_shapes_rollout_h31.npz"
    if not data.exists(): generate_benchmark_npz(data,DatasetConfig())
    need_len=max(args.horizons)+1
    regenerate=True
    if rollout_data.exists():
        try:
            import numpy as np
            with np.load(rollout_data,allow_pickle=True) as d: regenerate=d["test_images"].shape[1] < need_len
        except Exception: regenerate=True
    if regenerate:
        # Separate held-out long sequences: training data remain unchanged.
        generate_benchmark_npz(rollout_data,DatasetConfig(seq_len=need_len,train_sequences=1,val_sequences=1,test_sequences=64,seed=7007))
    device=("cuda" if torch.cuda.is_available() else "cpu") if args.device=="auto" else args.device
    if device.startswith("cuda") and not torch.cuda.is_available(): raise RuntimeError("CUDA requested but unavailable")
    out.mkdir(parents=True,exist_ok=True); grouped={m:[] for m in args.models}; raw=[]
    for seed in args.seeds:
        sd=out/f"seed_{seed}"
        for m in args.models:
            print(f"\n=== seed={seed} model={m} horizons={args.horizons} ===")
            r=train_one_model(m,data,sd,epochs=args.epochs,batch_size=args.batch_size,seed=seed,device=device,rollout_data_path=rollout_data,rollout_horizons=tuple(args.horizons))
            grouped[m].append(r); row={"seed":seed,"model":m,"params":r["params"]}
            for k,g in BASE_METRICS: row[k]=g(r)
            for h in args.horizons:
                rr=r["rollout_by_horizon"][str(h)]; row[f"rollout_mse_h{h}"]=rr["rollout_mse"]; row[f"time_per_rollout_step_ms_h{h}"]=rr["time_per_rollout_step_ms"]
            raw.append(row)
    with (out/"metrics_by_seed.csv").open("w",newline="") as f:
        w=csv.DictWriter(f,fieldnames=list(raw[0])); w.writeheader(); w.writerows(raw)
    summary=[]
    for m,runs in grouped.items():
        row={"model":m,"params":runs[0]["params"],"n_seeds":len(runs)}
        for k,g in BASE_METRICS:
            mu,sd=stats([float(g(r)) for r in runs]); row[k+"_mean"]=mu; row[k+"_std"]=sd
        for h in args.horizons:
            vals=[float(r["rollout_by_horizon"][str(h)]["rollout_mse"]) for r in runs]; mu,sd=stats(vals); row[f"rollout_mse_h{h}_mean"]=mu; row[f"rollout_mse_h{h}_std"]=sd
        summary.append(row)
    with (out/"metrics_summary.csv").open("w",newline="") as f:
        w=csv.DictWriter(f,fieldnames=list(summary[0])); w.writeheader(); w.writerows(summary)
    with (out/"run_config.json").open("w") as f: json.dump(vars(args)|{"device_resolved":device,"rollout_data_resolved":str(rollout_data)},f,indent=2)
    print("\n=== Long-horizon summary ===")
    for r in summary: print(r["model"]+": "+" | ".join(f'H{h} {r[f"rollout_mse_h{h}_mean"]:.6f} ± {r[f"rollout_mse_h{h}_std"]:.6f}' for h in args.horizons))
    print(f"\nWrote: {out/'metrics_summary.csv'}")
if __name__=="__main__": main()