#!/usr/bin/env python3 """Measure real training seconds/step for colony configs, so compute-matched comparisons use measurements rather than parameter-count guesses. Usage (on the box): python scripts/time_colony.py --widths 64 128 --agents 1 2 4 8 --steps 12 """ import argparse import os import sys import time sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) import torch # noqa: E402 from dsem.device import get_device # noqa: E402 from dsem.models.losses import ACTLossHead # noqa: E402 from dsem.models.recursive_reasoning.colony import ColonyModel_ACTV1 # noqa: E402 from dsem.optim.adam_atan2 import AdamATan2 # noqa: E402 SEQ, VOCAB = 81, 11 def cfg(width: int, k: int, batch: int): return dict( batch_size=batch, seq_len=SEQ, vocab_size=VOCAB, 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="bfloat16", mlp_t=True, no_ACT_continue=True, causal=False, num_agents=k, blackboard="tokens", standpoints="sudoku", aggregate="mean", halt_rule="quorum", role_embeddings=False, ) def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--widths", type=int, nargs="+", default=[64, 128]) ap.add_argument("--agents", type=int, nargs="+", default=[1, 2, 4, 8]) ap.add_argument("--batch", type=int, default=768) ap.add_argument("--steps", type=int, default=12) ap.add_argument("--warmup", type=int, default=3) args = ap.parse_args() device = get_device() g = torch.Generator().manual_seed(0) batch = { "inputs": torch.randint(1, VOCAB, (args.batch, SEQ), generator=g).to(device), "labels": torch.randint(1, VOCAB, (args.batch, SEQ), generator=g).to(device), "puzzle_identifiers": torch.zeros(args.batch, dtype=torch.int32, device=device), } print(f"{'width':>6} {'K':>3} {'params':>10} {'s/step':>8} {'vs D=512 TRM (2.40s)':>22}") for width in args.widths: for k in args.agents: with torch.device(device): model = ACTLossHead(ColonyModel_ACTV1(cfg(width, k, args.batch)), loss_type="stablemax_cross_entropy") carry = model.initial_carry(batch) opt = AdamATan2(model.parameters(), lr=1e-4, betas=(0.9, 0.95), weight_decay=1.0) n = sum(p.numel() for p in model.parameters()) times = [] for i in range(args.warmup + args.steps): if device.type == "cuda": torch.cuda.synchronize() t0 = time.time() carry, loss, _metrics, _preds, _fin = model(carry=carry, batch=batch, return_keys=[]) ((1 / args.batch) * loss).backward() opt.step() opt.zero_grad() if device.type == "cuda": torch.cuda.synchronize() if i >= args.warmup: times.append(time.time() - t0) sps = sorted(times)[len(times) // 2] print(f"{width:>6} {k:>3} {n:>10,} {sps:>8.3f} {sps/2.40:>21.2f}x") del model, opt, carry if device.type == "cuda": torch.cuda.empty_cache() if __name__ == "__main__": main()