bgc-setnet / source /src /bgc_retrieval /external_evaluation.py
whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
9.65 kB
"""External product-structure benchmark with overlap quarantine and block bootstrap."""
from __future__ import annotations
from collections.abc import Mapping
import numpy as np
import pandas as pd
import torch
from scipy import stats
from torch.nn import functional as F
from .metrics import expected_tie_aware_metrics
def load_structure_matrix(path: str) -> pd.DataFrame:
matrix = pd.read_csv(path, index_col=0)
matrix.index = matrix.index.astype(str)
matrix.columns = matrix.columns.astype(str)
if matrix.shape[0] != matrix.shape[1] or set(matrix.index) != set(matrix.columns):
raise ValueError("Product-structure similarity matrix must be square")
matrix = matrix.loc[matrix.index, matrix.index]
values = matrix.to_numpy(dtype=float)
if not np.allclose(values, values.T, equal_nan=False):
raise ValueError("Product-structure similarity matrix must be symmetric")
if np.nanmin(values) < 0 or np.nanmax(values) > 1:
raise ValueError("Product-structure similarities must be in [0, 1]")
return matrix
def eligible_external_ids(
structure_matrix: pd.DataFrame,
embedding_ids: set[str],
training_split: pd.DataFrame,
) -> list[str]:
blocked = set(
training_split.loc[
training_split["split"].isin(["train", "validation"]), "group_id"
].astype(str)
)
identifiers = set(structure_matrix.index).intersection(embedding_ids).difference(blocked)
if identifiers.intersection(blocked):
raise AssertionError("Training/validation MIBiG references leaked into external evaluation")
return sorted(identifiers)
def all_pair_scores(
embeddings: Mapping[str, torch.Tensor], identifiers: list[str]
) -> pd.DataFrame:
values = torch.stack([F.normalize(embeddings[identifier].float(), dim=0) for identifier in identifiers])
similarities = (values @ values.T).cpu().numpy()
left, right = np.triu_indices(len(identifiers), k=1)
return pd.DataFrame(
{
"record_a": [identifiers[index] for index in left],
"record_b": [identifiers[index] for index in right],
"score": similarities[left, right],
}
)
def score_requested_pairs(
embeddings: Mapping[str, torch.Tensor], pairs: pd.DataFrame
) -> pd.DataFrame:
required = {"record_a", "record_b"}
if missing := required.difference(pairs.columns):
raise ValueError(f"Requested pairs are missing columns: {sorted(missing)}")
result = pairs[["record_a", "record_b"]].copy()
result["score"] = [
float(F.cosine_similarity(embeddings[left], embeddings[right], dim=0))
for left, right in zip(result["record_a"], result["record_b"])
]
return result
def attach_structural_truth(edges: pd.DataFrame, matrix: pd.DataFrame) -> pd.DataFrame:
required = {"record_a", "record_b", "score"}
if missing := required.difference(edges.columns):
raise ValueError(f"Pair scores are missing columns: {sorted(missing)}")
valid = edges["record_a"].isin(matrix.index) & edges["record_b"].isin(matrix.index)
result = edges[valid].copy()
result["structural_similarity"] = [
float(matrix.loc[left, right])
for left, right in zip(result["record_a"], result["record_b"])
]
return result
def spearman_summary(edges: pd.DataFrame) -> dict[str, float | int]:
if len(edges) < 3:
raise ValueError("At least three scored pairs are required")
correlation, p_value = stats.spearmanr(edges["score"], edges["structural_similarity"])
return {"pairs": len(edges), "spearman_r": float(correlation), "p_value": float(p_value)}
def anchor_block_bootstrap(
edges: pd.DataFrame,
samples: int,
confidence: float,
seed: int,
) -> tuple[float, float]:
"""Two-endpoint BGC cluster bootstrap on fixed full-sample ranks.
Spearman correlation is Pearson correlation of ranks. Ranking once and
resampling endpoint-level sufficient statistics avoids materializing a
million-row pair table for every replicate. Each dyad contributes half of
its weight to each endpoint, so every BGC is represented as a dependence
block instead of assigning pairs to the lexicographically smaller ID.
"""
if samples < 1:
raise ValueError("Bootstrap samples must be positive")
if not 0.0 < confidence < 1.0:
raise ValueError("Bootstrap confidence must be between zero and one")
if len(edges) < 3:
raise ValueError("At least three scored pairs are required")
score_rank = stats.rankdata(edges["score"].to_numpy(dtype=float), method="average")
truth_rank = stats.rankdata(
edges["structural_similarity"].to_numpy(dtype=float), method="average"
)
endpoint_frame = pd.DataFrame(
{
"anchor": np.concatenate(
[
edges["record_a"].astype(str).to_numpy(),
edges["record_b"].astype(str).to_numpy(),
]
),
"weight": 0.5,
"x": np.tile(score_rank, 2),
"y": np.tile(truth_rank, 2),
}
)
endpoint_frame["x2"] = endpoint_frame["x"] ** 2
endpoint_frame["y2"] = endpoint_frame["y"] ** 2
endpoint_frame["xy"] = endpoint_frame["x"] * endpoint_frame["y"]
for column in ("x", "y", "x2", "y2", "xy"):
endpoint_frame[column] *= endpoint_frame["weight"]
blocks = (
endpoint_frame.groupby("anchor", sort=True)[
["weight", "x", "y", "x2", "y2", "xy"]
]
.sum()
.to_numpy(dtype=float)
)
if len(blocks) < 2:
raise ValueError("At least two BGC endpoint blocks are required")
random_state = np.random.default_rng(seed)
correlations: list[float] = []
for _ in range(samples):
selected = random_state.integers(0, len(blocks), size=len(blocks))
weight, sum_x, sum_y, sum_x2, sum_y2, sum_xy = blocks[selected].sum(axis=0)
covariance = sum_xy - (sum_x * sum_y / weight)
variance_x = sum_x2 - (sum_x * sum_x / weight)
variance_y = sum_y2 - (sum_y * sum_y / weight)
denominator = np.sqrt(max(variance_x, 0.0) * max(variance_y, 0.0))
if denominator > 0.0:
correlations.append(float(covariance / denominator))
if not correlations:
raise ValueError("No finite block-bootstrap correlations could be calculated")
tail = (1.0 - confidence) / 2.0
return tuple(float(value) for value in np.quantile(correlations, [tail, 1.0 - tail]))
def mark_cross_genus(edges: pd.DataFrame, metadata: pd.DataFrame) -> pd.DataFrame:
genus = metadata.loc[metadata["genus_count"] == 1].set_index("bgc_id")["genera"].to_dict()
result = edges.copy()
result["cross_genus"] = [
left in genus and right in genus and genus[left].lower() != genus[right].lower()
for left, right in zip(result.record_a, result.record_b)
]
return result
def exact_product_retrieval(
embeddings: Mapping[str, torch.Tensor],
gold_mapping: pd.DataFrame,
eligible_ids: set[str],
cutoff: int = 50,
) -> pd.DataFrame:
mapping = gold_mapping[gold_mapping["bgc_id"].isin(eligible_ids)].copy()
sizes = mapping.groupby("product_group_id")["bgc_id"].nunique()
mapping = mapping[mapping["product_group_id"].isin(sizes[sizes >= 2].index)]
universe = sorted(mapping["bgc_id"].unique())
group_by_bgc = mapping.set_index("bgc_id")["product_group_id"].to_dict()
genus_by_bgc = mapping.set_index("bgc_id")["genus"].astype(str).to_dict()
rows = []
for reference in universe:
candidates = [identifier for identifier in universe if identifier != reference]
relevant = {
identifier for identifier in candidates
if group_by_bgc[identifier] == group_by_bgc[reference]
}
if not relevant:
continue
scores = {
identifier: float(F.cosine_similarity(embeddings[reference], embeddings[identifier], dim=0))
for identifier in candidates
}
metrics = expected_tie_aware_metrics(scores, relevant, recall_at=(cutoff,), ndcg_at=(cutoff,))
cross_genus_relevant = {
identifier for identifier in relevant
if genus_by_bgc[identifier].lower() != genus_by_bgc[reference].lower()
}
rows.append(
{
"reference_id": reference,
"product_group_id": group_by_bgc[reference],
"reference_genus": genus_by_bgc[reference],
"cross_genus_positive_count": len(cross_genus_relevant),
**metrics,
}
)
return pd.DataFrame(rows)
def evaluate_similarity_method(
name: str,
edges: pd.DataFrame,
structure_matrix: pd.DataFrame,
metadata: pd.DataFrame,
bootstrap_samples: int,
confidence: float,
seed: int,
) -> tuple[pd.DataFrame, list[dict[str, object]]]:
scored = mark_cross_genus(attach_structural_truth(edges, structure_matrix), metadata)
summaries: list[dict[str, object]] = []
for subset_name, subset in (("all", scored), ("cross_genus", scored[scored["cross_genus"]])):
if len(subset) < 3:
continue
summary = spearman_summary(subset)
lower, upper = anchor_block_bootstrap(
subset, bootstrap_samples, confidence, seed
)
summaries.append(
{"method": name, "subset": subset_name, **summary,
"ci_lower": lower, "ci_upper": upper,
"bootstrap_unit": "two_endpoint_bgc_cluster_fixed_ranks"}
)
scored["method"] = name
return scored, summaries