qgate-checkpoints / code /train_wslbridge.py
kolbrian's picture
code: model/train variants + minimal champion
9affe33 verified
Raw
History Blame Contribute Delete
12 kB
"""
Training script with qv_variant choices extended for Phase 7 normed Q-only variants.
"""
import os
import time
import math
import pickle
import argparse
from contextlib import nullcontext
import numpy as np
import torch
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.distributed import init_process_group, destroy_process_group
from model import GPTConfig, GPT
out_dir = 'out-shakespeare-char'
eval_interval = 250
log_interval = 10
eval_iters = 200
eval_only = False
always_save_checkpoint = False
wandb_log = False
wandb_project = 'shakespeare-char'
wandb_run_name = 'mini-gpt'
dataset = 'shakespeare_char'
gradient_accumulation_steps = 1
batch_size = 64
block_size = 256
n_layer = 6
n_head = 6
n_embd = 384
dropout = 0.2
bias = False
learning_rate = 1e-3
max_iters = 5000
lr_decay_iters = 5000
min_lr = 1e-4
beta1 = 0.9
beta2 = 0.99
warmup_iters = 100
device = 'cuda'
dtype = 'bfloat16' if torch.cuda.is_available() and torch.cuda.is_bf16_supported() else 'float16'
compile = False
config_keys = [k for k, v in globals().items() if not k.startswith('_') and isinstance(v, (int, float, bool, str))]
parser = argparse.ArgumentParser(description='Train a mini-GPT')
parser.add_argument('--out_dir', type=str, default=out_dir)
parser.add_argument('--eval_interval', type=int, default=eval_interval)
parser.add_argument('--log_interval', type=int, default=log_interval)
parser.add_argument('--eval_iters', type=int, default=eval_iters)
parser.add_argument('--eval_only', action='store_true')
parser.add_argument('--always_save_checkpoint', action='store_true')
parser.add_argument('--wandb_log', action='store_true')
parser.add_argument('--wandb_project', type=str, default=wandb_project)
parser.add_argument('--wandb_run_name', type=str, default=wandb_run_name)
parser.add_argument('--dataset', type=str, default=dataset)
parser.add_argument('--gradient_accumulation_steps', type=int, default=gradient_accumulation_steps)
parser.add_argument('--batch_size', type=int, default=batch_size)
parser.add_argument('--block_size', type=int, default=block_size)
parser.add_argument('--n_layer', type=int, default=n_layer)
parser.add_argument('--n_head', type=int, default=n_head)
parser.add_argument('--n_embd', type=int, default=n_embd)
parser.add_argument('--dropout', type=float, default=dropout)
parser.add_argument('--bias', action='store_true')
parser.add_argument('--learning_rate', type=float, default=learning_rate)
parser.add_argument('--max_iters', type=int, default=max_iters)
parser.add_argument('--lr_decay_iters', type=int, default=lr_decay_iters)
parser.add_argument('--min_lr', type=float, default=min_lr)
parser.add_argument('--beta1', type=float, default=beta1)
parser.add_argument('--beta2', type=float, default=beta2)
parser.add_argument('--warmup_iters', type=int, default=warmup_iters)
parser.add_argument('--device', type=str, default=device)
parser.add_argument('--dtype', type=str, default=dtype)
parser.add_argument('--compile', action='store_true')
parser.add_argument('--seed', type=int, default=1337, help='random seed')
parser.add_argument('--qv_variant', type=str, default='none',
choices=[
# baseline
'none',
# normalization variants
'vnorm',
'qvnorm',
# static gates
'static_gate',
'static_gate_prehead',
# post-attn normalization
'post_rmsnorm_y',
# Q-conditioned gates
'dynamic',
'dynamic_swiglu',
'dynamic_qconditioned_mlp128',
'dynamic_qconditioned_normed',
'dynamic_qconditioned_mlp128_normed',
'dynamic_qconditioned_mlp192',
'dynamic_qconditioned_fullwidth',
'dynamic_qconditioned_fullwidth_headspecific',
'dynamic_q_headshared_elementwise',
# X-conditioned gates
'dynamic_xconditioned_g1',
'dynamic_xconditioned_fullwidth_headspecific',
'dynamic_xconditioned_bottleneck',
'dynamic_x_g1_headspecific_elementwise',
'dynamic_x_g1_headspecific_headwise',
'dynamic_x_g1_headshared_elementwise',
# Random / ablation gates
'dynamic_random_gate',
'dynamic_ones_gate',
'dynamic_random_normal',
'dynamic_bernoulli_gate',
# Zero-parameter dot product gates
'dynamic_dot_scalar',
'dynamic_dot_elementwise',
# A-only gate
'dynamic_a_conditioned',
# Q+A dual-signal gates
'dynamic_qa_conditioned',
'dynamic_qa_conditioned_headspecific',
'dynamic_qa_conditioned_mlp128',
'dynamic_qa_headshared_elementwise',
'dynamic_qa_bilinear_diag',
'dynamic_qa_conditioned_normed',
'dynamic_qa_conditioned_lowrank16',
# legacy
'dynamic_postconcat_matched',
],
help='QV experiment variant')
args = parser.parse_args()
for k in config_keys:
if hasattr(args, k):
globals()[k] = getattr(args, k)
config = {k: globals()[k] for k in config_keys}
ddp = int(os.environ.get('RANK', -1)) != -1
if ddp:
init_process_group(backend='nccl')
ddp_rank = int(os.environ['RANK'])
ddp_local_rank = int(os.environ['LOCAL_RANK'])
ddp_world_size = int(os.environ['WORLD_SIZE'])
device = f'cuda:{ddp_local_rank}'
torch.cuda.set_device(device)
master_process = ddp_rank == 0
seed_offset = ddp_rank
gradient_accumulation_steps //= ddp_world_size
else:
master_process = True
seed_offset = 0
ddp_world_size = 1
tokens_per_iter = gradient_accumulation_steps * ddp_world_size * batch_size * block_size
print(f"tokens per iteration will be: {tokens_per_iter:,}")
if master_process:
os.makedirs(out_dir, exist_ok=True)
torch.manual_seed(args.seed + seed_offset)
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
device_type = 'cuda' if 'cuda' in device else 'cpu'
ptdtype = {'float32': torch.float32, 'bfloat16': torch.bfloat16, 'float16': torch.float16}[dtype]
ctx = nullcontext() if device_type == 'cpu' else torch.amp.autocast(device_type=device_type, dtype=ptdtype)
data_dir = os.path.join('data', dataset)
# OPEN MEMMAPS ONCE, REUSE FOREVER
train_data = np.memmap(os.path.join(data_dir, 'train.bin'), dtype=np.uint16, mode='r')
val_data = np.memmap(os.path.join(data_dir, 'val.bin'), dtype=np.uint16, mode='r')
def get_batch(split):
data = train_data if split == 'train' else val_data
ix = torch.randint(len(data) - block_size, (batch_size,))
x = torch.stack([torch.from_numpy((data[i:i+block_size]).astype(np.int64)) for i in ix])
y = torch.stack([torch.from_numpy((data[i+1:i+1+block_size]).astype(np.int64)) for i in ix])
if device_type == 'cuda':
x, y = x.pin_memory().to(device, non_blocking=True), y.pin_memory().to(device, non_blocking=True)
else:
x, y = x.to(device), y.to(device)
return x, y
iter_num = 0
best_val_loss = 1e9
meta_path = os.path.join(data_dir, 'meta.pkl')
meta_vocab_size = None
if os.path.exists(meta_path):
with open(meta_path, 'rb') as f:
meta = pickle.load(f)
meta_vocab_size = meta['vocab_size']
print(f"found vocab_size = {meta_vocab_size} (inside {meta_path})")
model_args = dict(n_layer=n_layer, n_head=n_head, n_embd=n_embd, block_size=block_size,
bias=bias, vocab_size=None, dropout=dropout)
if meta_vocab_size is None:
print("defaulting to vocab_size of GPT-2 to 50304 (50257 rounded up for efficiency)")
model_args['vocab_size'] = meta_vocab_size if meta_vocab_size is not None else 50304
gptconf = GPTConfig(**model_args)
gptconf.qv_variant = args.qv_variant
model = GPT(gptconf)
model.to(device)
scaler = torch.cuda.amp.GradScaler(enabled=(dtype == 'float16'))
optimizer = model.configure_optimizers(1e-1, learning_rate, (beta1, beta2), device_type)
if compile:
print("compiling the model... (takes a ~minute)")
model = torch.compile(model)
@torch.no_grad()
def estimate_loss():
out = {}
model.eval()
for split in ['train', 'val']:
losses = torch.zeros(eval_iters)
for k in range(eval_iters):
X, Y = get_batch(split)
with ctx:
logits, loss = model(X, Y)
losses[k] = loss.item()
out[split] = losses.mean()
model.train()
return out
def get_lr(it):
if it < warmup_iters:
return learning_rate * (it + 1) / (warmup_iters + 1)
if it > lr_decay_iters:
return min_lr
decay_ratio = (it - warmup_iters) / (lr_decay_iters - warmup_iters)
coeff = 0.5 * (1.0 + math.cos(math.pi * decay_ratio))
return min_lr + coeff * (learning_rate - min_lr)
if wandb_log and master_process:
import wandb
wandb.init(project=wandb_project, name=wandb_run_name, config=config)
X, Y = get_batch('train')
t0 = time.time()
local_iter_num = 0
raw_model = model.module if hasattr(model, 'module') else model
running_mfu = -1.0
while True:
lr = get_lr(iter_num)
for param_group in optimizer.param_groups:
param_group['lr'] = lr
if iter_num % eval_interval == 0 and master_process:
losses = estimate_loss()
print(f"step {iter_num}: train loss {losses['train']:.4f}, val loss {losses['val']:.4f}")
if losses['val'] < best_val_loss or always_save_checkpoint:
best_val_loss = losses['val']
if iter_num > 0:
checkpoint = {
'model': raw_model.state_dict(),
'optimizer': optimizer.state_dict(),
'model_args': model_args,
'iter_num': iter_num,
'best_val_loss': best_val_loss,
'config': config,
}
print(f"saving checkpoint to {out_dir}")
torch.save(checkpoint, os.path.join(out_dir, 'ckpt.pt'))
if iter_num == 0 and eval_only:
break
for micro_step in range(gradient_accumulation_steps):
with ctx:
logits, loss = model(X, Y)
loss = loss / gradient_accumulation_steps
X, Y = get_batch('train')
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
optimizer.zero_grad(set_to_none=True)
t1 = time.time()
dt = t1 - t0
t0 = t1
if iter_num % log_interval == 0 and master_process:
lossf = loss.item() * gradient_accumulation_steps
if local_iter_num >= 5:
mfu = raw_model.estimate_mfu(batch_size * gradient_accumulation_steps, dt)
running_mfu = mfu if running_mfu == -1.0 else 0.9 * running_mfu + 0.1 * mfu
print(f"iter {iter_num}: loss {lossf:.4f}, time {dt*1000:.2f}ms, mfu {running_mfu*100:.2f}%")
iter_num += 1
local_iter_num += 1
if iter_num > max_iters:
break
if ddp:
destroy_process_group()