Team Ai
Datasetpublic

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.

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes101downloads
train_autoencoder.py394 linesDownload Raw Back to src
1#!/usr/bin/env python32"""3Training script for AST Autoencoder using Graph Neural Networks.4 5This script implements the training loop for the ASTAutoencoder model that6reconstructs Ruby method ASTs from learned embeddings. It uses a frozen encoder7and only trains the decoder weights.8"""9 10import sys11import os12import time13import argparse14import torch15import torch.nn.functional as F16from torch_geometric.data import Batch17 18# Add src directory to path19sys.path.insert(0, os.path.join(os.path.dirname(__file__), 'src'))20 21from torch.optim.lr_scheduler import ReduceLROnPlateau22from data_processing import create_data_loaders23from models import ASTAutoencoder24from loss import (25    ast_reconstruction_loss_improved,26    ast_reconstruction_loss_comprehensive,27    ast_reconstruction_loss_simple,28    ast_reconstruction_loss,29)30 31# Performance optimization: Cache CUDA availability32CUDA_AVAILABLE = torch.cuda.is_available()33 34 35def train_epoch(model, train_loader, optimizer, device, type_weight, parent_weight, scaler, loss_fn=None):36    if loss_fn is None:37        loss_fn = ast_reconstruction_loss_improved38    model.train()39    total_loss = 0.040    num_graphs = 041    42    # Pre-compute autocast context for efficiency43    autocast_ctx = torch.autocast(device_type=device.type, dtype=torch.float16, enabled=CUDA_AVAILABLE)44    45    # Memory optimization: Enable memory efficient attention if available46    if hasattr(torch.backends.cuda, 'enable_math_sdp'):47        torch.backends.cuda.enable_math_sdp(True)48    49    for data in train_loader:50        # Early skip for empty batches51        if data.num_nodes == 0: 52            continue53 54        data = data.to(device, non_blocking=True)55        56        # Clear cache periodically to prevent OOM57        if CUDA_AVAILABLE and num_graphs % 100 == 0:58            torch.cuda.empty_cache()59            60        optimizer.zero_grad()61        62        # Use pre-computed autocast context63        with autocast_ctx:64            result = model(data)65            loss = loss_fn(66                data, 67                result['reconstruction'],68                type_weight=type_weight,69                parent_weight=parent_weight70            )71        72        # Scale the loss and backpropagate73        scaler.scale(loss).backward()74        75        # Gradient clipping (unscale gradients first)76        scaler.unscale_(optimizer)77        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)78        79        # Update weights80        scaler.step(optimizer)81        scaler.update()82        83        total_loss += loss.item() * data.num_graphs84        num_graphs += data.num_graphs85 86    return total_loss / num_graphs if num_graphs > 0 else 0.087 88 89def validate_epoch(model, val_loader, device, type_weight, parent_weight, loss_fn=None):90    if loss_fn is None:91        loss_fn = ast_reconstruction_loss_improved92    model.eval()93    total_loss = 0.094    num_graphs = 095    96    # Pre-compute autocast context for efficiency97    autocast_ctx = torch.autocast(device_type=device.type, dtype=torch.float16, enabled=CUDA_AVAILABLE)98    99    with torch.no_grad():100        for data in val_loader:101            # Early skip for empty batches102            if data.num_nodes == 0: 103                continue104                105            data = data.to(device, non_blocking=True)106 107            with autocast_ctx:108                result = model(data)109                loss = loss_fn(110                    data, 111                    result['reconstruction'],112                    type_weight=type_weight,113                    parent_weight=parent_weight114                )115            total_loss += loss.item() * data.num_graphs116            num_graphs += data.num_graphs117 118    return total_loss / num_graphs if num_graphs > 0 else 0.0119 120 121def save_decoder_weights(model, filepath, epoch, train_loss, val_loss):122    """123    Save decoder weights and training metadata.124    125    Args:126        model: The autoencoder model127        filepath: Path to save the decoder weights128        epoch: Current epoch number129        train_loss: Training loss130        val_loss: Validation loss131    """132    torch.save({133        'epoch': epoch,134        'decoder_state_dict': model.decoder.state_dict(),135        'train_loss': train_loss,136        'val_loss': val_loss,137        'model_config': {138            'embedding_dim': model.decoder.embedding_dim,139            'output_node_dim': model.decoder.output_node_dim,140            'hidden_dim': model.decoder.hidden_dim,141            'num_layers': model.decoder.num_layers,142            'max_nodes': model.decoder.max_nodes143        }144    }, filepath)145 146 147def parse_args():148    """Parse command line arguments."""149    parser = argparse.ArgumentParser(description='Train AST Autoencoder model')150    parser.add_argument('--dataset_path', type=str, default='dataset/',151                        help='Path to dataset directory (default: dataset/)')152    parser.add_argument('--epochs', type=int, default=100,153                        help='Number of training epochs (default: 100)')154    parser.add_argument('--output_path', type=str, default='models/best_decoder.pt',155                        help='Path to save the best decoder model (default: models/best_decoder.pt)')156    parser.add_argument('--encoder_weights_path', type=str, default='models/best_model.pt',157                        help='Path to pre-trained encoder weights (default: models/best_model.pt)')158    parser.add_argument('--batch_size', type=int, default=4096,159                        help='Batch size for pre-collation and training (default: 4096)')160    parser.add_argument('--learning_rate', type=float, default=0.001,161                        help='Learning rate (default: 0.001)')162    parser.add_argument('--hidden_dim', type=int, default=256,163                        help='Hidden dimension size (default: 256)')164    parser.add_argument('--num_layers', type=int, default=5,165                        help='Number of GNN layers (default: 5)')166    parser.add_argument('--conv_type', type=str, default='SAGE', choices=['GCN', 'SAGE'],167                        help='GNN convolution type for the ENCODER (default: SAGE)')168    parser.add_argument('--decoder_conv_type', type=str, default='GAT', choices=['GCN', 'SAGE', 'GAT', 'GIN', 'GraphConv'],169                        help='GNN convolution type for the DECODER (default: GAT)')170    parser.add_argument('--dropout', type=float, default=0.1,171                        help='Dropout rate (default: 0.1)')172    parser.add_argument('--type_weight', type=float, default=2.0,173                        help='Weight for the node type loss component.')174    parser.add_argument('--parent_weight', type=float, default=1.0,175                        help='Weight for the parent prediction loss component.')176    parser.add_argument('--loss_fn', type=str, default='improved',177                        choices=['improved', 'comprehensive', 'simple', 'original'],178                        help='Loss function variant (default: improved)')179    parser.add_argument('--decoder_edge_mode', type=str, default='chain',180                        choices=['chain', 'teacher_forced', 'iterative'],181                        help='Decoder edge construction: chain (legacy sequential), '182                             'teacher_forced (ground-truth AST edges), '183                             'iterative (predictβ†’refine). Default: chain')184    parser.add_argument('--profile', action='store_true',185                        help='Enable profiling for one epoch to identify performance bottlenecks.')186    return parser.parse_args()187 188 189def main():190    """Main training function."""191    args = parse_args()192    193    print("πŸš€ AST Autoencoder Training")194    print("=" * 50)195    196    # Training configuration from args197    config = {198        'epochs': args.epochs,199        'batch_size': args.batch_size,200        'learning_rate': args.learning_rate,201        'hidden_dim': args.hidden_dim,202        'num_layers': args.num_layers,203        'conv_type': args.conv_type,204        'dropout': args.dropout,205        'freeze_encoder': True,  # Key requirement: freeze encoder206        'encoder_weights_path': args.encoder_weights_path,207        'loss_fn': args.loss_fn,208    }209 210    # Select loss function variant211    LOSS_FUNCTIONS = {212        'improved': ast_reconstruction_loss_improved,213        'comprehensive': ast_reconstruction_loss_comprehensive,214        'simple': ast_reconstruction_loss_simple,215        'original': ast_reconstruction_loss,216    }217    loss_fn = LOSS_FUNCTIONS[args.loss_fn]218    219    print("πŸ“‹ Training Configuration:")220    for key, value in config.items():221        print(f"   {key}: {value}")222    print(f"   decoder_conv_type: {args.decoder_conv_type}")223    print(f"   decoder_edge_mode: {args.decoder_edge_mode}")224    print(f"   type_weight: {args.type_weight}")225    print(f"   parent_weight: {args.parent_weight}")226    print(f"   dataset_path: {args.dataset_path}")227    print(f"   output_path: {args.output_path}")228    print()229    230    # Setup device231    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')232    print(f"πŸ–₯️  Using device: {device}")233    234    # Create data loaders235    print("πŸ“‚ Loading datasets...")236    237    # Try pre-collated data first (most efficient), fall back to JSONL238    b_size = args.batch_size239    train_collated = os.path.join(args.dataset_path, f"train_collated_b{b_size}.pt")240    val_collated = os.path.join(args.dataset_path, f"validation_collated_b{b_size}.pt")241    242    if os.path.exists(train_collated) and os.path.exists(val_collated):243        print("   Using pre-collated batches (fastest)")244        train_loader, val_loader = create_data_loaders(245            train_collated, val_collated,246            batch_size=1, shuffle=True, num_workers=0, pre_collated=True,247        )248    else:249        print("   Pre-collated data not found, loading from JSONL (slower but works)")250        train_jsonl = os.path.join(args.dataset_path, "train.jsonl")251        val_jsonl = os.path.join(args.dataset_path, "val.jsonl")252        if not os.path.exists(val_jsonl):253            val_jsonl = os.path.join(args.dataset_path, "validation.jsonl")254        train_loader, val_loader = create_data_loaders(255            train_jsonl, val_jsonl,256            batch_size=b_size, shuffle=True, num_workers=0,257        )258    259    print(f"   Training batches: {len(train_loader)}")260    print(f"   Validation batches: {len(val_loader)}")261    print()262    263    # Initialize autoencoder model with performance optimizations264    print("🧠 Initializing AST Autoencoder...")265    model = ASTAutoencoder(266        encoder_input_dim=74,  # AST node feature dimension267        node_output_dim=74,    # Reconstruct same dimension268        hidden_dim=config['hidden_dim'],269        num_layers=config['num_layers'],270        conv_type=config['conv_type'],271        dropout=config['dropout'],272        freeze_encoder=config['freeze_encoder'],273        encoder_weights_path=config['encoder_weights_path'],274        decoder_conv_type=args.decoder_conv_type,275        gradient_checkpointing=True,  # Enable for memory efficiency276        decoder_edge_mode=args.decoder_edge_mode,277    ).to(device)278    279    # Count parameters280    total_params = sum(p.numel() for p in model.parameters())281    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)282    frozen_params = total_params - trainable_params283    284    print(f"   Model: {model.get_model_info()}")285    print(f"   Total parameters: {total_params:,}")286    print(f"   Trainable parameters: {trainable_params:,} (decoder only)")287    print(f"   Frozen parameters: {frozen_params:,} (encoder)")288    print()289    290    # Setup optimizer and scheduler291    optimizer = torch.optim.Adam(292        filter(lambda p: p.requires_grad, model.parameters()), 293        lr=config['learning_rate']294    )295    scheduler = ReduceLROnPlateau(optimizer, 'min', factor=0.5, patience=5)296    297    # Initialize GradScaler for Automatic Mixed Precision (AMP)298    scaler = torch.amp.GradScaler('cuda', enabled=CUDA_AVAILABLE)299    300    print("βš™οΈ  Training setup:")301    print(f"   Optimizer: Adam (lr={config['learning_rate']})")302    print(f"   Scheduler: ReduceLROnPlateau (patience=5)")303    print(f"   Loss function: Improved Reconstruction Loss")304    print(f"   AMP Enabled: {CUDA_AVAILABLE}")305    print()306    307    # Ensure output directory exists308    os.makedirs(os.path.dirname(args.output_path), exist_ok=True)309    310    # Training loop with Early Stopping311    print("πŸ‹οΈ  Starting training...")312    print("=" * 50)313    314    if args.profile:315        import cProfile, pstats316        profiler = cProfile.Profile()317        print("πŸ”¬ PROFILING ENABLED: Running for one epoch...")318        profiler.enable()319 320    best_val_loss = float('inf')321    epochs_no_improve = 0322    323    # Performance optimization: Enable optimized attention if available324    if CUDA_AVAILABLE and hasattr(torch.backends.cuda, 'enable_flash_sdp'):325        torch.backends.cuda.enable_flash_sdp(True)326    early_stopping_patience = 10327    start_time = time.time()328    329    for epoch in range(config['epochs']):330        epoch_start = time.time()331        332        train_loss = train_epoch(model, train_loader, optimizer, device, args.type_weight, args.parent_weight, scaler, loss_fn=loss_fn)333        334        # If profiling, stop after one training epoch and print results335        if args.profile:336            profiler.disable()337            print("πŸ“Š Profiling Results (top 20 functions by cumulative time):")338            stats = pstats.Stats(profiler).sort_stats('cumtime')339            stats.print_stats(20)340            break # Exit after profiling341            342        val_loss = validate_epoch(model, val_loader, device, args.type_weight, args.parent_weight, loss_fn=loss_fn)343        344        epoch_time = time.time() - epoch_start345        346        print(f"Epoch {epoch+1:2d}/{config['epochs']} | "347              f"Train Loss: {train_loss:.4f} | "348              f"Val Loss: {val_loss:.4f} | "349              f"LR: {optimizer.param_groups[0]['lr']:.1e} | "350              f"Time: {epoch_time:.2f}s")351        352        scheduler.step(val_loss)353        354        if val_loss < best_val_loss:355            best_val_loss = val_loss356            epochs_no_improve = 0357            save_decoder_weights(model, args.output_path, epoch, train_loss, val_loss)358            print(f"   πŸ’Ύ New best decoder saved (val_loss: {val_loss:.4f})")359        else:360            epochs_no_improve += 1361 362        if epochs_no_improve >= early_stopping_patience:363            print(f"   πŸ›‘ Early stopping triggered after {early_stopping_patience} epochs with no improvement.")364            break365    366    # This part will not be reached if profiling is enabled and successful367    if not args.profile:368        total_time = time.time() - start_time369        370        print("=" * 50)371        print("πŸŽ‰ Training completed successfully!")372        print(f"   Total time: {total_time:.2f}s")373        print(f"   Best validation loss: {best_val_loss:.4f}")374        print(f"   Best decoder weights saved to: {args.output_path}")375        376        # Final decoder save (optional, keeping for compatibility)377        final_path = args.output_path.replace('.pt', '_final.pt')378        save_decoder_weights(model, final_path, config['epochs']-1, train_loss, val_loss)379        print(f"   Final decoder weights saved to: {final_path}")380        381        # Verify training objectives382        print("\nβœ… Training Objectives Met:")383        print(f"   βœ“ Trained for {config['epochs']} epochs (β‰₯2 required)")384        print(f"   βœ“ Only decoder weights trained (encoder frozen)")385        print(f"   βœ“ Used AST reconstruction loss function")386        print(f"   βœ“ Input and target are same AST graph")387        print(f"   βœ“ Best decoder weights saved to {args.output_path}")388        if config['epochs'] > 1:389            print(f"   βœ“ Training completed successfully over multiple epochs")390 391 392if __name__ == "__main__":393    main()394