mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
launcher: make Slurm memory and user defaults configurable (#1791)
## Summary - Adds `SlurmConfig.mem` and a `SLURM_MEM` env override for launcher Slurm jobs. - Passes configured memory through to `nemo_run.SlurmExecutor` instead of always using `"0"`. - Lets `SLURM_USER` provide the launcher default user when local and cluster usernames differ. - Adds focused tests for Slurm memory defaults and overrides. ## Test plan - [x] `uv run pytest tools/launcher/tests/test_slurm_config.py tools/launcher/tests/test_slurm_executor.py` <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit ## Release Notes * **New Features** * Slurm job memory can now be configured via `SLURM_MEM`, with a default of `0` when unset. * Job launch username can now be set via `SLURM_USER`; if omitted, it falls back to the local login name. * **Bug Fixes** * Slurm executor now correctly uses the configured memory value from Slurm settings, falling back to `0` only when missing/empty. * **Tests** * Expanded unit tests to cover default, environment-driven, and executor parameter memory behavior (including missing/empty/`None` cases). <!-- end of auto-generated comment: release notes by coderabbit.ai --> Signed-off-by: Chenhan Yu <chenhany@nvidia.com>
This commit is contained in:
@@ -312,7 +312,7 @@ def build_slurm_executor(
|
||||
container_mounts=container_mounts,
|
||||
array=slurm_config.array,
|
||||
time=slurm_config.time,
|
||||
mem="0",
|
||||
mem=getattr(slurm_config, "mem", None) or "0",
|
||||
retries=0,
|
||||
packager=packager,
|
||||
srun_args=slurm_config.srun_args,
|
||||
|
||||
@@ -23,6 +23,7 @@ Environment variables:
|
||||
SLURM_HOST Slurm login node hostname (required for remote jobs)
|
||||
SLURM_ACCOUNT Slurm account/partition billing (default: from YAML)
|
||||
SLURM_JOB_DIR Remote directory for job artifacts
|
||||
SLURM_USER Remote Slurm/SSH username (default: local login name)
|
||||
SLURM_HF_LOCAL Path to HuggingFace model cache on the cluster
|
||||
HF_TOKEN HuggingFace API token
|
||||
NEMORUN_HOME NeMo Run home directory (default: current working directory)
|
||||
@@ -112,12 +113,15 @@ def launch(
|
||||
job_dir: str = os.environ.get("SLURM_JOB_DIR", os.path.expanduser("~/experiments")),
|
||||
pipeline: SandboxPipeline = None,
|
||||
hf_local: str = None, # noqa: RUF013
|
||||
user: str = getpass.getuser(),
|
||||
user: str | None = None,
|
||||
identity: str = None, # noqa: RUF013
|
||||
detach: bool = False,
|
||||
clean: bool = False,
|
||||
) -> None:
|
||||
"""Launch ModelOpt jobs on Slurm or locally with Docker."""
|
||||
if user is None:
|
||||
user = os.environ.get("SLURM_USER", getpass.getuser())
|
||||
|
||||
if clean:
|
||||
if _mo_symlink is None:
|
||||
raise ValueError("--clean requires a dev checkout; modelopt source not found.")
|
||||
|
||||
@@ -50,6 +50,7 @@ class SlurmConfig:
|
||||
ntasks_per_node: int = 1
|
||||
gpus_per_node: int = 1
|
||||
time: str = "04:00:00"
|
||||
mem: str = "0"
|
||||
local: bool = False
|
||||
# Slurm --segment=<N>: force the job's nodes into a single topology block.
|
||||
# On a topology/block cluster (e.g. GB200 NVL72, where one block = one NVLink
|
||||
@@ -77,6 +78,7 @@ def slurm_factory(
|
||||
array: Optional[str] = None,
|
||||
requeue: bool = False,
|
||||
time: str = "04:00:00",
|
||||
mem: str = os.environ.get("SLURM_MEM", "0"),
|
||||
segment: Optional[int] = None,
|
||||
) -> SlurmConfig:
|
||||
"""Generic Slurm factory — configure via environment variables or CLI overrides."""
|
||||
@@ -95,5 +97,6 @@ def slurm_factory(
|
||||
array=array,
|
||||
requeue=requeue,
|
||||
time=time,
|
||||
mem=mem,
|
||||
segment=segment,
|
||||
)
|
||||
|
||||
@@ -41,6 +41,7 @@ class TestSlurmConfig:
|
||||
assert cfg.nodes == 1
|
||||
assert cfg.ntasks_per_node == 1
|
||||
assert cfg.gpus_per_node == 1
|
||||
assert cfg.mem == "0"
|
||||
assert cfg.local is False
|
||||
assert cfg.container_mounts is None
|
||||
assert cfg.srun_args is None
|
||||
@@ -52,6 +53,7 @@ class TestSlurmConfig:
|
||||
account="my_account",
|
||||
nodes=4,
|
||||
gpus_per_node=8,
|
||||
mem="128G",
|
||||
container="nvcr.io/nvidia/pytorch:24.01-py3",
|
||||
container_mounts=["/data:/data"],
|
||||
srun_args=["--no-container-mount-home"],
|
||||
@@ -60,6 +62,7 @@ class TestSlurmConfig:
|
||||
assert cfg.account == "my_account"
|
||||
assert cfg.nodes == 4
|
||||
assert cfg.gpus_per_node == 8
|
||||
assert cfg.mem == "128G"
|
||||
assert cfg.container_mounts == ["/data:/data"]
|
||||
|
||||
|
||||
@@ -79,6 +82,10 @@ class TestSlurmFactory:
|
||||
cfg = slurm_factory()
|
||||
assert cfg.srun_args == ["--no-container-mount-home"]
|
||||
|
||||
def test_default_mem(self):
|
||||
cfg = slurm_factory()
|
||||
assert cfg.mem == "0"
|
||||
|
||||
def test_default_container_mounts_from_env(self, monkeypatch):
|
||||
monkeypatch.setenv("SLURM_HF_LOCAL", "/custom/hf-local")
|
||||
# Reload to pick up the env var — slurm_factory reads SLURM_HF_LOCAL at module-import
|
||||
@@ -102,3 +109,9 @@ class TestSlurmFactory:
|
||||
importlib.reload(slurm_config)
|
||||
cfg = slurm_config.slurm_factory()
|
||||
assert cfg.host == "test-host.example.com"
|
||||
|
||||
def test_env_var_mem(self, monkeypatch):
|
||||
monkeypatch.setenv("SLURM_MEM", "100G")
|
||||
importlib.reload(slurm_config)
|
||||
cfg = slurm_config.slurm_factory()
|
||||
assert cfg.mem == "100G"
|
||||
|
||||
@@ -167,6 +167,7 @@ class TestBuildSlurmExecutor:
|
||||
gpus_per_node=8,
|
||||
array="0-3",
|
||||
time="04:00:00",
|
||||
mem="128G",
|
||||
)
|
||||
|
||||
packager = MagicMock()
|
||||
@@ -191,8 +192,48 @@ class TestBuildSlurmExecutor:
|
||||
assert kw["array"] == "0-3"
|
||||
assert kw["packager"] is packager
|
||||
assert kw["time"] == "04:00:00"
|
||||
assert kw["mem"] == "128G"
|
||||
assert kw["retries"] == 0
|
||||
|
||||
@patch("core.run.SlurmExecutor")
|
||||
@patch("core.run.SSHTunnel")
|
||||
def test_default_mem_remains_all_node_memory(self, mock_tunnel, mock_executor):
|
||||
mock_tunnel.return_value = MagicMock()
|
||||
|
||||
for mem in ("missing", "", None):
|
||||
slurm_config = MagicMock(
|
||||
requeue=False,
|
||||
host="h",
|
||||
port=22,
|
||||
account="a",
|
||||
partition="b",
|
||||
container="c",
|
||||
modelopt_install_path="/m",
|
||||
container_mounts=[],
|
||||
srun_args=[],
|
||||
nodes=1,
|
||||
ntasks_per_node=1,
|
||||
gpus_per_node=1,
|
||||
array=None,
|
||||
)
|
||||
if mem == "missing":
|
||||
del slurm_config.mem
|
||||
else:
|
||||
slurm_config.mem = mem
|
||||
|
||||
build_slurm_executor(
|
||||
user="u",
|
||||
identity=None,
|
||||
slurm_config=slurm_config,
|
||||
experiment_id="e",
|
||||
job_dir="/j",
|
||||
task_name="t",
|
||||
packager=MagicMock(),
|
||||
)
|
||||
|
||||
assert mock_executor.call_args[1]["mem"] == "0"
|
||||
mock_executor.reset_mock()
|
||||
|
||||
@patch("core.run.SlurmExecutor")
|
||||
@patch("core.run.SSHTunnel")
|
||||
def test_none_container_mounts_handled(self, mock_tunnel, mock_executor):
|
||||
|
||||
Reference in New Issue
Block a user