Spaces:
Runtime error
Runtime error
| 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["<pad>"] | |
| 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() | |