WikifitaGitHub live67e8de5
outro · compute-optimization/ml-architectures/muonclip-optimizer

MuonClip Optimizer — Newton-Schulz Orthogonalization and QK-Clip

Custom MLX optimizer: Newton-Schulz matrix orthogonalization, QK-Clip for gradient stabilization, rank-1 matrix hazard, 1D variant for holographic weights.

Baixar raw

MuonClip Optimizer

A custom optimizer built on mlx.optimizers.Optimizer for the ring-transformer. Two implementations exist: a full MLX-native version for the neural network, and a 1D variant for the holographic weight array.


MLX-Native Implementation

File: src/chronobreaker/train/muonclip_mlx.py (110 lines)

Extends mlx.optimizers.Optimizer with two custom mechanisms: Newton-Schulz orthogonalization for weight matrices, and QK-Clip for attention gradient stabilization.

Newton-Schulz Orthogonalization

# muonclip_mlx.py, lines 5-14
def newton_schulz(X, steps=5):
    """X_{k+1} = 1.5 * X_k - 0.5 * X_k @ (X_k^T @ X_k)"""
    for _ in range(steps):
        X = 1.5 * X - 0.5 * X @ (X.T @ X)
    return X

This iterative algorithm approximates the matrix inverse square root (used in Shampoo/Muon-style optimizers). After 5 iterations, the result is approximately orthogonal — it preserves the direction of the gradient while normalizing its magnitude.

Purpose: Prevents gradient explosion in deep attention networks by ensuring update matrices are approximately orthogonal.

Momentum

# Default decay: 0.95
momentum = decay * momentum + (1 - decay) * grad

Standard exponential moving average momentum.

QK-Clip

# muonclip_mlx.py, lines 83-110
tau = 100  # Threshold
for layer in model.layers:
    attn = layer.attention
    qk_scores = attn.Q @ attn.K.T  # Attention logits
    max_logit = mx.max(mx.abs(qk_scores))
    if max_logit > tau:
        gamma = tau / max_logit  # Per-head scaling factor
        attn.Q.weight *= gamma
        attn.K.weight *= gamma

When attention logit magnitudes exceed the threshold tau=100, Q and K weights are scaled down proportionally. This prevents the "attention entropy collapse" where a few tokens dominate all attention.

Per-head gamma: Each attention head gets its own scaling factor based on its maximum logit magnitude.


The Rank-1 Matrix Hazard

Critical finding: Newton-Schulz orthogonalization destroys rank-1 matrices.

The output head of ring-transformer is nn.Linear(128, 1, bias=False) — a weight matrix of shape [128, 1], which is rank-1 by construction. Applying Newton-Schulz to this matrix would zero it out (the orthogonalization converges to a zero matrix for rank-deficient inputs).

Mitigation:

# trainer.py, line 150
use_muon_for_linears = False  # "Muon destroys rank-1 matrices like the [d, 1] head projection"

MuonClip is only applied to the attention weight matrices (Q, K, V, Out), not to linear projections. The standard MLX optimizer (Adam/SGD) handles the linear layers.

Lesson: Any custom optimizer using matrix orthogonalization must check matrix rank before applying the transformation. Rank-deficient matrices require fallback to standard update rules.


1D Variant (Holographic Weights)

File: src/chronobreaker/train/muon_clip.py (77 lines)

A simplified variant operating on the 19-element holographic WeightsArray:

# Domain splitting
sha256_rings = weights[0:12]    # SHA-256 carrier rings
context_rings = weights[12:16]  # Context rings
mdl_coeffs = weights[16:19]     # MDL coefficients

# Per-domain Newton-Schulz with different scaling
sha256_rings = newton_schulz_1d(sha256_rings, scale=0.4)
context_rings = newton_schulz_1d(context_rings, scale=0.2)
mdl_coeffs = newton_schulz_1d(mdl_coeffs, scale=0.1)

Selective Weight Decay

decay_mask = [1]*16 + [0]*3  # Decay only ring weights (0-15), not MDL (16-18)

MDL coefficients are protected from weight decay because they represent structural information (model complexity penalties), not learned features.


Integration with RingTransformer

Training Step
  ├── Forward pass (MLX)
  ├── Loss computation (cyclic loss)
  ├── Autograd: nn.value_and_grad()
  ├── MuonClip update (attention weights only)
  │    ├── Newton-Schulz on Q, K, V, Out weight matrices
  │    └── Momentum update
  ├── QK-Clip pass (post-update)
  │    └── Scale Q, K if max_logit > 100
  └── Standard optimizer update (linear layers)

Design Decisions

DecisionRationale
5 Newton-Schulz iterationsEmpirically sufficient for convergence
tau=100 for QK-ClipAbove this, attention entropy collapses
Momentum decay 0.95Standard for transformer training
No Muon on linearsRank-1 matrix destruction
No decay on MDLStructural information preservation

Cross-References