the-jashthakkar/CodeMode
0
1import torch2import torch.nn as nn3from torch.optim import AdamW4from torch.utils.data import DataLoader5from tqdm import tqdm6import os7import logging8from .model import CodeEmbedder9 10# Setup Logger11logging.basicConfig(level=logging.INFO)12logger = logging.getLogger(__name__)13 14class CodeTrainer:15 def __init__(16 self,17 model: CodeEmbedder,18 train_loader: DataLoader,19 val_loader: DataLoader = None,20 epochs: int = 3,21 learning_rate: float = 2e-5,22 accumulation_steps: int = 1,23 mixed_precision: bool = True,24 output_dir: str = "./output",25 device: str = "cuda" if torch.cuda.is_available() else "cpu"26 ):27 self.model = model.to(device)28 self.train_loader = train_loader29 self.val_loader = val_loader30 self.epochs = epochs31 self.lr = learning_rate32 self.accumulation_steps = accumulation_steps33 self.mixed_precision = mixed_precision34 self.output_dir = output_dir35 self.device = device36 37 # Optimizer38 self.optimizer = AdamW(self.model.parameters(), lr=self.lr)39 40 # Scheduler (Optional: constant for now, can transform to Linear later)41 # self.scheduler = ...42 43 # Mixed Precision Scaler44 self.scaler = torch.cuda.amp.GradScaler(enabled=self.mixed_precision)45 46 # Loss Function: Triplet Margin Loss (Standard for Sentence Embeddings)47 # Tries to maximize distance between Anchor-Negative and minimize Anchor-Positive48 self.criterion = nn.TripletMarginLoss(margin=1.0, p=2)49 50 def train_step(self, batch):51 """52 Runs one training step. Returns loss.53 """54 # Unpack the Triplet Batch55 # We assume the Dataset returns keys: 'anchor_input_ids', 'anchor_attention_mask', etc.56 57 # Helper to move dict to device58 to_device = lambda x: x.to(self.device)59 60 # Autocast for Mixed Precision61 with torch.cuda.amp.autocast(enabled=self.mixed_precision):62 # 1. Forward Pass for all 3 components63 anchor_emb = self.model(to_device(batch['anchor_input_ids']), to_device(batch['anchor_attention_mask']))64 positive_emb = self.model(to_device(batch['positive_input_ids']), to_device(batch['positive_attention_mask']))65 negative_emb = self.model(to_device(batch['negative_input_ids']), to_device(batch['negative_attention_mask']))66 67 # 2. Compute Triplet Loss68 loss = self.criterion(anchor_emb, positive_emb, negative_emb)69 70 return loss71 72 def train(self):73 logger.info(f"Starting training on {self.device}...")74 logger.info(f"Batch Size: {self.train_loader.batch_size}, Accumulation Steps: {self.accumulation_steps}")75 logger.info(f"Effective Batch Size: {self.train_loader.batch_size * self.accumulation_steps}")76 77 self.model.train()78 79 for epoch in range(self.epochs):80 total_loss = 081 self.optimizer.zero_grad()82 83 progress_bar = tqdm(self.train_loader, desc=f"Epoch {epoch+1}/{self.epochs}")84 85 for step, batch in enumerate(progress_bar):86 87 # Forward + Loss Calculation88 loss = self.train_step(batch)89 90 # Gradient Accumulation: Normalize loss91 loss = loss / self.accumulation_steps92 93 # Backward Pass (Scaled)94 self.scaler.scale(loss).backward()95 96 if (step + 1) % self.accumulation_steps == 0:97 # Update Weights98 self.scaler.step(self.optimizer)99 self.scaler.update()100 self.optimizer.zero_grad()101 102 total_loss += loss.item() * self.accumulation_steps103 progress_bar.set_postfix({'loss': total_loss / (step + 1)})104 105 # Save Checkpoint106 self.save_model(epoch+1)107 108 def save_model(self, epoch):109 save_path = os.path.join(self.output_dir, f"checkpoint-{epoch}")110 os.makedirs(save_path, exist_ok=True)111 112 logger.info(f"Saving model to {save_path}...")113 114 # Save explicitly as safetensors via transformers API115 self.model.encoder.save_pretrained(save_path, safe_serialization=True)116 self.model.config.save_pretrained(save_path)117 # Note: We save the 'encoder' which is the AutoModel, 118 # so it can be loaded easily by others.119 