Files
SkillCompiler/data/skills-bench/tasks-extra/mhc-layer-impl/oracle/solve.sh
T
2026-09-04 14:58:42 +08:00

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 ==="