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