Download source/scripts/evaluate_external.py from rustambekurokov/bgc-setnet: direct link, hf CLI and curl.
- Browser
- Download file 6.54 kB
-
https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/scripts/evaluate_external.py
- Command line
-
hf download hf://rustambekurokov/bgc-setnet/source/scripts/evaluate_external.py
-
curl -L -o evaluate_external.py https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/scripts/evaluate_external.py
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() | |