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..a0da377 100644 --- a/tests/test_slurm_render.py +++ b/tests/test_slurm_render.py @@ -139,6 +139,56 @@ 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 + + +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 purge" not in content + assert "module load" not in content + + # --------------------------------------------------------------------------- # --tokenize-only forwarding — TRL templates and backend dispatch # ---------------------------------------------------------------------------