Download source/src/bgc_retrieval/data.py from rustambekurokov/bgc-setnet: direct link, hf CLI and curl.
- Browser
- Download file 9.44 kB
-
https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/src/bgc_retrieval/data.py
- Command line
-
hf download hf://rustambekurokov/bgc-setnet/source/src/bgc_retrieval/data.py
-
curl -L -o data.py https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/src/bgc_retrieval/data.py
9.44 kB
| """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() | |
| 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 | |