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
0likes104downloads
data_processing.py1461 linesDownload Raw Back to src
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]

Showing the first 1,200 of 1461 lines. Download the file for the rest.