Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
8 changes: 8 additions & 0 deletions plugins/nemo-automodel/src/nemo_automodel_plugin/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,14 @@ class ParallelismSpec(AutomodelSchema):
context_parallel_size: int = Field(default=1, gt=0)
expert_parallel_size: int | None = Field(default=None, gt=0)
sequence_parallel: bool = Field(default=False, description="Enable sequence parallelism.")
activation_checkpointing: bool | Literal["full", "selective"] | None = Field(
default=None,
description=(
"Recompute activations during the backward pass to cut peak memory at the cost of "
"speed. 'selective' checkpoints only the most memory-heavy ops. Left unset, Automodel "
"defaults to disabled."
),
)


class OutputRequest(AutomodelSchema):
Expand Down
1 change: 1 addition & 0 deletions services/automodel/src/nmp/automodel/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,6 +80,7 @@ def _build_training_block(spec: dict[str, Any]) -> SFTTraining | DistillationTra
context_parallel_size=parallelism.get("context_parallel_size", 1),
expert_parallel_size=parallelism.get("expert_parallel_size"),
sequence_parallel=parallelism.get("sequence_parallel", False),
activation_checkpointing=parallelism.get("activation_checkpointing"),
),
"execution_profile": training.get("execution_profile"),
}
Expand Down
8 changes: 8 additions & 0 deletions services/automodel/src/nmp/automodel/api/v2/jobs/schemas.py
Original file line number Diff line number Diff line change
Expand Up @@ -136,6 +136,14 @@ class ParallelismParams(BaseModel):
context_parallel_size: int = Field(default=1, gt=0, description="Context parallel size.")
expert_parallel_size: Optional[int] = Field(default=None, gt=0, description="Expert parallel size (MoE models).")
sequence_parallel: bool = Field(default=False, description="Enable sequence parallelism.")
activation_checkpointing: Optional[Union[bool, Literal["full", "selective"]]] = Field(
default=None,
description=(
"Recompute activations during the backward pass to reduce peak memory at the cost of "
"speed. 'selective' checkpoints only the most memory-heavy ops. Unset leaves Automodel's "
"default (disabled)."
),
)


# ============================================================
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -178,6 +178,7 @@ def compile_training_step(
context_parallel_size=p.context_parallel_size,
expert_parallel_size=p.expert_parallel_size,
sequence_parallel=p.sequence_parallel,
activation_checkpointing=p.activation_checkpointing,
),
integrations=job_spec.integrations,
output_model=job_spec.output.name,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
# SPDX-License-Identifier: Apache-2.0

from enum import Enum
from typing import Optional
from typing import Literal, Optional, Union

