Download source/src/bgc_retrieval/cli.py from rustambekurokov/bgc-setnet: direct link, hf CLI and curl.
- Browser
- Download file 15.8 kB
-
https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/src/bgc_retrieval/cli.py
- Command line
-
hf download hf://rustambekurokov/bgc-setnet/source/src/bgc_retrieval/cli.py
-
curl -L -o cli.py https://huggingface.co/rustambekurokov/bgc-setnet/resolve/main/source/src/bgc_retrieval/cli.py
15.8 kB
| """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() | |