Download source/scripts/evaluate_residual.py from rustambekurokov/bgc-setnet: direct link, hf CLI and curl.
- Browser
- Download file 6.68 kB
-
https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/scripts/evaluate_residual.py
- Command line
-
hf download hf://rustambekurokov/bgc-setnet/source/scripts/evaluate_residual.py
-
curl -L -o evaluate_residual.py https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/scripts/evaluate_residual.py
6.68 kB
| #!/usr/bin/env python3 | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| from typing import Any | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| from bgc_retrieval.artifacts import create_run_directory, sha256_file, write_json_immutable | |
| from bgc_retrieval.baselines import cosine_scores, pfam_jaccard_scores | |
| from bgc_retrieval.checkpoints import load_checkpoint | |
| from bgc_retrieval.config import load_config | |
| from bgc_retrieval.data import BGCEmbeddingDataset | |
| from bgc_retrieval.evaluation import evaluate_retrieval | |
| from bgc_retrieval.model import ModelConfig | |
| from bgc_retrieval.reporting import write_paper_outputs | |
| from bgc_retrieval.residual import ( | |
| ResidualGeneWeightingEncoder, | |
| RetrievalScoreCache, | |
| encode_residual_components, | |
| select_validation_weights, | |
| validation_grid, | |
| ) | |
| from bgc_retrieval.splits import load_split | |
| from bgc_retrieval.training import choose_device | |
| def main() -> None: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--config", default="configs/residual_griseus.yaml") | |
| parser.add_argument("--checkpoint", action="append", required=True) | |
| parser.add_argument("--run-id", required=True) | |
| parser.add_argument("--split", default="data/manifests/silver_split.csv") | |
| args = parser.parse_args() | |
| config = load_config(args.config) | |
| values = config.values | |
| if values["scope"]["organism"] != "Streptomyces griseus": | |
| raise ValueError("Residual evaluation is locked to Streptomyces griseus") | |
| split_path = Path(args.split).resolve() | |
| assignments = load_split(split_path) | |
| atlas_path = config.resolve_path("data", "atlas_csv") | |
| embeddings_path = config.resolve_path("data", "embeddings_h5") | |
| model_config = ModelConfig.from_dict(values["model"]) | |
| dataset = BGCEmbeddingDataset( | |
| embeddings_path, atlas_path, assignments, model_config.esm_dimension | |
| ) | |
| device = choose_device() | |
| raw_embeddings: dict[str, torch.Tensor] | None = None | |
| learned_models: list[dict[str, torch.Tensor]] = [] | |
| for checkpoint in args.checkpoint: | |
| model = ResidualGeneWeightingEncoder(model_config) | |
| load_checkpoint(checkpoint, model, split_path) | |
| model.to(device) | |
| raw, learned = encode_residual_components( | |
| model, | |
| dataset, | |
| device, | |
| int(values["training"]["num_workers"]), | |
| ) | |
| if raw_embeddings is None: | |
| raw_embeddings = raw | |
| learned_models.append(learned) | |
| if raw_embeddings is None: | |
| raise RuntimeError("No residual checkpoints were loaded") | |
| atlas = pd.read_csv(atlas_path, usecols=["bgc_id", "pfam_ids"]) | |
| pfam_sets = { | |
| str(row.bgc_id): set(str(row.pfam_ids).split(";")) | |
| if pd.notna(row.pfam_ids) | |
| else set() | |
| for row in atlas.itertuples(index=False) | |
| } | |
| evaluation = values["evaluation"] | |
| residual = values["residual"] | |
| evaluation_seed = int(evaluation["seed"]) | |
| alphas = [float(value) for value in residual["alpha_grid"]] | |
| betas = [float(value) for value in residual["pfam_beta_grid"]] | |
| validation_cache = RetrievalScoreCache(raw_embeddings, learned_models, pfam_sets) | |
| validation = validation_grid( | |
| assignments, | |
| validation_cache, | |
| alphas, | |
| betas, | |
| int(evaluation["reference_size"]), | |
| int(evaluation["query_draws"]), | |
| evaluation_seed, | |
| evaluation["recall_at"], | |
| evaluation["ndcg_at"], | |
| ) | |
| metric = str(evaluation["primary_metric"]) | |
| residual_alpha, _ = select_validation_weights(validation, metric, "residual_") | |
| hybrid_alpha, hybrid_beta = select_validation_weights(validation, metric, "hybrid_") | |
| if hybrid_beta is None: | |
| raise RuntimeError("Hybrid validation did not select beta") | |
| test_cache = RetrievalScoreCache(raw_embeddings, learned_models, pfam_sets) | |
| def raw_score(candidates: list[str], references: list[str]) -> dict[str, float]: | |
| return cosine_scores(candidates, references, raw_embeddings, "mean") | |
| def learned_score(candidates: list[str], references: list[str]) -> dict[str, float]: | |
| _, learned, _ = test_cache.components(candidates, references) | |
| return learned | |
| def pfam_score(candidates: list[str], references: list[str]) -> dict[str, float]: | |
| return pfam_jaccard_scores(candidates, references, pfam_sets, "max") | |
| methods: dict[str, Any] = { | |
| "pfam_jaccard_max": pfam_score, | |
| "raw_esm_mean": raw_score, | |
| "weighted_gene_esm": learned_score, | |
| f"residual_validation_alpha_{residual_alpha:g}": ( | |
| lambda candidates, references: test_cache.residual( | |
| candidates, references, residual_alpha | |
| ) | |
| ), | |
| f"residual_pfam_validation_a{hybrid_alpha:g}_b{hybrid_beta:g}": ( | |
| lambda candidates, references: test_cache.hybrid( | |
| candidates, references, hybrid_alpha, hybrid_beta | |
| ) | |
| ), | |
| } | |
| test_results = evaluate_retrieval( | |
| assignments, | |
| "test", | |
| methods, | |
| int(evaluation["reference_size"]), | |
| int(evaluation["query_draws"]), | |
| evaluation_seed, | |
| evaluation["recall_at"], | |
| evaluation["ndcg_at"], | |
| ) | |
| run_dir = create_run_directory( | |
| config.resolve_path("project", "run_root"), | |
| args.run_id, | |
| ) | |
| metadata = { | |
| "schema_version": 1, | |
| "organism_scope": values["scope"]["organism"], | |
| "task_scope": values["scope"]["task"], | |
| "analysis_status": "post_hoc_redesign_pilot", | |
| "selected_residual_alpha": residual_alpha, | |
| "selected_hybrid_alpha": hybrid_alpha, | |
| "selected_pfam_beta": hybrid_beta, | |
| "weights_selected_on": "validation_only", | |
| "checkpoint_sha256": { | |
| str(path): sha256_file(path) for path in args.checkpoint | |
| }, | |
| "split_sha256": sha256_file(split_path), | |
| "label_tier": "silver", | |
| "publication_eligible": False, | |
| "scope_warning": ( | |
| "The atlas provenance states Streptomyces griseus, but the atlas table " | |
| "does not contain a machine-verifiable species column." | |
| ), | |
| } | |
| write_paper_outputs( | |
| run_dir, | |
| test_results, | |
| [metric, "mrr", "map", "ndcg@50", "tie_fraction"], | |
| int(evaluation["bootstrap_samples"]), | |
| float(evaluation["confidence_level"]), | |
| evaluation_seed, | |
| metadata, | |
| ) | |
| validation.to_csv(run_dir / "validation_weight_search.csv", index=False) | |
| write_json_immutable(run_dir / "config.json", values) | |
| print(json.dumps(metadata, indent=2, sort_keys=True)) | |
| if __name__ == "__main__": | |
| main() | |