Robotics
multilingual
ternary
multimodal
pretraining
jirack
ternarytransformer
kgrabko commited on
Commit
59162db
·
verified ·
1 Parent(s): f85fe90

Create convert_to_14b.py

Browse files
Files changed (1) hide show
  1. convert_to_14b.py +117 -0
convert_to_14b.py ADDED
@@ -0,0 +1,117 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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()