Spaces:
Runtime error
Runtime error
File size: 4,340 Bytes
3781007 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 | 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()
|