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