ounce100m-code / config /param_count.py
Cion-lab's picture
Phase 2: source inventory probe + d576/49152 param candidates
fc350c9 verified
Raw History Blame Contribute Delete
7.22 kB
"""Exact parameter accounting for a Llama-style decoder-only model.
The parameter budget (§2: 90-110M *including* embeddings) cannot be met by eyeballing hidden size
and layer count: at a 50k vocab the embedding alone is 25-50% of the model, and whether the
embeddings are tied changes the total by the whole embedding matrix. So count it arithmetically,
then confirm against the real object.
Two independent counts are printed deliberately:
formula -- closed form, so the number is explainable without instantiating the model
torch -- sum of p.numel() over a constructed LlamaForCausalLM, which is the number a third
party can reproduce from the published config
They must agree. If they do not, the formula is wrong (usually GQA head packing or tying), and the
torch count is authoritative because that is what `AutoModelForCausalLM` will report.
Run: python code/config/param_count.py # candidate table
python code/config/param_count.py --verify # also build each model (needs transformers)
"""
import argparse
import json
# Measured on Kaggle (docs/00-platform-notes.md section 9): this shape is 112.9M, i.e. OVER budget.
PROBE_SHAPE = dict(hidden_size=768, num_hidden_layers=12, num_attention_heads=12,
num_key_value_heads=3, intermediate_size=2048, vocab_size=50257,
tie_word_embeddings=True)
CANDIDATES = {
# Shrinking the probe shape toward the window from several directions, so the choice is between
# real options rather than a single lucky guess.
"A_768_h_11l_gqa3_2048": dict(PROBE_SHAPE, num_hidden_layers=11),
"B_768_h_12l_gqa3_1792": dict(PROBE_SHAPE, intermediate_size=1792),
"C_768_h_12l_gqa2_1824": dict(PROBE_SHAPE, intermediate_size=1824, num_key_value_heads=2),
"D_640_h_16l_gqa4_2048": dict(PROBE_SHAPE, hidden_size=640, num_hidden_layers=16,
num_attention_heads=10, num_key_value_heads=4,
intermediate_size=2048),
"E_768_h_12l_gqa3_32k": dict(PROBE_SHAPE, vocab_size=32768),
"F_896_h_9l_gqa2_2432": dict(PROBE_SHAPE, hidden_size=896, num_hidden_layers=9,
num_attention_heads=14, num_key_value_heads=2,
intermediate_size=2432),
"G_1024_h_8l_gqa2_2816": dict(PROBE_SHAPE, hidden_size=1024, num_hidden_layers=8,
num_attention_heads=8, num_key_value_heads=2,
intermediate_size=2816),
# The family the design actually converges on (docs/01-plan.md section 2): deep-thin at d576 with
# the SmolLM2 tokenizer's 49,152 vocab, which is what makes the embedding tax affordable. Depth is
# the knob that moves the total, so size the whole band rather than guessing one value.
"H_576_20l_gqa3_smol": dict(hidden_size=576, num_hidden_layers=20, num_attention_heads=9,
num_key_value_heads=3, intermediate_size=1536, vocab_size=49152,
tie_word_embeddings=True),
"I_576_22l_gqa3_smol": dict(hidden_size=576, num_hidden_layers=22, num_attention_heads=9,
num_key_value_heads=3, intermediate_size=1536, vocab_size=49152,
tie_word_embeddings=True),
"J_576_24l_gqa3_smol": dict(hidden_size=576, num_hidden_layers=24, num_attention_heads=9,
num_key_value_heads=3, intermediate_size=1536, vocab_size=49152,
tie_word_embeddings=True),
"K_576_26l_gqa3_smol": dict(hidden_size=576, num_hidden_layers=26, num_attention_heads=9,
num_key_value_heads=3, intermediate_size=1536, vocab_size=49152,
tie_word_embeddings=True),
"L_640_20l_gqa4_smol": dict(hidden_size=640, num_hidden_layers=20, num_attention_heads=10,
num_key_value_heads=4, intermediate_size=1728, vocab_size=49152,
tie_word_embeddings=True),
}
def count_formula(v, h, nl, nh, nkv, ff, tied):
"""Closed-form Llama param count. All terms explained inline so the arithmetic is checkable."""
embed = v * h
d_head = h // nh
# GQA: k/v are d_head * nkv wide, not h, so the ratio nkv/nh matters.
attn = h * h + 2 * d_head * nkv * h + h * h # q, k, v, o
mlp = 3 * h * ff # SwiGLU gate + up + down
norms = 2 * h * nl # input + post-attention RMSNorm
per_layer = attn + mlp
total = embed * (1 if tied else 2) + per_layer * nl + norms + h
return {
"embedding_one_copy": embed,
"per_layer_attn": attn, "per_layer_mlp": mlp, "layers": nl,
"transformer_blocks": per_layer * nl,
"norms": norms + h,
"tied_embeddings": tied,
"total": total,
"embedding_pct_of_total": round(100.0 * embed / total, 1),
}
def count_torch(cfg):
import torch
from transformers import LlamaConfig, LlamaForCausalLM
m = LlamaForCausalLM(LlamaConfig(**cfg))
n = sum(p.numel() for p in m.parameters())
emb = m.get_input_embeddings().weight.numel()
tied = m.config.tie_word_embeddings and (
m.get_input_embeddings().weight.data_ptr() == m.lm_head.weight.data_ptr())
del m
torch.cuda.empty_cache() if torch.cuda.is_available() else None
return {"total": n, "embedding": emb, "actually_tied": bool(tied)}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--verify", action="store_true",
help="build each model with transformers and compare (slow, needs the lib)")
args = ap.parse_args()
rows = []
for name, cfg in CANDIDATES.items():
f = count_formula(cfg["vocab_size"], cfg["hidden_size"], cfg["num_hidden_layers"],
cfg["num_attention_heads"], cfg["num_key_value_heads"],
cfg["intermediate_size"], cfg["tie_word_embeddings"])
row = {"name": name, "formula_total": f["total"], "emb_pct": f["embedding_pct_of_total"],
"in_90_110M": 90_000_000 <= f["total"] <= 110_000_000, **{
k: cfg[k] for k in ("hidden_size", "num_hidden_layers", "num_attention_heads",
"num_key_value_heads", "intermediate_size", "vocab_size")}}
if args.verify:
t = count_torch(cfg)
row["torch_total"] = t["total"]
row["matches_formula"] = t["total"] == f["total"]
row["torch_reports_tied"] = t["actually_tied"]
rows.append(row)
p = count_formula(**{ # the probe shape, for the record
"v": PROBE_SHAPE["vocab_size"], "h": PROBE_SHAPE["hidden_size"],
"nl": PROBE_SHAPE["num_hidden_layers"], "nh": PROBE_SHAPE["num_attention_heads"],
"nkv": PROBE_SHAPE["num_key_value_heads"], "ff": PROBE_SHAPE["intermediate_size"],
"tied": PROBE_SHAPE["tie_word_embeddings"]})
print("probe shape (measured 112,934,400 in the image):", json.dumps(p))
for r in rows:
print(json.dumps(r))
if __name__ == "__main__":
main()