Robotics
multilingual
ternary
multimodal
pretraining
jirack
ternarytransformer
kgrabko commited on
Commit
dd17e7f
·
verified ·
1 Parent(s): 80e3270

Delete convert_to_14b.py

Browse files
Files changed (1) hide show
  1. convert_to_14b.py +0 -117
convert_to_14b.py DELETED
@@ -1,117 +0,0 @@
1
- #!/usr/bin/env python3
2
- # ==============================================================================
3
- # COPYRIGHT (C) 2026 KONSTANTIN VLADIMIROVICH GRABKO. ALL RIGHTS RESERVED.
4
- # PATENT PENDING | CMS MANHATTAN JIRACK TECHNOLOGY
5
- # ==============================================================================
6
- #
7
- # size_planner.py
8
- # ---------------
9
- # Compute exact parameter counts for a JiRack config, and suggest
10
- # (num_layers, intermediate_size) combinations that hit a target size.
11
- # No torch needed - pure arithmetic, so you can plan before touching a GPU.
12
- #
13
- # Examples:
14
- # python size_planner.py --hidden 3072 --layers 36 --intermediate 16384
15
- # python size_planner.py --hidden 3072 --target-b 14 # suggest combos
16
- # ==============================================================================
17
-
18
- import argparse
19
-
20
-
21
- def param_count(hidden, layers, intermediate, vocab=128256,
22
- num_heads=24, num_kv_heads=8, tied_embeddings=False):
23
- head_dim = hidden // num_heads
24
- kv_dim = num_kv_heads * head_dim
25
-
26
- attn = (hidden * hidden # q_proj
27
- + hidden * kv_dim # k_proj
28
- + hidden * kv_dim # v_proj
29
- + hidden * hidden) # out_proj
30
- ffn = 3 * intermediate * hidden # w1, w3, w2 (SwiGLU)
31
- per_layer = attn + ffn
32
-
33
- emb = vocab * hidden * (1 if tied_embeddings else 2) # token_emb (+ lm_head)
34
- total = emb + layers * per_layer
35
-
36
- return {
37
- "total": total,
38
- "embeddings": emb,
39
- "per_layer": per_layer,
40
- "attn_per_layer": attn,
41
- "ffn_per_layer": ffn,
42
- "layers_total": layers * per_layer,
43
- }
44
-
45
-
46
- def intermediate_for_target(hidden, layers, target_params, vocab=128256,
47
- num_heads=24, num_kv_heads=8, tied_embeddings=False):
48
- """Solve for the intermediate size that hits target_params at fixed layers."""
49
- head_dim = hidden // num_heads
50
- kv_dim = num_kv_heads * head_dim
51
- attn = 2 * hidden * hidden + 2 * hidden * kv_dim
52
- emb = vocab * hidden * (1 if tied_embeddings else 2)
53
- # target = emb + layers*(attn + 3*I*hidden) -> solve I
54
- budget_per_layer = (target_params - emb) / layers
55
- ffn_budget = budget_per_layer - attn
56
- if ffn_budget <= 0:
57
- return None
58
- inter = ffn_budget / (3 * hidden)
59
- # round to a multiple of 256 for clean shapes
60
- return int(round(inter / 256) * 256)
61
-
62
-
63
- def fmt(n):
64
- return f"{n/1e9:.3f} B" if n >= 1e9 else f"{n/1e6:.1f} M"
65
-
66
-
67
- def main():
68
- ap = argparse.ArgumentParser()
69
- ap.add_argument("--hidden", type=int, default=3072)
70
- ap.add_argument("--vocab", type=int, default=128256)
71
- ap.add_argument("--heads", type=int, default=24)
72
- ap.add_argument("--kv-heads", type=int, default=8)
73
- ap.add_argument("--tied", action="store_true", help="tie lm_head to token_emb")
74
- ap.add_argument("--layers", type=int, help="report a single exact config")
75
- ap.add_argument("--intermediate", type=int, help="report a single exact config")
76
- ap.add_argument("--target-b", type=float, help="suggest combos for this size (in B)")
77
- args = ap.parse_args()
78
-
79
- if args.layers and args.intermediate:
80
- r = param_count(args.hidden, args.layers, args.intermediate, args.vocab,
81
- args.heads, args.kv_heads, args.tied)
82
- print(f"hidden={args.hidden} layers={args.layers} "
83
- f"intermediate={args.intermediate} vocab={args.vocab}")
84
- print(f" embeddings : {fmt(r['embeddings'])} "
85
- f"({100*r['embeddings']/r['total']:.1f}% of total)")
86
- print(f" per layer : {fmt(r['per_layer'])} "
87
- f"(attn {fmt(r['attn_per_layer'])} + ffn {fmt(r['ffn_per_layer'])})")
88
- print(f" all layers : {fmt(r['layers_total'])}")
89
- print(f" TOTAL : {fmt(r['total'])}")
90
- ratio = args.intermediate / args.hidden
91
- print(f" ffn expansion : {ratio:.2f}x hidden "
92
- f"({'normal' if 2.0 <= ratio <= 3.0 else 'unusual - very fat' if ratio > 4 else 'narrow'})")
93
- return
94
-
95
- if args.target_b:
96
- target = args.target_b * 1e9
97
- print(f"Configs hitting ~{args.target_b}B (hidden={args.hidden}, "
98
- f"vocab={args.vocab}):\n")
99
- print(f" {'layers':>6} {'intermediate':>12} {'ffn_ratio':>9} {'total':>9}")
100
- print(" " + "-" * 46)
101
- for L in range(24, 81, 4):
102
- I = intermediate_for_target(args.hidden, L, target, args.vocab,
103
- args.heads, args.kv_heads, args.tied)
104
- if I is None or I < args.hidden: # skip infeasible / too-narrow FFN
105
- continue
106
- r = param_count(args.hidden, L, I, args.vocab,
107
- args.heads, args.kv_heads, args.tied)
108
- ratio = I / args.hidden
109
- flag = " <- fat FFN" if ratio > 4 else (" <- balanced" if ratio <= 3 else "")
110
- print(f" {L:>6} {I:>12} {ratio:>7.2f}x {r['total']/1e9:>7.2f}B{flag}")
111
- return
112
-
113
- ap.error("provide either --layers with --intermediate, or --target-b")
114
-
115
-
116
- if __name__ == "__main__":
117
- main()