#!/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()