bozdaglab's picture
Add ZeroGPU-enabled BioLM-NET workbench
31376a7 verified
Raw
History Blame Contribute Delete
7.38 kB
from __future__ import annotations
import math
from dataclasses import dataclass
import torch
from torch import Tensor, nn
from torch.nn import functional as F
def _activation(name: str) -> nn.Module:
choices: dict[str, nn.Module] = {
"relu": nn.ReLU(),
"gelu": nn.GELU(),
"sigmoid": nn.Sigmoid(),
"tanh": nn.Tanh(),
}
try:
return choices[name.lower()]
except KeyError as exc:
raise ValueError(f"Unsupported activation: {name}") from exc
class BioMaskedLinear(nn.Module):
"""Trainable linear layer whose weights are constrained by a biological mask.
The paper defines W_masked = W ⊙ M. The public notebook approximates this
with a dense projection followed by a fixed matrix multiplication; this
layer implements the paper's equation directly.
"""
def __init__(self, mask: Tensor, bias: bool = True) -> None:
super().__init__()
if mask.ndim != 2:
raise ValueError("Biological mask must be [input_genes, hidden_genes].")
if mask.shape[0] == 0 or mask.shape[1] == 0:
raise ValueError("Biological mask cannot be empty.")
input_features, output_features = mask.shape
self.input_features = int(input_features)
self.output_features = int(output_features)
self.weight = nn.Parameter(torch.empty(output_features, input_features))
self.bias = nn.Parameter(torch.empty(output_features)) if bias else None
self.register_buffer("mask", mask.T.to(dtype=torch.float32).contiguous())
self.reset_parameters()
def reset_parameters(self) -> None:
nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))
if self.bias is not None:
fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
bound = 1 / math.sqrt(fan_in) if fan_in > 0 else 0
nn.init.uniform_(self.bias, -bound, bound)
def forward(self, inputs: Tensor) -> Tensor:
return F.linear(inputs, self.weight * self.mask, self.bias)
class AttentionPathwayLayer(nn.Module):
"""GenePT-guided attention from biological hidden genes to pathways."""
def __init__(self, gene_embeddings: Tensor, pathway_mask: Tensor) -> None:
super().__init__()
if gene_embeddings.ndim != 2:
raise ValueError("Gene embeddings must be [hidden_genes, embedding_dim].")
if pathway_mask.ndim != 2:
raise ValueError("Pathway mask must be [hidden_genes, pathways].")
if gene_embeddings.shape[0] != pathway_mask.shape[0]:
raise ValueError("Gene embeddings and pathway mask must share gene order.")
if torch.any(pathway_mask.sum(dim=0) == 0):
raise ValueError("Every pathway must contain at least one retained gene.")
embeddings = gene_embeddings.to(dtype=torch.float32)
membership = pathway_mask.to(dtype=torch.bool)
self.register_buffer("gene_embeddings", embeddings)
self.register_buffer("pathway_mask", membership)
self.query = nn.Parameter(
torch.empty(pathway_mask.shape[1], gene_embeddings.shape[1])
)
nn.init.xavier_uniform_(self.query)
def attention_weights(self) -> Tensor:
scale = math.sqrt(self.gene_embeddings.shape[1])
# [hidden genes, pathways]
scores = self.gene_embeddings @ self.query.T / scale
scores = scores.masked_fill(~self.pathway_mask, torch.finfo(scores.dtype).min)
return torch.softmax(scores, dim=0)
def forward(self, hidden_gene_signal: Tensor) -> Tensor:
return hidden_gene_signal @ self.attention_weights()
class BioBranch(nn.Module):
def __init__(
self,
biological_mask: Tensor,
gene_embeddings: Tensor,
pathway_mask: Tensor,
projection_dim: int,
dropout: float,
biological_activation: str,
projection_activation: str,
) -> None:
super().__init__()
self.biological = BioMaskedLinear(biological_mask)
self.biological_activation = _activation(biological_activation)
self.dropout = nn.Dropout(dropout)
self.pathway_attention = AttentionPathwayLayer(
gene_embeddings, pathway_mask
)
self.projection = nn.Linear(pathway_mask.shape[1], projection_dim)
self.projection_activation = _activation(projection_activation)
def forward(self, inputs: Tensor) -> Tensor:
hidden = self.dropout(self.biological_activation(self.biological(inputs)))
pathways = self.pathway_attention(hidden)
return self.projection_activation(self.projection(pathways))
@dataclass(frozen=True)
class ModelDimensions:
gene_inputs: int
dna_inputs: int
gene_hidden: int
dna_hidden: int
gene_pathways: int
dna_pathways: int
classes: int
class BioLMNet(nn.Module):
"""Dual-omics BioLM-NET classifier."""
def __init__(
self,
gene_biological_mask: Tensor,
dna_biological_mask: Tensor,
gene_embeddings: Tensor,
dna_embeddings: Tensor,
gene_pathway_mask: Tensor,
dna_pathway_mask: Tensor,
n_classes: int,
projection_dim: int = 64,
fusion_dim: int = 12,
dropout: float = 0.3,
biological_activation: str = "relu",
projection_activation: str = "sigmoid",
fusion_activation: str = "tanh",
) -> None:
super().__init__()
if n_classes < 2:
raise ValueError("BioLM-NET requires at least two label classes.")
self.gene_branch = BioBranch(
gene_biological_mask,
gene_embeddings,
gene_pathway_mask,
projection_dim,
dropout,
biological_activation,
projection_activation,
)
self.dna_branch = BioBranch(
dna_biological_mask,
dna_embeddings,
dna_pathway_mask,
projection_dim,
dropout,
biological_activation,
projection_activation,
)
self.fusion = nn.Linear(projection_dim * 2, fusion_dim)
self.fusion_activation = _activation(fusion_activation)
self.fusion_dropout = nn.Dropout(dropout)
self.output = nn.Linear(fusion_dim, n_classes)
self.dimensions = ModelDimensions(
gene_inputs=gene_biological_mask.shape[0],
dna_inputs=dna_biological_mask.shape[0],
gene_hidden=gene_biological_mask.shape[1],
dna_hidden=dna_biological_mask.shape[1],
gene_pathways=gene_pathway_mask.shape[1],
dna_pathways=dna_pathway_mask.shape[1],
classes=n_classes,
)
def forward(self, gene_expression: Tensor, dna_methylation: Tensor) -> Tensor:
gene_projection = self.gene_branch(gene_expression)
dna_projection = self.dna_branch(dna_methylation)
fused = torch.cat([gene_projection, dna_projection], dim=1)
fused = self.fusion_dropout(
self.fusion_activation(self.fusion(fused))
)
return self.output(fused)
def pathway_attention(self) -> dict[str, Tensor]:
return {
"gene_expression": self.gene_branch.pathway_attention.attention_weights(),
"dna_methylation": self.dna_branch.pathway_attention.attention_weights(),
}