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