import argparse import json import random from pathlib import Path import torch from torch.utils.data import DataLoader, Dataset from superlillm.model import ModelConfig, SuperLilLM from superlillm.tokenizer import WordTokenizer DATA_PATH = Path("data/superlillm_dataset.json") CHECKPOINT_DIR = Path("checkpoints") def chat_text(example): return f"User: {example['input']}\nAssistant: {example['output']}" class ChatDataset(Dataset): def __init__(self, examples, tokenizer, block_size, sft=False): self.rows = [] for ex in examples: prompt = f"User: {ex['input']}\nAssistant:" full = chat_text(ex) ids = tokenizer.encode(full, add_bos=True, add_eos=True) if len(ids) > block_size: ids = ids[:block_size] labels = ids[1:] + [-100] labels = labels[: len(ids)] if sft: prompt_len = len(tokenizer.encode(prompt, add_bos=True)) for i in range(max(0, prompt_len - 1)): if i < len(labels): labels[i] = -100 self.rows.append((ids, labels)) self.block_size = block_size self.pad_id = tokenizer.token_to_id[""] def __len__(self): return len(self.rows) def __getitem__(self, idx): ids, labels = self.rows[idx] x = ids + [self.pad_id] * (self.block_size - len(ids)) y = labels + [-100] * (self.block_size - len(labels)) return torch.tensor(x, dtype=torch.long), torch.tensor(y, dtype=torch.long) def train_phase(model, loader, optimizer, device, epochs, phase_name): model.train() for epoch in range(1, epochs + 1): losses = [] for x, y in loader: x, y = x.to(device), y.to(device) _, loss = model(x, y) optimizer.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() losses.append(loss.item()) avg = sum(losses) / len(losses) print(f"{phase_name} epoch {epoch:02d}/{epochs} loss {avg:.4f}") def main(): parser = argparse.ArgumentParser() parser.add_argument("--epochs-pretrain", type=int, default=18) parser.add_argument("--epochs-sft", type=int, default=35) parser.add_argument("--batch-size", type=int, default=32) parser.add_argument("--block-size", type=int, default=160) parser.add_argument("--seed", type=int, default=7) args = parser.parse_args() random.seed(args.seed) torch.manual_seed(args.seed) with DATA_PATH.open("r", encoding="utf-8") as f: examples = json.load(f) random.shuffle(examples) tokenizer = WordTokenizer() tokenizer.build([chat_text(ex) for ex in examples]) config = ModelConfig( vocab_size=len(tokenizer.token_to_id), block_size=args.block_size, n_embd=128, n_head=4, n_layer=4, dropout=0.1, ) device = "mps" if torch.backends.mps.is_available() else "cuda" if torch.cuda.is_available() else "cpu" print(f"Training on {device} with {len(examples)} examples and vocab size {config.vocab_size}") model = SuperLilLM(config).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=0.01) pretrain_data = ChatDataset(examples, tokenizer, args.block_size, sft=False) sft_data = ChatDataset(examples, tokenizer, args.block_size, sft=True) pretrain_loader = DataLoader(pretrain_data, batch_size=args.batch_size, shuffle=True) sft_loader = DataLoader(sft_data, batch_size=args.batch_size, shuffle=True) train_phase(model, pretrain_loader, optimizer, device, args.epochs_pretrain, "pretrain") for group in optimizer.param_groups: group["lr"] = 1e-4 train_phase(model, sft_loader, optimizer, device, args.epochs_sft, "sft") CHECKPOINT_DIR.mkdir(parents=True, exist_ok=True) tokenizer.save(CHECKPOINT_DIR / "tokenizer.json") torch.save( { "model_state": model.state_dict(), "config": config.__dict__, "examples": len(examples), }, CHECKPOINT_DIR / "superlillm.pt", ) print(f"Saved checkpoint to {CHECKPOINT_DIR / 'superlillm.pt'}") if __name__ == "__main__": main()