BioXMol / modeling_bioxmol.py
arslanmasood's picture
Upload modeling_bioxmol.py
85851de verified
Raw History Blame Contribute Delete
5.22 kB
"""
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")