bozdaglab's picture
Add ZeroGPU-enabled BioLM-NET workbench
31376a7 verified
Raw
History Blame Contribute Delete
16 kB
from __future__ import annotations
import copy
import random
from dataclasses import asdict, dataclass
from typing import Callable
import numpy as np
import pandas as pd
import torch
from sklearn.metrics import (
accuracy_score,
confusion_matrix,
f1_score,
precision_score,
recall_score,
)
from sklearn.model_selection import train_test_split
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from .data import PreparedWorkspace
from .model import BioLMNet
ProgressCallback = Callable[[float, str], None]
@dataclass(frozen=True)
class Hyperparameters:
epochs: int = 50
batch_size: int = 16
learning_rate: float = 0.001
weight_decay: float = 0.01
dropout: float = 0.3
projection_dim: int = 64
fusion_dim: int = 12
validation_fraction: float = 0.2
optimizer: str = "Adam"
class_weighting: bool = True
early_stopping_patience: int = 12
seed: int = 42
def validate(self) -> None:
if not 1 <= self.epochs <= 1000:
raise ValueError("Epochs must be between 1 and 1,000.")
if not 2 <= self.batch_size <= 1024:
raise ValueError("Batch size must be between 2 and 1,024.")
if not 0 < self.learning_rate <= 1:
raise ValueError("Learning rate must be in (0, 1].")
if not 0 <= self.dropout < 1:
raise ValueError("Dropout must be in [0, 1).")
if not 0.05 <= self.validation_fraction <= 0.5:
raise ValueError("Validation fraction must be between 0.05 and 0.5.")
if self.optimizer.lower() not in {"adam", "sgd"}:
raise ValueError("Optimizer must be Adam or SGD.")
@dataclass
class ModelBundle:
model: BioLMNet
gene_features: list[str]
dna_features: list[str]
label_names: list[str]
gene_mean: np.ndarray
gene_scale: np.ndarray
dna_mean: np.ndarray
dna_scale: np.ndarray
config: dict
metrics: dict
history: list[dict[str, float]]
@dataclass
class TrainingResult:
bundle: ModelBundle
validation_predictions: pd.DataFrame
confusion: np.ndarray
def set_reproducible_seed(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def _fit_scaler(values: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
mean = values.mean(axis=0, dtype=np.float64).astype(np.float32)
scale = values.std(axis=0, dtype=np.float64).astype(np.float32)
scale[scale < 1e-8] = 1.0
return mean, scale
def _scale(values: np.ndarray, mean: np.ndarray, scale: np.ndarray) -> np.ndarray:
return ((values.astype(np.float32) - mean) / scale).astype(np.float32)
def _make_model(
workspace: PreparedWorkspace, hyperparameters: Hyperparameters
) -> BioLMNet:
gene = workspace.gene_branch
dna = workspace.dna_branch
if (
gene.embeddings is None
or dna.embeddings is None
or gene.pathway_mask is None
or dna.pathway_mask is None
):
raise ValueError("Workspace priors are incomplete; run data preparation first.")
return BioLMNet(
gene_biological_mask=torch.from_numpy(gene.biological_mask),
dna_biological_mask=torch.from_numpy(dna.biological_mask),
gene_embeddings=torch.from_numpy(gene.embeddings),
dna_embeddings=torch.from_numpy(dna.embeddings),
gene_pathway_mask=torch.from_numpy(gene.pathway_mask),
dna_pathway_mask=torch.from_numpy(dna.pathway_mask),
n_classes=len(workspace.label_names),
projection_dim=hyperparameters.projection_dim,
fusion_dim=hyperparameters.fusion_dim,
dropout=hyperparameters.dropout,
)
def _evaluate(
model: BioLMNet,
gene_values: np.ndarray,
dna_values: np.ndarray,
labels: np.ndarray,
device: torch.device,
) -> tuple[float, np.ndarray, np.ndarray]:
model.eval()
with torch.no_grad():
logits = model(
torch.from_numpy(gene_values).to(device),
torch.from_numpy(dna_values).to(device),
)
loss = nn.functional.cross_entropy(
logits, torch.from_numpy(labels).to(device)
).item()
probabilities = torch.softmax(logits, dim=1).cpu().numpy()
predictions = probabilities.argmax(axis=1)
return float(loss), probabilities, predictions
def train(
workspace: PreparedWorkspace,
hyperparameters: Hyperparameters,
progress: ProgressCallback | None = None,
) -> TrainingResult:
hyperparameters.validate()
set_reproducible_seed(hyperparameters.seed)
labels = workspace.labels
classes, class_counts = np.unique(labels, return_counts=True)
if class_counts.min() < 2:
raise ValueError(
"Each class needs at least two samples for a stratified train/validation split."
)
indices = np.arange(len(labels))
train_index, validation_index = train_test_split(
indices,
test_size=hyperparameters.validation_fraction,
random_state=hyperparameters.seed,
stratify=labels,
)
gene_mean, gene_scale = _fit_scaler(
workspace.gene_expression[train_index]
)
dna_mean, dna_scale = _fit_scaler(workspace.dna_methylation[train_index])
gene_train = _scale(
workspace.gene_expression[train_index], gene_mean, gene_scale
)
gene_validation = _scale(
workspace.gene_expression[validation_index], gene_mean, gene_scale
)
dna_train = _scale(
workspace.dna_methylation[train_index], dna_mean, dna_scale
)
dna_validation = _scale(
workspace.dna_methylation[validation_index], dna_mean, dna_scale
)
y_train = labels[train_index]
y_validation = labels[validation_index]
dataset = TensorDataset(
torch.from_numpy(gene_train),
torch.from_numpy(dna_train),
torch.from_numpy(y_train),
)
generator = torch.Generator().manual_seed(hyperparameters.seed)
loader = DataLoader(
dataset,
batch_size=min(hyperparameters.batch_size, len(dataset)),
shuffle=True,
generator=generator,
)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = _make_model(workspace, hyperparameters).to(device)
if hyperparameters.optimizer.lower() == "sgd":
optimizer = torch.optim.SGD(
model.parameters(),
lr=hyperparameters.learning_rate,
weight_decay=hyperparameters.weight_decay,
)
else:
optimizer = torch.optim.Adam(
model.parameters(),
lr=hyperparameters.learning_rate,
weight_decay=hyperparameters.weight_decay,
)
class_weights: torch.Tensor | None = None
if hyperparameters.class_weighting:
count_by_class = np.bincount(
y_train, minlength=len(workspace.label_names)
).astype(np.float32)
weights = len(y_train) / (len(count_by_class) * count_by_class)
class_weights = torch.from_numpy(weights).to(device)
criterion = nn.CrossEntropyLoss(weight=class_weights)
history: list[dict[str, float]] = []
best_state: dict[str, torch.Tensor] | None = None
best_loss = float("inf")
patience = 0
for epoch in range(hyperparameters.epochs):
model.train()
running_loss = 0.0
seen = 0
for gene_batch, dna_batch, label_batch in loader:
gene_batch = gene_batch.to(device)
dna_batch = dna_batch.to(device)
label_batch = label_batch.to(device)
optimizer.zero_grad(set_to_none=True)
logits = model(gene_batch, dna_batch)
loss = criterion(logits, label_batch)
loss.backward()
optimizer.step()
running_loss += loss.item() * len(label_batch)
seen += len(label_batch)
validation_loss, _, validation_predictions = _evaluate(
model,
gene_validation,
dna_validation,
y_validation,
device,
)
epoch_row = {
"epoch": float(epoch + 1),
"training_loss": float(running_loss / max(seen, 1)),
"validation_loss": validation_loss,
"validation_accuracy": float(
accuracy_score(y_validation, validation_predictions)
),
"validation_f1_macro": float(
f1_score(
y_validation,
validation_predictions,
average="macro",
zero_division=0,
)
),
}
history.append(epoch_row)
if progress:
progress(
(epoch + 1) / hyperparameters.epochs,
(
f"Epoch {epoch + 1}/{hyperparameters.epochs} · "
f"validation F1 {epoch_row['validation_f1_macro']:.3f}"
),
)
if validation_loss < best_loss - 1e-5:
best_loss = validation_loss
best_state = copy.deepcopy(model.state_dict())
patience = 0
else:
patience += 1
if patience >= hyperparameters.early_stopping_patience:
break
if best_state is not None:
model.load_state_dict(best_state)
validation_loss, probabilities, predictions = _evaluate(
model,
gene_validation,
dna_validation,
y_validation,
device,
)
confusion = confusion_matrix(
y_validation,
predictions,
labels=np.arange(len(workspace.label_names)),
)
metrics = {
"validation_loss": validation_loss,
"accuracy": float(accuracy_score(y_validation, predictions)),
"f1_macro": float(
f1_score(
y_validation, predictions, average="macro", zero_division=0
)
),
"f1_weighted": float(
f1_score(
y_validation, predictions, average="weighted", zero_division=0
)
),
"precision_macro": float(
precision_score(
y_validation, predictions, average="macro", zero_division=0
)
),
"recall_macro": float(
recall_score(
y_validation, predictions, average="macro", zero_division=0
)
),
"epochs_completed": len(history),
"training_samples": int(len(train_index)),
"validation_samples": int(len(validation_index)),
"device": str(device),
}
predictions_frame = pd.DataFrame(
{
"sample_row": validation_index,
"observed": [
workspace.label_names[value] for value in y_validation
],
"predicted": [
workspace.label_names[value] for value in predictions
],
}
)
for index, label in enumerate(workspace.label_names):
predictions_frame[f"P({label})"] = probabilities[:, index]
gene = workspace.gene_branch
dna = workspace.dna_branch
config = {
"format_version": 1,
"source_name": workspace.source_name,
"gene_features": gene.input_genes,
"dna_features": dna.input_genes,
"label_names": workspace.label_names,
"gene_hidden_genes": gene.hidden_genes,
"dna_hidden_genes": dna.hidden_genes,
"gene_pathways": gene.pathways,
"dna_pathways": dna.pathways,
"hyperparameters": asdict(hyperparameters),
"architecture": {
"projection_dim": hyperparameters.projection_dim,
"fusion_dim": hyperparameters.fusion_dim,
"dropout": hyperparameters.dropout,
"biological_activation": "relu",
"projection_activation": "sigmoid",
"fusion_activation": "tanh",
},
}
bundle = ModelBundle(
model=model.cpu(),
gene_features=gene.input_genes,
dna_features=dna.input_genes,
label_names=workspace.label_names,
gene_mean=gene_mean,
gene_scale=gene_scale,
dna_mean=dna_mean,
dna_scale=dna_scale,
config=config,
metrics=metrics,
history=history,
)
return TrainingResult(
bundle=bundle,
validation_predictions=predictions_frame,
confusion=confusion,
)
def validate_prediction_frames(
gene_frame: pd.DataFrame,
dna_frame: pd.DataFrame,
bundle: ModelBundle,
) -> tuple[np.ndarray, np.ndarray]:
if len(gene_frame) != len(dna_frame):
raise ValueError(
"Prediction gene-expression and DNA-methylation files must have "
"the same number of rows."
)
missing_gene = sorted(set(bundle.gene_features) - set(gene_frame.columns))
missing_dna = sorted(set(bundle.dna_features) - set(dna_frame.columns))
if missing_gene or missing_dna:
details = []
if missing_gene:
details.append(
"gene-expression: " + ", ".join(missing_gene[:8])
)
if missing_dna:
details.append("DNA-methylation: " + ", ".join(missing_dna[:8]))
raise ValueError(
"Prediction files are missing trained features (" + "; ".join(details) + ")."
)
gene_values = gene_frame.loc[:, bundle.gene_features].apply(
pd.to_numeric, errors="coerce"
)
dna_values = dna_frame.loc[:, bundle.dna_features].apply(
pd.to_numeric, errors="coerce"
)
if gene_values.isna().any().any() or dna_values.isna().any().any():
raise ValueError("Prediction inputs contain missing or non-numeric values.")
return (
_scale(gene_values.to_numpy(), bundle.gene_mean, bundle.gene_scale),
_scale(dna_values.to_numpy(), bundle.dna_mean, bundle.dna_scale),
)
def predict(
gene_frame: pd.DataFrame,
dna_frame: pd.DataFrame,
bundle: ModelBundle,
) -> pd.DataFrame:
gene_values, dna_values = validate_prediction_frames(
gene_frame, dna_frame, bundle
)
bundle.model.eval()
with torch.no_grad():
logits = bundle.model(
torch.from_numpy(gene_values), torch.from_numpy(dna_values)
)
probabilities = torch.softmax(logits, dim=1).numpy()
predicted = probabilities.argmax(axis=1)
output = pd.DataFrame(
{
"sample_row": np.arange(len(gene_frame)),
"predicted_class": [
bundle.label_names[index] for index in predicted
],
"confidence": probabilities.max(axis=1),
}
)
for index, label in enumerate(bundle.label_names):
output[f"P({label})"] = probabilities[:, index]
return output
def pathway_importance(bundle: ModelBundle) -> pd.DataFrame:
bundle.model.eval()
with torch.no_grad():
attention = bundle.model.pathway_attention()
records: list[dict[str, float | str]] = []
for branch_name, config_key in (
("Gene expression", "gene_pathways"),
("DNA methylation", "dna_pathways"),
):
key = "gene_expression" if branch_name == "Gene expression" else "dna_methylation"
weights = attention[key].cpu().numpy()
pathways = bundle.config[config_key]
peak_attention = weights.max(axis=0)
entropy = -(
weights * np.log(np.clip(weights, 1e-12, None))
).sum(axis=0)
for pathway, peak, entropy_value in zip(
pathways, peak_attention, entropy, strict=True
):
records.append(
{
"branch": branch_name,
"pathway": pathway,
"peak_gene_attention": float(peak),
"attention_entropy": float(entropy_value),
}
)
return (
pd.DataFrame(records)
.sort_values("peak_gene_attention", ascending=False)
.reset_index(drop=True)
)