Something went wrong. Try again.
This repository has no description
Something went wrong. Try again.
1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889#!/usr/bin/env python3"""Measure real training seconds/step for colony configs, so compute-matchedcomparisons 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 argparseimport osimport sysimport 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: E402from dsem.models.losses import ACTLossHead # noqa: E402from dsem.models.recursive_reasoning.colony import ColonyModel_ACTV1 # noqa: E402from 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()