whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
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()