CoolFace
Apppublic

the-jashthakkar/CodeMode

sourceHugging Facemitupdated 8mo agoView on Hugging Face
0likes
trainer.py119 linesDownload Raw Back to training
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