638 lines
25 KiBLFS
Bash
638 lines
25 KiBLFS
Bash
#!/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 ==="
|