whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
3.21 kB
"""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