timlawrenz/gnn-ruby-code-study
GNN Ruby Code Study Systematic study of Graph Neural Network architectures for Ruby code complexity prediction and generation. Paper: Graph Neural Networks for Ruby Code Complexity Prediction and Generation: A Systematic Architecture Study Dataset 22,452 Ruby methods parsed into AST graphs with 74-dimensional node features. Split Samples File Train 19,084 dataset/train.jsonl Validation 3,368 dataset/val.jsonl Each JSONL record contains:โฆ See the full description on the dataset page: https://huggingface.co/datasets/timlawrenz/gnn-ruby-code-study.
0101
1#!/usr/bin/env python32"""3Training script for Ruby complexity prediction using Graph Neural Networks.4 5This script implements the main training and validation loop for the GNN model6that predicts Ruby method complexity based on AST structure.7"""8 9import sys10import os11import time12import argparse13import torch14import torch.nn.functional as F15from torch_geometric.data import Data16 17# Add src directory to path18sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'src'))19 20from data_processing import create_data_loaders21from models import RubyComplexityGNN22 23 24def train_epoch(model, train_loader, optimizer, criterion, device):25 """26 Train the model for one epoch.27 28 Args:29 model: The GNN model30 train_loader: Training data loader31 optimizer: Optimizer instance32 criterion: Loss function33 device: Device to run on34 35 Returns:36 Average training loss for the epoch37 """38 model.train()39 total_loss = 0.040 num_batches = 041 42 for batch in train_loader:43 # Convert to PyTorch tensors and move to device44 x = torch.tensor(batch['x'], dtype=torch.float).to(device)45 edge_index = torch.tensor(batch['edge_index'], dtype=torch.long).to(device)46 y = torch.tensor(batch['y'], dtype=torch.float).to(device)47 batch_idx = torch.tensor(batch['batch'], dtype=torch.long).to(device)48 49 # Create PyTorch Geometric Data object50 data = Data(x=x, edge_index=edge_index, batch=batch_idx)51 52 # Forward pass53 optimizer.zero_grad()54 predictions = model(data)55 loss = criterion(predictions.squeeze(), y)56 57 # Backward pass58 loss.backward()59 optimizer.step()60 61 total_loss += loss.item()62 num_batches += 163 64 return total_loss / num_batches if num_batches > 0 else 0.065 66 67def validate_epoch(model, val_loader, criterion, device):68 """69 Validate the model for one epoch.70 71 Args:72 model: The GNN model73 val_loader: Validation data loader74 criterion: Loss function75 device: Device to run on76 77 Returns:78 Average validation loss for the epoch79 """80 model.eval()81 total_loss = 0.082 num_batches = 083 84 with torch.no_grad():85 for batch in val_loader:86 # Convert to PyTorch tensors and move to device87 x = torch.tensor(batch['x'], dtype=torch.float).to(device)88 edge_index = torch.tensor(batch['edge_index'], dtype=torch.long).to(device)89 y = torch.tensor(batch['y'], dtype=torch.float).to(device)90 batch_idx = torch.tensor(batch['batch'], dtype=torch.long).to(device)91 92 # Create PyTorch Geometric Data object93 data = Data(x=x, edge_index=edge_index, batch=batch_idx)94 95 # Forward pass96 predictions = model(data)97 loss = criterion(predictions.squeeze(), y)98 99 total_loss += loss.item()100 num_batches += 1101 102 return total_loss / num_batches if num_batches > 0 else 0.0103 104 105def save_model(model, filepath, epoch, train_loss, val_loss):106 """107 Save model weights and training metadata.108 109 Args:110 model: The model to save111 filepath: Path to save the model112 epoch: Current epoch number113 train_loss: Training loss114 val_loss: Validation loss115 """116 torch.save({117 'epoch': epoch,118 'model_state_dict': model.state_dict(),119 'train_loss': train_loss,120 'val_loss': val_loss,121 'model_config': {122 'input_dim': 74,123 'hidden_dim': model.convs[0].out_channels if hasattr(model.convs[0], 'out_channels') else 64,124 'num_layers': model.num_layers,125 'conv_type': model.conv_type,126 'dropout': model.dropout127 }128 }, filepath)129 130 131def parse_args():132 """Parse command line arguments."""133 parser = argparse.ArgumentParser(description='Train Ruby complexity prediction GNN model')134 parser.add_argument('--dataset_path', type=str, default='dataset/',135 help='Path to dataset directory (default: dataset/)')136 parser.add_argument('--epochs', type=int, default=100,137 help='Number of training epochs (default: 100)')138 parser.add_argument('--output_path', type=str, default='models/best_model.pt',139 help='Path to save the best model (default: models/best_model.pt)')140 parser.add_argument('--batch_size', type=int, default=32,141 help='Batch size for training (default: 32)')142 parser.add_argument('--learning_rate', type=float, default=0.001,143 help='Learning rate (default: 0.001)')144 parser.add_argument('--hidden_dim', type=int, default=64,145 help='Hidden dimension size (default: 64)')146 parser.add_argument('--num_layers', type=int, default=3,147 help='Number of GNN layers (default: 3)')148 parser.add_argument('--conv_type', type=str, default='SAGE',149 choices=['GCN', 'SAGE', 'GAT', 'GIN', 'GraphConv'],150 help='GNN convolution type (default: SAGE)')151 parser.add_argument('--dropout', type=float, default=0.1,152 help='Dropout rate (default: 0.1)')153 parser.add_argument('--num_workers', type=int, default=0,154 help='DataLoader workers (default: 0 for Docker compat)')155 return parser.parse_args()156 157 158def main():159 """Main training function."""160 args = parse_args()161 162 print("๐ Ruby Complexity GNN Training")163 print("=" * 50)164 165 # Training configuration from args166 config = {167 'epochs': args.epochs,168 'batch_size': args.batch_size,169 'learning_rate': args.learning_rate,170 'hidden_dim': args.hidden_dim,171 'num_layers': args.num_layers,172 'conv_type': args.conv_type,173 'dropout': args.dropout174 }175 176 print("๐ Training Configuration:")177 for key, value in config.items():178 print(f" {key}: {value}")179 print(f" dataset_path: {args.dataset_path}")180 print(f" output_path: {args.output_path}")181 print()182 183 # Setup device184 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')185 print(f"๐ฅ๏ธ Using device: {device}")186 187 # Create data loaders188 print("๐ Loading datasets...")189 190 # Handle sample dataset naming convention191 if args.dataset_path.rstrip('/').endswith('samples'):192 train_data_path = os.path.join(args.dataset_path, "train_sample.jsonl")193 val_data_path = os.path.join(args.dataset_path, "validation_sample.jsonl")194 else:195 train_data_path = os.path.join(args.dataset_path, "train.jsonl")196 val_data_path = os.path.join(args.dataset_path, "validation.jsonl")197 198 train_loader, val_loader = create_data_loaders(199 train_data_path,200 val_data_path,201 batch_size=config['batch_size'],202 shuffle=True,203 num_workers=args.num_workers204 )205 206 print(f" Training batches: {len(train_loader)}")207 print(f" Validation batches: {len(val_loader)}")208 print()209 210 # Initialize model211 print("๐ง Initializing model...")212 model = RubyComplexityGNN(213 input_dim=74, # AST node feature dimension214 hidden_dim=config['hidden_dim'],215 num_layers=config['num_layers'],216 conv_type=config['conv_type'],217 dropout=config['dropout']218 ).to(device)219 220 param_count = sum(p.numel() for p in model.parameters())221 print(f" Model: {model.get_model_info()}")222 print(f" Parameters: {param_count:,}")223 print()224 225 # Setup optimizer and loss function226 optimizer = torch.optim.Adam(model.parameters(), lr=config['learning_rate'])227 criterion = torch.nn.MSELoss()228 229 print("โ๏ธ Training setup:")230 print(f" Optimizer: Adam (lr={config['learning_rate']})")231 print(f" Loss function: MSELoss")232 print()233 234 # Ensure output directory exists235 os.makedirs(os.path.dirname(args.output_path), exist_ok=True)236 237 # Training loop238 print("๐๏ธ Starting training...")239 print("=" * 50)240 241 best_val_loss = float('inf')242 start_time = time.time()243 244 for epoch in range(config['epochs']):245 epoch_start = time.time()246 247 # Train for one epoch248 train_loss = train_epoch(model, train_loader, optimizer, criterion, device)249 250 # Validate251 val_loss = validate_epoch(model, val_loader, criterion, device)252 253 epoch_time = time.time() - epoch_start254 255 # Print results for each epoch (required by Definition of Done)256 print(f"Epoch {epoch+1:2d}/{config['epochs']} | "257 f"Train Loss: {train_loss:.4f} | "258 f"Val Loss: {val_loss:.4f} | "259 f"Time: {epoch_time:.2f}s")260 261 # Save best model (required by Definition of Done)262 if val_loss < best_val_loss:263 best_val_loss = val_loss264 save_model(model, args.output_path, epoch, train_loss, val_loss)265 print(f" ๐พ New best model saved (val_loss: {val_loss:.4f})")266 267 total_time = time.time() - start_time268 269 print("=" * 50)270 print("๐ Training completed successfully!")271 print(f" Total time: {total_time:.2f}s")272 print(f" Best validation loss: {best_val_loss:.4f}")273 print(f" Best model saved to: {args.output_path}")274 275 # Final model save (optional, keeping for compatibility)276 final_path = args.output_path.replace('.pt', '_final.pt')277 save_model(model, final_path, config['epochs']-1, train_loss, val_loss)278 print(f" Final model saved to: {final_path}")279 280 281if __name__ == "__main__":282 main()