bgc-setnet / source /scripts /evaluate_residual.py
whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
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()