File size: 7,222 Bytes
e6aee19
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fc350c9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e6aee19
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
121
122
123
124
125
126
127
128
129
130
131
132
133
134
"""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()