FlowRes-1 / train.py
arpecious's picture
Upload 18 files
dbd41fe verified
Raw
History Blame Contribute Delete
1.8 kB
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")