modmul-challenge / coppola_pretraining.py
alstrup's picture
modmul router v0: t12 + composed-t3b members (pre-t3g)
14bef4a verified
Raw
History Blame Contribute Delete
24.5 kB
#!/usr/bin/env python3
"""Utilities for full Coppola-typed pretraining in nanochat-style trainers.
This module is intentionally trainer-agnostic. It provides:
- basis extraction from current attention-output and MLP down-project weights
- typed gradient decomposition and projection
- a depth-aware scale policy
- a small controller that can refresh bases from a nanochat-style model
The intended integration is:
1. instantiate a controller from the current model weights
2. before each matrix optimizer step, project the raw gradient by family
3. apply Muon / Newton-Schulz or another matrix optimizer
4. re-project the transformed gradient before the final step
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Callable, Dict, List, Optional, Sequence, Tuple
import torch
Tensor = torch.Tensor
@dataclass(frozen=True)
class RankPolicy:
"""Energy-based rank truncation policy."""
energy: float = 0.99
max_rank: Optional[int] = None
min_rank: int = 1
@dataclass(frozen=True)
class FamilyScales:
"""Typed gradient multipliers for a layer."""
routing: float = 1.0
and_transport: float = 1.0
or_update: float = 1.0
remainder: float = 1.0
coupling: Optional[float] = None
def coupling_scale(self) -> float:
if self.coupling is not None:
return self.coupling
return 0.5 * (self.and_transport + self.or_update)
@dataclass
class LayerBases:
"""Typed bases for one transformer layer."""
routing_head_bases: List[Tensor] = field(default_factory=list)
q_head_bases: List[Tensor] = field(default_factory=list)
k_head_bases: List[Tensor] = field(default_factory=list)
and_output_basis: Optional[Tensor] = None
or_input_basis: Optional[Tensor] = None
@dataclass
class ComponentNorms:
"""Norms of the typed gradient components before scaling."""
routing: float = 0.0
and_transport: float = 0.0
or_update: float = 0.0
coupling: float = 0.0
remainder: float = 0.0
def _select_rank(singular_values: Tensor, policy: RankPolicy) -> int:
if singular_values.numel() == 0:
return policy.min_rank
total = float(singular_values.square().sum().item())
if total <= 0.0:
return policy.min_rank
cumulative = singular_values.square().cumsum(0) / total
rank = int(torch.searchsorted(cumulative, torch.tensor(policy.energy, device=cumulative.device)).item()) + 1
rank = max(policy.min_rank, rank)
if policy.max_rank is not None:
rank = min(rank, policy.max_rank)
return min(rank, singular_values.numel())
def _orthonormal_rows(basis: Tensor) -> Tensor:
"""Return an orthonormal row basis spanning the same row space."""
if basis.ndim != 2:
raise ValueError(f"basis must be rank-2, got {tuple(basis.shape)}")
if basis.shape[0] == 0:
return basis
q, _ = torch.linalg.qr(basis.T, mode="reduced")
return q.T.contiguous()
def _project_left(grad: Tensor, basis: Tensor) -> Tensor:
basis = basis.to(device=grad.device, dtype=grad.dtype)
proj = basis.T @ basis
return proj @ grad
def _project_right(grad: Tensor, basis: Tensor) -> Tensor:
basis = basis.to(device=grad.device, dtype=grad.dtype)
proj = basis.T @ basis
return grad @ proj
def _zero_like(grad: Tensor) -> Tensor:
return torch.zeros_like(grad)
def compute_attention_output_bases(
weight: Tensor,
n_head: int,
policy: RankPolicy = RankPolicy(),
) -> List[Tensor]:
"""Build one output-space basis per attention head block.
Args:
weight: output projection weight of shape [hidden, hidden]
n_head: number of attention heads
policy: rank selection policy inside each head block
"""
if weight.ndim != 2:
raise ValueError(f"attention output weight must be rank-2, got {tuple(weight.shape)}")
hidden, width = weight.shape
if width % n_head != 0:
raise ValueError(f"weight width {width} not divisible by n_head={n_head}")
head_dim = width // n_head
bases: List[Tensor] = []
with torch.no_grad():
w = weight.detach().float().cpu()
for h in range(n_head):
cols = w[:, h * head_dim:(h + 1) * head_dim]
u, s, _ = torch.linalg.svd(cols, full_matrices=False)
rank = _select_rank(s, policy)
basis = u[:, :rank].T.contiguous()
bases.append(_orthonormal_rows(basis))
return bases
def compute_attention_score_bases(
weight: Tensor,
n_head: int,
policy: RankPolicy = RankPolicy(),
) -> List[Tensor]:
"""Build one row-space basis per attention head for Q/K score-side blocks."""
if weight.ndim != 2:
raise ValueError(f"attention score weight must be rank-2, got {tuple(weight.shape)}")
width, hidden = weight.shape
if width % n_head != 0:
raise ValueError(f"weight height {width} not divisible by n_head={n_head}")
head_dim = width // n_head
bases: List[Tensor] = []
with torch.no_grad():
w = weight.detach().float().cpu()
for h in range(n_head):
rows = w[h * head_dim:(h + 1) * head_dim, :]
u, s, _ = torch.linalg.svd(rows, full_matrices=False)
rank = _select_rank(s, policy)
basis = u[:, :rank].T.contiguous()
bases.append(_orthonormal_rows(basis))
return bases
def compute_down_proj_bases(
weight: Tensor,
output_policy: RankPolicy = RankPolicy(),
input_policy: RankPolicy = RankPolicy(),
) -> Tuple[Tensor, Tensor]:
"""Build output-side AND basis and input-side OR basis for down_proj."""
if weight.ndim != 2:
raise ValueError(f"down projection weight must be rank-2, got {tuple(weight.shape)}")
with torch.no_grad():
w = weight.detach().float().cpu()
u, s, vh = torch.linalg.svd(w, full_matrices=False)
out_rank = _select_rank(s, output_policy)
in_rank = _select_rank(s, input_policy)
and_basis = _orthonormal_rows(u[:, :out_rank].T.contiguous())
or_basis = _orthonormal_rows(vh[:in_rank, :].contiguous())
return and_basis, or_basis
def decompose_down_proj_gradient(
grad: Tensor,
and_output_basis: Optional[Tensor],
or_input_basis: Optional[Tensor],
) -> Dict[str, Tensor]:
"""Decompose an MLP down_proj gradient into typed components.
The components are:
- and_transport_only
- or_update_only
- coupling: in both the output and input typed subspaces
- remainder
"""
if grad.ndim != 2:
raise ValueError(f"gradient must be rank-2, got {tuple(grad.shape)}")
g_and = _zero_like(grad)
g_or = _zero_like(grad)
g_both = _zero_like(grad)
if and_output_basis is not None:
g_and = _project_left(grad, and_output_basis)
if or_input_basis is not None:
g_or = _project_right(grad, or_input_basis)
if and_output_basis is not None and or_input_basis is not None:
g_both = _project_right(g_and, or_input_basis)
g_and_only = g_and - g_both
g_or_only = g_or - g_both
g_remainder = grad - g_and_only - g_or_only - g_both
return {
"and_transport_only": g_and_only,
"or_update_only": g_or_only,
"coupling": g_both,
"remainder": g_remainder,
}
def project_down_proj_gradient(
grad: Tensor,
and_output_basis: Optional[Tensor],
or_input_basis: Optional[Tensor],
scales: FamilyScales,
) -> Tensor:
parts = decompose_down_proj_gradient(grad, and_output_basis, or_input_basis)
return (
scales.and_transport * parts["and_transport_only"]
+ scales.or_update * parts["or_update_only"]
+ scales.coupling_scale() * parts["coupling"]
+ scales.remainder * parts["remainder"]
)
def support_project_down_proj_gradient(
grad: Tensor,
and_output_basis: Optional[Tensor],
or_input_basis: Optional[Tensor],
) -> Tensor:
"""Project a down-proj update onto typed Coppola support, dropping remainder."""
parts = decompose_down_proj_gradient(grad, and_output_basis, or_input_basis)
return parts["and_transport_only"] + parts["or_update_only"] + parts["coupling"]
def down_proj_component_norms(
grad: Tensor,
and_output_basis: Optional[Tensor],
or_input_basis: Optional[Tensor],
) -> ComponentNorms:
parts = decompose_down_proj_gradient(grad, and_output_basis, or_input_basis)
return ComponentNorms(
and_transport=float(parts["and_transport_only"].norm().item()),
or_update=float(parts["or_update_only"].norm().item()),
coupling=float(parts["coupling"].norm().item()),
remainder=float(parts["remainder"].norm().item()),
)
def project_attention_output_gradient(
grad: Tensor,
routing_head_bases: Sequence[Tensor],
scales: FamilyScales,
) -> Tensor:
"""Project an attention output gradient head-by-head on the left."""
if grad.ndim != 2:
raise ValueError(f"gradient must be rank-2, got {tuple(grad.shape)}")
if not routing_head_bases:
return grad
hidden, width = grad.shape
if width % len(routing_head_bases) != 0:
raise ValueError(f"gradient width {width} incompatible with {len(routing_head_bases)} head bases")
head_dim = width // len(routing_head_bases)
out = torch.empty_like(grad)
for h, basis in enumerate(routing_head_bases):
cols = grad[:, h * head_dim:(h + 1) * head_dim]
routed = _project_left(cols, basis)
remainder = cols - routed
out[:, h * head_dim:(h + 1) * head_dim] = scales.routing * routed + scales.remainder * remainder
return out
def project_attention_qk_gradient(
grad: Tensor,
q_head_bases: Sequence[Tensor],
k_head_bases: Sequence[Tensor],
scales: FamilyScales,
) -> Tensor:
"""Project fused Q/K/V gradient on Q and K row blocks only, leaving V unchanged."""
if grad.ndim != 2:
raise ValueError(f"gradient must be rank-2, got {tuple(grad.shape)}")
if not q_head_bases and not k_head_bases:
return grad
hidden3, width = grad.shape
if hidden3 % 3 != 0:
raise ValueError(f"expected fused qkv height divisible by 3, got {hidden3}")
hidden = hidden3 // 3
if q_head_bases and hidden % len(q_head_bases) != 0:
raise ValueError(f"q block height {hidden} incompatible with {len(q_head_bases)} q bases")
if k_head_bases and hidden % len(k_head_bases) != 0:
raise ValueError(f"k block height {hidden} incompatible with {len(k_head_bases)} k bases")
out = grad.clone()
if q_head_bases:
head_dim = hidden // len(q_head_bases)
for h, basis in enumerate(q_head_bases):
rows = grad[h * head_dim:(h + 1) * head_dim, :]
routed = _project_left(rows, basis)
remainder = rows - routed
out[h * head_dim:(h + 1) * head_dim, :] = scales.routing * routed + scales.remainder * remainder
if k_head_bases:
head_dim = hidden // len(k_head_bases)
offset = hidden
for h, basis in enumerate(k_head_bases):
start = offset + h * head_dim
stop = offset + (h + 1) * head_dim
rows = grad[start:stop, :]
routed = _project_left(rows, basis)
remainder = rows - routed
out[start:stop, :] = scales.routing * routed + scales.remainder * remainder
return out
def support_project_attention_output_gradient(
grad: Tensor,
routing_head_bases: Sequence[Tensor],
) -> Tensor:
"""Project an attention-output update onto routing support, dropping remainder."""
if grad.ndim != 2:
raise ValueError(f"gradient must be rank-2, got {tuple(grad.shape)}")
if not routing_head_bases:
return grad
hidden, width = grad.shape
if width % len(routing_head_bases) != 0:
raise ValueError(f"gradient width {width} incompatible with {len(routing_head_bases)} head bases")
head_dim = width // len(routing_head_bases)
out = torch.empty_like(grad)
for h, basis in enumerate(routing_head_bases):
cols = grad[:, h * head_dim:(h + 1) * head_dim]
out[:, h * head_dim:(h + 1) * head_dim] = _project_left(cols, basis)
return out
def support_project_attention_qk_gradient(
grad: Tensor,
q_head_bases: Sequence[Tensor],
k_head_bases: Sequence[Tensor],
) -> Tensor:
"""Project fused Q/K/V update on Q and K support only, leaving V unchanged."""
if grad.ndim != 2:
raise ValueError(f"gradient must be rank-2, got {tuple(grad.shape)}")
if not q_head_bases and not k_head_bases:
return grad
hidden3, width = grad.shape
if hidden3 % 3 != 0:
raise ValueError(f"expected fused qkv height divisible by 3, got {hidden3}")
hidden = hidden3 // 3
out = grad.clone()
if q_head_bases:
head_dim = hidden // len(q_head_bases)
for h, basis in enumerate(q_head_bases):
rows = grad[h * head_dim:(h + 1) * head_dim, :]
out[h * head_dim:(h + 1) * head_dim, :] = _project_left(rows, basis)
if k_head_bases:
head_dim = hidden // len(k_head_bases)
offset = hidden
for h, basis in enumerate(k_head_bases):
start = offset + h * head_dim
stop = offset + (h + 1) * head_dim
rows = grad[start:stop, :]
out[start:stop, :] = _project_left(rows, basis)
return out
def attention_component_norms(grad: Tensor, routing_head_bases: Sequence[Tensor]) -> ComponentNorms:
if not routing_head_bases:
return ComponentNorms(routing=float(grad.norm().item()))
hidden, width = grad.shape
head_dim = width // len(routing_head_bases)
routed_norm_sq = 0.0
remainder_norm_sq = 0.0
for h, basis in enumerate(routing_head_bases):
cols = grad[:, h * head_dim:(h + 1) * head_dim]
routed = _project_left(cols, basis)
remainder = cols - routed
routed_norm_sq += float(routed.square().sum().item())
remainder_norm_sq += float(remainder.square().sum().item())
return ComponentNorms(
routing=routed_norm_sq**0.5,
remainder=remainder_norm_sq**0.5,
)
def attention_qk_component_norms(
grad: Tensor,
q_head_bases: Sequence[Tensor],
k_head_bases: Sequence[Tensor],
) -> ComponentNorms:
if not q_head_bases and not k_head_bases:
return ComponentNorms(routing=float(grad.norm().item()))
hidden3, width = grad.shape
if hidden3 % 3 != 0:
raise ValueError(f"expected fused qkv height divisible by 3, got {hidden3}")
hidden = hidden3 // 3
routed_norm_sq = 0.0
remainder_norm_sq = 0.0
if q_head_bases:
head_dim = hidden // len(q_head_bases)
for h, basis in enumerate(q_head_bases):
rows = grad[h * head_dim:(h + 1) * head_dim, :]
routed = _project_left(rows, basis)
remainder = rows - routed
routed_norm_sq += float(routed.square().sum().item())
remainder_norm_sq += float(remainder.square().sum().item())
if k_head_bases:
head_dim = hidden // len(k_head_bases)
offset = hidden
for h, basis in enumerate(k_head_bases):
start = offset + h * head_dim
stop = offset + (h + 1) * head_dim
rows = grad[start:stop, :]
routed = _project_left(rows, basis)
remainder = rows - routed
routed_norm_sq += float(routed.square().sum().item())
remainder_norm_sq += float(remainder.square().sum().item())
return ComponentNorms(
routing=routed_norm_sq**0.5,
remainder=remainder_norm_sq**0.5,
)
def default_zone_scales(layer_idx: int, n_layer: int) -> FamilyScales:
"""Default full-Coppola depth policy for from-scratch pretraining."""
frac = layer_idx / max(n_layer - 1, 1)
if frac < 0.20:
return FamilyScales(routing=1.40, and_transport=1.10, or_update=0.70, remainder=0.0, coupling=0.90)
if frac < 0.65:
return FamilyScales(routing=0.90, and_transport=1.40, or_update=0.90, remainder=0.0, coupling=1.15)
if frac < 0.90:
return FamilyScales(routing=0.75, and_transport=0.95, or_update=1.40, remainder=0.0, coupling=1.15)
return FamilyScales(routing=0.60, and_transport=0.85, or_update=1.15, remainder=0.0, coupling=1.00)
@dataclass
class CoppolaPretrainingConfig:
n_head: int
attn_output_policy: RankPolicy = field(default_factory=RankPolicy)
attn_qk_policy: RankPolicy = field(default_factory=RankPolicy)
mlp_output_policy: RankPolicy = field(default_factory=RankPolicy)
mlp_input_policy: RankPolicy = field(default_factory=RankPolicy)
basis_update_interval: int = 250
uniform_scales: FamilyScales = field(default_factory=FamilyScales)
scale_fn: Callable[[int, int], FamilyScales] = default_zone_scales
class CoppolaPretrainingController:
"""Refreshes typed bases and projects gradients for nanochat-style models."""
def __init__(self, config: CoppolaPretrainingConfig):
self.config = config
self.layer_bases: Dict[int, LayerBases] = {}
def refresh_from_model(self, model) -> None:
layers = self._resolve_layers(model)
n_layer = len(layers)
bases: Dict[int, LayerBases] = {}
for layer_idx, block in enumerate(layers):
attn_weight = self._resolve_attn_out_weight(block)
q_weight, k_weight = self._resolve_attn_qk_weights(block)
mlp_weight = self._resolve_mlp_down_weight(block)
routing = compute_attention_output_bases(
attn_weight, self.config.n_head, self.config.attn_output_policy
)
q_bases: List[Tensor] = []
k_bases: List[Tensor] = []
if q_weight is not None and k_weight is not None:
q_bases = compute_attention_score_bases(
q_weight, self.config.n_head, self.config.attn_qk_policy
)
k_bases = compute_attention_score_bases(
k_weight, self.config.n_head, self.config.attn_qk_policy
)
and_basis, or_basis = compute_down_proj_bases(
mlp_weight,
output_policy=self.config.mlp_output_policy,
input_policy=self.config.mlp_input_policy,
)
bases[layer_idx] = LayerBases(
routing_head_bases=routing,
q_head_bases=q_bases,
k_head_bases=k_bases,
and_output_basis=and_basis,
or_input_basis=or_basis,
)
self.layer_bases = bases
self._n_layer = n_layer
def scales_for_layer(self, layer_idx: int, mode: str = "zoned") -> FamilyScales:
if mode == "uniform":
return self.config.uniform_scales
if not hasattr(self, "_n_layer"):
raise RuntimeError("refresh_from_model() must be called before requesting zoned scales")
return self.config.scale_fn(layer_idx, self._n_layer)
def project_attn_out_grad(self, layer_idx: int, grad: Tensor, mode: str = "zoned") -> Tensor:
bases = self.layer_bases[layer_idx]
scales = self.scales_for_layer(layer_idx, mode=mode)
return project_attention_output_gradient(grad, bases.routing_head_bases, scales)
def project_attn_qk_grad(self, layer_idx: int, grad: Tensor, mode: str = "zoned") -> Tensor:
bases = self.layer_bases[layer_idx]
scales = self.scales_for_layer(layer_idx, mode=mode)
return project_attention_qk_gradient(grad, bases.q_head_bases, bases.k_head_bases, scales)
def project_mlp_down_grad(self, layer_idx: int, grad: Tensor, mode: str = "zoned") -> Tensor:
bases = self.layer_bases[layer_idx]
scales = self.scales_for_layer(layer_idx, mode=mode)
return project_down_proj_gradient(grad, bases.and_output_basis, bases.or_input_basis, scales)
def project_attn_out_support(self, layer_idx: int, grad: Tensor) -> Tensor:
bases = self.layer_bases[layer_idx]
return support_project_attention_output_gradient(grad, bases.routing_head_bases)
def project_attn_qk_support(self, layer_idx: int, grad: Tensor) -> Tensor:
bases = self.layer_bases[layer_idx]
return support_project_attention_qk_gradient(grad, bases.q_head_bases, bases.k_head_bases)
def project_mlp_down_support(self, layer_idx: int, grad: Tensor) -> Tensor:
bases = self.layer_bases[layer_idx]
return support_project_down_proj_gradient(grad, bases.and_output_basis, bases.or_input_basis)
def project_after_matrix_transform(
self,
layer_idx: int,
grad: Tensor,
param_kind: str,
transform: Callable[[Tensor], Tensor],
mode: str = "zoned",
) -> Tensor:
"""Project -> transform -> re-project for Muon-style optimizers."""
g = self.project_gradient(layer_idx, grad, param_kind=param_kind, mode=mode)
g = transform(g)
return self.project_gradient(layer_idx, g, param_kind=param_kind, mode=mode)
def project_gradient(self, layer_idx: int, grad: Tensor, param_kind: str, mode: str = "zoned") -> Tensor:
if param_kind == "attn_out":
return self.project_attn_out_grad(layer_idx, grad, mode=mode)
if param_kind == "attn_qk":
return self.project_attn_qk_grad(layer_idx, grad, mode=mode)
if param_kind == "mlp_down":
return self.project_mlp_down_grad(layer_idx, grad, mode=mode)
raise ValueError(f"unknown param_kind={param_kind!r}")
def project_update_support(self, layer_idx: int, grad: Tensor, param_kind: str) -> Tensor:
if param_kind == "attn_out":
return self.project_attn_out_support(layer_idx, grad)
if param_kind == "attn_qk":
return self.project_attn_qk_support(layer_idx, grad)
if param_kind == "mlp_down":
return self.project_mlp_down_support(layer_idx, grad)
raise ValueError(f"unknown param_kind={param_kind!r}")
def gradient_component_norms(self, layer_idx: int, grad: Tensor, param_kind: str) -> ComponentNorms:
bases = self.layer_bases[layer_idx]
if param_kind == "attn_out":
return attention_component_norms(grad, bases.routing_head_bases)
if param_kind == "attn_qk":
return attention_qk_component_norms(grad, bases.q_head_bases, bases.k_head_bases)
if param_kind == "mlp_down":
return down_proj_component_norms(grad, bases.and_output_basis, bases.or_input_basis)
raise ValueError(f"unknown param_kind={param_kind!r}")
@staticmethod
def _resolve_layers(model):
if hasattr(model, "transformer") and hasattr(model.transformer, "h"):
return list(model.transformer.h)
if hasattr(model, "model") and hasattr(model.model, "layers"):
return list(model.model.layers)
raise ValueError("could not resolve transformer layers from model")
@staticmethod
def _resolve_attn_out_weight(block) -> Tensor:
if hasattr(block, "attn") and hasattr(block.attn, "c_proj"):
return block.attn.c_proj.weight
if hasattr(block, "self_attn") and hasattr(block.self_attn, "o_proj"):
return block.self_attn.o_proj.weight
raise ValueError("could not resolve attention output weight")
@staticmethod
def _resolve_attn_qk_weights(block) -> Tuple[Optional[Tensor], Optional[Tensor]]:
if hasattr(block, "attn") and hasattr(block.attn, "c_attn"):
weight = block.attn.c_attn.weight
if weight.shape[0] % 3 != 0:
raise ValueError(f"expected fused qkv height divisible by 3, got {weight.shape[0]}")
hidden = weight.shape[0] // 3
return weight[:hidden, :], weight[hidden:2 * hidden, :]
if hasattr(block, "self_attn") and hasattr(block.self_attn, "q_proj") and hasattr(block.self_attn, "k_proj"):
return block.self_attn.q_proj.weight, block.self_attn.k_proj.weight
return None, None
@staticmethod
def _resolve_mlp_down_weight(block) -> Tensor:
if hasattr(block, "mlp") and hasattr(block.mlp, "c_proj"):
return block.mlp.c_proj.weight
if hasattr(block, "mlp") and hasattr(block.mlp, "down_proj"):
return block.mlp.down_proj.weight
raise ValueError("could not resolve MLP down projection weight")