Snaseem2026/code-comment-classifier
132
1"""2Data loader utilities for Code Comment Quality Classifier3"""4import pandas as pd5from datasets import Dataset, DatasetDict6from sklearn.model_selection import train_test_split7from typing import Tuple, Dict, List, Optional8import yaml9import logging10import os11from pathlib import Path12 13 14def load_config(config_path: str = "config.yaml") -> dict:15 """Load configuration from YAML file."""16 with open(config_path, 'r') as f:17 config = yaml.safe_load(f)18 return config19 20 21def load_data(data_path: str) -> pd.DataFrame:22 """23 Load data from CSV file with validation.24 25 Expected format:26 - comment: str (the code comment text)27 - label: str (excellent, helpful, unclear, or outdated)28 29 Args:30 data_path: Path to the CSV file31 32 Returns:33 DataFrame with validated data34 35 Raises:36 FileNotFoundError: If data file doesn't exist37 ValueError: If data format is invalid38 """39 if not os.path.exists(data_path):40 raise FileNotFoundError(f"Data file not found: {data_path}")41 42 df = pd.read_csv(data_path)43 44 # Validate required columns45 required_columns = ['comment', 'label']46 missing_columns = [col for col in required_columns if col not in df.columns]47 if missing_columns:48 raise ValueError(f"Missing required columns: {missing_columns}")49 50 # Remove rows with missing values51 initial_len = len(df)52 df = df.dropna(subset=required_columns)53 if len(df) < initial_len:54 logging.warning(f"Removed {initial_len - len(df)} rows with missing values")55 56 # Remove empty comments57 df = df[df['comment'].str.strip().str.len() > 0]58 59 # Validate labels60 if df['label'].isna().any():61 logging.warning("Found NaN labels, removing those rows")62 df = df.dropna(subset=['label'])63 64 logging.info(f"Loaded {len(df)} samples from {data_path}")65 return df66 67 68def create_label_mapping(labels: list) -> Tuple[Dict[str, int], Dict[int, str]]:69 """Create bidirectional label mapping."""70 label2id = {label: idx for idx, label in enumerate(labels)}71 id2label = {idx: label for idx, label in enumerate(labels)}72 return label2id, id2label73 74 75def prepare_dataset(76 df: pd.DataFrame,77 label2id: Dict[str, int],78 train_size: float = 0.8,79 val_size: float = 0.1,80 test_size: float = 0.1,81 seed: int = 42,82 stratify: bool = True83) -> DatasetDict:84 """85 Prepare dataset splits for training.86 87 Args:88 df: DataFrame with 'comment' and 'label' columns89 label2id: Mapping from label names to IDs90 train_size: Proportion of training data91 val_size: Proportion of validation data92 test_size: Proportion of test data93 seed: Random seed for reproducibility94 stratify: Whether to maintain class distribution in splits95 96 Returns:97 DatasetDict with train, validation, and test splits98 """99 # Validate label distribution100 invalid_labels = set(df['label'].unique()) - set(label2id.keys())101 if invalid_labels:102 raise ValueError(f"Found invalid labels: {invalid_labels}. Expected: {list(label2id.keys())}")103 104 # Convert labels to IDs105 df['label_id'] = df['label'].map(label2id)106 107 # Check for missing mappings108 if df['label_id'].isna().any():109 missing_labels = df[df['label_id'].isna()]['label'].unique()110 raise ValueError(f"Labels not found in label2id mapping: {missing_labels}")111 112 # Validate split proportions113 total_size = train_size + val_size + test_size114 if abs(total_size - 1.0) > 1e-6:115 raise ValueError(f"Split sizes must sum to 1.0, got {total_size}")116 117 # Stratification column118 stratify_col = df['label_id'] if stratify else None119 120 # First split: separate test set121 train_val_df, test_df = train_test_split(122 df,123 test_size=test_size,124 random_state=seed,125 stratify=stratify_col126 )127 128 # Second split: separate train and validation129 val_size_adjusted = val_size / (train_size + val_size)130 stratify_col_train = train_val_df['label_id'] if stratify else None131 train_df, val_df = train_test_split(132 train_val_df,133 test_size=val_size_adjusted,134 random_state=seed,135 stratify=stratify_col_train136 )137 138 # Log distribution139 logging.info(f"Dataset splits - Train: {len(train_df)}, Val: {len(val_df)}, Test: {len(test_df)}")140 logging.info(f"Train label distribution:\n{train_df['label'].value_counts().sort_index()}")141 142 # Create datasets143 dataset_dict = DatasetDict({144 'train': Dataset.from_pandas(train_df[['comment', 'label_id']], preserve_index=False),145 'validation': Dataset.from_pandas(val_df[['comment', 'label_id']], preserve_index=False),146 'test': Dataset.from_pandas(test_df[['comment', 'label_id']], preserve_index=False)147 })148 149 return dataset_dict150 151 152def tokenize_function(examples, tokenizer, max_length: int = 512):153 """Tokenize the input text."""154 return tokenizer(155 examples['comment'],156 padding='max_length',157 truncation=True,158 max_length=max_length159 )160 161 162def prepare_datasets_for_training(config_path: str = "config.yaml"):163 """164 Complete pipeline to prepare datasets for training.165 166 Returns:167 Tuple of (tokenized_datasets, label2id, id2label, tokenizer)168 """169 from transformers import AutoTokenizer170 171 config = load_config(config_path)172 173 # Load data174 df = load_data(config['data']['data_path'])175 176 # Create label mappings177 labels = config['labels']178 label2id, id2label = create_label_mapping(labels)179 180 # Prepare dataset splits181 stratify = config['data'].get('stratify', True)182 dataset_dict = prepare_dataset(183 df,184 label2id,185 train_size=config['data']['train_size'],186 val_size=config['data']['val_size'],187 test_size=config['data']['test_size'],188 seed=config['training']['seed'],189 stratify=stratify190 )191 192 # Load tokenizer193 tokenizer = AutoTokenizer.from_pretrained(config['model']['name'])194 195 # Tokenize datasets196 tokenized_datasets = dataset_dict.map(197 lambda x: tokenize_function(x, tokenizer, config['model']['max_length']),198 batched=True,199 remove_columns=['comment']200 )201 202 # Rename label_id to labels for training203 tokenized_datasets = tokenized_datasets.rename_column('label_id', 'labels')204 205 return tokenized_datasets, label2id, id2label, tokenizer206 