feat: Initial Commit
This commit is contained in:
@@ -0,0 +1,156 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user