#!/usr/bin/env python3 from __future__ import annotations import argparse import json from pathlib import Path import pandas as pd from bgc_retrieval.artifacts import create_run_directory, write_json_immutable from bgc_retrieval.config import load_config from bgc_retrieval.data import BGCEmbeddingDataset from bgc_retrieval.model import ModelConfig from bgc_retrieval.residual import ( ResidualGeneWeightingEncoder, train_residual_gene_weighting, ) from bgc_retrieval.splits import load_split def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--config", default="configs/residual_griseus.yaml") 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 campaign 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"]) train_rows = assignments[assignments["split"] == "train"].copy() validation_rows = assignments[assignments["split"] == "validation"].copy() train_dataset = BGCEmbeddingDataset( embeddings_path, atlas_path, train_rows, model_config.esm_dimension ) validation_dataset = BGCEmbeddingDataset( embeddings_path, atlas_path, validation_rows, model_config.esm_dimension ) 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) } run_root = config.resolve_path("project", "run_root") run_dir = create_run_directory(run_root, args.run_id) write_json_immutable(run_dir / "config.json", values) model = ResidualGeneWeightingEncoder(model_config) checkpoint = train_residual_gene_weighting( model, model_config, train_dataset, validation_dataset, validation_rows, pfam_sets, split_path, [atlas_path, embeddings_path], run_dir, values["training"], values["residual"], int(values["project"]["seed"]), ) print(json.dumps({"checkpoint": str(checkpoint), "scope": values["scope"]}, indent=2)) if __name__ == "__main__": main()