CoolFace
Apppublic

the-jashthakkar/CodeMode

sourceHugging Facemitupdated 8mo agoView on Hugging Face
0likes
train.py146 linesDownload Raw Back to training
1import argparse2import os3import torch4from torch.utils.data import DataLoader, Dataset5from transformers import AutoTokenizer6 7from scripts.core.training.model import CodeEmbedder8from scripts.core.training.trainer import CodeTrainer9 10import json11 12# Real Dataset class for Triplet Training13class RealCodeDataset(Dataset):14    def __init__(self, jsonl_path, tokenizer, max_length=512):15        self.tokenizer = tokenizer16        self.max_length = max_length17        self.data = []18        19        print(f"Loading data from {jsonl_path}...")20        with open(jsonl_path, 'r', encoding='utf-8') as f:21            for line in f:22                if line.strip():23                    self.data.append(json.loads(line))24        print(f"Loaded {len(self.data)} triplets.")25 26    def __len__(self):27        return len(self.data)28 29    def __getitem__(self, idx):30        item = self.data[idx]31        32        # Helper to tokenize33        def tokenize_text(text):34            return self.tokenizer(35                text,36                return_tensors='pt',37                padding='max_length',38                truncation=True,39                max_length=self.max_length40            )41        42        # Tokenize all three parts43        anchor = tokenize_text(item['anchor'])44        positive = tokenize_text(item['positive'])45        negative = tokenize_text(item['negative'])46        47        # Return a flat dict with prefixed keys48        return {49            'anchor_input_ids': anchor['input_ids'].squeeze(0),50            'anchor_attention_mask': anchor['attention_mask'].squeeze(0),51            'positive_input_ids': positive['input_ids'].squeeze(0),52            'positive_attention_mask': positive['attention_mask'].squeeze(0),53            'negative_input_ids': negative['input_ids'].squeeze(0),54            'negative_attention_mask': negative['attention_mask'].squeeze(0)55        }56 57# Dummy Dataset class for MVP testing without the robust data pipeline availability58class DummyCodeDataset(Dataset):59    def __init__(self, tokenizer, size=100):60        self.tokenizer = tokenizer61        self.size = size62        # Generate dummy triplet structure63        self.data = [{"anchor": "def hello(): return 'world'", "positive": "def hi(): return 'earth'", "negative": "class Foo: pass"}] * size64 65    def __len__(self):66        return self.size67 68    def __getitem__(self, idx):69        item = self.data[idx]70        71        # Helper to tokenize72        def tokenize_text(text):73            return self.tokenizer(74                text,75                return_tensors='pt',76                padding='max_length',77                truncation=True,78                max_length=12879            )80            81        anchor = tokenize_text(item['anchor'])82        positive = tokenize_text(item['positive'])83        negative = tokenize_text(item['negative'])84 85        return {86            'anchor_input_ids': anchor['input_ids'].squeeze(0),87            'anchor_attention_mask': anchor['attention_mask'].squeeze(0),88            'positive_input_ids': positive['input_ids'].squeeze(0),89            'positive_attention_mask': positive['attention_mask'].squeeze(0),90            'negative_input_ids': negative['input_ids'].squeeze(0),91            'negative_attention_mask': negative['attention_mask'].squeeze(0)92        }93 94def main():95    parser = argparse.ArgumentParser(description="Train CodeMode Embeddings")96    97    parser.add_argument("--model_name", type=str, default="microsoft/codebert-base", help="Hub model name")98    parser.add_argument("--data_path", type=str, required=False, help="Path to parsed chunks.jsonl")99    parser.add_argument("--output_dir", type=str, default="./output", help="Where to save checkpoints")100    parser.add_argument("--epochs", type=int, default=3)101    parser.add_argument("--batch_size", type=int, default=8)102    parser.add_argument("--accumulation_steps", type=int, default=4, help="Gradient Accumulation Steps")103    parser.add_argument("--lr", type=float, default=2e-5)104    parser.add_argument("--dry_run", action="store_true", help="Run with dummy data for 1 epoch")105 106    args = parser.parse_args()107    108    print(f"Initializing Training Pipeline...")109    print(f"   Model: {args.model_name}")110    print(f"   Output: {args.output_dir}")111    print(f"   Device: {'cuda' if torch.cuda.is_available() else 'cpu'}")112 113    # 1. Initialize Tokenizer114    tokenizer = AutoTokenizer.from_pretrained(args.model_name)115 116    # 2. Load Dataset (Real or Dummy)117    if args.data_path and os.path.exists(args.data_path):118        train_dataset = RealCodeDataset(args.data_path, tokenizer)119    else:120        print("No data path provided or file missing. Using DUMMY data for verification.")121        train_dataset = DummyCodeDataset(tokenizer, size=100)122 123    train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True)124 125    # 3. Initialize Model126    model = CodeEmbedder(model_name_or_path=args.model_name)127 128    # 4. Initialize Trainer129    trainer = CodeTrainer(130        model=model,131        train_loader=train_loader,132        epochs=args.epochs,133        learning_rate=args.lr,134        accumulation_steps=args.accumulation_steps,135        mixed_precision=True, # Hardcoded True for the "Zero-Cost" philosophy136        output_dir=args.output_dir137    )138 139    # 5. Connect and Train140    trainer.train()141    142    print("Training Complete.")143 144if __name__ == "__main__":145    main()146