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
loss.py596 linesDownload Raw Back to src
1"""2Loss functions for AST reconstruction tasks and contrastive learning.3 4This module provides loss functions for:51. Measuring the difference between original and reconstructed Abstract Syntax Trees62. Contrastive learning between code and text embeddings for alignment7"""8 9import torch10import torch.nn.functional as F11from torch_geometric.data import Data12from typing import Dict, Any, Union13 14 15def ast_reconstruction_loss_comprehensive(original: Data, reconstructed: Dict[str, Any], 16                                        node_weight: float = 1.0, parent_weight: float = 1.0) -> torch.Tensor:17    """18    Computes a comprehensive reconstruction loss for an AST.19 20    This loss combines:21    1. Node Type Loss: Cross-entropy for predicting the correct node types.22    2. Parent Prediction Loss: Cross-entropy for predicting the correct parent for each node.23    """24    # --- Node Type Loss ---25    recon_node_logits = reconstructed['node_features'].squeeze(0)26    27    # Numerical stability: Clamp values to a reasonable range to prevent overflow28    recon_node_logits = torch.clamp(recon_node_logits, min=-100, max=100)29    30    true_node_types = original.x.argmax(dim=1)31    32    num_nodes = min(recon_node_logits.size(0), true_node_types.size(0))33    if num_nodes == 0:34        return torch.tensor(0.0, device=original.x.device, requires_grad=True)35        36    node_loss = F.cross_entropy(37        recon_node_logits[:num_nodes], 38        true_node_types[:num_nodes]39    )40 41    # --- Parent Prediction Loss ---42    recon_parent_logits = reconstructed['parent_logits'].squeeze(0) # [num_nodes, max_nodes]43    44    # Numerical stability: Clamp values45    recon_parent_logits = torch.clamp(recon_parent_logits, min=-100, max=100)46    47    max_nodes = recon_parent_logits.size(1)48    49    # Create the true parent labels50    num_true_nodes = original.num_nodes51    # Initialize with an ignore_index52    ignore_index = -100 53    true_parents = torch.full((num_true_nodes,), ignore_index, dtype=torch.long, device=original.x.device)54    55    # Edge index is [parent, child], so edge_index[0] are parents and edge_index[1] are children56    children = original.edge_index[1]57    parents = original.edge_index[0]58 59    # Clamp parent indices to be within the prediction range [0, max_nodes-1]60    valid_parents = torch.clamp(parents, 0, max_nodes - 1)61    true_parents[children] = valid_parents62 63    # We only care about the first num_nodes predictions and labels64    num_nodes = min(recon_parent_logits.size(0), true_parents.size(0))65 66    # Check if there are any valid parent labels to compute loss on67    if (true_parents[:num_nodes] != ignore_index).any():68        parent_loss = F.cross_entropy(69            recon_parent_logits[:num_nodes],70            true_parents[:num_nodes],71            ignore_index=ignore_index72        )73    else:74        # No valid parents to compute loss on (e.g., single-node graph)75        parent_loss = torch.tensor(0.0, device=original.x.device, requires_grad=True)76 77    # --- Total Loss ---78    total_loss = (node_weight * node_loss) + (parent_weight * parent_loss)79    return total_loss80 81 82def ast_reconstruction_loss(original: Data, reconstructed: Dict[str, Any], 83                          node_weight: float = 1.0, edge_weight: float = 0.5) -> torch.Tensor:84    """85    Compute the reconstruction loss between original and reconstructed AST.86    87    This loss function combines:88    1. Node Type Loss: Cross-entropy loss for predicting correct node types89    2. Edge Prediction Loss: Loss for predicting correct graph connectivity90    91    Args:92        original: Original AST as torch_geometric.data.Data object93        reconstructed: Reconstructed AST from decoder containing:94            - 'node_features': Tensor of shape [batch_size, num_nodes, feature_dim]95            - 'edge_index': Edge connectivity (optional, for edge loss)96            - 'batch': Batch indices97            - 'num_nodes_per_graph': List of node counts per graph98        node_weight: Weight for node type loss component99        edge_weight: Weight for edge prediction loss component100        101    Returns:102        Scalar tensor representing the total reconstruction loss103    """104    # Extract original data105    original_x = original.x  # [total_nodes, feature_dim]106    original_edge_index = original.edge_index  # [2, total_edges]107    original_batch = original.batch  # [total_nodes]108    109    # Extract reconstructed data110    recon_node_features = reconstructed['node_features']111    if recon_node_features.dim() == 2:112        recon_node_features = recon_node_features.unsqueeze(0)113        114    batch_size = recon_node_features.size(0)115    max_nodes = recon_node_features.size(1)116    feature_dim = recon_node_features.size(2)117    118    # Compute node type loss119    node_loss = compute_node_type_loss(original_x, recon_node_features, original_batch)120    121    # Compute edge prediction loss (simplified version)122    edge_loss = compute_edge_prediction_loss(original_edge_index, original_batch, 123                                           reconstructed, batch_size)124    125    # Combine losses126    total_loss = node_weight * node_loss + edge_weight * edge_loss127    128    return total_loss129 130 131def compute_node_type_loss(original_x: torch.Tensor, 132                          recon_node_features: torch.Tensor,133                          original_batch: torch.Tensor) -> torch.Tensor:134    """135    Compute cross-entropy loss for node type prediction.136    137    Args:138        original_x: Original node features [total_nodes, feature_dim] (one-hot encoded)139        recon_node_features: Reconstructed features [batch_size, max_nodes, feature_dim] (logits)140        original_batch: Batch indices for original nodes [total_nodes]141        142    Returns:143        Average cross-entropy loss across all nodes144    """145    if recon_node_features.dim() == 2:146        recon_node_features = recon_node_features.unsqueeze(0)147 148    batch_size = recon_node_features.size(0)149    max_nodes = recon_node_features.size(1)150    feature_dim = recon_node_features.size(2)151    152    total_loss = 0.0153    total_nodes = 0154    155    # Process each graph in the batch156    for batch_idx in range(batch_size):157        # Get original nodes for this graph158        mask = (original_batch == batch_idx)159        if not mask.any():160            continue161            162        original_nodes = original_x[mask]  # [num_nodes_in_graph, feature_dim]163        num_original_nodes = original_nodes.size(0)164        165        # Get reconstructed nodes for this graph (up to actual node count)166        # Handle case where reconstruction has fewer nodes than original167        num_recon_nodes = min(num_original_nodes, max_nodes)168        recon_nodes = recon_node_features[batch_idx, :num_recon_nodes, :]  # [num_recon_nodes, feature_dim]169        170        # Numerical stability: Check for and handle NaN/Inf values in reconstructed logits171        if torch.isnan(recon_nodes).any() or torch.isinf(recon_nodes).any():172            # Replace NaN/Inf with safe values to prevent loss explosion173            recon_nodes = torch.where(torch.isnan(recon_nodes), torch.zeros_like(recon_nodes), recon_nodes)174            recon_nodes = torch.clamp(recon_nodes, min=-100, max=100)  # Clamp to reasonable range175        176        # Only use original nodes up to the number of reconstructed nodes177        original_nodes_subset = original_nodes[:num_recon_nodes, :]  # [num_recon_nodes, feature_dim]178        179        # Convert one-hot original to class indices for cross-entropy180        # Assumes original_x is one-hot encoded181        original_classes = torch.argmax(original_nodes_subset, dim=1)  # [num_recon_nodes]182        183        # Compute cross-entropy loss184        # recon_nodes are logits, original_classes are target class indices185        loss = F.cross_entropy(recon_nodes, original_classes, reduction='sum')186        187        total_loss += loss188        total_nodes += num_recon_nodes189    190    # Return average loss per node191    if total_nodes > 0:192        return total_loss / total_nodes193    else:194        return torch.tensor(0.0, device=original_x.device, requires_grad=True)195 196 197def compute_edge_prediction_loss(original_edge_index: torch.Tensor,198                                original_batch: torch.Tensor,199                                reconstructed: Dict[str, Any],200                                batch_size: int) -> torch.Tensor:201    """202    Compute edge prediction loss based on graph connectivity.203    204    This is a simplified version that compares the number of edges per graph205    rather than exact edge-to-edge matching, which would be more complex.206    207    Args:208        original_edge_index: Original edges [2, total_edges]209        original_batch: Batch indices for original nodes [total_nodes]210        reconstructed: Dictionary containing reconstruction info211        batch_size: Number of graphs in batch212        213    Returns:214        Loss based on edge count differences215    """216    if original_edge_index.size(1) == 0:217        # No edges in original, return zero loss218        return torch.tensor(0.0, device=original_edge_index.device, requires_grad=True)219    220    # --- Vectorized implementation to avoid CPU bottlenecks ---221    222    # 1. Get the batch index for the source node of each edge223    edge_batch_indices = original_batch[original_edge_index[0]]224 225    # 2. Count the number of edges for each graph in the batch226    # `bincount` is a highly optimized way to count occurrences of each index227    original_edge_counts = torch.bincount(edge_batch_indices, minlength=batch_size).float()228 229    # 3. Estimate reconstructed edge counts (maintaining original logic)230    # Get the number of nodes in each graph of the batch231    num_nodes_per_graph = torch.bincount(original_batch, minlength=batch_size).float()232    # Estimate edge count as num_nodes - 1 (for a tree-like structure)233    recon_edge_counts = torch.clamp(num_nodes_per_graph - 1, min=0)234 235    # 4. Compute the loss as the mean squared error between the counts236    # This is a single, fast, vectorized operation237    loss = F.mse_loss(recon_edge_counts, original_edge_counts)238    239    return loss240 241 242def ast_reconstruction_loss_improved(original: Data, reconstructed: Dict[str, Any],243                                   type_weight: float = 1.0, 244                                   parent_weight: float = 1.0) -> torch.Tensor:245    """246    Improved AST reconstruction loss with explicit parent prediction for batches.247    248    This loss function provides a strong structural learning signal by combining249    node type prediction with explicit parent prediction for each node across an250    entire batch of graphs.251    252    Args:253        original: A `torch_geometric.data.Batch` object containing a batch of original ASTs.254        reconstructed: Reconstructed AST from the decoder, containing batched 'node_features' 255                       and 'parent_logits'.256        type_weight: Weight for the node type prediction loss.257        parent_weight: Weight for the parent prediction loss.258        259    Returns:260        Scalar tensor representing the total weighted reconstruction loss for the batch.261    """262    # --- Component 1: Node Type Loss (Batched) ---263    recon_node_logits = reconstructed['node_features'] # Shape: [total_nodes, feature_dim]264    true_node_types = original.x.argmax(dim=1)265    266    # The number of nodes should match between the batched original and reconstruction.267    num_nodes = min(recon_node_logits.size(0), true_node_types.size(0))268    if num_nodes == 0:269        return torch.tensor(0.0, device=original.x.device, requires_grad=True)270        271    type_loss = F.cross_entropy(272        recon_node_logits[:num_nodes], 273        true_node_types[:num_nodes]274    )275 276    # --- Component 2: Parent Prediction Loss (Batched) ---277    recon_parent_logits = reconstructed['parent_logits'] # Shape: [total_nodes, max_nodes]278    max_nodes = recon_parent_logits.size(1)279    280    # Create the ground truth parent labels for the entire batch.281    num_true_nodes = original.num_nodes282    ignore_index = -100283    true_parents = torch.full((num_true_nodes,), ignore_index, dtype=torch.long, device=original.x.device)284    285    # To correctly handle parent indices in a batch, we need to offset them.286    # The parent of a node in graph `i` must be one of the nodes *within* graph `i`.287    # We first create a global offset for each node.288    num_nodes_per_graph = torch.bincount(original.batch)289    node_offsets = torch.cumsum(num_nodes_per_graph, dim=0) - num_nodes_per_graph290    291    # Offset the parent indices in the edge list.292    children = original.edge_index[1]293    parents = original.edge_index[0]294    295    # The parent prediction is local to each graph. The `parent_predictor` outputs logits296    # where the `j`-th logit corresponds to the `j`-th node *within that graph*.297    # Therefore, we need to calculate the local parent index.298    local_parents = parents - node_offsets[original.batch[parents]]299    300    # Populate the true_parents tensor with the local parent indices.301    # Clamp to ensure indices are within the prediction range [0, max_nodes-1].302    valid_parents = torch.clamp(local_parents, 0, max_nodes - 1)303    true_parents[children] = valid_parents304 305    # Check if there are any valid parent-child relationships to compute loss on.306    if (true_parents != ignore_index).any():307        parent_loss = F.cross_entropy(308            recon_parent_logits,309            true_parents,310            ignore_index=ignore_index311        )312    else:313        parent_loss = torch.tensor(0.0, device=original.x.device)314 315    # --- Total Loss ---316    total_loss = (type_weight * type_loss) + (parent_weight * parent_loss)317    return total_loss318 319 320def _compute_role_loss(original: Data, reconstructed: Dict[str, Any]) -> torch.Tensor:321    """322    Compute role loss component for improved AST reconstruction.323    324    This function computes a loss that encourages the model to understand the 325    functional role of identifiers (e.g., method argument, local variable).326    327    For backward compatibility with current one-hot node features, this implements328    a simplified role-aware loss based on node types and graph structure.329    In the future, this will use dedicated role embeddings.330    331    Args:332        original: Original AST data333        reconstructed: Reconstructed AST data334        335    Returns:336        Scalar tensor representing the role loss337    """338    recon_node_features = reconstructed['node_features']339    batch_size = recon_node_features.size(0)340    341    # For backward compatibility, derive role information from node types and graph structure342    # This is a simplified approach until dedicated role features are implemented343    344    total_loss = 0.0345    total_nodes = 0346    347    for batch_idx in range(batch_size):348        # Get original nodes for this graph349        mask = (original.batch == batch_idx)350        if not mask.any():351            continue352            353        original_nodes = original.x[mask]  # [num_nodes_in_graph, feature_dim]354        num_original_nodes = original_nodes.size(0)355        356        # Get node types for role inference357        original_node_types = torch.argmax(original_nodes, dim=1)358        359        # Simple role-based loss: encourage consistency in how similar node types are handled360        # This approximates role understanding until full role features are available361        if num_original_nodes > 1:362            # Create a simple role similarity matrix based on node types363            type_similarity = (original_node_types.unsqueeze(0) == original_node_types.unsqueeze(1)).float()364            365            # Get reconstructed features for this batch366            max_nodes = min(num_original_nodes, recon_node_features.size(1))367            recon_features = recon_node_features[batch_idx, :max_nodes, :]368            369            # Compute pairwise similarities in reconstructed space370            recon_normalized = F.normalize(recon_features, p=2, dim=1)371            recon_similarity = torch.matmul(recon_normalized, recon_normalized.t())372            373            # Encourage similar node types to have similar representations (role consistency)374            role_consistency_loss = F.mse_loss(recon_similarity, type_similarity[:max_nodes, :max_nodes])375            total_loss += role_consistency_loss376            total_nodes += 1377    378    # Return average loss379    if total_nodes > 0:380        avg_loss = total_loss / total_nodes381        if isinstance(avg_loss, torch.Tensor):382            return avg_loss.requires_grad_(True)383        else:384            return torch.tensor(avg_loss, device=original.x.device, requires_grad=True)385    else:386        return torch.tensor(0.0, device=original.x.device, requires_grad=True)387 388 389def _compute_name_loss(original: Data, reconstructed: Dict[str, Any]) -> torch.Tensor:390    """391    Compute name loss component for improved AST reconstruction.392    393    This function computes a loss that lightly encourages the model to use394    appropriate names while not penalizing heavily for choosing different395    but valid names.396    397    For backward compatibility with current features, this implements a 398    placeholder loss that encourages semantic consistency.399    In the future, this will use dedicated name embeddings.400    401    Args:402        original: Original AST data403        reconstructed: Reconstructed AST data404        405    Returns:406        Scalar tensor representing the name loss407    """408    recon_node_features = reconstructed['node_features']409    batch_size = recon_node_features.size(0)410    411    # For backward compatibility, implement a lightweight semantic consistency loss412    # This will be replaced with proper name embedding loss in the future413    414    total_loss = 0.0415    total_nodes = 0416    417    for batch_idx in range(batch_size):418        # Get original nodes for this graph419        mask = (original.batch == batch_idx)420        if not mask.any():421            continue422            423        original_nodes = original.x[mask]424        num_original_nodes = original_nodes.size(0)425        426        # Get reconstructed features427        max_nodes = min(num_original_nodes, recon_node_features.size(1))428        recon_features = recon_node_features[batch_idx, :max_nodes, :]429        430        # Lightweight semantic consistency: encourage reconstructed features to maintain431        # relative relationships present in original (approximates name consistency)432        if max_nodes > 1:433            # Compute cosine similarities in both spaces434            orig_normalized = F.normalize(original_nodes[:max_nodes], p=2, dim=1)435            recon_normalized = F.normalize(recon_features, p=2, dim=1)436            437            orig_similarities = torch.matmul(orig_normalized, orig_normalized.t())438            recon_similarities = torch.matmul(recon_normalized, recon_normalized.t())439            440            # Light penalty for changing semantic relationships (low weight applied externally)441            semantic_consistency_loss = F.mse_loss(recon_similarities, orig_similarities)442            total_loss += semantic_consistency_loss443            total_nodes += 1444    445    # Return average loss446    if total_nodes > 0:447        avg_loss = total_loss / total_nodes448        if isinstance(avg_loss, torch.Tensor):449            return avg_loss.requires_grad_(True)450        else:451            return torch.tensor(avg_loss, device=original.x.device, requires_grad=True)452    else:453        return torch.tensor(0.0, device=original.x.device, requires_grad=True)454 455 456def ast_reconstruction_loss_simple(original: Data, reconstructed: Dict[str, Any]) -> torch.Tensor:457    """458    Simplified version of AST reconstruction loss focusing primarily on node prediction.459    460    This version is easier to use and debug, focusing on the core node type prediction461    task which is the most important component for AST reconstruction.462    463    Args:464        original: Original AST as torch_geometric.data.Data object465        reconstructed: Reconstructed AST from decoder466        467    Returns:468        Scalar tensor representing the node type reconstruction loss469    """470    return compute_node_type_loss(original.x, reconstructed['node_features'], original.batch)471 472 473# ============================================================================474# Contrastive Loss Functions for Code-Text Alignment (Phase 5)475# ============================================================================476 477def info_nce_loss(code_embeddings: torch.Tensor, text_embeddings: torch.Tensor, 478                  temperature: float = 0.07) -> torch.Tensor:479    """480    InfoNCE (Information Noise Contrastive Estimation) loss for contrastive learning.481    482    This loss encourages correct (code, text) pairs to have high similarity while483    pushing incorrect pairs to have low similarity. It's commonly used in 484    contrastive learning and multimodal alignment.485    486    Args:487        code_embeddings: Code embeddings tensor of shape [batch_size, embedding_dim]488        text_embeddings: Text embeddings tensor of shape [batch_size, embedding_dim]489        temperature: Temperature parameter for scaling similarities (higher = softer)490        491    Returns:492        Scalar tensor representing the InfoNCE loss493        494    Note:495        Assumes that code_embeddings[i] and text_embeddings[i] form a positive pair,496        while all other combinations are negative pairs.497    """498    batch_size = code_embeddings.size(0)499    500    # Normalize embeddings to unit vectors for stable cosine similarity501    code_embeddings = F.normalize(code_embeddings, p=2, dim=1)502    text_embeddings = F.normalize(text_embeddings, p=2, dim=1)503    504    # Compute similarity matrix: [batch_size, batch_size]505    # similarity[i, j] = similarity between code[i] and text[j]506    similarity_matrix = torch.matmul(code_embeddings, text_embeddings.t()) / temperature507    508    # Create labels: positive pairs are on the diagonal509    labels = torch.arange(batch_size, device=code_embeddings.device)510    511    # InfoNCE loss is cross-entropy between similarity scores and correct indices512    # For each code embedding, we want the corresponding text embedding to have highest similarity513    loss_code_to_text = F.cross_entropy(similarity_matrix, labels)514    515    # Symmetric loss: for each text embedding, we want the corresponding code embedding to have highest similarity516    loss_text_to_code = F.cross_entropy(similarity_matrix.t(), labels)517    518    # Return average of both directions519    return (loss_code_to_text + loss_text_to_code) / 2.0520 521 522def cosine_embedding_loss(code_embeddings: torch.Tensor, text_embeddings: torch.Tensor,523                         margin: float = 0.2) -> torch.Tensor:524    """525    Simple cosine embedding loss for contrastive learning.526    527    This loss encourages positive pairs to have high cosine similarity (close to 1)528    and negative pairs to have low cosine similarity (below margin).529    530    Args:531        code_embeddings: Code embeddings tensor of shape [batch_size, embedding_dim]532        text_embeddings: Text embeddings tensor of shape [batch_size, embedding_dim]533        margin: Margin for negative pairs (similarity should be below this)534        535    Returns:536        Scalar tensor representing the cosine embedding loss537    """538    batch_size = code_embeddings.size(0)539    540    # Normalize embeddings for stable cosine similarity541    code_embeddings = F.normalize(code_embeddings, p=2, dim=1)542    text_embeddings = F.normalize(text_embeddings, p=2, dim=1)543    544    # Compute cosine similarities for all pairs545    similarity_matrix = torch.matmul(code_embeddings, text_embeddings.t())546    547    # Positive pairs: diagonal elements (code[i] with text[i])548    positive_similarities = torch.diag(similarity_matrix)549    550    # Loss for positive pairs: encourage high similarity (target = 1)551    positive_loss = F.mse_loss(positive_similarities, torch.ones_like(positive_similarities))552    553    # For negative pairs, only apply if we have more than one sample554    if batch_size > 1:555        # Negative pairs: off-diagonal elements556        mask = torch.eye(batch_size, device=code_embeddings.device).bool()557        negative_similarities = similarity_matrix[~mask]558        559        # Loss for negative pairs: encourage low similarity (below margin)560        # Only penalize if similarity is above margin561        negative_loss = F.relu(negative_similarities - margin).mean()562    else:563        # No negative pairs when batch size is 1564        negative_loss = torch.tensor(0.0, device=code_embeddings.device)565    566    # Combine losses567    return positive_loss + negative_loss568 569 570def simple_contrastive_loss(code_embeddings: torch.Tensor, text_embeddings: torch.Tensor,571                           temperature: float = 0.1) -> torch.Tensor:572    """573    Simplified contrastive loss using cosine similarity.574    575    This is a straightforward implementation that maximizes cosine similarity576    between correct pairs and minimizes it for incorrect pairs.577    578    Args:579        code_embeddings: Code embeddings tensor of shape [batch_size, embedding_dim]580        text_embeddings: Text embeddings tensor of shape [batch_size, embedding_dim]581        temperature: Temperature for scaling similarities582        583    Returns:584        Scalar tensor representing the contrastive loss585    """586    # Normalize embeddings587    code_embeddings = F.normalize(code_embeddings, p=2, dim=1)588    text_embeddings = F.normalize(text_embeddings, p=2, dim=1)589    590    # Compute cosine similarities591    similarities = F.cosine_similarity(code_embeddings, text_embeddings, dim=1)592    593    # Loss is simply negative mean similarity (we want to maximize similarity)594    # Scale by temperature for better gradient flow595    return -similarities.mean() / temperature596