the-jashthakkar/CodeMode
0
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 