Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions pipeline/config/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion projects/export/export/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion projects/export/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
2 changes: 1 addition & 1 deletion projects/online/online/cli.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))


Expand Down
2 changes: 1 addition & 1 deletion projects/online/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 1 addition & 1 deletion projects/plots/plots/vizapp/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))

Expand Down
2 changes: 1 addition & 1 deletion projects/plots/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
4 changes: 2 additions & 2 deletions projects/train/configs/bbh.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
4 changes: 2 additions & 2 deletions projects/train/configs/bns_time_spectrogram.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 2 additions & 2 deletions projects/train/configs/time_heterodyne.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
2 changes: 1 addition & 1 deletion projects/train/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
24 changes: 20 additions & 4 deletions projects/train/train.smk
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand All @@ -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}"
7 changes: 3 additions & 4 deletions projects/train/train.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion projects/train/train/remote.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)


Expand Down
14 changes: 7 additions & 7 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading