""" Self-contained molecular encoder for BioXMol. This is the molecule branch (a Gated Graph Neural Network) of the BioXMol soft-contrastive model. It maps a molecular graph to a fixed-length embedding. No part of the mocop training framework is required to run it. Architecture (as trained): n_edge=1, in_dim=75, n_conv=6, fc_dims=[1024, 128]. Embedding layers exposed by `embed(...)`: "GNN" -> 75-d (graph-level sum-pooled node features) "first_fc" -> 1024-d (after fc_layers[0]: Linear -> Tanh) "second_fc" -> 128-d (after fc_layers[1]: Linear; contrastive projection) """ import math from typing import Iterable, List import torch import torch.nn as nn import torch.nn.functional as F class GraphConvolution(nn.Module): def __init__(self, in_dim: int, out_dim: int, n_edge: int = 1, bias: bool = True): super().__init__() self.weight = nn.Parameter(torch.Tensor(n_edge, in_dim, out_dim)) if bias: self.bias = nn.Parameter(torch.Tensor(n_edge, 1, out_dim)) else: self.register_parameter("bias", None) self.reset_parameters() def reset_parameters(self): nn.init.xavier_uniform_(self.weight) if self.bias is not None: nn.init.zeros_(self.bias) def forward(self, adj: torch.Tensor, feat: torch.Tensor) -> torch.Tensor: if len(adj.size()) == 3: adj = adj.unsqueeze(1) feat = feat.unsqueeze(1) output = torch.matmul(adj, feat) output = torch.matmul(output, self.weight) if self.bias is not None: output = output + self.bias output = output.sum(dim=1) return F.relu(output) class GRU2D(nn.Module): """2D GRU cell used to gate information flow between conv layers.""" def __init__(self, in_dim: int, hidden_dim: int, bias: bool = True): super().__init__() self.x_to_intermediate = nn.Linear(in_dim, 3 * hidden_dim, bias=bias) self.h_to_intermediate = nn.Linear(in_dim, 3 * hidden_dim, bias=bias) self.reset_parameters() def reset_parameters(self): for k, v in self.state_dict().items(): if "weight" in k: std = math.sqrt(6.0 / (v.size(1) + v.size(0) / 3)) nn.init.uniform_(v, a=-std, b=std) elif "bias" in k: nn.init.zeros_(v) def forward(self, x: torch.Tensor, h_0: torch.Tensor) -> torch.Tensor: ix = self.x_to_intermediate(x) ih = self.h_to_intermediate(h_0) x_r, x_z, x_n = ix.chunk(3, -1) h_r, h_z, h_n = ih.chunk(3, -1) r = torch.sigmoid(x_r + h_r) z = torch.sigmoid(x_z + h_z) n = torch.tanh(x_n + (r * h_n)) return (1 - z) * n + z * h_0 class GatedGraphConvolution(nn.Module): def __init__(self, in_dim: int, out_dim: int, n_edge: int = 1, bias: bool = True): super().__init__() if in_dim != out_dim: raise ValueError(f"in_dim ({in_dim}) must equal out_dim ({out_dim}).") self.gc = GraphConvolution(in_dim=in_dim, out_dim=out_dim, n_edge=n_edge, bias=bias) self.gru = GRU2D(in_dim=out_dim, hidden_dim=out_dim, bias=bias) def forward(self, adj: torch.Tensor, h_0: torch.Tensor) -> torch.Tensor: h = self.gc(adj, h_0) return self.gru(h, h_0) class GatedGraphNeuralNetwork(nn.Module): def __init__( self, n_edge: int = 1, in_dim: int = 75, n_conv: int = 6, fc_dims: Iterable[int] = (1024, 128), p_dropout: float = 0.0, ): super().__init__() fc_dims = list(fc_dims) self.conv_layers = nn.ModuleList( [GatedGraphConvolution(in_dim=in_dim, out_dim=in_dim, n_edge=n_edge) for _ in range(n_conv)] ) num_fc = len(fc_dims) dims = [in_dim] + fc_dims fc_layers = [] for i, (a, b) in enumerate(zip(dims[:-1], dims[1:])): layer = nn.Linear(a, b) if i < num_fc - 2: layer = nn.Sequential(layer, nn.ReLU()) elif i == num_fc - 2: layer = nn.Sequential(layer, nn.Tanh()) fc_layers.append(layer) self.fc_layers = nn.ModuleList(fc_layers) self.dropout = nn.Dropout(p=p_dropout) def _pool(self, adj, node_feat, atom_vec): for conv in self.conv_layers: node_feat = conv(adj, node_feat) node_feat = self.dropout(node_feat) return (node_feat * atom_vec).sum(1) # graph-level, masks padded atoms def embed(self, adj, node_feat, atom_vec, layer: str = "first_fc") -> torch.Tensor: """Return embeddings from a chosen layer. Call model.eval() first.""" h = self._pool(adj, node_feat, atom_vec) if layer == "GNN": return h h = self.fc_layers[0](h) if layer == "first_fc": return h h = self.dropout(h) h = self.fc_layers[1](h) if layer == "second_fc": return h raise ValueError(f"Unknown layer '{layer}'. Use 'GNN', 'first_fc', or 'second_fc'.") def forward(self, adj, node_feat, atom_vec): return self.embed(adj, node_feat, atom_vec, layer="second_fc")