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