File size: 5,655 Bytes
17d5066
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
149
150
151
152
153
154
"""Emit the paper's analysis artifacts from a trained Sheaf-ADMM checkpoint.

Loads a ``checkpoint.pkl`` written by ``scripts/train.py`` (a dict with
``params`` / ``ema_params`` / ``config``), rebuilds the model + task from the
embedded config, pulls one validation batch, and writes the three artifacts:

* ``prediction_evolution.pdf`` — global prediction vs ADMM count ``k``;
* ``coordination_dynamics.pdf`` — primal/dual residual heatmaps + curves;
* ``xz_trajectories.pdf`` — per-agent x-vs-z 2-D trajectories.

    python -m scripts.visualize --checkpoint outputs/.../checkpoint.pkl --out-dir /tmp/viz

The ``--task`` / ``--data-dir`` flags override what the checkpoint config says
(useful for visualizing on an OOD split). Evaluation uses the EMA parameters by
default (``--params ema``), matching how the trainer reports metrics.
Only Sheaf-ADMM checkpoints are supported; MPNN baselines do not expose an ADMM
trajectory.
"""

from __future__ import annotations

import argparse
import pickle
import sys
from pathlib import Path

import numpy as np

PROJECT_ROOT = Path(__file__).resolve().parents[1]
if str(PROJECT_ROOT / "src") not in sys.path:
    sys.path.insert(0, str(PROJECT_ROOT / "src"))

import sheaf_admm as _sheaf_admm  # noqa: F401,E402  (sets matmul precision on import)
from sheaf_admm.data import ImageDataset, PuzzleDataset  # noqa: E402
from sheaf_admm.models import model_config_from_dict  # noqa: E402
from sheaf_admm.training import build_model, make_task  # noqa: E402
from sheaf_admm.viz import (  # noqa: E402
    plot_coordination_dynamics,
    plot_prediction_evolution,
    plot_xz_trajectories,
    run_trajectory,
)

# Default per-task ADMM counts for the prediction-evolution panels.
DEFAULT_KS = {"maze": (1, 3, 5, 10, 30), "sudoku": (1, 3, 5, 10, 20), "mnist": (1, 3, 5, 10, 30)}


def _first_batch(cfg: dict, task_name: str, data_dir: str, split: str):
    loader = cfg["data"].get("loader", "puzzle" if task_name in ("maze", "sudoku") else "image")
    if loader == "puzzle":
        ds = PuzzleDataset(data_dir, split)
        for _set, b in ds.iter_test_batches(batch_size=8):
            batch = {"inputs": np.asarray(b["inputs"]), "labels": np.asarray(b["labels"])}
            for key in ("height", "width"):
                if key in b:
                    batch[key] = b[key]
            return batch
    ds = ImageDataset(data_dir, split)
    for b in ds.iter_batches(batch_size=8, shuffle=False):
        return {"images": np.asarray(b["images"]), "labels": np.asarray(b["labels"])}
    raise RuntimeError("no batch available from the dataset")


def main() -> None:
    p = argparse.ArgumentParser(description=__doc__)
    p.add_argument("--checkpoint", required=True, help="path to checkpoint.pkl")
    p.add_argument("--out-dir", required=True, help="directory for the artifact files")
    p.add_argument("--task", default=None, help="override task (maze|sudoku|mnist)")
    p.add_argument("--data-dir", default=None, help="override dataset directory")
    p.add_argument("--split", default=None, help="dataset split (default: first val split)")
    p.add_argument("--params", choices=("ema", "raw"), default="ema", help="which weights to use")
    p.add_argument(
        "--num-iters", type=int, default=30, help="ADMM steps for the trajectory artifacts"
    )
    p.add_argument("--batch-index", type=int, default=0, help="example index within the batch")
    p.add_argument(
        "--ks", type=int, nargs="+", default=None, help="ADMM counts for prediction evolution"
    )
    args = p.parse_args()

    with open(args.checkpoint, "rb") as f:
        ckpt = pickle.load(f)
    cfg = ckpt["config"]
    if cfg.get("model_type") != "sheaf":
        raise ValueError(
            "scripts.visualize supports Sheaf-ADMM checkpoints only "
            f"(got model_type={cfg.get('model_type')!r})."
        )
    params = ckpt.get("ema_params") if args.params == "ema" else ckpt["params"]
    if params is None:
        params = ckpt["params"]

    task_name = args.task or cfg["task"]
    data_dir = args.data_dir or cfg["data"]["dir"]
    split = args.split or cfg["data"]["val_splits"][0]
    ks = tuple(args.ks) if args.ks else DEFAULT_KS[task_name]

    task = make_task(task_name, **cfg.get("task_cfg", {}))
    model_cfg = model_config_from_dict(cfg["model"])
    model = build_model(model_cfg, cfg["model_type"])

    batch = _first_batch(cfg, task_name, data_dir, split)
    fwd, targets, aux = task.prepare(batch)

    out_dir = Path(args.out_dir)
    out_dir.mkdir(parents=True, exist_ok=True)

    written = []
    written.append(
        plot_prediction_evolution(
            model,
            params,
            task,
            fwd,
            targets,
            aux,
            ks,
            str(out_dir / "prediction_evolution.pdf"),
            batch_index=args.batch_index,
            title=f"{task_name}: prediction vs k",
        )
    )

    centers = aux.get("centers")
    traj = run_trajectory(
        model,
        params,
        fwd,
        num_iters=args.num_iters,
        batch_index=args.batch_index,
        centers=centers,
    )
    written.append(
        plot_coordination_dynamics(
            traj,
            str(out_dir / "coordination_dynamics.pdf"),
            title=f"{task_name}: coordination dynamics (rho={traj.rho:.3g})",
        )
    )
    written.append(
        plot_xz_trajectories(
            traj,
            str(out_dir / "xz_trajectories.pdf"),
            title=f"{task_name}: x vs z trajectories",
        )
    )

    for path in written:
        print(f"[viz] wrote {path}")


if __name__ == "__main__":
    main()