bgc-setnet / source /scripts /train_residual.py
whiteh4t's picture
Release final BGC retrieval checkpoints and model card
c87881a verified
Raw
History Blame Contribute Delete
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()