"""Strict input schemas and a leakage-free BGC embedding dataset.""" from __future__ import annotations from pathlib import Path from typing import Iterable, Mapping import h5py import numpy as np import pandas as pd import torch from torch.utils.data import Dataset FORBIDDEN_MODEL_INPUTS = frozenset( { "pident", "qcovs", "evalue", "avg_mibig_identity", "deepbgc_score", "product_activity", "product_class", "antibacterial", "cytotoxic", "inhibitor", "antifungal", "Alkaloid", "NRP", "Other", "Polyketide", "RiPP", "Saccharide", "Terpene", } ) ALLOWED_MODEL_INPUTS = frozenset( {"gene_embeddings", "relative_positions", "padding_mask", "pfam_tokens"} ) def require_columns(frame: pd.DataFrame, required: Iterable[str], table_name: str) -> None: missing = set(required).difference(frame.columns) if missing: raise ValueError(f"{table_name} is missing columns: {sorted(missing)}") def validate_model_input_names(names: Iterable[str]) -> None: supplied = set(names) forbidden = supplied.intersection(FORBIDDEN_MODEL_INPUTS) unknown = supplied.difference(ALLOWED_MODEL_INPUTS) if forbidden: raise ValueError(f"Target-leaking model inputs are forbidden: {sorted(forbidden)}") if unknown: raise ValueError(f"Unknown model inputs: {sorted(unknown)}") def build_pfam_vocab(atlas_csv: str | Path, training_bgc_ids: Iterable[str]) -> dict[str, int]: """Build a Pfam vocabulary from training BGCs only.""" atlas = pd.read_csv(atlas_csv, usecols=["bgc_id", "pfam_ids"]) wanted = {str(value) for value in training_bgc_ids} tokens: set[str] = set() for row in atlas.itertuples(index=False): if str(row.bgc_id) not in wanted or pd.isna(row.pfam_ids): continue tokens.update(value for value in str(row.pfam_ids).split(";") if value) return {token: index for index, token in enumerate(sorted(tokens), start=2)} def load_legacy_labels(atlas_csv: str | Path) -> pd.DataFrame: atlas = pd.read_csv(atlas_csv) require_columns(atlas, ["bgc_id", "compound_family"], "legacy atlas") labels = atlas[["bgc_id", "compound_family"]].rename( columns={"compound_family": "mibig_reference_id"} ) labels = labels.dropna(subset=["mibig_reference_id"]).copy() labels["mibig_reference_id"] = labels["mibig_reference_id"].astype(str) labels["group_id"] = labels["mibig_reference_id"] labels["label_tier"] = "silver" return labels def load_gold_mapping(mapping_csv: str | Path) -> pd.DataFrame: mapping = pd.read_csv(mapping_csv) required = ["bgc_id", "product_group_id", "product_id", "source"] require_columns(mapping, required, "gold mapping") result = mapping.copy() result["group_id"] = result["product_group_id"].astype(str) result["label_tier"] = "gold" return result class BGCEmbeddingDataset(Dataset): """Load ESM embeddings in canonical atlas protein order. The alignment-expanded gene table is deliberately not accepted here: it can contain several hit rows per biological gene and therefore corrupt gene rank. """ def __init__( self, embeddings_h5: str | Path, atlas_csv: str | Path, assignments: pd.DataFrame, esm_dimension: int = 1280, pfam_vocab: Mapping[str, int] | None = None, ) -> None: require_columns(assignments, ["bgc_id", "group_id", "split", "label_tier"], "assignments") atlas = pd.read_csv(atlas_csv, usecols=["bgc_id", "protein_ids", "pfam_ids"]) if atlas["bgc_id"].duplicated().any(): raise ValueError("Atlas contains duplicate BGC identifiers") wanted = set(assignments["bgc_id"].astype(str)) atlas = atlas[atlas["bgc_id"].astype(str).isin(wanted)].copy() self.h5_path = str(Path(embeddings_h5).resolve()) self.esm_dimension = int(esm_dimension) self._h5: h5py.File | None = None with h5py.File(self.h5_path, "r") as handle: available = set(handle.keys()) missing_rows: list[dict[str, str]] = [] grouped: dict[str, list[tuple[str, float]]] = {} pfam_by_bgc: dict[str, list[str]] = {} for row in atlas.itertuples(index=False): bgc_id = str(row.bgc_id) protein_ids = ( [value for value in str(row.protein_ids).split(";") if value] if pd.notna(row.protein_ids) else [] ) if len(protein_ids) != len(set(protein_ids)): raise ValueError(f"Atlas protein order contains duplicate IDs for {bgc_id}") denominator = max(1, len(protein_ids) - 1) present: list[tuple[str, float]] = [] for rank, gene_id in enumerate(protein_ids): if gene_id in available: present.append((gene_id, rank / denominator)) else: missing_rows.append({"bgc_id": bgc_id, "gene_id": gene_id}) if present: grouped[bgc_id] = present if pd.isna(row.pfam_ids): pfam_by_bgc[bgc_id] = [] else: pfam_by_bgc[bgc_id] = sorted( {value for value in str(row.pfam_ids).split(";") if value} ) self.missing_gene_rows = pd.DataFrame(missing_rows, columns=["bgc_id", "gene_id"]) metadata = assignments.drop_duplicates("bgc_id").set_index("bgc_id") self.bgc_ids = [str(bgc_id) for bgc_id in metadata.index if str(bgc_id) in grouped] rejected = set(metadata.index.astype(str)).difference(self.bgc_ids) if rejected: raise ValueError(f"BGCs have no usable ESM embeddings: {sorted(rejected)[:10]}") self.bgc_to_genes = grouped self.pfam_vocab = dict(pfam_vocab or {}) self.pfam_tokens_by_bgc = { bgc_id: [self.pfam_vocab.get(token, 1) for token in pfam_by_bgc.get(bgc_id, [])] for bgc_id in self.bgc_ids } self.group_by_bgc = metadata["group_id"].astype(str).to_dict() self.tier_by_bgc = metadata["label_tier"].astype(str).to_dict() self.split_by_bgc = metadata["split"].astype(str).to_dict() @property def h5(self) -> h5py.File: if self._h5 is None: self._h5 = h5py.File(self.h5_path, "r") return self._h5 def __len__(self) -> int: return len(self.bgc_ids) def __getitem__(self, index: int) -> dict[str, object]: bgc_id = self.bgc_ids[index] embeddings: list[np.ndarray] = [] positions: list[float] = [] gene_ids: list[str] = [] for gene_id, position in self.bgc_to_genes[bgc_id]: embedding = np.asarray(self.h5[gene_id][()], dtype=np.float32) if embedding.shape != (self.esm_dimension,): raise ValueError(f"{gene_id} has shape {embedding.shape}; expected {(self.esm_dimension,)}") embeddings.append(embedding) positions.append(position) gene_ids.append(gene_id) return { "gene_embeddings": torch.from_numpy(np.stack(embeddings)), "relative_positions": torch.tensor(positions, dtype=torch.float32), "bgc_id": bgc_id, "gene_ids": gene_ids, "pfam_tokens": torch.tensor( self.pfam_tokens_by_bgc[bgc_id], dtype=torch.long ), "group_id": self.group_by_bgc[bgc_id], "label_tier": self.tier_by_bgc[bgc_id], "split": self.split_by_bgc[bgc_id], } def close(self) -> None: if self._h5 is not None: self._h5.close() self._h5 = None def __del__(self) -> None: self.close() def collate_bgcs(batch: list[dict[str, object]]) -> dict[str, object]: if not batch: raise ValueError("Cannot collate an empty batch") max_genes = max(item["gene_embeddings"].shape[0] for item in batch) max_pfams = max(1, max(item["pfam_tokens"].shape[0] for item in batch)) dimension = batch[0]["gene_embeddings"].shape[1] embeddings = torch.zeros(len(batch), max_genes, dimension, dtype=torch.float32) positions = torch.zeros(len(batch), max_genes, dtype=torch.float32) padding_mask = torch.ones(len(batch), max_genes, dtype=torch.bool) pfam_tokens = torch.zeros(len(batch), max_pfams, dtype=torch.long) for row, item in enumerate(batch): count = item["gene_embeddings"].shape[0] embeddings[row, :count] = item["gene_embeddings"] positions[row, :count] = item["relative_positions"] padding_mask[row, :count] = False pfam_count = item["pfam_tokens"].shape[0] if pfam_count: pfam_tokens[row, :pfam_count] = item["pfam_tokens"] result: dict[str, object] = { "gene_embeddings": embeddings, "relative_positions": positions, "padding_mask": padding_mask, "pfam_tokens": pfam_tokens, } result["bgc_ids"] = [item["bgc_id"] for item in batch] result["gene_ids"] = [item["gene_ids"] for item in batch] result["group_ids"] = [item["group_id"] for item in batch] result["label_tiers"] = [item["label_tier"] for item in batch] result["splits"] = [item["split"] for item in batch] validate_model_input_names( ["gene_embeddings", "relative_positions", "padding_mask", "pfam_tokens"] ) return result