whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
19.7 kB
"""One-GPU Phase 1 and Phase 2 training with immutable artifacts."""
from __future__ import annotations
import copy
import random
from collections.abc import Iterable
from pathlib import Path
from typing import Any
import numpy as np
import pandas as pd
import torch
from torch import nn
from torch.nn import functional as F
from torch.utils.data import DataLoader
from .artifacts import write_json_immutable
from .checkpoints import load_checkpoint, save_checkpoint
from .data import BGCEmbeddingDataset, collate_bgcs
from .losses import augment_gene_sets, masked_gene_loss, supervised_contrastive_loss
from .metrics import expected_tie_aware_metrics
from .model import LeakageFreeBGCSetNet, MaskedGenePredictionHead, ModelConfig, WeightedPfamJaccard
from .sampling import UniqueGroupBatchSampler
def set_reproducible(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)
torch.backends.cudnn.benchmark = False
torch.backends.cudnn.deterministic = True
def choose_device() -> torch.device:
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
def _move_batch(batch: dict[str, object], device: torch.device) -> tuple[torch.Tensor, ...]:
return (
batch["gene_embeddings"].to(device),
batch["relative_positions"].to(device),
batch["padding_mask"].to(device),
batch["pfam_tokens"].to(device),
)
def _encode_genes(
model: nn.Module,
embeddings: torch.Tensor,
positions: torch.Tensor,
padding_mask: torch.Tensor,
pfam_tokens: torch.Tensor,
) -> torch.Tensor:
if getattr(model, "uses_pfam", False):
return model.encode_genes(embeddings, positions, padding_mask, pfam_tokens)
return model.encode_genes(embeddings, positions, padding_mask)
def _encode_bgc(
model: nn.Module,
embeddings: torch.Tensor,
positions: torch.Tensor,
padding_mask: torch.Tensor,
pfam_tokens: torch.Tensor,
) -> torch.Tensor:
if getattr(model, "uses_pfam", False):
return model(embeddings, positions, padding_mask, pfam_tokens)
return model(embeddings, positions, padding_mask)
def _mask_gene_batch(
embeddings: torch.Tensor,
padding_mask: torch.Tensor,
probability: float,
seed: int,
) -> tuple[torch.Tensor, torch.Tensor]:
generator = torch.Generator(device="cpu").manual_seed(seed)
random_values = torch.rand(padding_mask.shape, generator=generator)
masked = (random_values < probability) & ~padding_mask.cpu()
for row in range(masked.shape[0]):
if not masked[row].any():
valid = torch.nonzero(~padding_mask[row].cpu(), as_tuple=False).flatten()
masked[row, valid[torch.randint(len(valid), (1,), generator=generator)]] = True
masked = masked.to(embeddings.device)
model_input = embeddings.clone()
model_input[masked] = 0.0
return model_input, masked
def _phase1_epoch(
model: LeakageFreeBGCSetNet,
head: MaskedGenePredictionHead,
loader: DataLoader,
device: torch.device,
mask_probability: float,
seed: int,
optimizer: torch.optim.Optimizer | None,
mixed_precision: bool,
) -> float:
training = optimizer is not None
model.train(training)
head.train(training)
losses: list[float] = []
scaler = torch.amp.GradScaler("cuda", enabled=mixed_precision and device.type == "cuda")
context = torch.enable_grad() if training else torch.no_grad()
with context:
for batch_index, batch in enumerate(loader):
embeddings, positions, padding_mask, pfam_tokens = _move_batch(batch, device)
model_input, masked = _mask_gene_batch(
embeddings, padding_mask, mask_probability, seed + batch_index
)
if optimizer:
optimizer.zero_grad(set_to_none=True)
with torch.amp.autocast(device_type=device.type, enabled=mixed_precision and device.type == "cuda"):
contextual = _encode_genes(
model, model_input, positions, padding_mask, pfam_tokens
)
prediction = head(contextual)
loss = masked_gene_loss(prediction, embeddings, masked)
if not torch.isfinite(loss):
raise FloatingPointError("Non-finite Phase 1 loss")
if optimizer:
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
nn.utils.clip_grad_norm_(list(model.parameters()) + list(head.parameters()), 1.0)
scaler.step(optimizer)
scaler.update()
losses.append(float(loss.detach().cpu()))
return float(np.mean(losses))
@torch.no_grad()
def encode_dataset(
model: LeakageFreeBGCSetNet,
dataset: BGCEmbeddingDataset,
device: torch.device,
num_workers: int = 0,
) -> dict[str, torch.Tensor]:
loader = DataLoader(
dataset, batch_size=32, shuffle=False, num_workers=num_workers, collate_fn=collate_bgcs
)
model.eval()
result: dict[str, torch.Tensor] = {}
for batch in loader:
embeddings, positions, padding_mask, pfam_tokens = _move_batch(batch, device)
encoded = _encode_bgc(model, embeddings, positions, padding_mask, pfam_tokens).cpu()
result.update(zip(batch["bgc_ids"], encoded))
return result
def embedding_validation_metric(
embeddings: dict[str, torch.Tensor], assignments: pd.DataFrame, cutoff: int = 50
) -> float:
group_by_bgc = assignments.set_index("bgc_id")["group_id"].astype(str).to_dict()
identifiers = sorted(embeddings)
values: list[float] = []
for reference in identifiers:
relevant = {
item for item in identifiers if item != reference and group_by_bgc[item] == group_by_bgc[reference]
}
if not relevant:
continue
scores = {
candidate: float(F.cosine_similarity(embeddings[reference], embeddings[candidate], dim=0))
for candidate in identifiers
if candidate != reference
}
values.append(expected_tie_aware_metrics(scores, relevant, recall_at=(cutoff,), ndcg_at=())[f"recall@{cutoff}"])
if not values:
raise ValueError("Validation split has no positive retrieval queries")
return float(np.mean(values))
def train_phase1(
model: LeakageFreeBGCSetNet,
model_config: ModelConfig,
train_dataset: BGCEmbeddingDataset,
validation_dataset: BGCEmbeddingDataset,
split_path: str | Path,
input_paths: list[str | Path],
output_dir: str | Path,
training_config: dict[str, Any],
seed: int,
) -> Path:
set_reproducible(seed)
device = choose_device()
model.to(device)
head = MaskedGenePredictionHead(model_config).to(device)
optimizer = torch.optim.AdamW(
list(model.parameters()) + list(head.parameters()),
lr=float(training_config["learning_rate"]),
weight_decay=float(training_config["weight_decay"]),
)
generator = torch.Generator().manual_seed(seed)
train_loader = DataLoader(
train_dataset,
batch_size=int(training_config["batch_groups"]),
shuffle=True,
generator=generator,
num_workers=int(training_config["num_workers"]),
collate_fn=collate_bgcs,
)
validation_loader = DataLoader(
validation_dataset,
batch_size=int(training_config["batch_groups"]),
shuffle=False,
num_workers=int(training_config["num_workers"]),
collate_fn=collate_bgcs,
)
history: list[dict[str, float | int]] = []
best_loss = float("inf")
best_state: dict[str, Any] | None = None
patience = 0
for epoch in range(int(training_config["phase1_epochs"])):
train_loss = _phase1_epoch(
model, head, train_loader, device, float(training_config["mask_probability"]),
seed + epoch * 10000, optimizer, bool(training_config["mixed_precision"]),
)
validation_loss = _phase1_epoch(
model, head, validation_loader, device, float(training_config["mask_probability"]),
seed + 900000, None, bool(training_config["mixed_precision"]),
)
history.append({"epoch": epoch, "train_loss": train_loss, "validation_loss": validation_loss})
if validation_loss < best_loss:
best_loss = validation_loss
best_state = {
"model": copy.deepcopy(model.state_dict()),
"head": copy.deepcopy(head.state_dict()),
"epoch": epoch,
}
patience = 0
else:
patience += 1
if patience >= int(training_config["patience"]):
break
if best_state is None:
raise RuntimeError("Phase 1 did not produce a valid checkpoint")
model.load_state_dict(best_state["model"])
output = Path(output_dir)
checkpoint_path = output / "phase1_best.pt"
save_checkpoint(
checkpoint_path, model, model_config, split_path, input_paths,
{"stage": "phase1", "best_epoch": best_state["epoch"], "best_loss": best_loss,
"head_state": best_state["head"]}, optimizer,
)
write_json_immutable(output / "phase1_history.json", history)
return checkpoint_path
def weighted_jaccard_contrastive_loss(
similarities: torch.Tensor,
group_ids: list[str],
temperature: float,
) -> torch.Tensor:
identity = torch.eye(len(group_ids), dtype=torch.bool, device=similarities.device)
positives = torch.tensor(
[[left == right for right in group_ids] for left in group_ids],
dtype=torch.bool,
device=similarities.device,
) & ~identity
positive_counts = positives.sum(dim=1)
if torch.any(positive_counts == 0):
raise ValueError("Every weighted-Pfam batch item must have a positive")
logits = similarities / temperature
logits = logits.masked_fill(identity, float("-inf"))
log_probabilities = logits - torch.logsumexp(logits, dim=1, keepdim=True)
positive_log_probability = log_probabilities.masked_fill(~positives, 0.0).sum(dim=1)
positive_log_probability = positive_log_probability / positive_counts
return -positive_log_probability.mean()
@torch.no_grad()
def weighted_jaccard_validation_metric(
model: WeightedPfamJaccard,
dataset: BGCEmbeddingDataset,
assignments: pd.DataFrame,
cutoff: int = 50,
) -> float:
device = next(model.parameters()).device
vocabulary = model.raw_weights.shape[0]
presence = torch.zeros(len(dataset), vocabulary, device=device)
for index in range(len(dataset)):
tokens = dataset[index]["pfam_tokens"].to(device)
if len(tokens):
presence[index, tokens] = 1.0
presence[:, 0] = 0.0
weighted = presence * model.domain_weights().to(device)
totals = weighted.sum(dim=1)
intersection = weighted @ presence.T
union = totals[:, None] + totals[None, :] - intersection
similarities = (intersection / union.clamp_min(1e-8)).cpu()
group_by_bgc = assignments.set_index("bgc_id")["group_id"].astype(str).to_dict()
values: list[float] = []
for index, identifier in enumerate(dataset.bgc_ids):
relevant = {
other_index
for other_index, other_id in enumerate(dataset.bgc_ids)
if other_index != index and group_by_bgc[other_id] == group_by_bgc[identifier]
}
if not relevant:
continue
scores = {
other_id: float(similarities[index, other_index])
for other_index, other_id in enumerate(dataset.bgc_ids)
if other_index != index
}
values.append(
expected_tie_aware_metrics(
scores, {dataset.bgc_ids[item] for item in relevant}, recall_at=(cutoff,), ndcg_at=()
)[f"recall@{cutoff}"]
)
if not values:
raise ValueError("Validation split has no positive weighted-Pfam queries")
return float(np.mean(values))
def train_weighted_pfam(
model: WeightedPfamJaccard,
model_config: ModelConfig,
train_dataset: BGCEmbeddingDataset,
validation_dataset: BGCEmbeddingDataset,
validation_assignments: pd.DataFrame,
split_path: str | Path,
input_paths: list[str | Path],
output_dir: str | Path,
training_config: dict[str, Any],
seed: int,
) -> Path:
set_reproducible(seed)
device = choose_device()
model.to(device)
optimizer = torch.optim.AdamW(
model.parameters(),
lr=float(training_config["learning_rate"]),
weight_decay=float(training_config["weight_decay"]),
)
sampler = UniqueGroupBatchSampler(
[train_dataset.group_by_bgc[item] for item in train_dataset.bgc_ids],
groups_per_batch=int(training_config["batch_groups"]),
examples_per_group=int(training_config["examples_per_group"]),
seed=seed,
)
loader = DataLoader(
train_dataset,
batch_sampler=sampler,
num_workers=int(training_config["num_workers"]),
collate_fn=collate_bgcs,
)
best_metric = -float("inf")
best_state: dict[str, Any] | None = None
history: list[dict[str, float | int]] = []
patience = 0
regularization = float(training_config.get("pfam_weight_regularization", 0.001))
for epoch in range(int(training_config["phase2_epochs"])):
sampler.set_epoch(epoch)
model.train()
epoch_losses: list[float] = []
for batch in loader:
tokens = batch["pfam_tokens"].to(device)
optimizer.zero_grad(set_to_none=True)
similarities = model.pairwise_jaccard(tokens)
loss = weighted_jaccard_contrastive_loss(
similarities, batch["group_ids"], float(training_config["temperature"])
)
log_weights = torch.log(model.domain_weights()[1:].clamp_min(1e-8))
loss = loss + regularization * torch.mean(log_weights.square())
if not torch.isfinite(loss):
raise FloatingPointError("Non-finite weighted-Pfam loss")
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
epoch_losses.append(float(loss.detach().cpu()))
validation_recall = weighted_jaccard_validation_metric(
model, validation_dataset, validation_assignments, cutoff=50
)
history.append(
{"epoch": epoch, "train_loss": float(np.mean(epoch_losses)),
"validation_recall@50": validation_recall}
)
if validation_recall > best_metric:
best_metric = validation_recall
best_state = {"model": copy.deepcopy(model.state_dict()), "epoch": epoch}
patience = 0
else:
patience += 1
if patience >= int(training_config["patience"]):
break
if best_state is None:
raise RuntimeError("Weighted Pfam training did not produce a valid checkpoint")
model.load_state_dict(best_state["model"])
output = Path(output_dir)
checkpoint_path = output / "phase2_best.pt"
save_checkpoint(
checkpoint_path, model, model_config, split_path, input_paths,
{"stage": "weighted_pfam", "best_epoch": best_state["epoch"],
"best_validation_recall@50": best_metric}, optimizer,
)
write_json_immutable(output / "phase2_history.json", history)
return checkpoint_path
def train_phase2(
model: LeakageFreeBGCSetNet,
model_config: ModelConfig,
train_dataset: BGCEmbeddingDataset,
validation_dataset: BGCEmbeddingDataset,
validation_assignments: pd.DataFrame,
split_path: str | Path,
input_paths: list[str | Path],
output_dir: str | Path,
training_config: dict[str, Any],
seed: int,
phase1_checkpoint: str | Path | None = None,
) -> Path:
set_reproducible(seed)
device = choose_device()
if phase1_checkpoint:
load_checkpoint(phase1_checkpoint, model, split_path)
model.to(device)
optimizer = torch.optim.AdamW(
model.parameters(), lr=float(training_config["learning_rate"]),
weight_decay=float(training_config["weight_decay"]),
)
sampler = UniqueGroupBatchSampler(
[train_dataset.group_by_bgc[item] for item in train_dataset.bgc_ids],
groups_per_batch=int(training_config["batch_groups"]),
examples_per_group=int(training_config["examples_per_group"]), seed=seed,
)
loader = DataLoader(
train_dataset, batch_sampler=sampler, num_workers=int(training_config["num_workers"]),
collate_fn=collate_bgcs,
)
scaler = torch.amp.GradScaler(
"cuda", enabled=bool(training_config["mixed_precision"]) and device.type == "cuda"
)
best_metric = -float("inf")
best_state: dict[str, Any] | None = None
history: list[dict[str, float | int]] = []
patience = 0
for epoch in range(int(training_config["phase2_epochs"])):
sampler.set_epoch(epoch)
model.train()
epoch_losses: list[float] = []
for batch in loader:
embeddings, positions, padding_mask, pfam_tokens = _move_batch(batch, device)
embeddings, positions, padding_mask = augment_gene_sets(
embeddings, positions, padding_mask,
float(training_config["gene_dropout"]), float(training_config["position_jitter"]),
)
optimizer.zero_grad(set_to_none=True)
with torch.amp.autocast(
device_type=device.type,
enabled=bool(training_config["mixed_precision"]) and device.type == "cuda",
):
encoded = _encode_bgc(
model, embeddings, positions, padding_mask, pfam_tokens
)
loss = supervised_contrastive_loss(
encoded, batch["group_ids"], float(training_config["temperature"])
)
if not torch.isfinite(loss):
raise FloatingPointError("Non-finite Phase 2 loss")
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
epoch_losses.append(float(loss.detach().cpu()))
validation_embeddings = encode_dataset(
model, validation_dataset, device, int(training_config["num_workers"])
)
validation_recall = embedding_validation_metric(
validation_embeddings, validation_assignments, cutoff=50
)
history.append(
{"epoch": epoch, "train_loss": float(np.mean(epoch_losses)),
"validation_recall@50": validation_recall}
)
if validation_recall > best_metric:
best_metric = validation_recall
best_state = {"model": copy.deepcopy(model.state_dict()), "epoch": epoch}
patience = 0
else:
patience += 1
if patience >= int(training_config["patience"]):
break
if best_state is None:
raise RuntimeError("Phase 2 did not produce a valid checkpoint")
model.load_state_dict(best_state["model"])
output = Path(output_dir)
checkpoint_path = output / "phase2_best.pt"
save_checkpoint(
checkpoint_path, model, model_config, split_path, input_paths,
{"stage": "phase2", "best_epoch": best_state["epoch"],
"best_validation_recall@50": best_metric, "phase1_checkpoint": str(phase1_checkpoint)},
optimizer,
)
write_json_immutable(output / "phase2_history.json", history)
return checkpoint_path