bgc-setnet / source /scripts /evaluate_external.py
whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
6.54 kB
#!/usr/bin/env python3
from __future__ import annotations
import argparse
import json
from pathlib import Path
import pandas as pd
import torch
from bgc_retrieval.artifacts import create_run_directory, sha256_file, write_json_immutable
from bgc_retrieval.baselines import aggregate_raw_esm
from bgc_retrieval.checkpoints import load_checkpoint
from bgc_retrieval.config import load_config
from bgc_retrieval.data import BGCEmbeddingDataset
from bgc_retrieval.external import load_released_embeddings
from bgc_retrieval.external_evaluation import (
all_pair_scores,
eligible_external_ids,
evaluate_similarity_method,
exact_product_retrieval,
load_structure_matrix,
score_requested_pairs,
)
from bgc_retrieval.model import LeakageFreeBGCSetNet, ModelConfig
from bgc_retrieval.splits import load_split
from bgc_retrieval.training import choose_device, encode_dataset
def parse_released(values: list[str]) -> dict[str, dict[str, torch.Tensor]]:
result = {}
for value in values:
if "=" not in value:
raise ValueError("Released embeddings use NAME=PATH syntax")
name, path = value.split("=", 1)
result[name] = load_released_embeddings(path)
return result
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", default="configs/main.yaml")
parser.add_argument("--checkpoint", required=True)
parser.add_argument("--run-id", required=True)
parser.add_argument("--split", default="data/manifests/silver_split.csv")
parser.add_argument("--processed-dir", default="data/external/processed")
parser.add_argument(
"--structure-matrix",
default="data/external/bgc-clustering-benchmark/tanimoto_results/NPAtlas_bm_v1.tsv",
)
parser.add_argument(
"--bigscape-edges",
default="data/external/bgc-clustering-benchmark/bgc_similarities/bigscape_similarity_score_1.csv",
)
parser.add_argument("--released-embedding", action="append", default=[])
parser.add_argument("--bootstrap-samples", type=int, default=1000)
args = parser.parse_args()
config = load_config(args.config)
split = load_split(args.split)
processed = Path(args.processed_dir)
external_atlas = pd.read_csv(processed / "external_atlas.csv")
assignments = pd.DataFrame(
{
"bgc_id": external_atlas["bgc_id"].astype(str),
"group_id": external_atlas["bgc_id"].astype(str),
"split": "test",
"label_tier": "gold",
}
)
model_config = ModelConfig.from_dict(config.values["model"])
model = LeakageFreeBGCSetNet(model_config)
load_checkpoint(args.checkpoint, model, args.split)
device = choose_device()
model.to(device)
dataset = BGCEmbeddingDataset(
processed / "external_esm2.h5", processed / "external_atlas.csv",
assignments, model_config.esm_dimension,
)
learned = encode_dataset(model, dataset, device, int(config.values["training"]["num_workers"]))
raw_esm = {
dataset[index]["bgc_id"]: aggregate_raw_esm(dataset[index]["gene_embeddings"], "mean")
for index in range(len(dataset))
}
methods = {"setnet": learned, "raw_esm_mean": raw_esm, **parse_released(args.released_embedding)}
matrix = load_structure_matrix(args.structure_matrix)
common_ids = set.intersection(*(set(values) for values in methods.values()))
identifiers = eligible_external_ids(matrix, common_ids, split)
metadata = pd.read_csv(processed / "all_bgc_product_metadata.csv")
gold = pd.read_csv(processed / "gold_bgc_product_mapping.csv")
evaluation_config = config.values["evaluation"]
evaluation_seed = int(
evaluation_config.get("seed", config.values["project"]["seed"])
)
run_root = config.resolve_path("project", "run_root")
run_dir = create_run_directory(run_root, args.run_id)
summaries = []
pair_frames = []
for name, embeddings in methods.items():
edges = all_pair_scores(embeddings, identifiers)
scored, summary = evaluate_similarity_method(
name, edges, matrix, metadata, args.bootstrap_samples,
float(evaluation_config["confidence_level"]), evaluation_seed,
)
pair_frames.append(scored)
summaries.extend(summary)
retrieval = exact_product_retrieval(embeddings, gold, set(identifiers), cutoff=50)
retrieval.to_csv(run_dir / f"{name}_exact_product_retrieval.csv", index=False)
bigscape = pd.read_csv(
args.bigscape_edges, header=None, names=["record_a", "record_b", "score"]
)
bigscape = bigscape[
bigscape["record_a"].isin(identifiers) & bigscape["record_b"].isin(identifiers)
]
scored, summary = evaluate_similarity_method(
"bigscape", bigscape, matrix, metadata, args.bootstrap_samples,
float(evaluation_config["confidence_level"]), evaluation_seed,
)
pair_frames.append(scored)
summaries.extend(summary)
for name, embeddings in methods.items():
same_edges = score_requested_pairs(embeddings, bigscape)
scored, summary = evaluate_similarity_method(
f"{name}_on_bigscape_edges", same_edges, matrix, metadata,
args.bootstrap_samples, float(evaluation_config["confidence_level"]),
evaluation_seed,
)
pair_frames.append(scored)
summaries.extend(summary)
pd.concat(pair_frames, ignore_index=True).to_csv(run_dir / "external_pair_scores.csv", index=False)
pd.DataFrame(summaries).to_csv(run_dir / "external_similarity_summary.csv", index=False)
lineage = {
"schema_version": 1,
"eligible_external_bgcs": len(identifiers),
"blocked_train_validation_references": int(
split[split["split"].isin(["train", "validation"])]["group_id"].nunique()
),
"checkpoint_sha256": sha256_file(args.checkpoint),
"split_sha256": sha256_file(args.split),
"structure_matrix_sha256": sha256_file(args.structure_matrix),
"bigscape_edges_sha256": sha256_file(args.bigscape_edges),
"primary_external_endpoint": "Spearman correlation with product-structure Tanimoto",
"pair_uncertainty_unit": (
"two-endpoint BGC cluster bootstrap on fixed full-sample ranks"
),
"bootstrap_samples": args.bootstrap_samples,
}
write_json_immutable(run_dir / "external_metadata.json", lineage)
print(json.dumps(lineage, indent=2, sort_keys=True))
if __name__ == "__main__":
main()