File size: 1,795 Bytes
dbd41fe | 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 | 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")
|