thefinalboss commited on
Commit
ec293e9
·
verified ·
1 Parent(s): 07c235f

Upload fractus/grow.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. fractus/grow.py +234 -0
fractus/grow.py ADDED
@@ -0,0 +1,234 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Fractus Progressive Growth — grow a model's width, depth, and expert count.
2
+
3
+ THE INNOVATION. Instead of training a large model from scratch (impossible on
4
+ CPU), we grow it palier by palier. Each palier inherits the previous model's
5
+ weights via zero-padding (for width) or copying (for depth/experts), then
6
+ trains briefly. The model never starts from random — it starts "warm".
7
+
8
+ Growth axes:
9
+ - WIDTH (d_model): zero-pad every d-coupled matrix to the new dimension.
10
+ - DEPTH (n_layers): copy old blocks, init new ones with standard scheme.
11
+ - EXPERTS (n_experts): copy old experts, zero-init new ones.
12
+ - RANK (siren_rank): zero-pad the rank dimension of U/V factors.
13
+ """
14
+
15
+ from __future__ import annotations
16
+ import torch
17
+ import torch.nn as nn
18
+ from typing import Dict, Any
19
+
20
+
21
+ def _pad_dim0(tensor: torch.Tensor, new_size: int) -> torch.Tensor:
22
+ """Grow a tensor along dim 0, zero-padding the new rows."""
23
+ old = tensor.shape[0]
24
+ if old >= new_size:
25
+ return tensor
26
+ pad_shape = list(tensor.shape)
27
+ pad_shape[0] = new_size - old
28
+ pad = torch.zeros(pad_shape, dtype=tensor.dtype)
29
+ return torch.cat([tensor, pad], dim=0)
30
+
31
+
32
+ def _pad_last_dim(tensor: torch.Tensor, new_size: int) -> torch.Tensor:
33
+ """Grow a tensor along the LAST dim, zero-padding the new columns."""
34
+ old = tensor.shape[-1]
35
+ if old >= new_size:
36
+ return tensor
37
+ pad_shape = list(tensor.shape)
38
+ pad_shape[-1] = new_size - old
39
+ pad = torch.zeros(pad_shape, dtype=tensor.dtype)
40
+ return torch.cat([tensor, pad], dim=-1)
41
+
42
+
43
+ def _transfer_block_weights(old_blk, new_blk, old_d: int, new_d: int):
44
+ """Transfer weights from old CTEBlock to new CTEBlock via zero-padding.
45
+
46
+ Copies old knowledge into the top-left block of every matrix. New dims
47
+ are zero (neutral start). New experts/blocks warm up during training.
48
+ """
49
+
50
+ # 1. Attention w_qkv: list of 3 tensors, each (d_model, d_model).
51
+ if hasattr(old_blk.attn, "w_qkv"):
52
+ for i in range(min(len(old_blk.attn.w_qkv), len(new_blk.attn.w_qkv))):
53
+ old_w = old_blk.attn.w_qkv[i].data
54
+ new_w = new_blk.attn.w_qkv[i].data
55
+ r_copy = min(old_w.shape[0], new_w.shape[0])
56
+ c_copy = min(old_w.shape[1], new_w.shape[1])
57
+ new_w.zero_()
58
+ new_w[:r_copy, :c_copy] = old_w[:r_copy, :c_copy]
59
+ # Attention biases.
60
+ if hasattr(old_blk.attn, "b_qkv"):
61
+ for i in range(min(len(old_blk.attn.b_qkv), len(new_blk.attn.b_qkv))):
62
+ old_b = old_blk.attn.b_qkv[i].data
63
+ new_b = new_blk.attn.b_qkv[i].data
64
+ c_copy = min(old_b.shape[0], new_b.shape[0])
65
+ new_b[:c_copy] = old_b[:c_copy]
66
+ # Attention w_out.
67
+ if hasattr(old_blk.attn, "w_out"):
68
+ old_wo = old_blk.attn.w_out.data
69
+ new_wo = new_blk.attn.w_out.data
70
+ r_copy = min(old_wo.shape[0], new_wo.shape[0])
71
+ c_copy = min(old_wo.shape[1], new_wo.shape[1])
72
+ new_wo.zero_()
73
+ new_wo[:r_copy, :c_copy] = old_wo[:r_copy, :c_copy]
74
+ if hasattr(old_blk.attn, "b_out"):
75
+ old_bo = old_blk.attn.b_out.data
76
+ new_bo = new_blk.attn.b_out.data
77
+ c_copy = min(old_bo.shape[0], new_bo.shape[0])
78
+ new_bo[:c_copy] = old_bo[:c_copy]
79
+ # Level offsets.
80
+ if hasattr(old_blk.attn, "level_offsets") and hasattr(new_blk.attn, "level_offsets"):
81
+ old_lo = old_blk.attn.level_offsets.data
82
+ new_lo = new_blk.attn.level_offsets.data
83
+ n_copy = min(old_lo.shape[0], new_lo.shape[0])
84
+ new_lo[:n_copy] = old_lo[:n_copy]
85
+ # Level logits.
86
+ if hasattr(old_blk.attn, "level_logits") and hasattr(new_blk.attn, "level_logits"):
87
+ old_ll = old_blk.attn.level_logits.data
88
+ new_ll = new_blk.attn.level_logits.data
89
+ n_copy = min(old_ll.shape[0], new_ll.shape[0])
90
+ new_ll[:n_copy] = old_ll[:n_copy]
91
+
92
+ # 2. LayerNorms: copy old dims, gamma=1/beta=0 for new.
93
+ for (old_norm, new_norm) in [
94
+ (old_blk.norm_attn, new_blk.norm_attn),
95
+ (old_blk.norm_kur, new_blk.norm_kur),
96
+ (old_blk.norm_moe, new_blk.norm_moe),
97
+ ]:
98
+ old_g = old_norm.weight.data
99
+ new_g = new_norm.weight.data
100
+ d_copy = min(old_g.shape[0], new_g.shape[0])
101
+ new_g[:d_copy] = old_g[:d_copy]
102
+ old_b = old_norm.bias.data
103
+ new_b = new_norm.bias.data
104
+ new_b[:d_copy] = old_b[:d_copy]
105
+
106
+ # 3. Kuramoto: grow oscillators + coupling rank.
107
+ if hasattr(old_blk.kuramoto, "omega"):
108
+ old_om = old_blk.kuramoto.omega.data
109
+ new_om = new_blk.kuramoto.omega.data
110
+ n_copy = min(old_om.shape[0], new_om.shape[0])
111
+ new_om[:n_copy] = old_om[:n_copy]
112
+ if hasattr(old_blk.kuramoto, "coupling_u"):
113
+ old_cu = old_blk.kuramoto.coupling_u.data
114
+ new_cu = new_blk.kuramoto.coupling_u.data
115
+ n_copy = min(old_cu.shape[0], new_cu.shape[0])
116
+ r_copy = min(old_cu.shape[1], new_cu.shape[1])
117
+ new_cu[:n_copy, :r_copy] = old_cu[:n_copy, :r_copy]
118
+ if hasattr(old_blk.kuramoto, "coupling_lambda"):
119
+ old_cl = old_blk.kuramoto.coupling_lambda.data
120
+ new_cl = new_blk.kuramoto.coupling_lambda.data
121
+ r_copy = min(old_cl.shape[0], new_cl.shape[0])
122
+ new_cl[:r_copy] = old_cl[:r_copy]
123
+
124
+ # 4. MoE experts (low-rank): U1/V1/U2/V2 are (E, ..., r).
125
+ moe_old = old_blk.moe
126
+ moe_new = new_blk.moe
127
+ e_copy = min(moe_old.n_experts, moe_new.n_experts)
128
+
129
+ if moe_old.expert_rank is not None and moe_new.expert_rank is not None:
130
+ for param_name in ["U1", "V1", "U2", "V2", "scale1", "scale2", "b1", "b2"]:
131
+ old_p = getattr(moe_old, param_name).data
132
+ new_p = getattr(moe_new, param_name).data
133
+ new_p.zero_() # neutral start for ALL
134
+ if old_p.dim() == 3:
135
+ dd = min(old_p.shape[1], new_p.shape[1])
136
+ rr = min(old_p.shape[2], new_p.shape[2])
137
+ new_p[:e_copy, :dd, :rr] = old_p[:e_copy, :dd, :rr]
138
+ elif old_p.dim() == 2:
139
+ dd = min(old_p.shape[1], new_p.shape[1])
140
+ new_p[:e_copy, :dd] = old_p[:e_copy, :dd]
141
+ elif old_p.dim() == 1:
142
+ new_p[:e_copy] = old_p[:e_copy]
143
+ # Restore scale=1 for OLD experts.
144
+ if param_name in ("scale1", "scale2"):
145
+ old_scale = getattr(moe_old, param_name).data[:e_copy]
146
+ new_p[:e_copy] = old_scale
147
+
148
+
149
+ def grow_cte(old_engine, new_config: Dict[str, Any]):
150
+ """Grow a ContinuousThoughtEngine to a larger config.
151
+
152
+ Copies old weights into the new model via zero-padding. Supports growth
153
+ in width (d_model), depth (n_layers), experts (n_experts), and rank
154
+ (siren_rank). Old knowledge is preserved; new capacity is neutral.
155
+
156
+ Args:
157
+ old_engine: a trained ContinuousThoughtEngine.
158
+ new_config: dict with any of:
159
+ d_model, n_layers, n_experts, top_k, expert_d_ff, siren_rank,
160
+ n_heads, d_head, n_levels, n_oscillators, coupling_rank, vocab_size.
161
+ """
162
+ from .continuous_engine import ContinuousThoughtEngine
163
+
164
+ old_d = old_engine.d_model
165
+ old_vocab = old_engine.vocab_size
166
+ old_n_layers = len(old_engine.blocks)
167
+
168
+ new_d = new_config.get("d_model", old_d)
169
+ new_vocab = new_config.get("vocab_size", old_vocab)
170
+ new_n_layers = new_config.get("n_layers", old_n_layers)
171
+ new_n_experts = new_config.get("n_experts", old_engine.blocks[0].moe.n_experts)
172
+ new_rank = new_config.get("siren_rank", old_engine.blocks[0].moe.expert_rank or 32)
173
+ new_d_ff = new_config.get("expert_d_ff", old_engine.blocks[0].moe.d_ff)
174
+ new_top_k = new_config.get("top_k", old_engine.blocks[0].moe.top_k)
175
+ new_n_heads = new_config.get("n_heads", old_engine.blocks[0].attn.n_heads)
176
+ new_d_head = new_config.get("d_head", old_engine.blocks[0].attn.d_head)
177
+ new_n_levels = new_config.get("n_levels", old_engine.blocks[0].attn.n_levels)
178
+ new_n_osc = new_config.get("n_oscillators", old_engine.blocks[0].kuramoto.N)
179
+ new_coupling_rank = new_config.get("coupling_rank", old_engine.blocks[0].kuramoto.rank)
180
+
181
+ # Build the new engine.
182
+ new_engine = ContinuousThoughtEngine(
183
+ vocab_size=new_vocab, d_model=new_d,
184
+ n_heads=new_n_heads, d_head=new_d_head, n_levels=new_n_levels,
185
+ n_oscillators=new_n_osc, coupling_rank=new_coupling_rank,
186
+ n_experts=new_n_experts, top_k=new_top_k,
187
+ expert_d_ff=new_d_ff, siren_rank=(new_rank if new_rank else None),
188
+ n_layers=new_n_layers,
189
+ )
190
+
191
+ # --- Transfer embedding (shared across all blocks) ---
192
+ old_emb = old_engine.observe.weight.data
193
+ new_emb = new_engine.observe.weight.data
194
+ v_copy = min(old_vocab, new_vocab)
195
+ d_copy = min(old_d, new_d)
196
+ new_emb.zero_()
197
+ new_emb[:v_copy, :d_copy] = old_emb[:v_copy, :d_copy]
198
+
199
+ # --- Transfer per-block weights (loop over ALL old blocks) ---
200
+ blocks_to_copy = min(old_n_layers, new_n_layers)
201
+ for blk_idx in range(blocks_to_copy):
202
+ _transfer_block_weights(
203
+ old_engine.blocks[blk_idx],
204
+ new_engine.blocks[blk_idx],
205
+ old_d, new_d,
206
+ )
207
+ # New blocks (indices old_n_layers..new_n_layers-1) keep random init —
208
+ # they warm up during training (like new brain regions developing).
209
+
210
+ # --- Transfer heads (top-level, not per-block) ---
211
+ for head_name in ["confidence_head", "salience_head"]:
212
+ old_h = getattr(old_engine, head_name)
213
+ new_h = getattr(new_engine, head_name)
214
+ old_w = old_h.weight.data
215
+ new_w = new_h.weight.data
216
+ d_copy = min(old_w.shape[1], new_w.shape[1])
217
+ new_w[:, :d_copy] = old_w[:, :d_copy]
218
+ new_h.bias.data[:] = old_h.bias.data[:]
219
+
220
+ # Output head is TIED with embedding — already handled above.
221
+
222
+ return new_engine
223
+
224
+
225
+ def grow_summary(old_engine, new_engine) -> dict:
226
+ """Report what changed between old and new engine."""
227
+ return {
228
+ "d_model": f"{old_engine.d_model} → {new_engine.d_model}",
229
+ "n_layers": f"{len(old_engine.blocks)} → {len(new_engine.blocks)}",
230
+ "n_experts": f"{old_engine.blocks[0].moe.n_experts} → {new_engine.blocks[0].moe.n_experts}",
231
+ "expert_rank": f"{old_engine.blocks[0].moe.expert_rank} → {new_engine.blocks[0].moe.expert_rank}",
232
+ "params": f"{sum(p.numel() for p in old_engine.parameters()):,} → {sum(p.numel() for p in new_engine.parameters()):,}",
233
+ "n_oscillators": f"{old_engine.blocks[0].kuramoto.N} → {new_engine.blocks[0].kuramoto.N}",
234
+ }