File size: 4,508 Bytes
038acee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Exact parameter counts for committed model configs.

Instantiates each config through the repo's own build path (build_model +
init_params) and prints name, architecture fields and the exact parameter
count. With --verify-checkpoint, additionally restores an Orbax checkpoint
via load_checkpoint and asserts its parameter count equals the config's.

Usage:
  uv run python scripts/count_params.py
  uv run python scripts/count_params.py --verify-checkpoint \
      checkpoints/offline/Craftax-Classic-Symbolic-v1-Offline-Diffusion-BC-100M/100000000 \
      --config configs/final_craftax_classic_gpu_24gb.yaml
"""
from __future__ import annotations

import argparse
import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))

import jax
import jax.numpy as jnp  # noqa: F401  (kept for parity with repo imports)
import yaml

from craftax.craftax_env import make_craftax_env_from_name
from src.planners.model import build_model, init_params, load_checkpoint

DEFAULT_CONFIGS = [
    "configs/classic_exp_d_100K_model.yaml",
    "configs/classic_exp_d_250K_model.yaml",
    "configs/classic_exp_d_850K_model.yaml",
    "configs/classic_exp_d_3M_model.yaml",
    "configs/craftax_exp_d_500K_model.yaml",
    "configs/craftax_exp_d_1M_model.yaml",
    "configs/craftax_exp_d_3M_model.yaml",
    "configs/craftax_exp_d_7M_model.yaml",
    "configs/final_craftax_classic_gpu_24gb.yaml",
    "experiments/rl_finetuning/configs/ablations_final_craftax_gpu_24gb.yaml",
]

_ENV_CACHE: dict[str, tuple] = {}


def _env_dims(env_name: str) -> tuple[int, int]:
    if env_name not in _ENV_CACHE:
        env = make_craftax_env_from_name(env_name, auto_reset=True)
        env_params = env.default_params
        _ENV_CACHE[env_name] = (
            env.action_space(env_params).n,
            env.observation_space(env_params).shape[0],
        )
    return _ENV_CACHE[env_name]


def count_config(cfg_path: str) -> dict:
    raw = yaml.safe_load(open(cfg_path))
    cfg = {k.upper(): v for k, v in raw.items()}
    env_name = cfg.get("ENV_NAME", "Craftax-Classic-Symbolic-v1")
    num_actions, obs_dim = _env_dims(env_name)
    model = build_model(cfg, num_actions)
    params = init_params(
        model, jax.random.PRNGKey(0), obs_dim, int(cfg.get("PLAN_HORIZON", 32))
    )
    n = int(sum(int(x.size) for x in jax.tree_util.tree_leaves(params)))
    return {
        "config": cfg_path,
        "env_name": env_name,
        "d_model": cfg.get("D_MODEL"),
        "n_heads": cfg.get("N_HEADS"),
        "n_layers": cfg.get("N_LAYERS"),
        "d_ff": cfg.get("D_FF"),
        "obs_dim": obs_dim,
        "params": n,
    }


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument("--configs", nargs="+", default=DEFAULT_CONFIGS)
    ap.add_argument("--verify-checkpoint", type=str, default=None,
                    help="Orbax checkpoint dir to restore and count against the config.")
    ap.add_argument("--config", type=str, default=None,
                    help="Config matching --verify-checkpoint.")
    a = ap.parse_args()

    print(f"{'config':52s} {'arch (d/h/l/ff)':>18s} {'obs':>6s} {'params':>12s}")
    rows = []
    for c in a.configs:
        r = count_config(c)
        rows.append(r)
        arch = f"{r['d_model']}/{r['n_heads']}/{r['n_layers']}/{r['d_ff']}"
        print(f"{r['config']:52s} {arch:>18s} {r['obs_dim']:>6d} {r['params']:>12,d}")

    if a.verify_checkpoint:
        if not a.config:
            raise SystemExit("--config is required with --verify-checkpoint")
        r = count_config(a.config)
        raw = yaml.safe_load(open(a.config))
        cfg = {k.upper(): v for k, v in raw.items()}
        num_actions, obs_dim = _env_dims(cfg.get("ENV_NAME", "Craftax-Classic-Symbolic-v1"))
        model = build_model(cfg, num_actions)
        restored = load_checkpoint(
            model, jax.random.PRNGKey(0), obs_dim,
            int(cfg.get("PLAN_HORIZON", 32)), a.verify_checkpoint,
        )
        n_ckpt = int(sum(int(x.size) for x in jax.tree_util.tree_leaves(restored)))
        print(f"\ncheckpoint {a.verify_checkpoint}: {n_ckpt:,d} params")
        print(f"config     {a.config}: {r['params']:,d} params")
        if n_ckpt == r["params"]:
            print("PARAM COUNT GATE: PASS (checkpoint parameter count equals config count)")
        else:
            print("PARAM COUNT GATE: FAIL (counts differ)")
            raise SystemExit(1)


if __name__ == "__main__":
    main()