File size: 6,544 Bytes
c87881a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 | #!/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()
|