"""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()