Something went wrong. Try again.
This repository has no description
Something went wrong. Try again.
12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970#!/usr/bin/env python3"""Report parameter counts for TRM at various hidden widths (CPU-only, no training).
Used to plan the capacity-floor sweep: how small does a unit actually get?Note SwiGLU's inner dim is rounded up to a multiple of 256, so params do NOTscale as D^2 all the way down — small widths are floored by that rounding.
Usage: DSEM_DEVICE=cpu python scripts/param_counts.py [--widths 512 256 128 64 32]"""
import argparseimport osimport sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))os.environ.setdefault("DSEM_DEVICE", "cpu")
import torch # noqa: E402
from dsem.models.recursive_reasoning.trm import ( # noqa: E402 TinyRecursiveReasoningModel_ACTV1,)
def build(width: int, seq_len: int = 81, vocab_size: int = 11) -> torch.nn.Module: cfg = dict( batch_size=8, seq_len=seq_len, vocab_size=vocab_size, num_puzzle_identifiers=1, puzzle_emb_ndim=width, puzzle_emb_len=16, H_cycles=3, L_cycles=6, H_layers=0, L_layers=2, hidden_size=width, expansion=4, num_heads=max(1, width // 64), pos_encodings="none", halt_max_steps=16, halt_exploration_prob=0.1, forward_dtype="float32", mlp_t=True, no_ACT_continue=True, causal=False, ) return TinyRecursiveReasoningModel_ACTV1(cfg)
def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--widths", type=int, nargs="+", default=[512, 384, 256, 192, 128, 96, 64, 32]) args = ap.parse_args()
base = None print(f"{'width':>6} {'params':>12} {'vs D=512':>9} {'rel FLOPs*':>11}") for w in args.widths: model = build(w) n = sum(p.numel() for p in model.parameters()) + sum(b.numel() for b in model.buffers()) if base is None: base = n print(f"{w:>6} {n:>12,} {n/base:>8.1%} {(w/args.widths[0])**2:>10.1%}") print("\n* rel FLOPs is the naive D^2 expectation; actual params flatten out at small D") print(" because SwiGLU's inner dimension is rounded up to a multiple of 256.")
if __name__ == "__main__": main()