Team Ai
Datasetpublic

timlawrenz/gnn-ruby-code-study

GNN Ruby Code Study Systematic study of Graph Neural Network architectures for Ruby code complexity prediction and generation. Paper: Graph Neural Networks for Ruby Code Complexity Prediction and Generation: A Systematic Architecture Study Dataset 22,452 Ruby methods parsed into AST graphs with 74-dimensional node features. Split Samples File Train 19,084 dataset/train.jsonl Validation 3,368 dataset/val.jsonl Each JSONL record contains:… See the full description on the dataset page: https://huggingface.co/datasets/timlawrenz/gnn-ruby-code-study.

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes101downloads
models.py1351 linesDownload Raw Back to src
1"""2Graph Neural Network models for Ruby code complexity prediction.3 4This module contains PyTorch Geometric models for learning from5Ruby AST structures with performance optimizations.6"""7 8import torch9import torch.nn.functional as F10from torch_geometric.nn import GCNConv, SAGEConv, GATConv, GINConv, GraphConv, global_mean_pool11from torch_geometric.data import Data, Batch12import torch_geometric13from typing import Dict14try:15    from sentence_transformers import SentenceTransformer16    SENTENCE_TRANSFORMERS_AVAILABLE = True17except ImportError:18    SENTENCE_TRANSFORMERS_AVAILABLE = False19 20# Performance optimization: Cache CUDA availability21CUDA_AVAILABLE = torch.cuda.is_available()22 23 24class RubyComplexityGNN(torch.nn.Module):25    """26    Graph Neural Network for predicting Ruby method complexity.27    28    This model uses Graph Convolutional Networks (GCN) or GraphSAGE layers29    to learn from Abstract Syntax Tree representations of Ruby methods.30    """31    32    def __init__(self, input_dim: int, hidden_dim: int = 64, num_layers: int = 3, 33                 conv_type: str = 'GCN', dropout: float = 0.1):34        """35        Initialize the GNN model.36        37        Args:38            input_dim: Dimension of input node features39            hidden_dim: Hidden layer dimension40            num_layers: Number of convolutional layers41            conv_type: Type of convolution ('GCN', 'SAGE', 'GAT', 'GIN', 'GraphConv')42            dropout: Dropout probability for regularization43        """44        super().__init__()45        46        supported = ['GCN', 'SAGE', 'GAT', 'GIN', 'GraphConv']47        if conv_type not in supported:48            raise ValueError(f"conv_type must be one of {supported}")49        50        self.num_layers = num_layers51        self.conv_type = conv_type52        self.dropout = dropout53        self.convs = torch.nn.ModuleList()54        55        def _make_conv(in_dim, out_dim):56            if conv_type == 'GCN':57                return GCNConv(in_dim, out_dim)58            elif conv_type == 'SAGE':59                return SAGEConv(in_dim, out_dim)60            elif conv_type == 'GAT':61                return GATConv(in_dim, out_dim, heads=1)62            elif conv_type == 'GIN':63                mlp = torch.nn.Sequential(64                    torch.nn.Linear(in_dim, out_dim),65                    torch.nn.ReLU(),66                    torch.nn.Linear(out_dim, out_dim),67                )68                return GINConv(mlp)69            elif conv_type == 'GraphConv':70                return GraphConv(in_dim, out_dim)71        72        # First layer73        self.convs.append(_make_conv(input_dim, hidden_dim))74        75        # Hidden layers76        for _ in range(num_layers - 2):77            self.convs.append(_make_conv(hidden_dim, hidden_dim))78            79        # Last layer80        if num_layers > 1:81            self.convs.append(_make_conv(hidden_dim, hidden_dim))82        83        # Output layer for complexity prediction84        self.predictor = torch.nn.Linear(hidden_dim, 1)85        86    def forward(self, data: Data, return_embedding: bool = False) -> torch.Tensor:87        """88        Forward pass through the network.89        90        Args:91            data: PyTorch Geometric Data object containing graph92            return_embedding: If True, return graph embedding instead of prediction93            94        Returns:95            Complexity prediction tensor of shape (batch_size, 1) or96            Graph embedding tensor of shape (batch_size, hidden_dim) if return_embedding=True97        """98        x, edge_index, batch = data.x, data.edge_index, data.batch99        100        # Apply convolution layers with ReLU activation and dropout101        for i, conv in enumerate(self.convs):102            x = conv(x, edge_index)103            if i < len(self.convs) - 1:  # No activation after last layer104                x = F.relu(x)105                x = F.dropout(x, p=self.dropout, training=self.training)106        107        # Global pooling to get graph-level representation108        embedding = global_mean_pool(x, batch)109        110        if return_embedding:111            return embedding112        113        # Predict complexity114        return self.predictor(embedding)115    116    def get_model_info(self) -> str:117        """118        Get information about the model configuration.119        120        Returns:121            String describing the model architecture122        """123        return (f"RubyComplexityGNN({self.conv_type}, "124                f"layers={self.num_layers}, "125                f"dropout={self.dropout})")126 127 128class ASTDecoder(torch.nn.Module):129    """130    GNN-based decoder for reconstructing Abstract Syntax Trees from embeddings.131    132    This module takes a graph embedding and autoregressively generates node features133    and edge structure to reconstruct an AST.134    """135    136    def __init__(self, embedding_dim: int, output_node_dim: int, hidden_dim: int = 256, 137                 num_layers: int = 5, max_nodes: int = 100, conv_type: str = 'GCN',138                 gradient_checkpointing: bool = False):139        """140        Initialize the AST decoder.141        142        Args:143            embedding_dim: Dimension of input graph embedding144            output_node_dim: Dimension of output node features145            hidden_dim: Hidden layer dimension for GNN layers.146            num_layers: Number of decoder GNN layers.147            max_nodes: Maximum number of nodes to generate.148            conv_type: The type of GNN layer to use ('GCN', 'SAGE', 'GAT', 'GIN', 'GraphConv').149            gradient_checkpointing: Whether to use gradient checkpointing for memory efficiency.150        """151        super().__init__()152        153        self.embedding_dim = embedding_dim154        self.output_node_dim = output_node_dim155        self.hidden_dim = hidden_dim156        self.num_layers = num_layers157        self.max_nodes = max_nodes158        self.gradient_checkpointing = gradient_checkpointing159        160        self.embedding_transform = torch.nn.Linear(embedding_dim, hidden_dim)161        162        self.convs = torch.nn.ModuleList()163        current_dim = hidden_dim164 165        for i in range(num_layers):166            if conv_type == 'GAT':167                heads = 4168                conv = GATConv(current_dim, hidden_dim, heads=heads)169                current_dim = hidden_dim * heads170            elif conv_type == 'GIN':171                mlp = torch.nn.Sequential(172                    torch.nn.Linear(current_dim, current_dim),173                    torch.nn.ReLU(),174                    torch.nn.Linear(current_dim, current_dim)175                )176                conv = GINConv(mlp)177            elif conv_type == 'SAGE':178                conv = SAGEConv(current_dim, current_dim)179            elif conv_type == 'GCN':180                conv = GCNConv(current_dim, current_dim)181            elif conv_type == 'GraphConv':182                conv = GraphConv(current_dim, current_dim)183            else:184                raise ValueError(f"Unsupported conv_type: {conv_type}")185            186            self.convs.append(conv)187 188        self.node_output = torch.nn.Linear(current_dim, output_node_dim)189        self.parent_predictor = torch.nn.Linear(current_dim, max_nodes)190        191    def forward(self, embedding: torch.Tensor, num_nodes_per_graph: torch.Tensor) -> dict:192        """193        Forward pass to decode a batch of embeddings into AST structures.194        195        Args:196            embedding: Graph embedding tensor of shape [batch_size, embedding_dim].197            num_nodes_per_graph: Tensor of shape [batch_size] with the number of nodes for each graph.198            199        Returns:200            Dictionary containing batched node features and parent predictions.201        """202        batch_size = embedding.size(0)203        device = embedding.device204        205        # Use torch.repeat_interleave to expand each graph's embedding206        # to match the number of nodes in that graph.207        # This is the core of the batch-aware processing.208        node_features = self.embedding_transform(embedding)209        node_features = node_features.repeat_interleave(num_nodes_per_graph, dim=0)210 211        # Vectorized edge construction for sequential edges within each graph.212        # This approach avoids loops over graphs in the batch, creating all edges213        # at once for efficiency.214        num_edges_per_graph = torch.clamp(num_nodes_per_graph - 1, min=0)215        total_edges = torch.sum(num_edges_per_graph).item()216 217        if total_edges == 0:218            edge_index = torch.empty((2, 0), dtype=torch.long, device=device)219        else:220            # Calculate node offsets for each graph221            node_offsets = torch.cat([torch.zeros(1, device=device, dtype=num_nodes_per_graph.dtype), 222                                    torch.cumsum(num_nodes_per_graph[:-1], dim=0)])223            224            # Efficient edge index computation for sequential nodes225            # Pre-allocate tensors to avoid repeated allocations226            total_edges = num_edges_per_graph.sum().item()227            228            # Determine which graph each edge belongs to229            graph_indices = torch.repeat_interleave(torch.arange(len(num_nodes_per_graph), device=device), num_edges_per_graph)230            231            # Calculate the starting edge index for each graph232            edge_offsets = torch.cat([torch.zeros(1, device=device, dtype=num_edges_per_graph.dtype), 233                                    torch.cumsum(num_edges_per_graph[:-1], dim=0)])234            235            # Compute local (within-graph) source indices more efficiently236            src_in_graph = torch.arange(total_edges, device=device) - edge_offsets[graph_indices]237            238            # Get the starting node index for each edge's graph239            edge_node_offsets = node_offsets[graph_indices]240            241            # Compute global source and destination indices242            src = edge_node_offsets + src_in_graph243            dst = src + 1244            245            edge_index = torch.stack([src, dst], dim=0)246 247        # GNNs are typically undirected, so we add reverse edges.248        edge_index = torch_geometric.utils.to_undirected(edge_index)249 250        # Apply GNN layers with optional gradient checkpointing251        x = node_features252        if self.gradient_checkpointing and self.training:253            # Use gradient checkpointing for memory efficiency during training254            def create_custom_forward(module):255                def custom_forward(*inputs):256                    return module(*inputs)257                return custom_forward258            259            for conv in self.convs:260                x = torch.utils.checkpoint.checkpoint(261                    create_custom_forward(conv), x, edge_index, use_reentrant=False262                )263                x = F.relu(x)264        else:265            # Standard forward pass266            for conv in self.convs:267                x = conv(x, edge_index)268                x = F.relu(x)  # In-place for memory efficiency269        270        # Predict the final node features and parent logits for all nodes in the batch.271        output_node_features = self.node_output(x)272        parent_logits = self.parent_predictor(x)273        274        return {275            'node_features': output_node_features, # Shape: [total_nodes, feature_dim]276            'parent_logits': parent_logits         # Shape: [total_nodes, max_nodes]277        }278 279 280class TreeAwareASTDecoder(torch.nn.Module):281    """282    Tree-topology-aware AST decoder.283 284    Unlike ASTDecoder which constructs sequential chain edges (0→1→2→…),285    this decoder uses the actual AST tree structure for GNN message passing.286 287    Three edge modes:288      - 'chain':          Legacy sequential edges (same as ASTDecoder).289      - 'teacher_forced': Uses ground-truth AST edges during training.290      - 'iterative':      Two-pass: chain edges → predict parents → rebuild291                          tree edges → refine predictions.  Fully feed-forward.292    """293 294    def __init__(self, embedding_dim: int, output_node_dim: int,295                 hidden_dim: int = 256, num_layers: int = 5,296                 max_nodes: int = 100, conv_type: str = 'GCN',297                 edge_mode: str = 'teacher_forced',298                 gradient_checkpointing: bool = False):299        super().__init__()300        self.embedding_dim = embedding_dim301        self.output_node_dim = output_node_dim302        self.hidden_dim = hidden_dim303        self.num_layers = num_layers304        self.max_nodes = max_nodes305        self.edge_mode = edge_mode306        self.gradient_checkpointing = gradient_checkpointing307 308        self.embedding_transform = torch.nn.Linear(embedding_dim, hidden_dim)309 310        # Primary GNN stack311        self.convs = torch.nn.ModuleList()312        current_dim = hidden_dim313        for _ in range(num_layers):314            conv, current_dim = self._make_conv(conv_type, current_dim, hidden_dim)315            self.convs.append(conv)316 317        self.node_output = torch.nn.Linear(current_dim, output_node_dim)318        self.parent_predictor = torch.nn.Linear(current_dim, max_nodes)319 320        # Refinement GNN stack (only used in iterative mode)321        if edge_mode == 'iterative':322            self.refine_convs = torch.nn.ModuleList()323            ref_dim = current_dim324            for _ in range(max(num_layers // 2, 1)):325                conv, ref_dim = self._make_conv(conv_type, ref_dim, hidden_dim)326                self.refine_convs.append(conv)327            self.refine_node_output = torch.nn.Linear(ref_dim, output_node_dim)328            self.refine_parent_predictor = torch.nn.Linear(ref_dim, max_nodes)329 330    @staticmethod331    def _make_conv(conv_type: str, in_dim: int, hidden_dim: int):332        if conv_type == 'GAT':333            heads = 4334            return GATConv(in_dim, hidden_dim, heads=heads), hidden_dim * heads335        elif conv_type == 'GIN':336            mlp = torch.nn.Sequential(337                torch.nn.Linear(in_dim, in_dim),338                torch.nn.ReLU(),339                torch.nn.Linear(in_dim, in_dim),340            )341            return GINConv(mlp), in_dim342        elif conv_type == 'SAGE':343            return SAGEConv(in_dim, in_dim), in_dim344        elif conv_type == 'GCN':345            return GCNConv(in_dim, in_dim), in_dim346        elif conv_type == 'GraphConv':347            return GraphConv(in_dim, in_dim), in_dim348        else:349            raise ValueError(f"Unsupported conv_type: {conv_type}")350 351    # ------------------------------------------------------------------352    # Edge construction helpers353    # ------------------------------------------------------------------354 355    @staticmethod356    def _build_chain_edges(num_nodes_per_graph: torch.Tensor) -> torch.Tensor:357        """Build sequential chain edges (legacy behaviour)."""358        device = num_nodes_per_graph.device359        num_edges_per_graph = torch.clamp(num_nodes_per_graph - 1, min=0)360        total_edges = num_edges_per_graph.sum().item()361        if total_edges == 0:362            return torch.empty((2, 0), dtype=torch.long, device=device)363 364        node_offsets = torch.cat([365            torch.zeros(1, device=device, dtype=num_nodes_per_graph.dtype),366            torch.cumsum(num_nodes_per_graph[:-1], dim=0),367        ])368        graph_indices = torch.repeat_interleave(369            torch.arange(len(num_nodes_per_graph), device=device),370            num_edges_per_graph,371        )372        edge_offsets = torch.cat([373            torch.zeros(1, device=device, dtype=num_edges_per_graph.dtype),374            torch.cumsum(num_edges_per_graph[:-1], dim=0),375        ])376        src_in_graph = torch.arange(total_edges, device=device) - edge_offsets[graph_indices]377        edge_node_offsets = node_offsets[graph_indices]378        src = edge_node_offsets + src_in_graph379        dst = src + 1380        return torch.stack([src, dst], dim=0)381 382    @staticmethod383    def _parents_to_edges(parent_logits: torch.Tensor,384                          num_nodes_per_graph: torch.Tensor) -> torch.Tensor:385        """Convert per-node parent logits to a hard edge_index (argmax)."""386        device = parent_logits.device387        total_nodes = parent_logits.size(0)388        max_nodes = parent_logits.size(1)389 390        # Compute graph membership and node offsets391        batch_vec = torch.repeat_interleave(392            torch.arange(len(num_nodes_per_graph), device=device),393            num_nodes_per_graph,394        )395        node_offsets = torch.cat([396            torch.zeros(1, device=device, dtype=num_nodes_per_graph.dtype),397            torch.cumsum(num_nodes_per_graph[:-1], dim=0),398        ])399 400        # Mask out logits for positions beyond each graph's node count401        mask = torch.arange(max_nodes, device=device).unsqueeze(0).expand(total_nodes, -1)402        graph_sizes = num_nodes_per_graph[batch_vec].unsqueeze(1)403        parent_logits = parent_logits.clone()404        parent_logits[mask >= graph_sizes] = float('-inf')405 406        # Local parent index → global parent index407        local_parent = parent_logits.argmax(dim=1)  # [total_nodes]408        global_parent = local_parent + node_offsets[batch_vec]409 410        # Node 0 of each graph (the root) has no parent — remove those edges411        local_idx = torch.arange(total_nodes, device=device) - node_offsets[batch_vec]412        is_root = local_idx == 0413        src = global_parent[~is_root]414        dst = torch.arange(total_nodes, device=device)[~is_root]415        return torch.stack([src, dst], dim=0).long()416 417    # ------------------------------------------------------------------418    # Forward419    # ------------------------------------------------------------------420 421    def _apply_convs(self, x, edge_index, convs):422        edge_index = torch_geometric.utils.to_undirected(edge_index)423        if self.gradient_checkpointing and self.training:424            def _make_fn(module):425                def fn(*inputs):426                    return module(*inputs)427                return fn428            for conv in convs:429                x = torch.utils.checkpoint.checkpoint(430                    _make_fn(conv), x, edge_index, use_reentrant=False,431                )432                x = F.relu(x)433        else:434            for conv in convs:435                x = conv(x, edge_index)436                x = F.relu(x)437        return x438 439    def forward(self, embedding: torch.Tensor,440                num_nodes_per_graph: torch.Tensor,441                gt_edge_index: torch.Tensor | None = None) -> dict:442        """443        Args:444            embedding: [batch_size, embedding_dim]445            num_nodes_per_graph: [batch_size]446            gt_edge_index: [2, num_edges] ground-truth AST edges (optional).447                           Required for teacher_forced mode during training.448        """449        device = embedding.device450        node_features = self.embedding_transform(embedding)451        node_features = node_features.repeat_interleave(num_nodes_per_graph, dim=0)452 453        # ---- choose edges for the first GNN pass ----454        if self.edge_mode == 'teacher_forced' and gt_edge_index is not None:455            first_pass_edges = gt_edge_index456        else:457            first_pass_edges = self._build_chain_edges(num_nodes_per_graph)458 459        x = self._apply_convs(node_features, first_pass_edges, self.convs)460        output_node_features = self.node_output(x)461        parent_logits = self.parent_predictor(x)462 463        # ---- optional second (refinement) pass ----464        if self.edge_mode == 'iterative':465            predicted_edges = self._parents_to_edges(parent_logits, num_nodes_per_graph)466            if predicted_edges.size(1) > 0:467                x2 = self._apply_convs(x, predicted_edges, self.refine_convs)468                output_node_features = self.refine_node_output(x2)469                parent_logits = self.refine_parent_predictor(x2)470 471        return {472            'node_features': output_node_features,473            'parent_logits': parent_logits,474        }475 476 477class AutoregressiveASTDecoder(torch.nn.Module):478    """479    Autoregressive decoder for generating Abstract Syntax Trees sequentially.480    481    This decoder generates AST nodes one by one, maintaining state across generation482    steps and considering both text description and current partial graph context.483    """484    485    def __init__(self, 486                 text_embedding_dim: int = 64,487                 graph_hidden_dim: int = 64,488                 state_hidden_dim: int = 128,489                 node_types: int = 74,490                 max_nodes: int = 100,491                 sequence_model: str = 'GRU'):  # Options: 'GRU', 'LSTM', 'Transformer'492        """493        Initialize the AutoregressiveASTDecoder.494        495        Args:496            text_embedding_dim: Dimension of text embeddings (from alignment model)497            graph_hidden_dim: Hidden dimension for graph encoding498            state_hidden_dim: Hidden dimension for sequential state499            node_types: Number of possible node types (also node feature dimension)500            max_nodes: Maximum number of nodes for connection prediction501            sequence_model: Type of sequence model ('GRU', 'LSTM', 'Transformer')502        """503        super().__init__()504        505        self.text_embedding_dim = text_embedding_dim506        self.graph_hidden_dim = graph_hidden_dim507        self.state_hidden_dim = state_hidden_dim508        self.node_types = node_types509        self.max_nodes = max_nodes510        self.sequence_model = sequence_model511        512        # Graph Context Encoder - GNN for processing partial graph structure513        # Note: Node features are node_types dimensional (one-hot encoded)514        self.graph_gnn_layers = torch.nn.ModuleList([515            GCNConv(node_types, graph_hidden_dim),516            GCNConv(graph_hidden_dim, graph_hidden_dim)517        ])518        self.graph_layer_norm = torch.nn.LayerNorm(graph_hidden_dim)519        self.graph_dropout = torch.nn.Dropout(0.1)520        521        # Sequential State Encoder - maintains state across generation steps522        input_size = text_embedding_dim + graph_hidden_dim523        524        if sequence_model == 'GRU':525            self.state_encoder = torch.nn.GRU(526                input_size=input_size,527                hidden_size=state_hidden_dim,528                num_layers=2,529                batch_first=True,530                dropout=0.1531            )532        elif sequence_model == 'LSTM':533            self.state_encoder = torch.nn.LSTM(534                input_size=input_size,535                hidden_size=state_hidden_dim,536                num_layers=2,537                batch_first=True,538                dropout=0.1539            )540        elif sequence_model == 'Transformer':541            # For transformer, we'll use a transformer encoder layer542            encoder_layer = torch.nn.TransformerEncoderLayer(543                d_model=state_hidden_dim,544                nhead=8,545                dim_feedforward=256,546                dropout=0.1,547                batch_first=True548            )549            self.state_encoder = torch.nn.TransformerEncoder(550                encoder_layer=encoder_layer,551                num_layers=4552            )553            # For transformer, we need to project input to state_hidden_dim554            self.input_projection = torch.nn.Linear(input_size, state_hidden_dim)555        else:556            raise ValueError(f"Unknown sequence model: {sequence_model}. Choose from 'GRU', 'LSTM', 'Transformer'")557        558        # Dual Prediction Heads559        560        # Predict next node type561        self.node_type_predictor = torch.nn.Linear(state_hidden_dim, node_types)562        563        # Predict connection to existing nodes564        self.connection_predictor = torch.nn.Sequential(565            torch.nn.Linear(state_hidden_dim, max_nodes),566            torch.nn.Sigmoid()  # Probability of connection to each existing node567        )568        569    def forward(self, text_embedding, partial_graph=None, hidden_state=None):570        """571        Forward pass for autoregressive AST generation.572        573        Args:574            text_embedding: (batch_size, text_embedding_dim) - Text description embedding575            partial_graph: Dict with keys 'x', 'edge_index', 'batch' - Current partial AST (optional)576            hidden_state: Previous hidden state for sequence model (optional)577            578        Returns:579            Dictionary containing:580                - node_type_logits: (batch_size, node_types) - Probabilities for next node type581                - connection_probs: (batch_size, max_nodes) - Connection probabilities582                - hidden_state: Updated hidden state583        """584        batch_size = text_embedding.size(0)585        device = text_embedding.device586        587        # 1. Encode current graph state using GNN588        if partial_graph is not None and 'x' in partial_graph and len(partial_graph['x']) > 0:589            # We have a non-empty partial graph - process it with GNN590            591            # Convert partial graph to tensor if needed592            graph_features = partial_graph['x']593            if isinstance(graph_features, list):594                # Convert list of features to tensor595                if graph_features and isinstance(graph_features[0], list):596                    graph_features = torch.tensor(graph_features, dtype=torch.float32, device=device)597                else:598                    # Empty or malformed graph599                    graph_encoded = torch.zeros(batch_size, self.graph_hidden_dim, device=device)600            else:601                graph_features = graph_features.to(device)602            603            if len(graph_features.shape) == 2 and graph_features.size(0) > 0:604                # Get edge information for GNN processing605                edge_index = partial_graph.get('edge_index', None)606                if edge_index is None:607                    # Create simple sequential edges if no edges provided608                    num_nodes = graph_features.size(0)609                    if num_nodes > 1:610                        edge_list = []611                        for i in range(num_nodes - 1):612                            edge_list.extend([[i, i + 1], [i + 1, i]])  # Bidirectional edges613                        edge_index = torch.tensor(edge_list, dtype=torch.long, device=device).t()614                    else:615                        # Single node - no edges616                        edge_index = torch.empty((2, 0), dtype=torch.long, device=device)617                else:618                    if isinstance(edge_index, list):619                        edge_index = torch.tensor(edge_index, dtype=torch.long, device=device)620                    else:621                        edge_index = edge_index.to(device)622                623                # Apply GNN layers for structural encoding624                x = graph_features625                for i, gnn_layer in enumerate(self.graph_gnn_layers):626                    x = gnn_layer(x, edge_index)627                    if i < len(self.graph_gnn_layers) - 1:  # Apply activation for all but last layer628                        x = F.relu(x)629                        x = self.graph_dropout(x)630                631                # Apply layer normalization to final GNN output632                x = self.graph_layer_norm(x)633                634                # Global pooling to get graph-level representation per batch635                if 'batch' in partial_graph and partial_graph['batch'] is not None:636                    # Use batch indices for proper pooling637                    batch_indices = partial_graph['batch']638                    if isinstance(batch_indices, list):639                        batch_indices = torch.tensor(batch_indices, dtype=torch.long, device=device)640                    else:641                        batch_indices = batch_indices.to(device)642                    643                    # Use global_mean_pool for proper batched pooling644                    graph_encoded = global_mean_pool(x, batch_indices, size=batch_size)645                    646                    # Ensure we have the right batch size647                    if graph_encoded.size(0) < batch_size:648                        # Pad with zeros for missing batches649                        padding = torch.zeros(batch_size - graph_encoded.size(0), self.graph_hidden_dim, device=device)650                        graph_encoded = torch.cat([graph_encoded, padding], dim=0)651                    elif graph_encoded.size(0) > batch_size:652                        # Trim if too many653                        graph_encoded = graph_encoded[:batch_size]654                else:655                    # Single graph case - use mean pooling656                    graph_encoded = x.mean(dim=0).unsqueeze(0).expand(batch_size, -1)657            else:658                # Unexpected shape or empty, use zeros659                graph_encoded = torch.zeros(batch_size, self.graph_hidden_dim, device=device)660        else:661            # Empty graph - start with zero representation662            graph_encoded = torch.zeros(batch_size, self.graph_hidden_dim, device=device)663        664        # 2. Combine text and graph context665        combined_input = torch.cat([text_embedding, graph_encoded], dim=-1)666        667        # 3. Update sequential state668        if self.sequence_model == 'Transformer':669            # For transformer, project input and treat as sequence670            sequence_input = self.input_projection(combined_input.unsqueeze(1))  # (batch_size, 1, state_hidden_dim)671            sequence_output = self.state_encoder(sequence_input)  # (batch_size, 1, state_hidden_dim)672            sequence_output = sequence_output.squeeze(1)  # (batch_size, state_hidden_dim)673            new_hidden_state = None  # Transformers don't maintain hidden state in the same way674        else:675            # For RNN/GRU/LSTM676            sequence_input = combined_input.unsqueeze(1)  # (batch_size, 1, input_size)677            sequence_output, new_hidden_state = self.state_encoder(sequence_input, hidden_state)678            sequence_output = sequence_output.squeeze(1)  # (batch_size, state_hidden_dim)679        680        # 4. Predict next step681        node_type_logits = self.node_type_predictor(sequence_output)682        connection_probs = self.connection_predictor(sequence_output)683        684        return {685            'node_type_logits': node_type_logits,686            'connection_probs': connection_probs,687            'hidden_state': new_hidden_state688        }689    690    def get_model_info(self) -> str:691        """692        Get information about the autoregressive decoder configuration.693        694        Returns:695            String describing the model architecture696        """697        return (f"AutoregressiveASTDecoder(\n"698                f"  text_dim={self.text_embedding_dim}, "699                f"  graph_dim={self.graph_hidden_dim}, "700                f"  state_dim={self.state_hidden_dim}\n"701                f"  node_types={self.node_types}, "702                f"  sequence_model={self.sequence_model}\n"703                f")")704 705 706class ASTAutoencoder(torch.nn.Module):707    """708    Autoencoder for Abstract Syntax Trees using Graph Neural Networks.709    710    Combines the existing RubyComplexityGNN (as encoder) with the new ASTDecoder711    to create an autoencoder that can reconstruct ASTs from learned embeddings.712    """713    714    def __init__(self, encoder_input_dim: int, node_output_dim: int, 715                 hidden_dim: int = 64, num_layers: int = 3, 716                 conv_type: str = 'GCN', dropout: float = 0.1,717                 freeze_encoder: bool = False, encoder_weights_path: str = None,718                 max_nodes: int = 100, decoder_conv_type: str = 'GCN',719                 gradient_checkpointing: bool = False,720                 decoder_edge_mode: str = 'chain'):721        """722        Initialize the AST autoencoder.723        724        Args:725            encoder_input_dim: Input dimension for encoder (node feature dimension)726            node_output_dim: Output dimension for decoder node features727            hidden_dim: Hidden dimension for both encoder and decoder728            num_layers: Number of layers in both encoder and decoder729            conv_type: Type of convolution for encoder ('GCN' or 'SAGE')730            dropout: Dropout rate for encoder731            freeze_encoder: Whether to freeze encoder weights732            encoder_weights_path: Path to pre-trained encoder weights733            max_nodes: Maximum number of nodes for the decoder.734            decoder_conv_type: The GNN layer type for the decoder.735            gradient_checkpointing: Whether to enable gradient checkpointing for memory efficiency.736            decoder_edge_mode: Edge construction strategy for the decoder.737                'chain' uses the original ASTDecoder with sequential edges.738                'teacher_forced' or 'iterative' uses TreeAwareASTDecoder.739        """740        super().__init__()741        742        self.decoder_edge_mode = decoder_edge_mode743        # Initialize encoder (RubyComplexityGNN without prediction head)744        self.encoder = RubyComplexityGNN(745            input_dim=encoder_input_dim,746            hidden_dim=hidden_dim,747            num_layers=num_layers,748            conv_type=conv_type,749            dropout=dropout750        )751        752        # Load pre-trained weights if provided and adjust encoder config if needed753        self.encoder_weights_path = encoder_weights_path754        if encoder_weights_path is not None:755            try:756                checkpoint = torch.load(encoder_weights_path, map_location='cpu', weights_only=True)757                # Check if checkpoint contains model config and use it to create compatible encoder758                if 'model_config' in checkpoint:759                    saved_config = checkpoint['model_config']760                    # Recreate encoder with saved configuration if it differs from current761                    if (saved_config.get('conv_type', conv_type) != conv_type or762                        saved_config.get('hidden_dim', hidden_dim) != hidden_dim or763                        saved_config.get('num_layers', num_layers) != num_layers or764                        saved_config.get('dropout', dropout) != dropout):765                        print(f"Adjusting encoder config to match saved model: conv_type={saved_config.get('conv_type', conv_type)}")766                        self.encoder = RubyComplexityGNN(767                            input_dim=encoder_input_dim,768                            hidden_dim=saved_config.get('hidden_dim', hidden_dim),769                            num_layers=saved_config.get('num_layers', num_layers),770                            conv_type=saved_config.get('conv_type', conv_type),771                            dropout=saved_config.get('dropout', dropout)772                        )773                        # Update hidden_dim for decoder compatibility774                        hidden_dim = saved_config.get('hidden_dim', hidden_dim)775                776                self.encoder.load_state_dict(checkpoint['model_state_dict'])777                print(f"Loaded encoder weights from {encoder_weights_path}")778            except FileNotFoundError:779                print(f"Warning: Could not find encoder weights at {encoder_weights_path}")780            except Exception as e:781                print(f"Warning: Could not load encoder weights: {e}")782        783        # Freeze encoder if requested784        if freeze_encoder:785            for param in self.encoder.parameters():786                param.requires_grad = False787            print("Encoder weights frozen")788        789        # Initialize decoder790        if decoder_edge_mode in ('teacher_forced', 'iterative'):791            self.decoder = TreeAwareASTDecoder(792                embedding_dim=hidden_dim,793                output_node_dim=node_output_dim,794                hidden_dim=hidden_dim,795                num_layers=num_layers,796                max_nodes=max_nodes,797                conv_type=decoder_conv_type,798                edge_mode=decoder_edge_mode,799                gradient_checkpointing=gradient_checkpointing,800            )801        else:802            self.decoder = ASTDecoder(803                embedding_dim=hidden_dim,804                output_node_dim=node_output_dim,805                hidden_dim=hidden_dim,806                num_layers=num_layers,807                max_nodes=max_nodes,808                conv_type=decoder_conv_type,809                gradient_checkpointing=gradient_checkpointing,810            )811        812        self.hidden_dim = hidden_dim813        self.freeze_encoder = freeze_encoder814        815    def forward(self, data: Data) -> dict:816        """817        Forward pass through the autoencoder.818        819        Args:820            data: PyTorch Geometric Data object containing a batch of input ASTs.821            822        Returns:823            Dictionary containing reconstructed AST information for the batch.824        """825        # Encode: Batch of ASTs -> Batch of embeddings826        embedding = self.encoder(data, return_embedding=True)827        828        # Get the number of nodes in each graph of the batch829        num_nodes_per_graph = torch.bincount(data.batch)830        831        # Decode: Batch of embeddings -> Batch of reconstructed ASTs832        # Pass ground-truth edges for tree-aware decoders833        if self.decoder_edge_mode != 'chain':834            reconstruction = self.decoder(835                embedding, num_nodes_per_graph,836                gt_edge_index=data.edge_index,837            )838        else:839            reconstruction = self.decoder(embedding, num_nodes_per_graph)840        841        return {842            'embedding': embedding,843            'reconstruction': reconstruction844        }845    846    def get_model_info(self) -> str:847        """848        Get information about the autoencoder configuration.849        850        Returns:851            String describing the model architecture852        """853        encoder_info = self.encoder.get_model_info()854        decoder_info = f"ASTDecoder(embedding_dim={self.hidden_dim})"855        freeze_status = " [FROZEN]" if self.freeze_encoder else ""856        857        return (f"ASTAutoencoder(\n"858                f"  encoder: {encoder_info}{freeze_status}\n"859                f"  decoder: {decoder_info}\n"860                f")")861 862 863class SimpleTextEncoder(torch.nn.Module):864    """865    Simple text encoder as fallback when sentence-transformers is not available.866    867    This provides a basic text encoding mechanism using character-level features868    and a simple neural network. Used as fallback for testing when internet869    access is not available.870    """871    872    def __init__(self, output_dim: int = 384, max_length: int = 100):873        """874        Initialize the simple text encoder.875        876        Args:877            output_dim: Output embedding dimension878            max_length: Maximum text length to consider879        """880        super().__init__()881        self.output_dim = output_dim882        self.max_length = max_length883        884        # Character embedding (256 ASCII characters)885        self.char_embedding = torch.nn.Embedding(256, 64)886        887        # Simple RNN for text processing888        self.rnn = torch.nn.LSTM(64, 128, batch_first=True, bidirectional=True)889        890        # Output projection891        self.output_proj = torch.nn.Linear(256, output_dim)892        893    def encode(self, texts: list, convert_to_tensor: bool = True) -> torch.Tensor:894        """895        Encode texts to embeddings.896        897        Args:898            texts: List of text strings899            convert_to_tensor: Whether to return tensor (for compatibility)900            901        Returns:902            Text embeddings tensor903        """904        batch_size = len(texts)905        906        # Convert texts to character indices907        char_sequences = []908        for text in texts:909            # Convert to lowercase and get character codes910            chars = [min(ord(c), 255) for c in text.lower()[:self.max_length]]911            # Pad to max_length912            chars.extend([0] * (self.max_length - len(chars)))913            char_sequences.append(chars[:self.max_length])914        915        # Convert to tensor and move to same device as model916        char_tensor = torch.tensor(char_sequences, dtype=torch.long)917        char_tensor = char_tensor.to(next(self.parameters()).device)918        919        # Embed characters920        embedded = self.char_embedding(char_tensor)  # (batch, seq_len, embed_dim)921        922        # Process with RNN923        rnn_output, (hidden, _) = self.rnn(embedded)924        925        # Use last hidden state (concatenated forward and backward)926        final_hidden = torch.cat([hidden[0], hidden[1]], dim=1)  # (batch, 256)927        928        # Project to output dimension929        embeddings = self.output_proj(final_hidden)930        931        return embeddings932    933    def get_sentence_embedding_dimension(self) -> int:934        """Get embedding dimension for compatibility."""935        return self.output_dim936 937 938class AlignmentModel(torch.nn.Module):939    """940    Dual-encoder model for aligning text descriptions with code embeddings.941    942    This model combines a frozen RubyComplexityGNN (code encoder) with a 943    sentence-transformers text encoder to create aligned embeddings in the944    same 64-dimensional space.945    """946    947    def __init__(self, input_dim: int, hidden_dim: int = 64, num_layers: int = 3,948                 conv_type: str = 'GCN', dropout: float = 0.1,949                 text_model_name: str = 'all-MiniLM-L6-v2',950                 code_encoder_weights_path: str = 'models/best_model.pt'):951        """952        Initialize the alignment model.953        954        Args:955            input_dim: Input dimension for code encoder (node feature dimension)956            hidden_dim: Hidden dimension for both encoders (default: 64)957            num_layers: Number of layers in code encoder958            conv_type: Type of convolution for code encoder ('GCN' or 'SAGE')959            dropout: Dropout rate for code encoder960            text_model_name: Name of the sentence-transformers model to use961            code_encoder_weights_path: Path to pre-trained code encoder weights (default: 'models/best_encoder_model.pt')962        """963        super().__init__()964        965        self.hidden_dim = hidden_dim966        967        # Initialize frozen code encoder (RubyComplexityGNN without prediction head)968        self.code_encoder = RubyComplexityGNN(969            input_dim=input_dim,970            hidden_dim=hidden_dim,971            num_layers=num_layers,972            conv_type=conv_type,973            dropout=dropout974        )975        976        # Load pre-trained weights if provided977        if code_encoder_weights_path is not None:978            try:979                checkpoint = torch.load(code_encoder_weights_path, map_location='cpu', weights_only=True)980                # Handle both direct state dict and checkpoint format981                if 'model_state_dict' in checkpoint:982                    state_dict = checkpoint['model_state_dict']983                else:984                    state_dict = checkpoint985                986                # Load state dict, ignoring predictor weights if present987                model_state = {}988                for key, value in state_dict.items():989                    if not key.startswith('predictor'):990                        model_state[key] = value991                992                self.code_encoder.load_state_dict(model_state, strict=False)993                print(f"Loaded code encoder weights from {code_encoder_weights_path}")994            except FileNotFoundError:995                print(f"Warning: Could not find code encoder weights at {code_encoder_weights_path}")996            except Exception as e:997                print(f"Warning: Could not load code encoder weights: {e}")998        999        # Freeze code encoder parameters1000        for param in self.code_encoder.parameters():1001            param.requires_grad = False1002        print("Code encoder weights frozen")1003        1004        # Initialize text encoder1005        if SENTENCE_TRANSFORMERS_AVAILABLE:1006            try:1007                self.text_encoder = SentenceTransformer(text_model_name)1008                self.text_encoder_type = "sentence_transformers"1009                print(f"Using SentenceTransformer: {text_model_name}")1010            except Exception as e:1011                print(f"Warning: Could not load SentenceTransformer ({e}), using fallback")1012                self.text_encoder = SimpleTextEncoder(output_dim=384)1013                self.text_encoder_type = "simple"1014        else:1015            print("SentenceTransformers not available, using simple text encoder")1016            self.text_encoder = SimpleTextEncoder(output_dim=384)1017            self.text_encoder_type = "simple"1018        1019        # Get text encoder output dimension1020        text_dim = self.text_encoder.get_sentence_embedding_dimension()1021        1022        # Projection head to align text embeddings to code embedding space1023        # Small MLP for better capacity: Linear(384 -> 256) -> ReLU() -> Linear(256 -> 64)1024        self.text_projection = torch.nn.Sequential(1025            torch.nn.Linear(text_dim, 256),1026            torch.nn.ReLU(),1027            torch.nn.Linear(256, hidden_dim)1028        )1029        1030        print(f"Text encoder output dim: {text_dim}, projecting to: {hidden_dim}")1031        1032    def encode_code(self, data: Data) -> torch.Tensor:1033        """1034        Encode graph data to embeddings using the frozen code encoder.1035        1036        Args:1037            data: PyTorch Geometric Data object containing graph1038            1039        Returns:1040            Code embeddings tensor of shape (batch_size, hidden_dim)1041        """1042        with torch.no_grad():  # Code encoder is frozen1043            return self.code_encoder(data, return_embedding=True)1044    1045    def encode_text(self, texts: list) -> torch.Tensor:1046        """1047        Encode text descriptions to embeddings using the text encoder.1048        1049        Args:1050            texts: List of text descriptions1051            1052        Returns:1053            Text embeddings tensor of shape (batch_size, hidden_dim)1054        """1055        # Get text embeddings from sentence transformer1056        text_embeddings = self.text_encoder.encode(texts, convert_to_tensor=True)1057        1058        # Clone tensor to create a normal tensor for autograd (SentenceTransformer creates inference tensors)1059        text_embeddings = text_embeddings.clone()1060        1061        # Project to code embedding space1062        projected_embeddings = self.text_projection(text_embeddings)1063        1064        return projected_embeddings1065    1066    def forward(self, data: Data, texts: list) -> dict:1067        """1068        Forward pass through both encoders.1069        1070        Args:1071            data: PyTorch Geometric Data object containing graphs1072            texts: List of text descriptions (same length as batch size)1073            1074        Returns:1075            Dictionary containing:1076                - 'code_embeddings': Code embeddings (batch_size, hidden_dim)1077                - 'text_embeddings': Text embeddings (batch_size, hidden_dim)1078        """1079        # Encode code1080        code_embeddings = self.encode_code(data)1081        1082        # Encode text1083        text_embeddings = self.encode_text(texts)1084        1085        # Ensure embeddings are on the same device1086        if code_embeddings.device != text_embeddings.device:1087            text_embeddings = text_embeddings.to(code_embeddings.device)1088        1089        return {1090            'code_embeddings': code_embeddings,1091            'text_embeddings': text_embeddings1092        }1093    1094    def get_model_info(self) -> str:1095        """1096        Get information about the alignment model configuration.1097        1098        Returns:1099            String describing the model architecture1100        """1101        code_info = self.code_encoder.get_model_info()1102        1103        if self.text_encoder_type == "sentence_transformers":1104            # Try to get model name from _model_config, fallback to transformer config, or use generic name1105            model_name = self.text_encoder._model_config.get('_name_or_path')1106            if model_name is None:1107                # Try to get from transformer module config1108                try:1109                    model_name = self.text_encoder[0].auto_model.config._name_or_path1110                except (AttributeError, IndexError):1111                    model_name = "SentenceTransformer"1112            text_info = f"SentenceTransformer({model_name})"1113        else:1114            text_info = f"SimpleTextEncoder(dim={self.text_encoder.output_dim})"1115            1116        # Handle Sequential projection (MLP) vs single Linear layer1117        if isinstance(self.text_projection, torch.nn.Sequential):1118            first_layer = self.text_projection[0]1119            last_layer = self.text_projection[2]1120            projection_info = f"MLP({first_layer.in_features} -> 256 -> {last_layer.out_features})"1121        else:1122            projection_info = f"Linear({self.text_projection.in_features} -> {self.text_projection.out_features})"1123        1124        return (f"AlignmentModel(\n"1125                f"  code_encoder: {code_info} [FROZEN]\n"1126                f"  text_encoder: {text_info}\n"1127                f"  projection: {projection_info}\n"1128                f")")1129 1130 1131class HierarchicalASTDecoder(torch.nn.Module):1132    """1133    Hierarchical, coarse-to-fine decoder for generating ASTs level by level.1134 1135    This model takes a text embedding and progressively generates an AST from the1136    root down, with each stage adding one level of depth to the tree. Uses proper1137    GNN layers to process graph structures at each level.1138    """1139 1140    def __init__(self, embedding_dim: int, hidden_dim: int, num_levels: int, node_feature_dim: int, conv_type: str = 'GCN'):1141        """1142        Initialize the HierarchicalASTDecoder.1143 1144        Args:1145            embedding_dim: Dimension of the input text embedding.1146            hidden_dim: Hidden dimension for the GNN layers.1147            num_levels: The maximum depth of the AST to generate (number of stages).1148            node_feature_dim: The dimension of the node features to be predicted.1149            conv_type: The type of GNN convolution to use ('GCN' or 'SAGE').1150        """1151        super().__init__()1152        self.embedding_dim = embedding_dim1153        self.hidden_dim = hidden_dim1154        self.num_levels = num_levels1155        self.node_feature_dim = node_feature_dim1156        self.conv_type = conv_type1157        self.register_buffer('device_indicator', torch.empty(0))1158 1159        # Select GNN layer type1160        if conv_type == 'GCN':1161            ConvLayer = GCNConv1162        elif conv_type == 'SAGE':1163            ConvLayer = SAGEConv1164        else:1165            raise ValueError(f"Unsupported conv_type: {conv_type}. Use 'GCN' or 'SAGE'.")1166 1167        # A ModuleList to hold the generator for each level of the AST.1168        self.level_generators = torch.nn.ModuleList()1169 1170        for i in range(num_levels):1171            # Level 0 takes embedding as input, subsequent levels take hidden state1172            # which has the same dimension as the output of the previous level's GNN1173            if i == 0:1174                input_dim = self.embedding_dim1175            else:1176                input_dim = self.hidden_dim1177            1178            # Each level generator uses proper GNN layers1179            level_gnn = ConvLayer(input_dim, self.hidden_dim)1180            node_predictor = torch.nn.Linear(self.hidden_dim, node_feature_dim)1181            adjacency_predictor = torch.nn.Linear(self.hidden_dim, self.hidden_dim)1182 1183            self.level_generators.append(torch.nn.ModuleDict({1184                'gnn': level_gnn,1185                'node_predictor': node_predictor,1186                'adjacency_predictor': adjacency_predictor,1187            }))1188 1189    @property1190    def device(self):1191        """Returns the device the model is on."""1192        return self.device_indicator.device1193 1194    def forward(self, input_data: Data, target_level: int) -> Dict[str, torch.Tensor]:1195        """1196        Performs a forward pass for a single level of generation.1197 1198        Args:1199            input_data: PyG Data object with node features (x) and edge indices (edge_index).1200                        For level 0, x should be the text embedding repeated for initial node(s).

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