File size: 2,656 Bytes
c87881a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
#!/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()