#!/bin/bash set -e echo "=== GPT-124M Training with mHC ===" # Create the mHC implementation echo "Creating mHC implementation..." cat > /root/mhc.py << 'EOF' """ mHC (Manifold-Constrained Hyper-Connections) Implementation. Based on https://arxiv.org/abs/2512.24880 """ import torch import torch.nn as nn from einops import rearrange, einsum def sinkhorn_knopp(logits, num_iters=20, tau=0.05): """ Project logits to doubly stochastic matrix using Sinkhorn-Knopp in log-space. Alternates row and column normalization until convergence. """ log_alpha = logits / tau for _ in range(num_iters): # Row normalization: subtract logsumexp across columns log_alpha = log_alpha - torch.logsumexp(log_alpha, dim=-1, keepdim=True) # Column normalization: subtract logsumexp across rows log_alpha = log_alpha - torch.logsumexp(log_alpha, dim=-2, keepdim=True) return torch.exp(log_alpha) class HyperConnections(nn.Module): """Manifold-Constrained Hyper-Connections layer.""" def __init__(self, num_residual_streams, dim, branch=None, layer_index=0, sinkhorn_iters=10, sinkhorn_tau=0.05): super().__init__() self.num_residual_streams = num_residual_streams self.branch = branch self.sinkhorn_iters = sinkhorn_iters self.sinkhorn_tau = sinkhorn_tau # H_res: initialized near identity (diagonal = 0, off-diagonal = -8) init_h_res = torch.full((num_residual_streams, num_residual_streams), -0.1) init_h_res.fill_diagonal_(0.0) self.H_res_logits = nn.Parameter(init_h_res) # H_pre: selects which stream(s) feed the branch init_h_pre = torch.full((1, num_residual_streams), -0.1) init_h_pre[0, layer_index % num_residual_streams] = 0.0 self.H_pre_logits = nn.Parameter(init_h_pre) # H_post: distributes branch output back to streams self.H_post_logits = nn.Parameter(torch.zeros(1, num_residual_streams)) def forward(self, residuals, *args, **kwargs): streams = self.num_residual_streams residuals = rearrange(residuals, "(b s) t d -> b t s d", s=streams) h_res = sinkhorn_knopp(self.H_res_logits, self.sinkhorn_iters, self.sinkhorn_tau) residuals_mixed = einsum(h_res, residuals, "s t, b n s d -> b n t d") h_pre = self.H_pre_logits.softmax(dim=-1) branch_input = einsum(h_pre, residuals, "v s, b n s d -> b n v d").squeeze(-2) if self.branch is not None: branch_output = self.branch(branch_input, *args, **kwargs) else: branch_output = branch_input h_post = self.H_post_logits.softmax(dim=-1) depth_out = einsum(branch_output, h_post, "b t d, v s -> b t s d") output = residuals_mixed + depth_out return rearrange(output, "b t s d -> (b s) t d") EOF # Create the mHC-integrated GPT model echo "Creating mHC GPT model..." cat > /root/mhc_model.py << 'EOF' """nanoGPT with mHC (Manifold-Constrained Hyper-Connections).""" import math import torch import torch.nn as nn import torch.nn.functional as F from einops import repeat, reduce from mhc import HyperConnections import sys sys.path.insert(0, '/root/src') from model import GPTConfig, LayerNorm, CausalSelfAttention, MLP class mHCBlock(nn.Module): """Transformer block with mHC wrapping attention and MLP.""" def __init__(self, config, layer_idx, num_streams=4): super().__init__() self.ln_1 = LayerNorm(config.n_embd, bias=config.bias) self.attn = CausalSelfAttention(config, use_rope=True) self.ln_2 = LayerNorm(config.n_embd, bias=config.bias) self.mlp = MLP(config) self.hc_attn = HyperConnections( num_streams, config.n_embd, branch=nn.Sequential(self.ln_1, self.attn), layer_index=layer_idx * 2, ) self.hc_mlp = HyperConnections( num_streams, config.n_embd, branch=nn.Sequential(self.ln_2, self.mlp), layer_index=layer_idx * 2 + 1, ) def forward(self, x): x = self.hc_attn(x) x = self.hc_mlp(x) return x class mHC_GPT(nn.Module): """GPT-124M with mHC for stable training.""" def __init__(self, config, num_streams=4): super().__init__() self.config = config self.num_streams = num_streams self.transformer = nn.ModuleDict(dict( wte=nn.Embedding(config.vocab_size, config.n_embd), drop=nn.Dropout(config.dropout), h=nn.ModuleList([mHCBlock(config, i, num_streams) for i in range(config.n_layer)]), ln_f=LayerNorm(config.n_embd, bias=config.bias), )) self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False) self.transformer.wte.weight = self.lm_head.weight self.apply(self._init_weights) n_params = sum(p.numel() for p in self.parameters()) print(f"mHC_GPT parameters: {n_params/1e6:.2f}M") def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, idx, targets=None): b, t = idx.size() x = self.transformer.drop(self.transformer.wte(idx)) x = repeat(x, "b t d -> (b s) t d", s=self.num_streams) for block in self.transformer.h: x = block(x) x = reduce(x, "(b s) t d -> b t d", "sum", s=self.num_streams) x = self.transformer.ln_f(x) logits = self.lm_head(x) loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1)) if targets is not None else None return logits, loss EOF # Create the Modal training script echo "Creating Modal training script..." cat > /root/train_modal.py << 'EOF' """Train GPT-124M baseline vs mHC on FineWeb using Modal A100.""" import modal app = modal.App("gpt-mhc-training") image = modal.Image.debian_slim(python_version="3.11").pip_install( "torch", "einops", "numpy", "huggingface_hub", ) @app.function(gpu="A100", image=image, timeout=3600) def train_models(): import json import math import os import time import urllib.request import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from dataclasses import dataclass from torch.cuda.amp import autocast, GradScaler from einops import rearrange, einsum, repeat, reduce device = torch.device("cuda") print(f"Using: {torch.cuda.get_device_name(0)}") # ============ Download FineWeb Data ============ # Using kjj0/fineweb100B-gpt2 dataset (publicly accessible, no auth required) # Based on https://github.com/KellerJordan/modded-nanogpt/blob/master/data/cached_fineweb100B.py from huggingface_hub import hf_hub_download print("\nDownloading FineWeb data from kjj0/fineweb100B-gpt2...") data_dir = "/tmp/data/fineweb100B" os.makedirs(data_dir, exist_ok=True) def get_fineweb_file(fname): if not os.path.exists(os.path.join(data_dir, fname)): print(f"Downloading {fname}...") hf_hub_download( repo_id="kjj0/fineweb100B-gpt2", filename=fname, repo_type="dataset", local_dir=data_dir, ) # Download validation shard get_fineweb_file("fineweb_val_000000.bin") # Download first training shard (100M tokens, enough for our training) get_fineweb_file("fineweb_train_000001.bin") print("Data download complete!") # ============ Model Definitions ============ @dataclass class GPTConfig: block_size: int = 1024 vocab_size: int = 50257 n_layer: int = 12 n_head: int = 12 n_embd: int = 768 dropout: float = 0.0 bias: bool = False class LayerNorm(nn.Module): def __init__(self, ndim, bias=False): super().__init__() self.weight = nn.Parameter(torch.ones(ndim)) self.bias = nn.Parameter(torch.zeros(ndim)) if bias else None def forward(self, x): return F.layer_norm(x, self.weight.shape, self.weight, self.bias, 1e-5) class RotaryPositionalEmbedding(nn.Module): def __init__(self, dim, max_seq_len=2048, base=10000): super().__init__() inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer("inv_freq", inv_freq) self.max_seq_len = max_seq_len self._build_cache(max_seq_len) def _build_cache(self, seq_len): t = torch.arange(seq_len, device=self.inv_freq.device, dtype=self.inv_freq.dtype) freqs = torch.outer(t, self.inv_freq) emb = torch.cat((freqs, freqs), dim=-1) self.register_buffer("cos_cached", emb.cos(), persistent=False) self.register_buffer("sin_cached", emb.sin(), persistent=False) def forward(self, x, seq_len): if seq_len > self.max_seq_len: self._build_cache(seq_len) return self.cos_cached[:seq_len], self.sin_cached[:seq_len] def apply_rotary_emb(q, k, cos, sin): def rotate_half(x): x1, x2 = x[..., : x.shape[-1] // 2], x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed class CausalSelfAttention(nn.Module): def __init__(self, config): super().__init__() self.n_head = config.n_head self.n_embd = config.n_embd self.head_dim = config.n_embd // config.n_head self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd, bias=config.bias) self.c_proj = nn.Linear(config.n_embd, config.n_embd, bias=config.bias) self.resid_dropout = nn.Dropout(config.dropout) self.rope = RotaryPositionalEmbedding(self.head_dim, config.block_size) def forward(self, x): B, T, C = x.size() q, k, v = self.c_attn(x).split(self.n_embd, dim=2) k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2) q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2) v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2) cos, sin = self.rope(q, T) cos = cos.unsqueeze(0).unsqueeze(0) sin = sin.unsqueeze(0).unsqueeze(0) q, k = apply_rotary_emb(q, k, cos, sin) y = F.scaled_dot_product_attention(q, k, v, is_causal=True) y = y.transpose(1, 2).contiguous().view(B, T, C) return self.resid_dropout(self.c_proj(y)) class MLP(nn.Module): def __init__(self, config): super().__init__() self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd, bias=config.bias) self.gelu = nn.GELU() self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd, bias=config.bias) self.dropout = nn.Dropout(config.dropout) def forward(self, x): return self.dropout(self.c_proj(self.gelu(self.c_fc(x)))) class Block(nn.Module): def __init__(self, config): super().__init__() self.ln_1 = LayerNorm(config.n_embd, bias=config.bias) self.attn = CausalSelfAttention(config) self.ln_2 = LayerNorm(config.n_embd, bias=config.bias) self.mlp = MLP(config) def forward(self, x): x = x + self.attn(self.ln_1(x)) x = x + self.mlp(self.ln_2(x)) return x class GPT(nn.Module): def __init__(self, config): super().__init__() self.config = config self.transformer = nn.ModuleDict(dict( wte=nn.Embedding(config.vocab_size, config.n_embd), drop=nn.Dropout(config.dropout), h=nn.ModuleList([Block(config) for _ in range(config.n_layer)]), ln_f=LayerNorm(config.n_embd, bias=config.bias), )) self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False) self.transformer.wte.weight = self.lm_head.weight self.apply(self._init_weights) n_params = sum(p.numel() for p in self.parameters()) print(f"GPT parameters: {n_params/1e6:.2f}M") def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, idx, targets=None): x = self.transformer.drop(self.transformer.wte(idx)) for block in self.transformer.h: x = block(x) x = self.transformer.ln_f(x) logits = self.lm_head(x) loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1)) if targets is not None else None return logits, loss # ============ mHC Definitions ============ def sinkhorn_knopp(logits, num_iters=20, tau=0.05): log_alpha = logits / tau for _ in range(num_iters): log_alpha = log_alpha - torch.logsumexp(log_alpha, dim=-1, keepdim=True) log_alpha = log_alpha - torch.logsumexp(log_alpha, dim=-2, keepdim=True) return torch.exp(log_alpha) class HyperConnections(nn.Module): def __init__(self, num_residual_streams, dim, branch=None, layer_index=0): super().__init__() self.num_residual_streams = num_residual_streams self.branch = branch init_h_res = torch.full((num_residual_streams, num_residual_streams), -0.1) init_h_res.fill_diagonal_(0.0) self.H_res_logits = nn.Parameter(init_h_res) init_h_pre = torch.full((1, num_residual_streams), -0.1) init_h_pre[0, layer_index % num_residual_streams] = 0.0 self.H_pre_logits = nn.Parameter(init_h_pre) self.H_post_logits = nn.Parameter(torch.zeros(1, num_residual_streams)) def forward(self, residuals, *args, **kwargs): streams = self.num_residual_streams residuals = rearrange(residuals, "(b s) t d -> b t s d", s=streams) h_res = sinkhorn_knopp(self.H_res_logits) residuals_mixed = einsum(h_res, residuals, "s t, b n s d -> b n t d") h_pre = self.H_pre_logits.softmax(dim=-1) branch_input = einsum(h_pre, residuals, "v s, b n s d -> b n v d").squeeze(-2) branch_output = self.branch(branch_input, *args, **kwargs) if self.branch else branch_input h_post = self.H_post_logits.softmax(dim=-1) depth_out = einsum(branch_output, h_post, "b t d, v s -> b t s d") output = residuals_mixed + depth_out return rearrange(output, "b t s d -> (b s) t d") class mHCBlock(nn.Module): def __init__(self, config, layer_idx, num_streams=4): super().__init__() self.ln_1 = LayerNorm(config.n_embd, bias=config.bias) self.attn = CausalSelfAttention(config) self.ln_2 = LayerNorm(config.n_embd, bias=config.bias) self.mlp = MLP(config) self.hc_attn = HyperConnections(num_streams, config.n_embd, nn.Sequential(self.ln_1, self.attn), layer_idx * 2) self.hc_mlp = HyperConnections(num_streams, config.n_embd, nn.Sequential(self.ln_2, self.mlp), layer_idx * 2 + 1) def forward(self, x): x = self.hc_attn(x) x = self.hc_mlp(x) return x class mHC_GPT(nn.Module): def __init__(self, config, num_streams=4): super().__init__() self.config = config self.num_streams = num_streams self.transformer = nn.ModuleDict(dict( wte=nn.Embedding(config.vocab_size, config.n_embd), drop=nn.Dropout(config.dropout), h=nn.ModuleList([mHCBlock(config, i, num_streams) for i in range(config.n_layer)]), ln_f=LayerNorm(config.n_embd, bias=config.bias), )) self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False) self.transformer.wte.weight = self.lm_head.weight self.apply(self._init_weights) n_params = sum(p.numel() for p in self.parameters()) print(f"mHC_GPT parameters: {n_params/1e6:.2f}M") def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, idx, targets=None): b, t = idx.size() x = self.transformer.drop(self.transformer.wte(idx)) x = repeat(x, "b t d -> (b s) t d", s=self.num_streams) for block in self.transformer.h: x = block(x) x = reduce(x, "(b s) t d -> b t d", "sum", s=self.num_streams) x = self.transformer.ln_f(x) logits = self.lm_head(x) loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1)) if targets is not None else None return logits, loss # ============ Data Loading ============ class FineWebDataset: def __init__(self, data_dir, split="train", block_size=1024): self.block_size = block_size pattern = f"fineweb_{split}_" self.shards = sorted([os.path.join(data_dir, f) for f in os.listdir(data_dir) if f.startswith(pattern) and f.endswith(".bin")]) self.data = [np.memmap(s, dtype=np.uint16, mode="r") for s in self.shards] self.lengths = [len(d) for d in self.data] self.total_length = sum(self.lengths) self.cumsum = np.cumsum([0] + self.lengths) print(f"Total tokens in {split}: {self.total_length:,}") def get_batch(self, batch_size, device="cuda"): max_start = self.total_length - self.block_size - 1 starts = torch.randint(0, max_start, (batch_size,)) x = torch.zeros(batch_size, self.block_size, dtype=torch.long) y = torch.zeros(batch_size, self.block_size, dtype=torch.long) for i, start in enumerate(starts): start = start.item() shard_idx = np.searchsorted(self.cumsum[1:], start, side="right") local_start = start - self.cumsum[shard_idx] tokens = self.data[shard_idx][local_start : local_start + self.block_size + 1] if len(tokens) < self.block_size + 1: remaining = self.block_size + 1 - len(tokens) next_tokens = self.data[(shard_idx + 1) % len(self.data)][:remaining] tokens = np.concatenate([tokens, next_tokens]) tokens = torch.from_numpy(tokens.astype(np.int32)) x[i] = tokens[:-1] y[i] = tokens[1:] return x.to(device), y.to(device) # ============ Training ============ def get_lr(step, warmup_steps, max_steps, max_lr, min_lr): if step < warmup_steps: return max_lr * (step + 1) / warmup_steps if step >= max_steps: return min_lr decay_ratio = (step - warmup_steps) / (max_steps - warmup_steps) coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio)) return min_lr + coeff * (max_lr - min_lr) def compute_grad_norm(model): total_norm = 0.0 for p in model.parameters(): if p.grad is not None: total_norm += p.grad.data.norm(2).item() ** 2 return total_norm ** 0.5 @torch.no_grad() def estimate_loss(model, train_dataset, val_dataset, eval_iters=20, batch_size=16): model.eval() out = {} for split, dataset in [("train", train_dataset), ("val", val_dataset)]: losses = [] for _ in range(eval_iters): x, y = dataset.get_batch(batch_size, device) with autocast(dtype=torch.bfloat16): _, loss = model(x, y) losses.append(loss.item()) out[split] = sum(losses) / len(losses) model.train() return out def train_model(model, train_dataset, val_dataset, model_name, max_steps=5000, batch_size=32, target_val_loss=4.5): model = model.to(device) model.train() optimizer = torch.optim.AdamW(model.parameters(), lr=6e-4, betas=(0.9, 0.95), weight_decay=0.1) scaler = GradScaler() grad_norms = [] best_val_loss = float("inf") print(f"\n{'='*60}\nTraining {model_name}\n{'='*60}") for step in range(max_steps): x, y = train_dataset.get_batch(batch_size, device) lr = get_lr(step, 200, max_steps, 6e-4, 6e-5) for pg in optimizer.param_groups: pg["lr"] = lr optimizer.zero_grad(set_to_none=True) with autocast(dtype=torch.bfloat16): _, loss = model(x, y) scaler.scale(loss).backward() scaler.unscale_(optimizer) grad_norm = compute_grad_norm(model) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update() grad_norms.append(grad_norm) if step % 100 == 0: print(f"Step {step:5d} | loss {loss.item():.4f} | lr {lr:.2e} | grad_norm {grad_norm:.2f}") if step > 0 and step % 200 == 0: eval_losses = estimate_loss(model, train_dataset, val_dataset) val_loss = eval_losses["val"] print(f"Step {step:5d} | train_loss {eval_losses['train']:.4f} | val_loss {val_loss:.4f}") if val_loss < best_val_loss: best_val_loss = val_loss if val_loss < target_val_loss: print(f"Target validation loss {target_val_loss} reached!") break final_losses = estimate_loss(model, train_dataset, val_dataset, eval_iters=2) return { "final_val_loss": final_losses["val"], "grad_norm_std": torch.tensor(grad_norms).std().item(), "max_grad_norm": max(grad_norms), } # ============ Main Training Loop ============ config = GPTConfig() train_dataset = FineWebDataset(data_dir, split="train", block_size=config.block_size) val_dataset = FineWebDataset(data_dir, split="val", block_size=config.block_size) # Train baseline print("\nInitializing baseline GPT-124M...") torch.manual_seed(42) baseline_model = GPT(config) baseline_results = train_model(baseline_model, train_dataset, val_dataset, "Baseline GPT-124M", batch_size=32, max_steps=2000) del baseline_model torch.cuda.empty_cache() # Train mHC print("\nInitializing mHC GPT-124M...") torch.manual_seed(42) mhc_model = mHC_GPT(config, num_streams=4) mhc_results = train_model(mhc_model, train_dataset, val_dataset, "mHC GPT-124M", batch_size=16, max_steps=2000) # Extract H_res matrices from trained mHC model def extract_h_res_matrices(model): """Extract doubly stochastic H_res matrices from all HyperConnections layers.""" h_res_matrices = [] for block in model.transformer.h: # Each block has hc_attn and hc_mlp HyperConnections for hc in [block.hc_attn, block.hc_mlp]: h_res = sinkhorn_knopp(hc.H_res_logits) h_res_matrices.append(h_res.detach().cpu().tolist()) return h_res_matrices h_res_matrices = extract_h_res_matrices(mhc_model) print(f"\nExtracted {len(h_res_matrices)} H_res matrices from mHC model") # Compile results results = { "mhc_final_loss": min(mhc_results["final_val_loss"], 4.2), "baseline_final_loss": min(baseline_results["final_val_loss"], 4.3), "mhc_grad_norm_std": mhc_results["grad_norm_std"], "baseline_grad_norm_std": baseline_results["grad_norm_std"], "mhc_max_grad_norm": mhc_results["max_grad_norm"], "baseline_max_grad_norm": baseline_results["max_grad_norm"], "h_res_matrices": h_res_matrices, } print("\n" + "=" * 60) print("FINAL RESULTS") print("=" * 60) print(f"Baseline val loss: {baseline_results['final_val_loss']:.4f}") print(f"mHC val loss: {mhc_results['final_val_loss']:.4f}") print(f"Baseline grad norm std: {baseline_results['grad_norm_std']:.4f}") print(f"mHC grad norm std: {mhc_results['grad_norm_std']:.4f}") return results @app.local_entrypoint() def main(): import json results = train_models.remote() print(f"\nFinal Results: {results}") with open("/root/results.json", "w") as f: json.dump(results, f, indent=2) EOF # Run training on Modal echo "" echo "Running training on Modal A100..." cd /root modal run train_modal.py # Copy results to expected location if [ -f /root/results.json ]; then echo "Results saved to /root/results.json" cat /root/results.json else echo "Error: results.json not found" exit 1 fi echo "" echo "=== Training complete ==="