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 # Kills 95% of RAM! ) 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")