"""Teach the base model to chat: supervised fine-tuning on conversations. python sft.py --base runs/base/ckpt_final.pt --chats data/chats.jsonl --out runs/chat A conversation becomes one line of tokens: <|endoftext|><|user|>hi!<|end|><|assistant|>Hello! How are you?<|end|><|user|>... The loss is counted only on the assistant's words and its closing <|end|>, so the model learns to answer, not to imitate the user. """ import argparse import json import os import random import time import torch from model import GPT, Config from muon import Muon from tokenizer import Tokenizer from train import get_device, lr_factor def render(tok, messages): """Token ids and a 0/1 mask saying which positions the model should learn to predict.""" ids, mask = [tok['<|endoftext|>']], [0] for m in messages: role = tok['<|user|>'] if m['role'] == 'user' else tok['<|assistant|>'] body = tok.encode(m['content'].strip(), allow_special=False) + [tok['<|end|>']] ids += [role] + body learn = 1 if m['role'] == 'assistant' else 0 mask += [0] + [learn] * len(body) return ids, mask def pack(examples, context): """Glue conversations into rows of exactly context+1 tokens (the rest is padding).""" rows, cur_ids, cur_mask = [], [], [] for ids, mask in examples: if len(ids) > context + 1: continue if len(cur_ids) + len(ids) > context + 1: rows.append((cur_ids, cur_mask)) cur_ids, cur_mask = [], [] cur_ids += ids cur_mask += mask if cur_ids: rows.append((cur_ids, cur_mask)) x = torch.zeros(len(rows), context + 1, dtype=torch.long) m = torch.zeros(len(rows), context + 1, dtype=torch.long) for i, (ids, mask) in enumerate(rows): x[i, :len(ids)] = torch.tensor(ids) m[i, :len(mask)] = torch.tensor(mask) return x, m def main(): ap = argparse.ArgumentParser() ap.add_argument('--base', required=True) ap.add_argument('--chats', default='data/chats.jsonl') ap.add_argument('--data', default='data') ap.add_argument('--out', default='runs/chat') ap.add_argument('--epochs', type=float, default=3) ap.add_argument('--batch', type=int, default=32) ap.add_argument('--lr', type=float, default=0.004) ap.add_argument('--lr-embed', type=float, default=0.0008) ap.add_argument('--val-frac', type=float, default=0.03) ap.add_argument('--seed', type=int, default=0) args = ap.parse_args() torch.manual_seed(args.seed) device = get_device() os.makedirs(args.out, exist_ok=True) tok = Tokenizer.load(f'{args.data}/tokenizer.json') ck = torch.load(args.base, map_location='cpu') cfg = Config(**ck['config']) model = GPT(cfg) model.load_state_dict(ck['model']) model.to(device) chats = [json.loads(line) for line in open(args.chats)] # some conversations appear several times on purpose; hold out whole conversations, # every copy of them, so validation text is really never trained on key = lambda c: json.dumps(c['messages'], sort_keys=True) unique = sorted({key(c) for c in chats}) random.Random(args.seed).shuffle(unique) held_out = set(unique[:int(len(unique) * args.val_frac)]) val = list({key(c): c for c in chats if key(c) in held_out}.values()) train = [c for c in chats if key(c) not in held_out] random.Random(args.seed).shuffle(train) val_x, val_m = pack([render(tok, c['messages']) for c in val], cfg.context) train_x, train_m = pack([render(tok, c['messages']) for c in train], cfg.context) print(f'{len(chats)} conversations -> {len(train_x)} train rows, {len(val_x)} val rows') matrices = [p for n, p in model.named_parameters() if p.ndim == 2 and 'embed' not in n] rest = [p for n, p in model.named_parameters() if not (p.ndim == 2 and 'embed' not in n)] muon = Muon(matrices, lr=args.lr, split={b.attn.qkv.weight: 3 for b in model.blocks}) adam = torch.optim.AdamW(rest, lr=args.lr_embed, betas=(0.9, 0.95), weight_decay=0.0) opts = [muon, adam] for opt in opts: for g in opt.param_groups: g['base_lr'] = g['lr'] ctx = torch.autocast(device_type=device, dtype=torch.bfloat16) @torch.no_grad() def evaluate(): model.eval() losses = [] for i in range(0, len(val_x), args.batch): x, m = val_x[i:i + args.batch].to(device), val_m[i:i + args.batch].to(device) with ctx: _, loss = model(x[:, :-1], x[:, 1:], loss_mask=m[:, 1:]) losses.append(loss.item()) model.train() return sum(losses) / len(losses) steps = int(args.epochs * len(train_x) / args.batch) log = open(f'{args.out}/log.jsonl', 'w') log.write(json.dumps({'config': cfg.to_dict(), 'args': vars(args), 'params': model.num_params(), 'steps': steps, 'conversations': len(chats), 'train_rows': len(train_x)}) + '\n') t0 = time.time() order = torch.randperm(len(train_x)) pos = 0 for step in range(steps + 1): if step % 50 == 0 or step == steps: rec = {'step': step, 'val': round(evaluate(), 4), 'time': round(time.time() - t0, 1)} print(json.dumps(rec), flush=True) log.write(json.dumps(rec) + '\n') log.flush() if step == steps: break if pos + args.batch > len(order): order, pos = torch.randperm(len(train_x)), 0 idx = order[pos:pos + args.batch] pos += args.batch x, m = train_x[idx].to(device), train_m[idx].to(device) with ctx: _, loss = model(x[:, :-1], x[:, 1:], loss_mask=m[:, 1:]) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) f = lr_factor(step, steps, warmup=20, cooldown=0.5) for opt in opts: for g in opt.param_groups: g['lr'] = g['base_lr'] * f opt.step() opt.zero_grad(set_to_none=True) if step % 10 == 0: log.write(json.dumps({'step': step, 'loss': round(loss.item(), 4), 'lr': round(f, 4)}) + '\n') torch.save({'config': cfg.to_dict(), 'model': model.state_dict(), 'step': steps}, f'{args.out}/ckpt_final.pt') log.close() if __name__ == '__main__': main()