"""Contracts for curated external mappings and released baseline outputs.""" from __future__ import annotations from pathlib import Path import numpy as np import pandas as pd import torch from .data import require_columns GOLD_COLUMNS = ( "bgc_id", "product_group_id", "product_id", "mibig_reference_id", "genus", "source", ) def validate_gold_mapping(mapping: pd.DataFrame) -> None: require_columns(mapping, GOLD_COLUMNS, "external gold mapping") if mapping["bgc_id"].duplicated().any(): raise ValueError("External BGC IDs must be unique") if mapping[list(GOLD_COLUMNS)].isna().any().any(): raise ValueError("Primary gold mapping fields may not be null") if (mapping["product_group_id"].astype(str).str.len() == 0).any(): raise ValueError("Product group IDs may not be empty") def quarantine_external_reference_overlap( training_assignments: pd.DataFrame, external_mapping: pd.DataFrame, ) -> tuple[pd.DataFrame, pd.DataFrame]: """Exclude training pseudo-labels linked to external MIBiG references.""" validate_gold_mapping(external_mapping) require_columns(training_assignments, ["mibig_reference_id", "split"], "training assignments") external_references = set(external_mapping["mibig_reference_id"].astype(str)) overlap = training_assignments["mibig_reference_id"].astype(str).isin(external_references) quarantined = training_assignments[overlap].copy() filtered = training_assignments[~overlap].copy() if set(filtered.loc[filtered["split"] == "train", "mibig_reference_id"]).intersection( external_references ): raise AssertionError("External reference leakage remains after quarantine") return filtered, quarantined def cross_genus_subset(mapping: pd.DataFrame, training_genera: set[str]) -> pd.DataFrame: validate_gold_mapping(mapping) normalized = {value.strip().lower() for value in training_genera} return mapping[~mapping["genus"].astype(str).str.strip().str.lower().isin(normalized)].copy() def load_released_embeddings(path: str | Path) -> dict[str, torch.Tensor]: """Load BGC-MLM or other released embeddings from an ID-keyed NPZ file.""" archive = np.load(path, allow_pickle=False) if "bgc_ids" not in archive or "embeddings" not in archive: raise ValueError("Released embedding NPZ requires 'bgc_ids' and 'embeddings' arrays") identifiers = archive["bgc_ids"].astype(str) embeddings = archive["embeddings"] if embeddings.ndim != 2 or len(identifiers) != len(embeddings): raise ValueError("Released embedding arrays have inconsistent shapes") return { identifier: torch.tensor(vector, dtype=torch.float32) for identifier, vector in zip(identifiers, embeddings) } def load_bigscape_edges(path: str | Path) -> pd.DataFrame: """Load a normalized BiG-SCAPE edge export without calling it Pfam Jaccard.""" edges = pd.read_csv(path) require_columns(edges, ["record_a", "record_b", "similarity"], "BiG-SCAPE edges") if ((edges["similarity"] < 0) | (edges["similarity"] > 1)).any(): raise ValueError("BiG-SCAPE similarities must be in [0, 1]") return edges