diff --git a/pipeline/config/config.yaml b/pipeline/config/config.yaml index 23a8fdd44..046a51ce0 100644 --- a/pipeline/config/config.yaml +++ b/pipeline/config/config.yaml @@ -115,8 +115,11 @@ resources: # Slurm GPU partition for training. train_partition: gpuA40x4 -# Number of GPUs to request for training. Must match trainer.devices in -# your train.yaml. +# GPUs for training, which override trainer.devices in the train config. +# On a shared node, pin specific devices with train_gpus (comma-separated +# IDs, e.g. "2" or "2,3"), since Lightning would otherwise take every +# visible GPU. Leave it null under slurm, which assigns train_num_gpus. +train_gpus: null train_num_gpus: 1 # Slurm GPU partition for export and inference. @@ -195,7 +198,7 @@ model_version: -1 # Resolved against $AFRAME_CONTAINER_ROOT unless given as an absolute path. triton_image: tritonserver_25.06.sif # Comma-separated GPU IDs for export and inference server only. -# Pin training GPUs via trainer.devices in the train_config. +# Training uses train_gpus. gpus: "0" # Snapshotter instances hosted per GPU. Baked into the model at export time diff --git a/projects/export/export/cli.py b/projects/export/export/cli.py index 955d3774a..8300ef669 100644 --- a/projects/export/export/cli.py +++ b/projects/export/export/cli.py @@ -18,7 +18,7 @@ def main(args=None): parser = build_parser() args = parser.parse_args(args) logfile = args.pop("logfile") - args = parser.instantiate_classes(args) + args = parser.instantiate(args) if logfile is not None: logdir = os.path.dirname(logfile) os.makedirs(logdir, exist_ok=True) diff --git a/projects/export/pyproject.toml b/projects/export/pyproject.toml index b12da2838..382c7f6e9 100644 --- a/projects/export/pyproject.toml +++ b/projects/export/pyproject.toml @@ -13,7 +13,7 @@ dependencies = [ "s3fs>=2024", "ml4gw-hermes[torch]", "utils[torch]", - "jsonargparse[signatures]>=4.29,<5", + "jsonargparse[signatures]>=4.49,<5", "h5py>=3.10,<4", "nvidia-cudnn-cu12", # CUDA-12 TensorRT build. The default `tensorrt` now resolves to CUDA-13, diff --git a/projects/online/online/cli.py b/projects/online/online/cli.py index 2177c0916..f0cb60f8f 100644 --- a/projects/online/online/cli.py +++ b/projects/online/online/cli.py @@ -30,7 +30,7 @@ def cli(args=None): parser = build_parser() args = parser.parse_args(args) args.pop("config") - args = parser.instantiate_classes(args) + args = parser.instantiate(args) main(**vars(args)) diff --git a/projects/online/pyproject.toml b/projects/online/pyproject.toml index 816bc4f94..a918486f9 100644 --- a/projects/online/pyproject.toml +++ b/projects/online/pyproject.toml @@ -31,7 +31,7 @@ dependencies = [ "ligo-gracedb[kafka]>=2.15.4", "tables>=3.9", "gwpy>=4.0,<5", - "jsonargparse[signatures]>=4.29,<5", + "jsonargparse[signatures]>=4.49,<5", "psutil>=5.9", "pytz", "certifi", diff --git a/projects/plots/plots/vizapp/main.py b/projects/plots/plots/vizapp/main.py index d416bd0ed..367876500 100644 --- a/projects/plots/plots/vizapp/main.py +++ b/projects/plots/plots/vizapp/main.py @@ -18,7 +18,7 @@ def cli(args=None): parser.add_argument("--config", action="config") parser.add_function_arguments(main) args = parser.parse_args() - args = parser.instantiate_classes(args) + args = parser.instantiate(args) args.pop("config", None) main(**vars(args)) diff --git a/projects/plots/pyproject.toml b/projects/plots/pyproject.toml index 7f8c067d6..21baf10ca 100644 --- a/projects/plots/pyproject.toml +++ b/projects/plots/pyproject.toml @@ -22,7 +22,7 @@ dependencies = [ # imports `lal` without declaring it, so lalsuite has to come along too. "igwn-ligolw>=2.1,<3", "lalsuite", - "jsonargparse[signatures]>=4.29,<5", + "jsonargparse[signatures]>=4.49,<5", ] # Make the vizapp dependencies fully optional. diff --git a/projects/train/configs/bbh.yaml b/projects/train/configs/bbh.yaml index 28ca602b3..a0236026c 100644 --- a/projects/train/configs/bbh.yaml +++ b/projects/train/configs/bbh.yaml @@ -14,9 +14,9 @@ model: init_args: layers: [3, 4, 6, 3] norm_layer: - class_path: ml4gw.nn.norm.GroupNorm1DGetter + class_path: ml4gw.nn.norm.GroupNorm1D init_args: - groups: 16 + num_groups: 16 metric: class_path: train.metrics.TimeSlideAUROC init_args: diff --git a/projects/train/configs/bns_time_spectrogram.yaml b/projects/train/configs/bns_time_spectrogram.yaml index b95a8dd8f..a73287b97 100644 --- a/projects/train/configs/bns_time_spectrogram.yaml +++ b/projects/train/configs/bns_time_spectrogram.yaml @@ -16,9 +16,9 @@ model: time_kernel_size: 25 time_classes: 1 time_norm_layer: - class_path: ml4gw.nn.norm.GroupNorm1DGetter + class_path: ml4gw.nn.norm.GroupNorm1D init_args: - groups: 16 + num_groups: 16 spec_layers: [3, 3, 2, 2] spec_kernel_size: 3 spec_classes: 1 diff --git a/projects/train/configs/time_heterodyne.yaml b/projects/train/configs/time_heterodyne.yaml index d2928f781..928ad26c5 100644 --- a/projects/train/configs/time_heterodyne.yaml +++ b/projects/train/configs/time_heterodyne.yaml @@ -15,9 +15,9 @@ model: layers: [3, 4, 6, 3] kernel_size: 25 norm_layer: - class_path: ml4gw.nn.norm.GroupNorm1DGetter + class_path: ml4gw.nn.norm.GroupNorm1D init_args: - groups: 16 + num_groups: 16 metric: class_path: train.metrics.TimeSlideAUROC init_args: diff --git a/projects/train/pyproject.toml b/projects/train/pyproject.toml index aefe94c01..c48755de6 100644 --- a/projects/train/pyproject.toml +++ b/projects/train/pyproject.toml @@ -14,7 +14,7 @@ dependencies = [ "torchmetrics>=1.0,<2", "torch==2.10.0", "lightning>=2.2.1", - "jsonargparse[signatures]>=4.29,<5", + "jsonargparse[signatures]>=4.49,<5", "wandb>=0.15", "ray[default, tune]>=2.8.0,<3", "lightray>=0.2.3", diff --git a/projects/train/train.smk b/projects/train/train.smk index c418ebad2..4eadd4a55 100644 --- a/projects/train/train.smk +++ b/projects/train/train.smk @@ -23,6 +23,17 @@ train_log_dir = log_dir / "train" TRAIN_CONTAINER = os.path.join(os.getenv("AFRAME_CONTAINER_ROOT", ""), "train.sif") +# GPUs for local training. `train_gpus` pins specific devices on a shared +# node; otherwise use `train_num_gpus` of whatever is visible, which under +# slurm is the allocation. +if config.get("train_gpus") is not None: + TRAIN_GPU_ENV = f"CUDA_VISIBLE_DEVICES={config['train_gpus']} " + TRAIN_NUM_GPUS = len(str(config["train_gpus"]).split(",")) +else: + TRAIN_GPU_ENV = "" + TRAIN_NUM_GPUS = config.get("train_num_gpus", 1) + + def _train_waveform_inputs(wildcards): """Pre-generated training waveforms, when enabled.""" if config.get("pregenerate_training_waveforms", False): @@ -98,7 +109,8 @@ else: snakemake is invoked. AFRAME_TRAIN_WAVEFORMS_DIR is exported so the config resolves - to this run's validation/training waveform files. + to this run's validation/training waveform files. The GPU count + overrides trainer.devices in the train config. """ input: background=get_train_background_files, @@ -116,12 +128,16 @@ else: resources: **rule_resources("train"), slurm_partition=config.get("train_partition", "gpuA40x4"), - gpu=config.get("train_num_gpus", 1), + gpu=TRAIN_NUM_GPUS, params: **train_data_params, + gpu_env=TRAIN_GPU_ENV, + num_gpus=TRAIN_NUM_GPUS, background_dir=str(train_bg), waveforms_dir=str(train_waveforms), save_dir=str(train_out), shell: - "AFRAME_TRAIN_WAVEFORMS_DIR={params.waveforms_dir}" - " python -m train fit" + train_cli_args + " &> {log}" + "{params.gpu_env}AFRAME_TRAIN_WAVEFORMS_DIR={params.waveforms_dir}" + " python -m train fit" + + train_cli_args + + " --trainer.devices {params.num_gpus} &> {log}" diff --git a/projects/train/train.yaml b/projects/train/train.yaml index 8d2e679ea..7b95714ca 100644 --- a/projects/train/train.yaml +++ b/projects/train/train.yaml @@ -14,9 +14,9 @@ model: init_args: layers: [3, 4, 6, 3] norm_layer: - class_path: ml4gw.nn.norm.GroupNorm1DGetter + class_path: ml4gw.nn.norm.GroupNorm1D init_args: - groups: 16 + num_groups: 16 metric: class_path: train.metrics.TimeSlideAUROC init_args: @@ -142,8 +142,7 @@ trainer: # class_path: lightning.pytorch.profilers.PyTorchProfiler # dict_kwargs: # profile_memory: true - # devices: - # strategy: set to ddp if len(devices) > 1 + # devices: set by the pipeline from train_gpus / train_num_gpus #precision: 16-mixed accelerator: auto max_epochs: 200 diff --git a/projects/train/train/remote.py b/projects/train/train/remote.py index 125336c3c..9e33d3943 100644 --- a/projects/train/train/remote.py +++ b/projects/train/train/remote.py @@ -287,7 +287,7 @@ def main(args=None): Path(cfg.logfile).parent.mkdir(parents=True, exist_ok=True) configure_logging(cfg.logfile, cfg.verbose) - cfg = parser.instantiate_classes(cfg) + cfg = parser.instantiate(cfg) cfg.trainer.run(cfg.train_args) diff --git a/uv.lock b/uv.lock index 03271d2b1..745d5beaa 100644 --- a/uv.lock +++ b/uv.lock @@ -1791,7 +1791,7 @@ requires-dist = [ { name = "boto3", specifier = "~=1.30" }, { name = "fsspec", extras = ["s3"], specifier = ">=2024" }, { name = "h5py", specifier = ">=3.10,<4" }, - { name = "jsonargparse", extras = ["signatures"], specifier = ">=4.29,<5" }, + { name = "jsonargparse", extras = ["signatures"], specifier = ">=4.49,<5" }, { name = "ml4gw", specifier = ">=0.8.0" }, { name = "ml4gw-hermes", extras = ["torch"], git = "https://github.com/ML4GW/hermes?branch=dev" }, { name = "nvidia-cudnn-cu12" }, @@ -2648,14 +2648,14 @@ wheels = [ [[package]] name = "jsonargparse" -version = "4.36.0" +version = "4.52.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "pyyaml" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/14/de/dd1e113182865f600d8b1f6c90ec5e584114c18e33f0794d1c43815f05f6/jsonargparse-4.36.0.tar.gz", hash = "sha256:95b9a5df843d0224acae1a9bb780745fa8f206784a87b2b30d5b1db11b00c125", size = 193845, upload-time = "2025-01-17T17:55:47.076Z" } +sdist = { url = "https://files.pythonhosted.org/packages/f1/1b/86ab3c48f74609de566e6c6fea5493a7bea4c7de485e5a56fa1239c698b5/jsonargparse-4.52.0.tar.gz", hash = "sha256:103168510fab9d45a980fe7560dcbf20d1cce38320a43904583d23869af3bee1", size = 164776, upload-time = "2026-09-01T05:53:07.575Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/e5/12/7134b957ca60fb10692c7701d1ed16928b8b2b331eeaa54bedd3a9dabaf3/jsonargparse-4.36.0-py3-none-any.whl", hash = "sha256:1140dc766957857282e81efbb88c51605c948aa3c3517ea98f3265d6df2df627", size = 214471, upload-time = "2025-01-17T17:55:42.075Z" }, + { url = "https://files.pythonhosted.org/packages/52/70/e9114d52e9e72daaaeec30ecafe1018a4fbe6288ec35e16822529676b671/jsonargparse-4.52.0-py3-none-any.whl", hash = "sha256:1878b16b6be90c3e838aac569784a110fe1eb1806c28e56175cdc7f2e563d78c", size = 174766, upload-time = "2026-09-01T05:53:04.257Z" }, ] [package.optional-dependencies] @@ -4117,7 +4117,7 @@ requires-dist = [ { name = "certifi" }, { name = "gwpy", specifier = ">=4.0,<5" }, { name = "h5py", specifier = ">=3.10,<4" }, - { name = "jsonargparse", extras = ["signatures"], specifier = ">=4.29,<5" }, + { name = "jsonargparse", extras = ["signatures"], specifier = ">=4.49,<5" }, { name = "lalsuite" }, { name = "ledger", editable = "libs/ledger" }, { name = "ligo-gracedb", extras = ["kafka"], specifier = ">=2.15.4" }, @@ -4595,7 +4595,7 @@ requires-dist = [ { name = "gwpy", specifier = ">=4.0,<5" }, { name = "h5py", specifier = ">=3.10,<4" }, { name = "igwn-ligolw", specifier = ">=2.1,<3" }, - { name = "jsonargparse", extras = ["signatures"], specifier = ">=4.29,<5" }, + { name = "jsonargparse", extras = ["signatures"], specifier = ">=4.49,<5" }, { name = "lalsuite" }, { name = "ledger", editable = "libs/ledger" }, { name = "matplotlib", marker = "extra == 'vizapp'", specifier = ">=3.9.4" }, @@ -6806,7 +6806,7 @@ requires-dist = [ { name = "filelock", specifier = ">=3.13.1,<5" }, { name = "fsspec", extras = ["s3"], specifier = ">=2024" }, { name = "h5py", specifier = ">=3.10,<4" }, - { name = "jsonargparse", extras = ["signatures"], specifier = ">=4.29,<5" }, + { name = "jsonargparse", extras = ["signatures"], specifier = ">=4.49,<5" }, { name = "kr8s", specifier = ">=0.20.0" }, { name = "ledger", editable = "libs/ledger" }, { name = "lightning", specifier = ">=2.2.1" },