chromatography-rt-prediction / src /neural_models.py
AI4deeperScience's picture
Add files using upload-large-folder tool
8f4ed7a verified
Raw
History Blame Contribute Delete
5.78 kB
"""Neural network architectures for descriptor and fingerprint baselines."""
from __future__ import annotations
from typing import List, Optional, Tuple, cast
import torch
import torch.nn as nn
class DescriptorNN(nn.Module):
"""Feed-forward network for molecular descriptors with optional lab embedding."""
def __init__(
self,
input_dim: int,
hidden_dims: List[int] | Tuple[int, ...] = (192, 96, 48),
dropout: float = 0.2,
use_batch_norm: bool = True,
num_labs: Optional[int] = None,
lab_embed_dim: int = 12,
input_dropout: float = 0.05,
) -> None:
super().__init__()
self.use_lab_embedding = num_labs is not None and num_labs > 0
lab_count = cast(int, num_labs) if self.use_lab_embedding else 0
self.lab_embedding = (
nn.Embedding(lab_count, lab_embed_dim)
if self.use_lab_embedding
else None
)
initial_dim = input_dim + (lab_embed_dim if self.use_lab_embedding else 0)
self.input_dropout = nn.Dropout(input_dropout) if input_dropout > 0 else None
layers: List[nn.Module] = []
prev_dim = initial_dim
for hidden_dim in hidden_dims:
layers.append(nn.Linear(prev_dim, hidden_dim))
if use_batch_norm:
layers.append(nn.BatchNorm1d(hidden_dim))
layers.append(nn.GELU())
layers.append(nn.Dropout(dropout))
prev_dim = hidden_dim
bottleneck_dim = max(32, prev_dim // 2)
layers.extend(
[
nn.Linear(prev_dim, bottleneck_dim),
nn.GELU(),
nn.Dropout(dropout * 0.5),
nn.Linear(bottleneck_dim, 1),
]
)
self.network = nn.Sequential(*layers)
self.apply(self._init_weights)
# placeholders populated during training for inverse scaling
self.target_mean: float = 0.0
self.target_std: float = 1.0
@staticmethod
def _init_weights(module: nn.Module) -> None:
if isinstance(module, nn.Linear):
nn.init.xavier_uniform_(module.weight)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.BatchNorm1d):
nn.init.ones_(module.weight)
nn.init.zeros_(module.bias)
def forward(self, x: torch.Tensor, lab_indices: Optional[torch.Tensor] = None) -> torch.Tensor:
if self.input_dropout is not None:
x = self.input_dropout(x)
if self.lab_embedding is not None:
if lab_indices is None:
raise ValueError("lab_indices must be provided when lab embedding is enabled.")
if lab_indices.dim() > 1:
lab_indices = lab_indices.squeeze(-1)
lab_embeddings = self.lab_embedding(lab_indices.long())
x = torch.cat([x, lab_embeddings], dim=-1)
return self.network(x).squeeze(-1)
class FingerprintNN(nn.Module):
"""Feed-forward network tailored for high-dimensional sparse fingerprints."""
def __init__(
self,
input_dim: int,
hidden_dims: List[int] | Tuple[int, ...] = (768, 384, 192, 96),
dropout: float = 0.15,
use_batch_norm: bool = True,
num_labs: Optional[int] = None,
lab_embed_dim: int = 16,
input_dropout: float = 0.1,
) -> None:
super().__init__()
self.use_lab_embedding = num_labs is not None and num_labs > 0
lab_count = cast(int, num_labs) if self.use_lab_embedding else 0
self.lab_embedding = (
nn.Embedding(lab_count, lab_embed_dim)
if self.use_lab_embedding
else None
)
initial_dim = input_dim + (lab_embed_dim if self.use_lab_embedding else 0)
self.input_dropout = nn.Dropout(input_dropout) if input_dropout > 0 else None
layers: List[nn.Module] = []
prev_dim = initial_dim
for idx, hidden_dim in enumerate(hidden_dims):
layers.append(nn.Linear(prev_dim, hidden_dim))
if use_batch_norm:
layers.append(nn.BatchNorm1d(hidden_dim))
layers.append(nn.GELU())
layers.append(nn.Dropout(dropout if idx < len(hidden_dims) - 1 else dropout * 0.5))
prev_dim = hidden_dim
layers.append(nn.Linear(prev_dim, 1))
self.network = nn.Sequential(*layers)
self.apply(self._init_weights)
self.target_mean: float = 0.0
self.target_std: float = 1.0
@staticmethod
def _init_weights(module: nn.Module) -> None:
if isinstance(module, nn.Linear):
nn.init.xavier_uniform_(module.weight)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.BatchNorm1d):
nn.init.ones_(module.weight)
nn.init.zeros_(module.bias)
def forward(self, x: torch.Tensor, lab_indices: Optional[torch.Tensor] = None) -> torch.Tensor:
if self.input_dropout is not None:
x = self.input_dropout(x)
if self.lab_embedding is not None:
if lab_indices is None:
raise ValueError("lab_indices must be provided when lab embedding is enabled.")
if lab_indices.dim() > 1:
lab_indices = lab_indices.squeeze(-1)
lab_embeddings = self.lab_embedding(lab_indices.long())
x = torch.cat([x, lab_embeddings], dim=-1)
return self.network(x).squeeze(-1)
__all__ = ["DescriptorNN", "FingerprintNN"]