From 2c6b401149906d38e49dc2f37d3d83ad2d120260 Mon Sep 17 00:00:00 2001 From: Ali Elganzory Date: Wed, 26 Aug 2026 10:59:07 +0200 Subject: [PATCH 1/3] feat: Support loading HPC modules for TRL container jobs. --- .../slurm/job_trl_container.sh.jinja | 10 +++++ src/post_training/slurm/launcher.py | 2 + tests/test_slurm_render.py | 39 +++++++++++++++++++ 3 files changed, 51 insertions(+) diff --git a/src/post_training/slurm/job_trl_container.sh.jinja b/src/post_training/slurm/job_trl_container.sh.jinja index 86699b0..989a947 100644 --- a/src/post_training/slurm/job_trl_container.sh.jinja +++ b/src/post_training/slurm/job_trl_container.sh.jinja @@ -30,6 +30,16 @@ set -euo pipefail +{% if modules %} +# ── Environment modules ──────────────────────────────────────────────────── +{% if module_purge %} +module purge +{% endif %} +{% for mod in modules %} +module load {{ mod }} +{% endfor %} +{% endif %} + # ── Source cluster environment ──────────────────────────────────────────── {% if env_file %}source "{{ env_file }}"{% endif %} diff --git a/src/post_training/slurm/launcher.py b/src/post_training/slurm/launcher.py index 000bce4..fd8338f 100644 --- a/src/post_training/slurm/launcher.py +++ b/src/post_training/slurm/launcher.py @@ -128,6 +128,8 @@ def render_trl_container_slurm_script( wall_time=config.slurm.wall_time, signal_time_seconds=config.slurm.signal_time_seconds, max_failures=config.slurm.max_failures, + modules=config.slurm.modules, + module_purge=config.slurm.module_purge, run_dir=str(run_dir.resolve()), config_path=config_path, tokenize_only=tokenize_only, diff --git a/tests/test_slurm_render.py b/tests/test_slurm_render.py index 6b2adba..262e0b4 100644 --- a/tests/test_slurm_render.py +++ b/tests/test_slurm_render.py @@ -139,6 +139,45 @@ def test_trl_container_qos_mem_absent_when_none(tmp_path, config): assert "--mem" not in content +# --------------------------------------------------------------------------- +# environment modules — TRL container template +# --------------------------------------------------------------------------- + + +def test_trl_container_modules_rendered(tmp_path, config): + """Configured host modules are loaded before the cluster environment file.""" + config.slurm.modules = ["singularity/4.1.0", "cuda/12.4"] + config.slurm.module_purge = True + run_dir = tmp_path / "outputs" / "my-run" + run_dir.mkdir(parents=True) + + content = render_trl_container_slurm_script(config, run_dir, "configs/trl/sft.yaml").read_text() + + commands = [ + "module purge", + "module load singularity/4.1.0", + "module load cuda/12.4", + 'source "/shared/env/cluster.env"', + ] + assert all(command in content for command in commands) + assert [content.index(command) for command in commands] == sorted( + content.index(command) for command in commands + ) + + +def test_trl_container_module_purge_omitted_when_disabled(tmp_path, config): + """Modules can be loaded without first purging the inherited module set.""" + config.slurm.modules = ["singularity/4.1.0"] + config.slurm.module_purge = False + run_dir = tmp_path / "outputs" / "my-run" + run_dir.mkdir(parents=True) + + content = render_trl_container_slurm_script(config, run_dir, "configs/trl/sft.yaml").read_text() + + assert "module load singularity/4.1.0" in content + assert "module purge" not in content + + # --------------------------------------------------------------------------- # --tokenize-only forwarding — TRL templates and backend dispatch # --------------------------------------------------------------------------- From 08576cb2419229bc4af3e3b41b08159a4eb5e8c5 Mon Sep 17 00:00:00 2001 From: Ali Elganzory Date: Wed, 26 Aug 2026 11:19:07 +0200 Subject: [PATCH 2/3] Test that no modules are rendered when no modules are configured. --- tests/test_slurm_render.py | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/tests/test_slurm_render.py b/tests/test_slurm_render.py index 262e0b4..4a7fc2b 100644 --- a/tests/test_slurm_render.py +++ b/tests/test_slurm_render.py @@ -178,6 +178,16 @@ def test_trl_container_module_purge_omitted_when_disabled(tmp_path, config): assert "module purge" not in content +def test_trl_container_host_setup_absent_when_unspecified(tmp_path, config): + """No module setup is rendered when the module list is empty.""" + run_dir = tmp_path / "outputs" / "my-run" + run_dir.mkdir(parents=True) + + content = render_trl_container_slurm_script(config, run_dir, "configs/trl/sft.yaml").read_text() + + assert "module" not in content + + # --------------------------------------------------------------------------- # --tokenize-only forwarding — TRL templates and backend dispatch # --------------------------------------------------------------------------- From 4a615a289e3f4c59b22139145bc8461d96918d48 Mon Sep 17 00:00:00 2001 From: Ali Elganzory Date: Wed, 26 Aug 2026 11:22:56 +0200 Subject: [PATCH 3/3] Make the test more specific. --- tests/test_slurm_render.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/test_slurm_render.py b/tests/test_slurm_render.py index 4a7fc2b..a0da377 100644 --- a/tests/test_slurm_render.py +++ b/tests/test_slurm_render.py @@ -185,7 +185,8 @@ def test_trl_container_host_setup_absent_when_unspecified(tmp_path, config): content = render_trl_container_slurm_script(config, run_dir, "configs/trl/sft.yaml").read_text() - assert "module" not in content + assert "module purge" not in content + assert "module load" not in content # ---------------------------------------------------------------------------