Upload fractus/grow.py with huggingface_hub
Browse files- 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 |
+
}
|