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[""] ) 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()