156 lines
3.2 KiB
Python
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() |