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"""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 