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
10 changes: 10 additions & 0 deletions src/post_training/slurm/job_trl_container.sh.jinja
Original file line number Diff line number Diff line change
Expand Up @@ -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 %}

Expand Down
2 changes: 2 additions & 0 deletions src/post_training/slurm/launcher.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
50 changes: 50 additions & 0 deletions tests/test_slurm_render.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
# ---------------------------------------------------------------------------
Expand Down
Loading