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
19 changes: 18 additions & 1 deletion miles/backends/megatron_utils/arguments.py
Original file line number Diff line number Diff line change
@@ -1,14 +1,31 @@
import logging
import os

from megatron.training.arguments import parse_args, validate_args
import torch
from megatron.training.arguments import parse_args
from megatron.training.arguments import validate_args as _megatron_validate_args
from megatron.training.tokenizer.tokenizer import _vocab_size_with_padding

__all__ = ["validate_args", "parse_args", "set_default_megatron_args"]

logger = logging.getLogger(__name__)


def validate_args(args):
args = _megatron_validate_args(args)

optimizer_state_dtypes = (args.main_params_dtype, args.exp_avg_dtype, args.exp_avg_sq_dtype)
if args.optimizer_cpu_offload and any(dtype != torch.float32 for dtype in optimizer_state_dtypes):
# HybridDeviceOptimizer currently creates FP32 master parameters and Adam states
# regardless of the precision-aware optimizer dtype settings.
raise ValueError(
"--optimizer-cpu-offload does not honor lower-precision --main-params-dtype, "
"--exp-avg-dtype, or --exp-avg-sq-dtype; remove these dtype overrides"
)

return args


def set_default_megatron_args(args):
# Muon currently owns its sharding path, and Megatron's distributed optimizer
# only supports Adam-family optimizers.
Expand Down
3 changes: 1 addition & 2 deletions tests/e2e/ckpt/test_glm47_flash_ckpt.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,7 @@

import miles.utils.external_utils.command_utils as U

# FIXME: need to modify megatron, better fix later.
register_cuda_ci(est_time=2400, suite="stage-c-8-gpu-h100", labels=["ckpt"], disabled="Disabled due to bugs.")
register_cuda_ci(est_time=2400, suite="stage-c-8-gpu-h200", labels=["ckpt"])

ENABLE_EVAL = 0
USE_DEEPEP = 0
Expand Down
7 changes: 0 additions & 7 deletions tests/e2e/megatron/test_glm47_flash/_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,6 @@
MODEL_NAME = "GLM-4.7-Flash"
MODEL_TYPE = "glm4.7-flash"

TIGHT_HOST_MEMORY = bool(int(os.environ.get("MILES_TEST_TIGHT_HOST_MEMORY", "1")))


@dataclass
class CaseConfig:
Expand Down Expand Up @@ -88,11 +86,6 @@ def build_train_args(case: CaseConfig, *, wandb_file: str) -> str:
f"--max-tokens-per-gpu {case.max_tokens_per_gpu} "
)

if TIGHT_HOST_MEMORY:
perf_args += "--exp-avg-dtype fp16 "
perf_args += "--exp-avg-sq-dtype fp16 "
perf_args += "--main-params-dtype fp16 "

grpo_args = (
"--advantage-estimator grpo "
"--use-kl-loss "
Expand Down
7 changes: 0 additions & 7 deletions tests/e2e/megatron/test_qwen3_30B_A3B/_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,6 @@
MODEL_NAME = "Qwen3-30B-A3B"
MODEL_TYPE = "qwen3-30B-A3B"

TIGHT_HOST_MEMORY = bool(int(os.environ.get("MILES_TEST_TIGHT_HOST_MEMORY", "1")))


@dataclass
class CaseConfig:
Expand Down Expand Up @@ -131,11 +129,6 @@ def build_train_args(case: CaseConfig, *, wandb_file: str) -> str:
f"--max-tokens-per-gpu {case.max_tokens_per_gpu} "
)

if TIGHT_HOST_MEMORY:
perf_args += "--exp-avg-dtype fp16 "
perf_args += "--exp-avg-sq-dtype fp16 "
perf_args += "--main-params-dtype fp16 "

# r3 path uses --use-rollout-routing-replay; non-r3 uses --use-routing-replay.
routing_flag = "--use-rollout-routing-replay" if case.use_r3 else "--use-routing-replay"
grpo_args = (
Expand Down
19 changes: 19 additions & 0 deletions tests/fast/test_megatron_cli_flags.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import sys
from types import SimpleNamespace

import pytest

Expand Down Expand Up @@ -54,3 +55,21 @@ def test_post_layernorm_flags_propagate_to_megatron(monkeypatch):

assert config.post_self_attn_layernorm is True
assert config.post_mlp_layernorm is True


def test_optimizer_cpu_offload_rejects_lower_precision_state_dtypes(monkeypatch):
torch = pytest.importorskip("torch")
pytest.importorskip("megatron.training.arguments")

import miles.backends.megatron_utils.arguments as megatron_arguments

args = SimpleNamespace(
optimizer_cpu_offload=True,
main_params_dtype=torch.float16,
exp_avg_dtype=torch.float16,
exp_avg_sq_dtype=torch.float16,
)
monkeypatch.setattr(megatron_arguments, "_megatron_validate_args", lambda args: args)

with pytest.raises(ValueError, match="--optimizer-cpu-offload does not honor lower-precision"):
megatron_arguments.validate_args(args)
Loading