| """Command-line entry points for the isolated rebuild.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| from pathlib import Path |
| from typing import Any |
|
|
| import pandas as pd |
| import torch |
|
|
| from .artifacts import build_file_manifest, create_run_directory, sha256_file, write_json_immutable |
| from .audit import audit_legacy_data |
| from .baselines import ( |
| aggregate_raw_esm, |
| cosine_scores, |
| ensemble_scores, |
| pfam_jaccard_scores, |
| select_ensemble_alpha, |
| weighted_pfam_jaccard_scores, |
| ) |
| from .checkpoints import load_checkpoint |
| from .config import LoadedConfig, load_config |
| from .data import BGCEmbeddingDataset, build_pfam_vocab, load_legacy_labels |
| from .evaluation import evaluate_retrieval |
| from .model import ModelConfig, build_model |
| from .reporting import write_paper_outputs |
| from .splits import build_group_split, load_split, split_summary |
| from .training import choose_device, encode_dataset, train_phase1, train_phase2, train_weighted_pfam |
|
|
|
|
| def _paths(config: LoadedConfig) -> dict[str, Path]: |
| return { |
| "atlas": config.resolve_path("data", "atlas_csv"), |
| "genes": config.resolve_path("data", "genes_csv"), |
| "h5": config.resolve_path("data", "embeddings_h5"), |
| "legacy_root": config.resolve_path("data", "legacy_root"), |
| "manifest_root": config.resolve_path("project", "manifest_root"), |
| "run_root": config.resolve_path("project", "run_root"), |
| } |
|
|
|
|
| def command_audit(args: argparse.Namespace) -> None: |
| config = load_config(args.config) |
| paths = _paths(config) |
| legacy_root = paths["legacy_root"] |
| audit = audit_legacy_data(paths["atlas"], paths["genes"], paths["h5"]) |
| candidates = [ |
| legacy_root / "community_atlas.csv", |
| legacy_root / "discovery_zone.csv", |
| legacy_root / "master_gene_annotations.csv", |
| legacy_root / "training_positive_pairs.csv", |
| legacy_root / "training_representatives.csv", |
| legacy_root / "esm2_embeddings.h5", |
| legacy_root / "paper/main.tex", |
| ] |
| candidates.extend(sorted((legacy_root / "models").glob("*.pt"))) |
| candidates.extend(sorted((legacy_root / "results/publication").glob("*"))) |
| manifest = build_file_manifest(candidates, legacy_root) |
| write_json_immutable(paths["manifest_root"] / "legacy_audit.json", audit) |
| write_json_immutable(paths["manifest_root"] / "legacy_manifest.json", manifest) |
| print(json.dumps(audit, indent=2, sort_keys=True)) |
|
|
|
|
| def command_build_splits(args: argparse.Namespace) -> None: |
| config = load_config(args.config) |
| paths = _paths(config) |
| values = config.values |
| labels = load_legacy_labels(paths["atlas"]) |
| data_config = values["data"] |
| assignments = build_group_split( |
| labels, |
| seed=int(values["project"]["seed"]), |
| minimum_group_size=int(data_config["minimum_group_size"]), |
| fractions=( |
| float(data_config["train_fraction"]), |
| float(data_config["validation_fraction"]), |
| float(data_config["test_fraction"]), |
| ), |
| ) |
| output = paths["manifest_root"] / "silver_split.csv" |
| summary_path = paths["manifest_root"] / "silver_split_summary.json" |
| output.parent.mkdir(parents=True, exist_ok=True) |
| if output.exists() or summary_path.exists(): |
| raise FileExistsError("Refusing to overwrite the frozen split; change its versioned name") |
| assignments.to_csv(output, index=False) |
| summary = split_summary(assignments) |
| summary["file_sha256"] = sha256_file(output) |
| summary["claim_scope"] = "development-only silver MIBiG-reference split" |
| write_json_immutable(summary_path, summary) |
| print(json.dumps(summary, indent=2, sort_keys=True)) |
|
|
|
|
| def _make_dataset( |
| config: LoadedConfig, |
| assignments: pd.DataFrame, |
| pfam_vocab: dict[str, int], |
| ) -> BGCEmbeddingDataset: |
| paths = _paths(config) |
| return BGCEmbeddingDataset( |
| paths["h5"], |
| paths["atlas"], |
| assignments, |
| esm_dimension=int(config.values["model"]["esm_dimension"]), |
| pfam_vocab=pfam_vocab, |
| ) |
|
|
|
|
| def _build_pfam_vocab(config: LoadedConfig, assignments: pd.DataFrame) -> dict[str, int]: |
| paths = _paths(config) |
| train_ids = assignments.loc[assignments["split"].eq("train"), "bgc_id"] |
| return build_pfam_vocab(paths["atlas"], train_ids) |
|
|
|
|
| def command_train(args: argparse.Namespace) -> None: |
| config = load_config(args.config) |
| paths = _paths(config) |
| split_path = Path(args.split).resolve() if args.split else paths["manifest_root"] / "silver_split.csv" |
| assignments = load_split(split_path) |
| train_assignments = assignments[assignments["split"] == "train"].copy() |
| validation_assignments = assignments[assignments["split"] == "validation"].copy() |
| pfam_vocab = _build_pfam_vocab(config, assignments) |
| train_dataset = _make_dataset(config, train_assignments, pfam_vocab) |
| validation_dataset = _make_dataset(config, validation_assignments, pfam_vocab) |
| run_dir = create_run_directory(paths["run_root"], args.run_id) |
| model_values = dict(config.values["model"]) |
| model_values["pfam_vocab_size"] = len(pfam_vocab) + 2 |
| config_values = dict(config.values) |
| config_values["model"] = model_values |
| write_json_immutable(run_dir / "config.json", config_values) |
| write_json_immutable(run_dir / "pfam_vocab.json", pfam_vocab) |
| model_config = ModelConfig.from_dict(model_values) |
| model = build_model(model_config) |
| input_paths = [paths["atlas"], paths["h5"]] |
| if model_config.architecture == "weighted_pfam_jaccard": |
| if args.stage != "phase2": |
| raise ValueError("Weighted Pfam Jaccard uses phase2 training directly") |
| checkpoint = train_weighted_pfam( |
| model, model_config, train_dataset, validation_dataset, validation_assignments, |
| split_path, input_paths, run_dir, config.values["training"], |
| int(config.values["project"]["seed"]), |
| ) |
| print(checkpoint) |
| return |
| if args.stage == "phase1": |
| checkpoint = train_phase1( |
| model, model_config, train_dataset, train_dataset, split_path, input_paths, |
| run_dir, config.values["training"], int(config.values["project"]["seed"]), |
| ) |
| else: |
| checkpoint = train_phase2( |
| model, model_config, train_dataset, validation_dataset, validation_assignments, |
| split_path, input_paths, run_dir, config.values["training"], |
| int(config.values["project"]["seed"]), args.phase1_checkpoint, |
| ) |
| print(checkpoint) |
|
|
|
|
| def _pfam_sets(atlas_path: Path) -> dict[str, set[str]]: |
| atlas = pd.read_csv(atlas_path, usecols=["bgc_id", "pfam_ids"]) |
| return { |
| 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) |
| } |
|
|
|
|
| def _raw_esm_embeddings( |
| dataset: BGCEmbeddingDataset, |
| ) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]: |
| mean: dict[str, torch.Tensor] = {} |
| maximum: dict[str, torch.Tensor] = {} |
| for index in range(len(dataset)): |
| item = dataset[index] |
| mean[item["bgc_id"]] = aggregate_raw_esm(item["gene_embeddings"], "mean") |
| maximum[item["bgc_id"]] = aggregate_raw_esm(item["gene_embeddings"], "max") |
| return mean, maximum |
|
|
|
|
| def command_evaluate(args: argparse.Namespace) -> None: |
| config = load_config(args.config) |
| paths = _paths(config) |
| split_path = Path(args.split).resolve() if args.split else paths["manifest_root"] / "silver_split.csv" |
| assignments = load_split(split_path) |
| pfam_vocab = _build_pfam_vocab(config, assignments) |
| model_values = dict(config.values["model"]) |
| model_values["pfam_vocab_size"] = len(pfam_vocab) + 2 |
| model_config = ModelConfig.from_dict(model_values) |
| model = build_model(model_config) |
| load_checkpoint(args.checkpoint, model, split_path) |
| device = choose_device() |
| model.to(device) |
| dataset = _make_dataset(config, assignments, pfam_vocab) |
| pfams = _pfam_sets(paths["atlas"]) |
| if model_config.architecture == "weighted_pfam_jaccard": |
| learned_weights = model.domain_weights().detach().cpu().tolist() |
| inverse_vocab = {index: token for token, index in pfam_vocab.items()} |
| weights_by_pfam = { |
| token: float(learned_weights[index]) |
| for index, token in inverse_vocab.items() |
| if index < len(learned_weights) |
| } |
| unknown_weight = float(learned_weights[1]) |
|
|
| def weighted_jaccard(candidates: list[str], references: list[str]) -> dict[str, float]: |
| return weighted_pfam_jaccard_scores( |
| candidates, references, pfams, weights_by_pfam, unknown_weight, "max" |
| ) |
|
|
| def plain_jaccard(candidates: list[str], references: list[str]) -> dict[str, float]: |
| return pfam_jaccard_scores(candidates, references, pfams, "max") |
|
|
| methods = { |
| "pfam_jaccard_max": plain_jaccard, |
| "weighted_pfam_jaccard": weighted_jaccard, |
| } |
| evaluation_config = config.values["evaluation"] |
| evaluation_seed = int(evaluation_config.get("seed", config.values["project"]["seed"])) |
| test_results = evaluate_retrieval( |
| assignments, "test", methods, int(evaluation_config["reference_size"]), |
| int(evaluation_config["query_draws"]), evaluation_seed, |
| evaluation_config["recall_at"], evaluation_config["ndcg_at"], |
| ) |
| run_dir = create_run_directory(paths["run_root"], args.run_id) |
| metrics = [evaluation_config["primary_metric"], "mrr", "map", "ndcg@50", "tie_fraction"] |
| metadata = { |
| "schema_version": 1, |
| "selected_alpha": None, |
| "alpha_selected_on": None, |
| "split_sha256": sha256_file(split_path), |
| "checkpoint_sha256": sha256_file(args.checkpoint), |
| "label_tiers": sorted(assignments["label_tier"].unique()), |
| "pfam_feature": "BGC-level pfam_ids inventory; vocabulary built from training BGCs only", |
| "weighted_pfam_feature": "trainable nonnegative domain weights optimized with differentiable weighted Jaccard", |
| "pfam_vocab_size": len(pfam_vocab) + 2, |
| "publication_eligible": bool((assignments["label_tier"] == "gold").any()), |
| "claim_warning": "Silver-only results are development evidence, not the primary biological result.", |
| } |
| write_paper_outputs( |
| run_dir, test_results, metrics, int(evaluation_config["bootstrap_samples"]), |
| float(evaluation_config["confidence_level"]), evaluation_seed, metadata, |
| ) |
| pd.DataFrame().to_csv(run_dir / "validation_alpha_search.csv", index=False) |
| print(json.dumps(metadata, indent=2, sort_keys=True)) |
| return |
| learned = encode_dataset(model, dataset, device, int(config.values["training"]["num_workers"])) |
| raw_mean, raw_max = _raw_esm_embeddings(dataset) |
| pfams = _pfam_sets(paths["atlas"]) |
| evaluation_config = config.values["evaluation"] |
| evaluation_seed = int( |
| evaluation_config.get("seed", config.values["project"]["seed"]) |
| ) |
|
|
| def jaccard(candidates: list[str], references: list[str]) -> dict[str, float]: |
| return pfam_jaccard_scores(candidates, references, pfams, "max") |
|
|
| def learned_cosine(candidates: list[str], references: list[str]) -> dict[str, float]: |
| return cosine_scores(candidates, references, learned, "mean") |
|
|
| validation_methods: dict[str, Any] = {"pfam_jaccard": jaccard, "setnet": learned_cosine} |
| for alpha in evaluation_config["ensemble_alphas"]: |
| validation_methods[f"ensemble_a{float(alpha):.1f}"] = ( |
| lambda candidates, references, weight=float(alpha): ensemble_scores( |
| learned_cosine(candidates, references), jaccard(candidates, references), weight |
| ) |
| ) |
| validation = evaluate_retrieval( |
| assignments, "validation", validation_methods, int(evaluation_config["reference_size"]), |
| int(evaluation_config["query_draws"]), evaluation_seed, |
| evaluation_config["recall_at"], evaluation_config["ndcg_at"], |
| ) |
| ensemble_validation = validation[validation["method"].str.startswith("ensemble_a")].copy() |
| ensemble_validation["alpha"] = ensemble_validation["method"].str.replace( |
| "ensemble_a", "", regex=False |
| ).astype(float) |
| alpha = select_ensemble_alpha(ensemble_validation, evaluation_config["primary_metric"]) |
|
|
| methods: dict[str, Any] = { |
| "pfam_jaccard_max": jaccard, |
| "pfam_jaccard_mean": lambda candidates, references: pfam_jaccard_scores( |
| candidates, references, pfams, "mean" |
| ), |
| "raw_esm_mean": lambda candidates, references: cosine_scores( |
| candidates, references, raw_mean, "mean" |
| ), |
| "raw_esm_max": lambda candidates, references: cosine_scores( |
| candidates, references, raw_max, "mean" |
| ), |
| "setnet": learned_cosine, |
| f"ensemble_validation_alpha_{alpha:.1f}": lambda candidates, references: ensemble_scores( |
| learned_cosine(candidates, references), jaccard(candidates, references), alpha |
| ), |
| } |
| test_results = evaluate_retrieval( |
| assignments, "test", methods, int(evaluation_config["reference_size"]), |
| int(evaluation_config["query_draws"]), evaluation_seed, |
| evaluation_config["recall_at"], evaluation_config["ndcg_at"], |
| ) |
| run_dir = create_run_directory(paths["run_root"], args.run_id) |
| metrics = [evaluation_config["primary_metric"], "mrr", "map", "ndcg@50", "tie_fraction"] |
| metadata = { |
| "schema_version": 1, |
| "selected_alpha": alpha, |
| "alpha_selected_on": "validation", |
| "split_sha256": sha256_file(split_path), |
| "checkpoint_sha256": sha256_file(args.checkpoint), |
| "label_tiers": sorted(assignments["label_tier"].unique()), |
| "pfam_feature": "BGC-level pfam_ids inventory; vocabulary built from training BGCs only", |
| "pfam_vocab_size": len(pfam_vocab) + 2, |
| "publication_eligible": bool((assignments["label_tier"] == "gold").any()), |
| "claim_warning": "Silver-only results are development evidence, not the primary biological result.", |
| } |
| write_paper_outputs( |
| run_dir, test_results, metrics, int(evaluation_config["bootstrap_samples"]), |
| float(evaluation_config["confidence_level"]), evaluation_seed, |
| metadata, |
| ) |
| validation.to_csv(run_dir / "validation_alpha_search.csv", index=False) |
| print(json.dumps(metadata, indent=2, sort_keys=True)) |
|
|
|
|
| def build_parser() -> argparse.ArgumentParser: |
| parser = argparse.ArgumentParser(prog="bgc-rebuild") |
| subparsers = parser.add_subparsers(dest="command", required=True) |
| for name, function in (("audit", command_audit), ("build-splits", command_build_splits)): |
| subparser = subparsers.add_parser(name) |
| subparser.add_argument("--config", default="configs/main.yaml") |
| subparser.set_defaults(function=function) |
| train = subparsers.add_parser("train") |
| train.add_argument("--config", default="configs/main.yaml") |
| train.add_argument("--stage", choices=("phase1", "phase2"), required=True) |
| train.add_argument("--run-id", required=True) |
| train.add_argument("--split") |
| train.add_argument("--phase1-checkpoint") |
| train.set_defaults(function=command_train) |
| evaluate = subparsers.add_parser("evaluate") |
| evaluate.add_argument("--config", default="configs/main.yaml") |
| evaluate.add_argument("--checkpoint", required=True) |
| evaluate.add_argument("--run-id", required=True) |
| evaluate.add_argument("--split") |
| evaluate.set_defaults(function=command_evaluate) |
| return parser |
|
|
|
|
| def main() -> None: |
| args = build_parser().parse_args() |
| args.function(args) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|