Skip to content
Open
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
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
11 changes: 11 additions & 0 deletions plugins/nemo-customizer/openapi/openapi.yaml

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

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