Files
2026-09-21 22:20:24 +02:00

156 lines
3.2 KiB
Python

import torch
import torch.nn as nn
import torch.optim as optim
from config import (
DATASET_PATH,
BATCH_SIZE,
LEARNING_RATE,
EPOCHS,
DEVICE,
CHECKPOINT_DIR,
)
from data.tokenizer import Tokenizer
from data.dataset import create_dataloader
from model.builder_model import BuilderModel
from utils.checkpoint import (
save_checkpoint,
latest_checkpoint,
load_checkpoint,
)
def format_params(n):
if n >= 1_000_000_000:
return f"{n:,} ({n/1_000_000_000:.2f}B)"
elif n >= 1_000_000:
return f"{n:,} ({n/1_000_000:.2f}M)"
elif n >= 1_000:
return f"{n:,} ({n/1_000:.2f}K)"
return str(n)
def train():
print("=" * 60)
print("Building tokenizer...")
print("=" * 60)
tokenizer = Tokenizer()
tokenizer.build(DATASET_PATH)
tokenizer.save("data")
print("Prompt vocabulary :", tokenizer.text_vocab_size)
print("Output vocabulary :", tokenizer.output_vocab_size)
print("=" * 60)
print("Loading dataset...")
print("=" * 60)
dataloader = create_dataloader(
DATASET_PATH,
tokenizer,
BATCH_SIZE,
)
print("Creating model...")
model = BuilderModel(tokenizer).to(DEVICE)
optimizer = optim.AdamW(
model.parameters(),
lr=LEARNING_RATE,
)
criterion = nn.CrossEntropyLoss(
ignore_index=tokenizer.output_to_id["<PAD>"]
)
start_epoch = 1
checkpoint = latest_checkpoint(CHECKPOINT_DIR)
if checkpoint:
print("Loading checkpoint:", checkpoint)
start_epoch, _ = load_checkpoint(
checkpoint,
model,
optimizer,
DEVICE,
)
start_epoch += 1
print("=" * 60)
print("Training")
print("=" * 60)
for epoch in range(start_epoch, EPOCHS + 1):
model.train()
total_loss = 0
for prompts, outputs in dataloader:
prompts = prompts.to(DEVICE)
outputs = outputs.to(DEVICE)
decoder_input = outputs[:, :-1]
targets = outputs[:, 1:]
optimizer.zero_grad()
logits = model(
prompts,
decoder_input,
)
loss = criterion(
logits.reshape(-1, logits.size(-1)),
targets.reshape(-1),
)
loss.backward()
torch.nn.utils.clip_grad_norm_(
model.parameters(),
1.0,
)
optimizer.step()
total_loss += loss.item()
avg_loss = total_loss / len(dataloader)
print(
f"Epoch {epoch}/{EPOCHS} | Loss: {avg_loss:.4f}"
)
print("\nTraining finished!")
total_params = sum(p.numel() for p in model.parameters())
trainable_params = sum(
p.numel() for p in model.parameters() if p.requires_grad
)
print("=" * 60)
print(f"Total parameters: {format_params(total_params)}")
print(f"Trainable parameters: {format_params(trainable_params)}")
print("=" * 60)
save_checkpoint(
model,
optimizer,
EPOCHS,
avg_loss,
CHECKPOINT_DIR / "final_model.pth",
)
if __name__ == "__main__":
train()