#!/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()