SuperLilLM / train.py
StarpowerTechnology's picture
Upload 19 files
3781007 verified
Raw
History Blame Contribute Delete
4.34 kB
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()