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.
0104
1"""2Data processing utilities for Ruby method datasets.3 4This module provides functions to load, preprocess, and prepare Ruby method5data for GNN training. Includes custom Dataset class for AST to graph conversion.6"""7 8import json9import random10import os11import logging12from pathlib import Path13from typing import List, Dict, Any, Tuple, Optional, Union14try:15 import torch16 from torch_geometric.data import Data17 TORCH_AVAILABLE = True18except ImportError:19 TORCH_AVAILABLE = False20 21 22def load_methods_json(filepath: str) -> List[Dict[str, Any]]:23 """24 Load Ruby methods from JSON file.25 26 Args:27 filepath: Path to the JSON file containing method data28 29 Returns:30 List of method dictionaries31 """32 with open(filepath, 'r') as f:33 return json.load(f)34 35 36def methods_to_dataframe(methods: List[Dict[str, Any]]) -> List[Dict[str, Any]]:37 """38 Convert list of method dictionaries to a structured format.39 40 Args:41 methods: List of method dictionaries42 43 Returns:44 List of method dictionaries (pass-through for compatibility)45 """46 return methods47 48 49def filter_methods_by_length(methods: List[Dict[str, Any]], min_lines: int = 5, max_lines: int = 100) -> List[Dict[str, Any]]:50 """51 Filter methods by source code length.52 53 Args:54 methods: List of method dictionaries55 min_lines: Minimum number of lines56 max_lines: Maximum number of lines57 58 Returns:59 Filtered list of methods60 """61 filtered = []62 for method in methods:63 if 'raw_source' in method:64 line_count = len(method['raw_source'].split('\n'))65 if min_lines <= line_count <= max_lines:66 method['line_count'] = line_count67 filtered.append(method)68 return filtered69 """70 Filter methods by source code length.71 72 Args:73 df: DataFrame containing method data74 min_lines: Minimum number of lines75 max_lines: Maximum number of lines76 77 Returns:78 Filtered DataFrame79 """80 df['line_count'] = df['raw_source'].apply(lambda x: len(x.split('\n')))81 return df[(df['line_count'] >= min_lines) & (df['line_count'] <= max_lines)]82 83 84class ASTNodeEncoder:85 """86 Encoder for mapping AST node types to feature vectors.87 88 This class maintains a vocabulary of AST node types found in Ruby code89 and maps them to dense feature vectors for GNN processing.90 """91 92 def __init__(self):93 """Initialize the node encoder with common Ruby AST node types."""94 # Common Ruby AST node types based on the parser gem95 self.node_types = [96 'def', 'defs', 'args', 'arg', 'begin', 'end', 'lvasgn', 'ivasgn', 'gvasgn',97 'cvasgn', 'send', 'block', 'if', 'unless', 'while', 'until', 'for', 'case',98 'when', 'rescue', 'ensure', 'retry', 'break', 'next', 'redo', 'return',99 'yield', 'super', 'zsuper', 'lambda', 'proc', 'and', 'or', 'not', 'true',100 'false', 'nil', 'self', 'int', 'float', 'str', 'sym', 'regexp', 'array',101 'hash', 'pair', 'splat', 'kwsplat', 'block_pass', 'const', 'cbase',102 'lvar', 'ivar', 'gvar', 'cvar', 'casgn', 'masgn', 'mlhs', 'op_asgn',103 'and_asgn', 'or_asgn', 'back_ref', 'nth_ref', 'class', 'sclass', 'module',104 'defined?', 'alias', 'undef', 'range', 'irange', 'erange', 'regopt'105 ]106 107 # Create mapping from node type to index108 self.type_to_idx = {node_type: idx for idx, node_type in enumerate(self.node_types)}109 self.unknown_idx = len(self.node_types) # Index for unknown node types110 self.vocab_size = len(self.node_types) + 1 # +1 for unknown111 112 def encode_node_type(self, node_type: str) -> int:113 """114 Encode a node type to its integer index.115 116 Args:117 node_type: The AST node type string118 119 Returns:120 Integer index for the node type121 """122 return self.type_to_idx.get(node_type, self.unknown_idx)123 124 def create_node_features(self, node_type: str) -> List[float]:125 """126 Create feature vector for a node type.127 128 Args:129 node_type: The AST node type string130 131 Returns:132 Feature vector as list of floats133 """134 # Simple one-hot encoding for now135 features = [0.0] * self.vocab_size136 idx = self.encode_node_type(node_type)137 features[idx] = 1.0138 return features139 140 141class ASTGraphConverter:142 """143 Converter for transforming AST JSON to graph representation.144 145 This class parses the AST JSON structure and converts it into146 a graph format suitable for GNN processing.147 """148 149 def __init__(self):150 """Initialize the AST to graph converter."""151 self.node_encoder = ASTNodeEncoder()152 self.reset()153 154 def reset(self):155 """Reset the converter state for processing a new AST."""156 self.nodes = [] # List of node features157 self.edges = [] # List of edge tuples (parent_idx, child_idx)158 self.edge_attrs = [] # List of edge attributes [child_index, depth, num_siblings]159 self.node_depths = [] # Depth of each node in the tree160 self.node_child_indices = [] # Position of each node among its siblings161 self.node_count = 0162 163 def parse_ast_json(self, ast_json: str) -> Dict[str, Any]:164 """165 Parse AST JSON string and convert to graph representation.166 167 Args:168 ast_json: JSON string representing the AST169 170 Returns:171 Dictionary containing node features, edge indices, and edge attributes.172 edge_attr contains [child_index, depth, num_siblings] per edge.173 node_pos contains [child_index, depth] per node for positional encoding.174 """175 self.reset()176 177 try:178 ast_data = json.loads(ast_json)179 self._process_node(ast_data, parent_idx=None, depth=0, child_index=0, num_siblings=1)180 181 # Convert to appropriate format182 if not self.nodes:183 # Handle empty AST case184 node_features = [[0.0] * self.node_encoder.vocab_size]185 edge_index = [[], []] # Empty edge list186 edge_attr = []187 node_pos = [[0, 0]]188 else:189 node_features = self.nodes190 if self.edges:191 # Transpose edge list to [2, num_edges] format192 edge_index = [[], []]193 for parent, child in self.edges:194 edge_index[0].append(parent)195 edge_index[1].append(child)196 else:197 edge_index = [[], []]198 edge_attr = self.edge_attrs199 node_pos = list(zip(self.node_child_indices, self.node_depths))200 201 return {202 'x': node_features,203 'edge_index': edge_index,204 'edge_attr': edge_attr,205 'node_pos': node_pos,206 'num_nodes': len(self.nodes) if self.nodes else 1207 }208 209 except (json.JSONDecodeError, Exception):210 # Handle malformed JSON or other errors gracefully211 return {212 'x': [[0.0] * self.node_encoder.vocab_size],213 'edge_index': [[], []],214 'edge_attr': [],215 'node_pos': [[0, 0]],216 'num_nodes': 1217 }218 219 def _process_node(self, node: Union[Dict, List, str, int, float, None],220 parent_idx: Optional[int] = None, depth: int = 0,221 child_index: int = 0, num_siblings: int = 1) -> int:222 """223 Recursively process an AST node and its children.224 225 Args:226 node: The AST node (dict, list, or primitive)227 parent_idx: Index of the parent node228 depth: Depth of the current node in the AST229 child_index: Position of this node among its siblings (0-based)230 num_siblings: Total number of siblings (including this node)231 232 Returns:233 Index of the current node234 """235 if isinstance(node, dict) and 'type' in node:236 # This is an AST node with a type237 node_type = node['type']238 current_idx = self.node_count239 self.node_count += 1240 241 # Create node features242 features = self.node_encoder.create_node_features(node_type)243 self.nodes.append(features)244 self.node_depths.append(depth)245 self.node_child_indices.append(child_index)246 247 # Add edge from parent to current node248 if parent_idx is not None:249 self.edges.append((parent_idx, current_idx))250 self.edge_attrs.append([child_index, depth, num_siblings])251 252 # Process children with positional information253 if 'children' in node:254 children = node['children']255 n_children = len(children)256 for i, child in enumerate(children):257 self._process_node(child, current_idx, depth=depth + 1,258 child_index=i, num_siblings=n_children)259 260 return current_idx261 262 elif isinstance(node, list):263 # Process list of nodes264 n_items = len(node)265 for i, child in enumerate(node):266 self._process_node(child, parent_idx, depth=depth,267 child_index=i, num_siblings=n_items)268 return parent_idx if parent_idx is not None else -1269 270 else:271 # Leaf node (string, int, float, None)272 if parent_idx is not None:273 current_idx = self.node_count274 self.node_count += 1275 276 # Create a generic leaf node277 leaf_type = 'leaf_' + type(node).__name__278 features = self.node_encoder.create_node_features(leaf_type)279 self.nodes.append(features)280 self.node_depths.append(depth)281 self.node_child_indices.append(child_index)282 283 # Add edge from parent to leaf284 self.edges.append((parent_idx, current_idx))285 self.edge_attrs.append([child_index, depth, num_siblings])286 287 return current_idx288 return -1289 290 291def load_jsonl_file(filepath: str, limit: Optional[int] = None) -> List[Dict[str, Any]]:292 """293 Load data from a JSONL file.294 295 Args:296 filepath: Path to the JSONL file297 limit: Optional maximum number of lines to load.298 299 Returns:300 List of dictionaries from the JSONL file301 """302 data = []303 with open(filepath, 'r', encoding='utf-8') as f:304 for i, line in enumerate(f):305 if limit is not None and i >= limit:306 break307 line = line.strip()308 if line:309 try:310 data.append(json.loads(line))311 except json.JSONDecodeError:312 continue # Skip malformed lines313 return data314 315 316class RubyASTDataset:317 """318 Dataset class for loading Ruby AST data and converting to graph format.319 320 This class loads JSONL files containing Ruby method data and converts321 the AST representations to graph objects suitable for GNN training.322 """323 324 def __init__(self, jsonl_path: str, transform=None, limit: Optional[int] = None):325 """326 Initialize the dataset.327 328 Args:329 jsonl_path: Path to the JSONL file containing method data330 transform: Optional transform to apply to each sample331 limit: Optional maximum number of samples to load.332 """333 self.jsonl_path = jsonl_path334 self.transform = transform335 self.converter = ASTGraphConverter()336 337 # Load the data338 self.data = load_jsonl_file(jsonl_path, limit=limit)339 340 print(f"Loaded {len(self.data)} samples from {jsonl_path}")341 342 def __len__(self) -> int:343 """Return the number of samples in the dataset."""344 return len(self.data)345 346 def __getitem__(self, idx: int) -> Dict[str, Any]:347 """348 Get a sample from the dataset.349 350 Args:351 idx: Index of the sample352 353 Returns:354 Dictionary containing graph data and target355 """356 if idx < 0 or idx >= len(self.data):357 raise IndexError(f"Index {idx} out of range for dataset of size {len(self.data)}")358 359 sample = self.data[idx]360 361 # Convert AST to graph362 graph_data = self.converter.parse_ast_json(sample['ast_json'])363 364 # Create the data object365 result = {366 'x': graph_data['x'],367 'edge_index': graph_data['edge_index'],368 'y': [sample.get('complexity_score', 5.0)], # Default complexity score if missing369 'num_nodes': graph_data['num_nodes'],370 'id': sample.get('id', f'sample_{idx}'),371 'repo_name': sample.get('repo_name', ''),372 'file_path': sample.get('file_path', '')373 }374 375 # Apply transform if provided376 if self.transform:377 result = self.transform(result)378 379 return result380 381 def get_feature_dim(self) -> int:382 """Return the dimension of node features."""383 return self.converter.node_encoder.vocab_size384 385 386def collate_graphs(batch: List[Dict[str, Any]]) -> Dict[str, Any]:387 """388 Collate function for batching graph data.389 390 Args:391 batch: List of graph data dictionaries392 393 Returns:394 Batched graph data395 """396 if not batch:397 raise ValueError("Cannot collate empty batch")398 399 # Collect all node features and edge indices400 all_x = []401 all_edge_index = [[], []] # [source_nodes, target_nodes]402 all_y = []403 batch_idx = []404 node_offset = 0405 406 metadata = {407 'ids': [],408 'repo_names': [],409 'file_paths': []410 }411 412 for i, sample in enumerate(batch):413 # Node features414 all_x.extend(sample['x'])415 416 # Edge indices (offset by current node count)417 edges = sample['edge_index']418 if len(edges[0]) > 0: # Only offset if there are edges419 for j in range(len(edges[0])):420 all_edge_index[0].append(edges[0][j] + node_offset)421 all_edge_index[1].append(edges[1][j] + node_offset)422 423 # Target values424 all_y.extend(sample['y'])425 426 # Batch indices for each node427 num_nodes = sample['num_nodes']428 batch_idx.extend([i] * num_nodes)429 node_offset += num_nodes430 431 # Metadata432 metadata['ids'].append(sample['id'])433 metadata['repo_names'].append(sample['repo_name'])434 metadata['file_paths'].append(sample['file_path'])435 436 return {437 'x': all_x,438 'edge_index': all_edge_index,439 'y': all_y,440 'batch': batch_idx,441 'num_graphs': len(batch),442 'metadata': metadata443 }444 445 446class SimpleDataLoader:447 """448 Simple DataLoader implementation for batching data.449 450 This provides a basic implementation that can be used when PyTorch451 DataLoader is not available, and can easily be replaced with the real452 PyTorch DataLoader when dependencies are installed.453 """454 455 def __init__(self, dataset, batch_size: int = 1, shuffle: bool = False, collate_fn=None):456 """457 Initialize the DataLoader.458 459 Args:460 dataset: Dataset to load from461 batch_size: Number of samples per batch462 shuffle: Whether to shuffle the data463 collate_fn: Function to collate samples into batches464 """465 self.dataset = dataset466 self.batch_size = batch_size467 self.shuffle = shuffle468 self.collate_fn = collate_fn or collate_graphs469 470 # Create indices471 self.indices = list(range(len(dataset)))472 if shuffle:473 import random474 random.shuffle(self.indices)475 476 def __len__(self) -> int:477 """Return number of batches."""478 return (len(self.dataset) + self.batch_size - 1) // self.batch_size479 480 def __iter__(self):481 """Iterate over batches."""482 for i in range(0, len(self.dataset), self.batch_size):483 batch_indices = self.indices[i:i + self.batch_size]484 batch = [self.dataset[idx] for idx in batch_indices]485 yield self.collate_fn(batch)486 487 488class PairedDataset:489 """490 Dataset class for loading paired Ruby AST and text description data.491 492 This class loads the paired_data.jsonl file containing Ruby method data 493 and converts AST representations to graph objects paired with text descriptions.494 For each method, it randomly samples one description from the available descriptions.495 """496 497 def __init__(self, jsonl_path: str, transform=None, seed: Optional[int] = None, limit: Optional[int] = None):498 """499 Initialize the paired dataset.500 501 Args:502 jsonl_path: Path to the paired_data.jsonl file503 transform: Optional transform to apply to each sample504 seed: Random seed for consistent description sampling505 limit: Optional maximum number of samples to load.506 """507 self.jsonl_path = jsonl_path508 self.transform = transform509 self.converter = ASTGraphConverter()510 511 if seed is not None:512 random.seed(seed)513 514 # Load the data515 self.data = load_jsonl_file(jsonl_path, limit=limit)516 517 print(f"Loaded {len(self.data)} samples from {jsonl_path}")518 519 def __len__(self) -> int:520 """Return the number of samples in the dataset."""521 return len(self.data)522 523 def __getitem__(self, idx: int) -> Tuple[Dict[str, Any], str]:524 """525 Get a sample from the dataset.526 527 Args:528 idx: Index of the sample529 530 Returns:531 Tuple of (graph_data, text_description)532 """533 if idx < 0 or idx >= len(self.data):534 raise IndexError(f"Index {idx} out of range for dataset of size {len(self.data)}")535 536 sample = self.data[idx]537 538 # Convert AST to graph539 graph_data = self.converter.parse_ast_json(sample['ast_json'])540 541 # Randomly sample one description542 descriptions = sample.get('descriptions', [])543 if descriptions:544 description = random.choice(descriptions)545 text_description = description['text']546 else:547 # Fallback to method name if no descriptions available548 text_description = sample.get('method_name', 'unknown_method')549 550 # Create the graph data object551 graph_result = {552 'x': graph_data['x'],553 'edge_index': graph_data['edge_index'],554 'num_nodes': graph_data['num_nodes'],555 'id': sample.get('id', f'sample_{idx}'),556 'repo_name': sample.get('repo_name', ''),557 'file_path': sample.get('file_path', '')558 }559 560 # Apply transform if provided561 if self.transform:562 graph_result = self.transform(graph_result)563 564 return graph_result, text_description565 566 def get_feature_dim(self) -> int:567 """Return the dimension of node features."""568 return self.converter.node_encoder.vocab_size569 570 571def collate_paired_data(batch: List[Tuple[Dict[str, Any], str]]) -> Tuple[Dict[str, Any], List[str]]:572 """573 Collate function for batching paired graph and text data.574 575 Args:576 batch: List of (graph_data, text_description) tuples577 578 Returns:579 Tuple of (batched_graph_data, list_of_text_descriptions)580 """581 if not batch:582 raise ValueError("Cannot collate empty batch")583 584 # Separate graph data and text descriptions585 graph_batch = [item[0] for item in batch]586 text_batch = [item[1] for item in batch]587 588 # Collate graph data manually (similar to collate_graphs but without 'y' field)589 all_x = []590 all_edge_index = [[], []] # [source_nodes, target_nodes]591 batch_idx = []592 node_offset = 0593 594 metadata = {595 'ids': [],596 'repo_names': [],597 'file_paths': []598 }599 600 for i, sample in enumerate(graph_batch):601 # Node features602 all_x.extend(sample['x'])603 604 # Edge indices (offset by current node count)605 edges = sample['edge_index']606 if len(edges[0]) > 0: # Only offset if there are edges607 for j in range(len(edges[0])):608 all_edge_index[0].append(edges[0][j] + node_offset)609 all_edge_index[1].append(edges[1][j] + node_offset)610 611 # Batch indices for each node612 num_nodes = sample['num_nodes']613 batch_idx.extend([i] * num_nodes)614 node_offset += num_nodes615 616 # Metadata617 metadata['ids'].append(sample['id'])618 metadata['repo_names'].append(sample['repo_name'])619 metadata['file_paths'].append(sample['file_path'])620 621 batched_graphs = {622 'x': all_x,623 'edge_index': all_edge_index,624 'batch': batch_idx,625 'num_graphs': len(batch),626 'metadata': metadata627 }628 629 return batched_graphs, text_batch630 631 632class PairedDataLoader:633 """634 DataLoader for paired graph and text data.635 636 Extends SimpleDataLoader to handle paired (graph, text) data.637 """638 639 def __init__(self, dataset, batch_size: int = 1, shuffle: bool = False):640 """641 Initialize the PairedDataLoader.642 643 Args:644 dataset: PairedDataset to load from645 batch_size: Number of samples per batch646 shuffle: Whether to shuffle the data647 """648 self.dataset = dataset649 self.batch_size = batch_size650 self.shuffle = shuffle651 652 # Create indices653 self.indices = list(range(len(dataset)))654 if shuffle:655 random.shuffle(self.indices)656 657 def __len__(self) -> int:658 """Return number of batches."""659 return (len(self.dataset) + self.batch_size - 1) // self.batch_size660 661 def __iter__(self):662 """Iterate over batches."""663 for i in range(0, len(self.dataset), self.batch_size):664 batch_indices = self.indices[i:i + self.batch_size]665 batch = [self.dataset[idx] for idx in batch_indices]666 yield collate_paired_data(batch)667 668 669 670class PrecomputedRubyASTDataset:671 """672 Dataset class for loading precomputed Ruby AST graph data.673 674 This class can load .pt files containing pre-converted PyTorch Geometric675 Data objects for speed, but also supports processing .jsonl files as a fallback.676 """677 678 def __init__(self, path: str, transform=None):679 """680 Initialize the dataset.681 682 Args:683 path: Path to the .pt or .jsonl file containing graph data.684 transform: Optional transform to apply to each sample.685 """686 self.path = path687 self.transform = transform688 689 if not TORCH_AVAILABLE:690 raise ImportError("PyTorch and PyG are required for this dataset.")691 692 if path.endswith('.pt'):693 # Load the precomputed data into RAM694 self.data = torch.load(path, weights_only=False)695 print(f"Loaded {len(self.data)} precomputed graphs from {path}")696 elif path.endswith('.jsonl'):697 print(f"Processing JSONL file into graphs: {path}")698 jsonl_data = load_jsonl_file(path)699 converter = ASTGraphConverter()700 self.data = []701 for sample in jsonl_data:702 graph_data = converter.parse_ast_json(sample['ast_json'])703 704 x = torch.tensor(graph_data['x'], dtype=torch.float)705 edge_index = torch.tensor(graph_data['edge_index'], dtype=torch.long)706 y = torch.tensor([sample.get('complexity_score', 5.0)], dtype=torch.float)707 708 data_obj = Data(x=x, edge_index=edge_index, y=y)709 710 # Add positional attributes — always set so PyG collation is consistent711 ea = graph_data.get('edge_attr', [])712 data_obj.edge_attr = torch.tensor(713 ea if ea else [], dtype=torch.float,714 ).reshape(-1, 3) if ea else torch.zeros((0, 3), dtype=torch.float)715 716 np_ = graph_data.get('node_pos', [])717 data_obj.node_pos = torch.tensor(718 np_ if np_ else [[0, 0]], dtype=torch.float,719 )720 721 self.data.append(data_obj)722 print(f"Converted {len(self.data)} graphs from {path}")723 else:724 raise ValueError(f"Unsupported file type: {path}. Please provide a .pt or .jsonl file.")725 726 def __len__(self) -> int:727 """Return the number of samples in the dataset."""728 return len(self.data)729 730 def __getitem__(self, idx: int):731 """732 Get a sample from the dataset.733 734 Args:735 idx: Index of the sample736 737 Returns:738 PyTorch Geometric Data object739 """740 if idx < 0 or idx >= len(self.data):741 raise IndexError(f"Index {idx} out of range for dataset of size {len(self.data)}")742 743 sample = self.data[idx]744 745 if self.transform:746 sample = self.transform(sample)747 748 return sample749 750 751class PreCollatedDataset:752 """753 Dataset class for loading pre-collated batches of graph data.754 755 This class loads a .pt file where each item is an already-collated756 `torch_geometric.data.Batch` object. This is the most efficient757 way to load data as it eliminates all real-time collation overhead.758 """759 def __init__(self, pt_path: str):760 """761 Initialize the dataset.762 763 Args:764 pt_path: Path to the .pt file containing pre-collated batches.765 """766 # Load the list of pre-collated batches into RAM767 self.batches = torch.load(pt_path, weights_only=False)768 print(f"Loaded {len(self.batches)} pre-collated batches from {pt_path}")769 770 def __len__(self):771 return len(self.batches)772 773 def __getitem__(self, idx):774 return self.batches[idx]775 776 777def create_data_loaders(train_path: str, val_path: str, batch_size: int = 32, shuffle: bool = True, num_workers: Optional[int] = None, pre_collated: bool = False):778 """779 Create train and validation data loaders.780 781 Supports two modes:782 1. Standard loading from a dataset of individual graphs (`pre_collated=False`).783 This uses a PyG DataLoader to perform real-time batching.784 2. Pre-collated loading from a dataset of pre-batched graphs (`pre_collated=True`).785 This is the most performant option, as it has near-zero CPU overhead.786 787 Args:788 train_path: Path to training .pt file.789 val_path: Path to validation .pt file.790 batch_size: Batch size (used only if `pre_collated=False`).791 shuffle: Whether to shuffle training data.792 num_workers: Number of workers for data loading (used only if `pre_collated=False`).793 pre_collated: Whether the dataset files contain pre-collated batches.794 795 Returns:796 Tuple of (train_loader, val_loader)797 """798 if not TORCH_AVAILABLE:799 raise ImportError("PyTorch is required to create data loaders.")800 801 if pre_collated:802 # --- Pre-collated path (most efficient) ---803 train_dataset = PreCollatedDataset(train_path)804 val_dataset = PreCollatedDataset(val_path)805 806 # The collate_fn simply returns the already-collated batch.807 # The input `batch` is a list of size 1 containing our pre-made Batch object.808 collate_fn = lambda x: x[0]809 810 # DataLoader is just a simple iterator here, no real collation work.811 # num_workers > 0 can actually be slower due to overhead of sending812 # already-large batches between processes.813 from torch.utils.data import DataLoader814 train_loader = DataLoader(train_dataset, batch_size=1, shuffle=shuffle, num_workers=0, collate_fn=collate_fn)815 val_loader = DataLoader(val_dataset, batch_size=1, shuffle=False, num_workers=0, collate_fn=collate_fn)816 817 print("✅ Using pre-collated data loader (maximum performance).")818 819 else:820 # --- Standard real-time collation path ---821 from torch_geometric.loader import DataLoader822 train_dataset = PrecomputedRubyASTDataset(train_path)823 val_dataset = PrecomputedRubyASTDataset(val_path)824 825 if num_workers is None:826 num_workers = os.cpu_count()827 828 train_loader = DataLoader(829 train_dataset, 830 batch_size=batch_size, 831 shuffle=shuffle,832 num_workers=num_workers,833 pin_memory=torch.cuda.is_available(),834 persistent_workers=num_workers > 0835 )836 val_loader = DataLoader(837 val_dataset, 838 batch_size=batch_size, 839 shuffle=False,840 num_workers=num_workers,841 pin_memory=torch.cuda.is_available(),842 persistent_workers=num_workers > 0843 )844 845 print(f"✅ Using standard PyG DataLoader with {num_workers} workers.")846 847 return train_loader, val_loader848 849 850 851def create_paired_data_loaders(paired_data_path: str, batch_size: int = 32, shuffle: bool = True, seed: Optional[int] = None):852 """853 Create data loader for paired graph and text data.854 855 Args:856 paired_data_path: Path to paired_data.jsonl file857 batch_size: Batch size for the loader858 shuffle: Whether to shuffle the data859 seed: Random seed for consistent description sampling860 861 Returns:862 PairedDataLoader instance863 """864 dataset = PairedDataset(paired_data_path, seed=seed)865 loader = PairedDataLoader(dataset, batch_size=batch_size, shuffle=shuffle)866 867 return loader868 869 870class AutoregressiveASTDataset:871 """872 Dataset class for autoregressive AST generation training.873 874 This class loads paired Ruby AST and text description data and converts875 each AST into a sequence of (partial_graph, target_node) pairs for 876 autoregressive training. Each method generates multiple training examples.877 """878 879 def __init__(self, paired_data_path: str, max_sequence_length: int = 50, seed: Optional[int] = None,880 precomputed_embeddings_path: Optional[str] = None):881 """882 Initialize the autoregressive dataset.883 884 Args:885 paired_data_path: Path to the paired_data.jsonl file886 max_sequence_length: Maximum number of nodes per sequence887 seed: Random seed for consistent description sampling888 precomputed_embeddings_path: Path to pre-computed text embeddings file (optional)889 """890 self.paired_data_path = paired_data_path891 self.max_sequence_length = max_sequence_length892 self.converter = ASTGraphConverter()893 894 if seed is not None:895 random.seed(seed)896 897 # Load pre-computed embeddings if available898 self.precomputed_embeddings = {}899 if precomputed_embeddings_path and os.path.exists(precomputed_embeddings_path):900 try:901 if TORCH_AVAILABLE:902 self.precomputed_embeddings = torch.load(precomputed_embeddings_path, map_location='cpu', weights_only=True)903 print(f"✅ Loaded {len(self.precomputed_embeddings)} pre-computed text embeddings")904 else:905 print("⚠️ PyTorch not available, skipping pre-computed embeddings")906 except Exception as e:907 print(f"⚠️ Warning: Could not load pre-computed embeddings: {e}")908 elif precomputed_embeddings_path:909 print(f"⚠️ Warning: Pre-computed embeddings file not found: {precomputed_embeddings_path}")910 911 # Load the paired data912 self.paired_data = load_jsonl_file(paired_data_path)913 914 # Generate sequential training pairs from all methods915 self.sequential_pairs = []916 self._generate_all_sequential_pairs()917 918 print(f"Loaded {len(self.paired_data)} methods from {paired_data_path}")919 print(f"Generated {len(self.sequential_pairs)} sequential training pairs")920 921 def _generate_all_sequential_pairs(self):922 """Generate sequential training pairs from all ASTs in the dataset."""923 for sample in self.paired_data:924 try:925 # Get text description926 descriptions = sample.get('descriptions', [])927 if descriptions:928 description = random.choice(descriptions)929 text_description = description['text']930 else:931 # Fallback to method name if no descriptions available932 text_description = sample.get('method_name', 'unknown_method')933 934 # Create sequential pairs for this AST935 sequential_pairs = self._create_sequential_pairs(936 sample['ast_json'], 937 text_description938 )939 940 # Add to global list941 self.sequential_pairs.extend(sequential_pairs)942 943 except Exception as e:944 # Skip malformed samples gracefully945 print(f"Warning: Skipping sample {sample.get('id', 'unknown')} due to error: {e}")946 continue947 948 def _create_sequential_pairs(self, ast_json: str, text_description: str) -> List[Dict[str, Any]]:949 """950 Convert single AST into sequence of (partial_graph, target_node) pairs.951 952 Args:953 ast_json: JSON string representing the AST954 text_description: Text description for this method955 956 Returns:957 List of sequential training pairs958 """959 pairs = []960 961 try:962 # Extract nodes in proper order along with their connections963 nodes, connections = self._extract_nodes_and_connections_in_order(ast_json)964 965 # Limit sequence length if needed966 if len(nodes) > self.max_sequence_length:967 nodes = nodes[:self.max_sequence_length]968 # Also limit connections to only include those within the sequence969 filtered_connections = []970 for src, tgt in connections:971 if src < self.max_sequence_length and tgt < self.max_sequence_length:972 filtered_connections.append((src, tgt))973 connections = filtered_connections974 975 # Get pre-computed text embedding if available, otherwise store text976 text_embedding = None977 if text_description in self.precomputed_embeddings:978 text_embedding = self.precomputed_embeddings[text_description]979 980 # Create sequential pairs981 for i in range(len(nodes)):982 # Build partial graph with nodes 0 to i-1983 partial_graph = self._build_partial_graph(nodes[:i])984 985 # Target is the i-th node986 target_node = nodes[i]987 988 # Create target connections for this step989 # This represents which existing nodes (0 to i-1) the new node i should connect to990 target_connections = self._create_target_connections(i, connections)991 992 pair = {993 'text_description': text_description,994 'text_embedding': text_embedding, # Pre-computed embedding if available995 'partial_graph': partial_graph,996 'target_node': target_node,997 'target_connections': target_connections,998 'step': i,999 'total_steps': len(nodes)1000 }1001 1002 pairs.append(pair)1003 1004 except Exception as e:1005 # Return empty list for malformed ASTs1006 print(f"Warning: Failed to create sequential pairs: {e}")1007 1008 return pairs1009 1010 def _extract_nodes_and_connections_in_order(self, ast_json: str) -> Tuple[List[Dict[str, Any]], List[Tuple[int, int]]]:1011 """1012 Extract nodes and their connections from AST in proper depth-first order.1013 1014 Args:1015 ast_json: JSON string representing the AST1016 1017 Returns:1018 Tuple of (nodes_list, connections_list) where connections are (parent_idx, child_idx) pairs1019 """1020 try:1021 ast_data = json.loads(ast_json)1022 nodes = []1023 connections = []1024 self._traverse_ast_nodes_with_connections(ast_data, nodes, connections, parent_idx=None)1025 return nodes, connections1026 except (json.JSONDecodeError, Exception):1027 # Return empty lists for malformed JSON1028 return [], []1029 1030 def _traverse_ast_nodes_with_connections(self, node: Union[Dict, List, str, int, float, None], 1031 nodes: List[Dict[str, Any]], 1032 connections: List[Tuple[int, int]],1033 parent_idx: Optional[int] = None):1034 """1035 Recursively traverse AST and collect nodes and connections in depth-first order.1036 1037 Args:1038 node: Current AST node1039 nodes: List to collect nodes1040 connections: List to collect connections as (parent_idx, child_idx) pairs1041 parent_idx: Index of parent node1042 """1043 if isinstance(node, dict) and 'type' in node:1044 # This is an AST node with a type1045 current_idx = len(nodes)1046 node_info = {1047 'node_type': node['type'],1048 'features': self.converter.node_encoder.create_node_features(node['type']),1049 'raw_node': node # Keep reference for debugging1050 }1051 nodes.append(node_info)1052 1053 # Add connection from parent to current node1054 if parent_idx is not None:1055 connections.append((parent_idx, current_idx))1056 1057 # Traverse children1058 if 'children' in node:1059 for child in node['children']:1060 self._traverse_ast_nodes_with_connections(child, nodes, connections, current_idx)1061 1062 elif isinstance(node, list):1063 # Process list of nodes1064 for child in node:1065 self._traverse_ast_nodes_with_connections(child, nodes, connections, parent_idx)1066 1067 def _create_target_connections(self, node_idx: int, all_connections: List[Tuple[int, int]]) -> List[float]:1068 """1069 Create target connection vector for a specific node being added.1070 1071 Args:1072 node_idx: Index of the node being added to the graph1073 all_connections: List of all connections in the full AST as (parent_idx, child_idx) pairs1074 1075 Returns:1076 Binary vector of length max_nodes indicating which existing nodes to connect to1077 """1078 # Initialize with zeros for all possible connections1079 target_vector = [0.0] * 100 # max_nodes = 100 from model1080 1081 # Find all connections where this node is the target (child)1082 # We want to know which existing nodes (with index < node_idx) should connect to this node1083 for parent_idx, child_idx in all_connections:1084 if child_idx == node_idx and parent_idx < node_idx and parent_idx < 100:1085 target_vector[parent_idx] = 1.01086 1087 return target_vector1088 1089 def _traverse_ast_nodes(self, node: Union[Dict, List, str, int, float, None], nodes: List[Dict[str, Any]]):1090 """1091 Recursively traverse AST and collect nodes in depth-first order.1092 1093 Args:1094 node: Current AST node1095 nodes: List to collect nodes1096 """1097 if isinstance(node, dict) and 'type' in node:1098 # This is an AST node with a type1099 node_info = {1100 'node_type': node['type'],1101 'features': self.converter.node_encoder.create_node_features(node['type']),1102 'raw_node': node # Keep reference for debugging1103 }1104 nodes.append(node_info)1105 1106 # Traverse children1107 if 'children' in node:1108 for child in node['children']:1109 self._traverse_ast_nodes(child, nodes)1110 1111 elif isinstance(node, list):1112 # Process list of nodes1113 for child in node:1114 self._traverse_ast_nodes(child, nodes)1115 1116 def _build_partial_graph(self, nodes: List[Dict[str, Any]]) -> Dict[str, Any]:1117 """1118 Build partial graph from first i nodes.1119 1120 Args:1121 nodes: List of nodes to include in partial graph1122 1123 Returns:1124 Partial graph representation1125 """1126 if not nodes:1127 # Empty graph case1128 return {1129 'x': [],1130 'edge_index': [[], []],1131 'num_nodes': 01132 }1133 1134 # Extract node features1135 node_features = [node['features'] for node in nodes]1136 1137 # Create simple sequential connections (each node connects to next)1138 # This is a simplified approach - in practice you'd want to preserve 1139 # the actual AST structure relationships1140 edge_list = []1141 for i in range(len(nodes) - 1):1142 edge_list.append([i, i + 1]) # Forward edge1143 edge_list.append([i + 1, i]) # Backward edge for undirected1144 1145 if edge_list:1146 edge_index = [[], []]1147 for source, target in edge_list:1148 edge_index[0].append(source)1149 edge_index[1].append(target)1150 else:1151 edge_index = [[], []]1152 1153 return {1154 'x': node_features,1155 'edge_index': edge_index,1156 'num_nodes': len(nodes)1157 }1158 1159 def __len__(self) -> int:1160 """Return the number of sequential training pairs."""1161 return len(self.sequential_pairs)1162 1163 def __getitem__(self, idx: int) -> Dict[str, Any]:1164 """1165 Get a sequential training pair.1166 1167 Args:1168 idx: Index of the training pair1169 1170 Returns:1171 Dictionary containing partial graph and target node data1172 """1173 if idx < 0 or idx >= len(self.sequential_pairs):1174 raise IndexError(f"Index {idx} out of range for dataset of size {len(self.sequential_pairs)}")1175 1176 return self.sequential_pairs[idx]1177 1178 def get_feature_dim(self) -> int:1179 """Return the dimension of node features."""1180 return self.converter.node_encoder.vocab_size1181 1182 1183def collate_autoregressive_data(batch: List[Dict[str, Any]]) -> Dict[str, Any]:1184 """1185 Collate function for batching autoregressive training data.1186 1187 Args:1188 batch: List of sequential training pairs1189 1190 Returns:1191 Batched autoregressive training data1192 """1193 if not batch:1194 raise ValueError("Cannot collate empty batch")1195 1196 # Separate different components1197 text_descriptions = [item['text_description'] for item in batch]1198 text_embeddings = [item.get('text_embedding') for item in batch]1199 steps = [item['step'] for item in batch]1200 total_steps = [item['total_steps'] for item in batch]