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