File size: 2,761 Bytes
986e0b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
#!/usr/bin/env python3
"""Train one or two Saluki heads with the bundled official trainer."""

from __future__ import annotations

import argparse
import json
from pathlib import Path

from _bootstrap import configure_runtime


ROOT = configure_runtime()


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(
        description="Train Saluki on one or two official RNA TFRecord datasets."
    )
    parser.add_argument(
        "data_dirs", nargs="+", type=Path,
        help="Dataset directories containing statistics.json and tfrecords/."
    )
    parser.add_argument(
        "--params", type=Path, default=ROOT / "conf" / "params.json",
        help="Model/training parameter JSON (default: conf/params.json)."
    )
    parser.add_argument(
        "--out-dir", type=Path, default=ROOT / "output" / "train",
        help="Checkpoint and log output directory."
    )
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    missing = [path for path in args.data_dirs if not (path / "statistics.json").is_file()]
    if missing:
        joined = ", ".join(str(path) for path in missing)
        raise FileNotFoundError(f"Missing statistics.json in: {joined}")
    if not args.params.is_file():
        raise FileNotFoundError(args.params)

    from model.basenji import dataset, rnann, trainer
    from model.saluki import load_params

    params = load_params(args.params)
    params_model = params["model"]
    params_train = params["train"]

    if len(args.data_dirs) != len(params_model["num_targets"]):
        raise ValueError(
            "The number of data directories must match the number of model "
            f"heads ({len(params_model['num_targets'])})."
        )

    args.out_dir.mkdir(parents=True, exist_ok=True)
    (args.out_dir / "params.json").write_text(
        json.dumps(params, indent=2) + "\n", encoding="utf-8"
    )

    train_data = [
        dataset.RnaDataset(
            str(data_dir), split_label="train",
            batch_size=params_train["batch_size"],
            shuffle_buffer=params_train.get("shuffle_buffer", 1024),
            mode="train"
        )
        for data_dir in args.data_dirs
    ]
    eval_data = [
        dataset.RnaDataset(
            str(data_dir), split_label="valid",
            batch_size=params_train["batch_size"], mode="eval"
        )
        for data_dir in args.data_dirs
    ]

    network = rnann.RnaNN(params_model)
    network_trainer = trainer.Trainer(
        params_train, train_data, eval_data, str(args.out_dir)
    )
    network_trainer.compile(network)
    if len(args.data_dirs) == 1:
        network_trainer.fit_tape(network)
    else:
        network_trainer.fit2(network)


if __name__ == "__main__":
    main()