| import torch |
| import torch.nn.functional as F |
| import time |
| import math |
| from ket_optimizer import KetOptimizer |
|
|
| from torch.utils.data import DataLoader |
|
|
| from config import * |
| from model import MiniTransformer |
| from dataset import TextDataset |
|
|
| tokens = torch.load("tokens.pt") |
|
|
| dataset = TextDataset(tokens, BLOCK_SIZE) |
|
|
| loader = DataLoader( |
| dataset, |
| batch_size=BATCH_SIZE, |
| shuffle=True |
| ) |
|
|
| device = "cuda" if torch.cuda.is_available() else "cpu" |
|
|
| model = MiniTransformer().to(device) |
|
|
| params = sum(p.numel() for p in model.parameters()) |
| print("Parameters:", params) |
|
|
| optimizer = KetOptimizer( |
| model.parameters(), |
| lr=LR, |
| rank_ratio=20 |
| ) |
|
|
| total_steps = EPOCHS * len(loader) |
|
|
| for epoch in range(EPOCHS): |
|
|
| for step, (x, y) in enumerate(loader): |
| start_time = time.time() |
|
|
| x = x.to(device) |
| y = y.to(device) |
|
|
| logits = model(x) |
|
|
| loss = F.cross_entropy( |
| logits.reshape(-1, VOCAB_SIZE), |
| y.reshape(-1) |
| ) |
|
|
| optimizer.zero_grad() |
| loss.backward() |
| optimizer.step() |
| |
| duration = time.time() - start_time |
| tps = (x.shape[0] * x.shape[1]) / duration |
| |
| global_step = epoch * len(loader) + step |
| steps_remaining = total_steps - global_step - 1 |
| eta_seconds = int(steps_remaining * duration) |
| |
| eta_mins, eta_secs = divmod(eta_seconds, 60) |
| eta_hours, eta_mins = divmod(eta_mins, 60) |
| eta_str = f"{eta_hours:02d}:{eta_mins:02d}:{eta_secs:02d}" |
|
|
| if step % 100 == 0: |
| print( |
| f"epoch={epoch+1}/{EPOCHS} step={step} loss={loss.item():.4f} tps={tps:.2f} Tokens/sec ETA={eta_str}" |
| ) |
|
|
| torch.save( |
| model.state_dict(), |
| "mini.pt" |
| ) |
|
|
| print("Saved model") |
|
|