from nemo_platform_plugin.integrations import IntegrationsSpec
from nmp.automodel.app.constants import (
Expand Down Expand Up @@ -204,6 +204,7 @@ class ParallelismConfig(BaseModel):
context_parallel_size: int = 1
expert_parallel_size: Optional[int] = None
sequence_parallel: bool = False
activation_checkpointing: Optional[Union[bool, Literal["full", "selective"]]] = None

# === Main Config Fields ===
model: ModelConfig
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -146,7 +146,11 @@ def compile_automodel_config(
"ep_size": p.expert_parallel_size,
"sequence_parallel": p.sequence_parallel,
}
if _is_embedding_model and embedding_config.do_gradient_checkpointing:
# Explicit parallelism setting wins; embedding jobs keep their own toggle as the
# fallback so `do_gradient_checkpointing` behaves as before when nothing is set.
if p.activation_checkpointing is not None:
cfg["distributed"]["activation_checkpointing"] = p.activation_checkpointing
elif _is_embedding_model and embedding_config.do_gradient_checkpointing:
cfg["distributed"]["activation_checkpointing"] = True
if p.pipeline_parallel_size > 1:
cfg["distributed"]["pipeline"] = {
Expand Down
56 changes: 56 additions & 0 deletions services/automodel/tests/tasks/training/backends/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -414,3 +414,59 @@ def test_compile_uses_fallback_packing_factor_for_schedule(tmp_path: Path) -> No
assert compiled["step_scheduler"]["max_steps"] == 2
assert compiled["step_scheduler"]["val_every_steps"] == 1
assert compiled["lr_scheduler"]["lr_warmup_steps"] == 1


def _compile_with_parallelism(tmp_path: Path, **parallelism: Any) -> dict[str, Any]:
"""Compile a known-good contract fixture with parallelism overrides applied."""
fixture = Path(__file__).parents[3] / "contract" / "input_configs" / "llama-3.2-1b" / "llama_3_2_1b_lora.json"
raw = json.loads(fixture.read_text())
raw.pop("backend", None)
config = TrainingStepConfig.model_validate(raw)
for key, value in parallelism.items():
setattr(config.parallelism, key, value)

prepared = PreparedDataset(
merged_dir=tmp_path,
train_file=tmp_path / "train.jsonl",
validation_file=tmp_path / "validation.jsonl",
train_samples=100,
validation_samples=10,
)

with (
patch(f"{CONFIG_MODULE}.prepare_dataset", return_value=prepared),
patch(f"{CONFIG_MODULE}.DatasetValidator"),
patch(f"{CONFIG_MODULE}.estimate_dataset_sequence_lengths", return_value=None),
patch(f"{CONFIG_MODULE}._configure_datasets"),
patch(f"{CONFIG_MODULE}._configure_moe_backend"),
patch(f"{CONFIG_MODULE}.build_wandb_config", return_value=None),
patch(f"{CONFIG_MODULE}.build_mlflow_config", return_value=None),
):
return compile_automodel_config(config, tmp_path, MagicMock())


class TestActivationCheckpointing:
"""Emission of `distributed.activation_checkpointing` for causal-LM jobs.

Automodel's FSDP2Config defaults this to False, so a key we never emit means
activation checkpointing is off — which is what every non-embedding job used to get.
"""

def test_omitted_when_unset(self, tmp_path: Path) -> None:
compiled = _compile_with_parallelism(tmp_path)
assert "activation_checkpointing" not in compiled["distributed"]

def test_emitted_when_enabled(self, tmp_path: Path) -> None:
compiled = _compile_with_parallelism(tmp_path, activation_checkpointing=True)
assert compiled["distributed"]["activation_checkpointing"] is True

@pytest.mark.parametrize("mode", ["full", "selective"])
def test_string_modes_pass_through(self, tmp_path: Path, mode: str) -> None:
# Automodel accepts bool | "full" | "selective"; the strings must not be coerced.
compiled = _compile_with_parallelism(tmp_path, activation_checkpointing=mode)
assert compiled["distributed"]["activation_checkpointing"] == mode

def test_explicit_false_is_emitted(self, tmp_path: Path) -> None:
# Distinct from unset: the user asked for it off, so say so rather than omit.
compiled = _compile_with_parallelism(tmp_path, activation_checkpointing=False)
assert compiled["distributed"]["activation_checkpointing"] is False
44 changes: 44 additions & 0 deletions services/automodel/tests/test_adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -225,3 +225,47 @@ def test_adapter_integrations_from_automodel_job_output() -> None:
assert spec.integrations is not None
assert spec.integrations.wandb is not None
assert spec.integrations.wandb.project == "plugin-project"


def test_adapter_plumbs_activation_checkpointing() -> None:
"""`parallelism.activation_checkpointing` must survive into the v2 training spec."""
spec = automodel_spec_to_compiler_output(
{
"model": "meta/llama",
"dataset": {"training": "default/train"},
"training": {"training_type": "sft", "finetuning_type": "lora"},
"parallelism": {"num_gpus_per_node": 8, "activation_checkpointing": True},
"output": {"name": "out", "type": "adapter", "fileset": "out-fs"},
},
)
assert isinstance(spec.training, SFTTraining)
assert spec.training.parallelism.activation_checkpointing is True


def test_adapter_activation_checkpointing_accepts_selective() -> None:
"""Automodel takes bool | 'full' | 'selective'; the string modes must pass through."""
spec = automodel_spec_to_compiler_output(
{
"model": "meta/llama",
"dataset": {"training": "default/train"},
"training": {"training_type": "sft", "finetuning_type": "lora"},
"parallelism": {"num_gpus_per_node": 8, "activation_checkpointing": "selective"},
"output": {"name": "out", "type": "adapter", "fileset": "out-fs"},
},
)
assert isinstance(spec.training, SFTTraining)
assert spec.training.parallelism.activation_checkpointing == "selective"


def test_adapter_activation_checkpointing_defaults_to_none() -> None:
"""Unset means "don't emit", preserving Automodel's own default of disabled."""
spec = automodel_spec_to_compiler_output(
{
"model": "meta/llama",
"dataset": {"training": "default/train"},
"training": {"training_type": "sft", "finetuning_type": "lora"},
"output": {"name": "out", "type": "adapter", "fileset": "out-fs"},
},
)
assert isinstance(spec.training, SFTTraining)
assert spec.training.parallelism.activation_checkpointing is None
Loading