| |
| 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() |
|
|