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
2 changes: 1 addition & 1 deletion docker/Dockerfile
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ ARG SGLANG_BRANCH=sglang-miles
ARG SGLANG_COMMIT=""

ARG MEGATRON_REPO=radixark/Megatron-LM
ARG MEGATRON_BRANCH=miles-main
ARG MEGATRON_BRANCH=miles-main-20260622

ARG ENABLE_CUDA_13=1

Expand Down
2 changes: 1 addition & 1 deletion miles/backends/megatron_utils/arguments.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,8 @@
import logging
import os

from megatron.core.tokenizers.utils.build_tokenizer import vocab_size_with_padding as _vocab_size_with_padding
from megatron.training.arguments import parse_args, validate_args
from megatron.training.tokenizer.tokenizer import _vocab_size_with_padding

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

Expand Down
2 changes: 1 addition & 1 deletion miles/backends/megatron_utils/initialize.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ def _initialize_distributed(args, get_embedding_ranks=None, get_position_embeddi
order="tp-cp-ep-dp-pp" if not args.use_tp_pp_dp_mapping else "tp-cp-ep-pp-dp",
get_embedding_ranks=get_embedding_ranks,
get_position_embedding_ranks=get_position_embedding_ranks,
create_gloo_process_groups=args.enable_gloo_process_groups,
create_gloo_process_groups=args.use_gloo_process_groups,
)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -149,7 +149,8 @@ def convert_deepseekv3_to_hf(args, name, param):
return [(f"model.layers.{layer_idx}.shared_head.norm.weight", param)]
else:
name = f"module.module.decoder.layers.{layer_idx}.{rest}"
name = name.replace("transformer_layer.", "")
# New Megatron renamed the MTP submodule transformer_layer -> mtp_model_layer.
name = name.replace("transformer_layer.", "").replace("mtp_model_layer.", "")
return convert_deepseekv3_to_hf(args, name, param)

raise ValueError(f"Unknown parameter name: {name}")
3 changes: 2 additions & 1 deletion miles/backends/megatron_utils/megatron_to_hf/glm4moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -133,7 +133,8 @@ def convert_glm4moe_to_hf(args, name, param):
return [(f"model.layers.{layer_idx}.shared_head.norm.weight", param)]
else:
name = f"module.module.decoder.layers.{layer_idx}.{rest}"
name = name.replace("transformer_layer.", "")
# New Megatron renamed the MTP submodule transformer_layer -> mtp_model_layer.
name = name.replace("transformer_layer.", "").replace("mtp_model_layer.", "")
return convert_glm4moe_to_hf(args, name, param)

raise ValueError(f"Unknown parameter name: {name}")
11 changes: 7 additions & 4 deletions miles/backends/megatron_utils/megatron_to_hf/mimo.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,10 +52,13 @@ def convert_mimo_mtp_param(args, name, param):
if component in direct_mappings:
return [(direct_mappings[component], param)]

# Handle transformer_layer components
if component.startswith("transformer_layer."):
# Remove "transformer_layer." prefix
transformer_component = component[len("transformer_layer.") :]
# Handle the wrapped transformer-layer components. New Megatron renamed the MTP
# submodule field `transformer_layer` -> `mtp_model_layer`; accept both.
transformer_layer_prefixes = ("transformer_layer.", "mtp_model_layer.")
matched_prefix = next((p for p in transformer_layer_prefixes if component.startswith(p)), None)
if matched_prefix is not None:
# Remove the wrapped-layer prefix
transformer_component = component[len(matched_prefix) :]

# Create proxy name for reusing existing Qwen2 conversion functions
proxy_name = f"module.module.decoder.layers.{layer_idx}.{transformer_component}"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@ def quantize_params_fp8(args, megatron_name, converted_named_params, quantizatio
if not match:
return converted_named_params
layer_idx, rest = match.groups()
rest = rest.replace("transformer_layer.", "")
rest = rest.replace("transformer_layer.", "").replace("mtp_model_layer.", "")
else:
layer_idx, rest = match.groups()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,7 @@ def quantize_params_mxfp8(args, megatron_name, converted_named_params, quantizat
if not match:
return converted_named_params
layer_idx, rest = match.groups()
rest = rest.replace("transformer_layer.", "")
rest = rest.replace("transformer_layer.", "").replace("mtp_model_layer.", "")
else:
layer_idx, rest = match.groups()

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,7 @@ def quantize_params_nvfp4(args, megatron_name, converted_named_params, quantizat
if not match:
return converted_named_params
layer_idx, rest = match.groups()
rest = rest.replace("transformer_layer.", "")
rest = rest.replace("transformer_layer.", "").replace("mtp_model_layer.", "")
else:
layer_idx, rest = match.groups()

Expand Down
5 changes: 3 additions & 2 deletions miles/backends/megatron_utils/megatron_to_hf/qwen3_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,9 @@ def _convert_mtp_layer(args, name, param, layer_idx):
if "eh_proj.weight" in name:
return [("mtp.fc.weight", param)]

if "transformer_layer" in name:
proxy_name = name.replace(f"mtp.layers.{layer_idx}.transformer_layer", f"decoder.layers.{layer_idx}")
_mtp_inner = next((p for p in ("transformer_layer", "mtp_model_layer") if p in name), None)
if _mtp_inner is not None:
proxy_name = name.replace(f"mtp.layers.{layer_idx}.{_mtp_inner}", f"decoder.layers.{layer_idx}")
mapped_params = convert_qwen3_5_to_hf(args, proxy_name, param)

final_params = []
Expand Down
7 changes: 4 additions & 3 deletions miles/backends/megatron_utils/megatron_to_hf/qwen3_next.py
Original file line number Diff line number Diff line change
Expand Up @@ -155,9 +155,10 @@ def convert_qwen3_next_to_hf(args, name, param):
elif rest == "final_layernorm.weight":
return [("mtp.norm.weight", param)]

# transformer_layer components → reuse decoder conversion with mtp prefix
if rest.startswith("transformer_layer."):
transformer_rest = rest[len("transformer_layer.") :]
# transformer_layer components → reuse decoder conversion with mtp prefix.
# New Megatron renamed the MTP submodule transformer_layer -> mtp_model_layer.
if rest.startswith(("transformer_layer.", "mtp_model_layer.")):
transformer_rest = rest.split(".", 1)[1]
proxy_name = f"module.module.decoder.layers.{layer_idx}.{transformer_rest}"
results = convert_qwen3_next_to_hf(args, proxy_name, param)
return [
Expand Down
27 changes: 16 additions & 11 deletions miles/backends/megatron_utils/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,14 +152,14 @@ def setup_model_and_optimizer(
optimizer = get_megatron_muon_optimizer(
config=config,
model_chunks=model,
use_gloo_process_groups=args.enable_gloo_process_groups,
use_gloo_process_groups=args.use_gloo_process_groups,
layer_wise_distributed_optimizer="dist" in config.optimizer.lower(),
)
else:
optimizer = get_megatron_optimizer(
config=config,
model_chunks=model,
use_gloo_process_groups=args.enable_gloo_process_groups,
use_gloo_process_groups=args.use_gloo_process_groups,
)
opt_param_scheduler = get_optimizer_param_scheduler(args, optimizer)
return model, optimizer, opt_param_scheduler
Expand Down Expand Up @@ -449,8 +449,9 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p
"loss_mask": batch["full_loss_masks"],
}

if args.enable_mtp_training:
forward_kwargs["mtp_kwargs"] = {"mtp_labels": batch["tokens"]}
# MTP labels: Megatron's process_mtp_loss derives them from input_ids
# (== batch["tokens"]) when labels is None, so no mtp_kwargs is needed.
# MTP head detach is configured via config.mtp_detach_heads (model_provider).

if (x := batch["multimodal_train_inputs"]) is not None:
forward_kwargs.update(x)
Expand Down Expand Up @@ -642,19 +643,23 @@ def train(
config.param_sync_func = param_sync_func
pre_hook_enabled = True

mtp_losses = None
if args.enable_mtp_training:
from megatron.core.transformer.multi_token_prediction import MTPLossLoggingHelper

mtp_loss_scale = 1 / num_microbatches[step_id]
# New Megatron tracks MTP loss as loss_sums/num_tokens (or loss_values) rather than a
# pre-divided "values" tensor. reduce_loss_in_tracker() all-reduces across ranks
# (collective: call on all ranks) and computes the per-token loss into tracker["values"].
MTPLossLoggingHelper.reduce_loss_in_tracker()
tracker = MTPLossLoggingHelper.tracker
# here we assume only one mtp layer
if "values" in tracker:
values = tracker["values"]
if (x := tracker.get("reduce_group")) is not None:
torch.distributed.all_reduce(values, group=x)
if (x := tracker.get("avg_group")) is not None:
torch.distributed.all_reduce(values, group=x, op=torch.distributed.ReduceOp.AVG)
# here we assume only one mtp layer
mtp_losses = (tracker["values"] * mtp_loss_scale).item()
elif "loss_values" in tracker:
mtp_losses = (tracker["loss_values"] * mtp_loss_scale).item()

if mtp_losses is not None:
MTPLossLoggingHelper.clean_loss_in_tracker()

# CI check: verify MTP loss is within expected bounds
Expand All @@ -670,7 +675,7 @@ def train(
role_tag = "" if role == "actor" else f"{role}-"

extra_metrics = {}
if args.enable_mtp_training:
if args.enable_mtp_training and mtp_losses is not None:
extra_metrics["mtp_loss"] = mtp_losses

if not disable_optimizer:
Expand Down
10 changes: 8 additions & 2 deletions miles/backends/megatron_utils/model_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,14 @@ def model_provider(
assert config is None, "miles builds the config from args, so it expects config to be None"
config = core_transformer_config_from_args(args)

# `enable_mtp_training` comes from miles' arg parser; megatron-only arg contexts
# (e.g. the run_megatron debug worker) won't have it, so default to False.
if getattr(args, "enable_mtp_training", False):
# RL MTP training: detach the MTP heads so MTP gradients do not flow into the
# shared output layer / embedding (Megatron implements this via mtp_detach_heads;
# previously miles patched Megatron directly). See bump_docs/01-cherry-pick.md (#6).
config.mtp_detach_heads = True

if args.spec is not None:
transformer_layer_spec = import_module(args.spec)
# Allow the spec to be a function so that user can use customized Megatron easier.
Expand All @@ -178,15 +186,13 @@ def model_provider(
moe_grouped_gemm=args.moe_grouped_gemm,
qk_layernorm=args.qk_layernorm,
multi_latent_attention=args.multi_latent_attention,
moe_use_legacy_grouped_gemm=args.moe_use_legacy_grouped_gemm,
)
else:
transformer_layer_spec = get_gpt_layer_local_spec(
num_experts=args.num_experts,
moe_grouped_gemm=args.moe_grouped_gemm,
qk_layernorm=args.qk_layernorm,
multi_latent_attention=args.multi_latent_attention,
moe_use_legacy_grouped_gemm=args.moe_use_legacy_grouped_gemm,
normalization=args.normalization,
use_kitchen=config.use_kitchen,
use_true_on_policy_backend=config.true_on_policy_contract is not None,
Expand Down
4 changes: 3 additions & 1 deletion miles/utils/debug_utils/run_megatron/worker/main.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,7 +132,9 @@ def _initialize_megatron(args: argparse.Namespace) -> None:

def _build_and_load_model(args: argparse.Namespace, script: WorkerScriptArgs) -> list[Any]:
model_provider: Callable[..., Any] = get_model_provider_func(args, role=script.role)
model: list[Any] = get_model(model_provider, ModelType.encoder_or_decoder)
# Forward-only runs skip DDP wrapping so the distributed-optimizer grad buffer
# (tens-to-hundreds of GB for large MoE models) is never allocated -> avoids OOM.
model: list[Any] = get_model(model_provider, ModelType.encoder_or_decoder, wrap_with_ddp=script.run_backward)

if args.load is not None:
load_checkpoint(
Expand Down
12 changes: 6 additions & 6 deletions miles_plugins/mbridge/deepseekv4.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,17 +107,17 @@ def _build_config(self):
config.dsv4_hc_sinkhorn_iters = getattr(self.hf_config, "hc_sinkhorn_iters", 20)
config.dsv4_hc_eps = getattr(self.hf_config, "hc_eps", 1e-6)

config.dsv4_compress_ratios = getattr(self.hf_config, "compress_ratios", None)
config.dsv4_compress_rope_theta = getattr(self.hf_config, "compress_rope_theta", 160000)
config.csa_compress_ratios = getattr(self.hf_config, "compress_ratios", None)
config.csa_compress_rotary_base = getattr(self.hf_config, "compress_rope_theta", 160000)

config.dsv4_swiglu_limit = getattr(self.hf_config, "swiglu_limit", 0.0)
if config.dsv4_swiglu_limit > 0:
config.bias_activation_fusion = False
config.activation_func_clamp_value = config.dsv4_swiglu_limit

config.dsv4_o_groups = getattr(self.hf_config, "o_groups", 8)
config.dsv4_o_lora_rank = getattr(self.hf_config, "o_lora_rank", 1024)
config.dsv4_n_hash_layers = getattr(self.hf_config, "n_hash_layers", 3)
config.dsv4_window_size = getattr(self.hf_config, "window_size", 128)
config.o_groups = getattr(self.hf_config, "o_groups", 8)
config.o_lora_rank = getattr(self.hf_config, "o_lora_rank", 1024)
config.moe_n_hash_layers = getattr(self.hf_config, "n_hash_layers", 3)
config.csa_window_size = getattr(self.hf_config, "window_size", 128)

return config
4 changes: 2 additions & 2 deletions miles_plugins/mbridge/glm4moe.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,14 +41,14 @@ def _weight_name_mapping_mtp(self, name: str, num_layers: int) -> str:
elif "mlp" in name:
mtp_layer_index = int(re.findall(r"mtp\.layers\.(\d+)\.", name)[0])
name_ = re.sub(
r"^mtp\.layers.\d+.transformer_layer", f"model.layers.{num_layers+mtp_layer_index}", name
r"^mtp\.layers.\d+.(?:transformer_layer|mtp_model_layer)", f"model.layers.{num_layers+mtp_layer_index}", name
)
convert_names = self._weight_name_mapping_mlp(name_)
break
elif "self_attention" in name:
mtp_layer_index = int(re.findall(r"mtp\.layers.(\d+)\.", name)[0])
name_ = re.sub(
r"^mtp\.layers.\d+.transformer_layer", f"model.layers.{num_layers+mtp_layer_index}", name
r"^mtp\.layers.\d+.(?:transformer_layer|mtp_model_layer)", f"model.layers.{num_layers+mtp_layer_index}", name
)
convert_names = self._weight_name_mapping_attention(name_)
break
Expand Down
7 changes: 5 additions & 2 deletions miles_plugins/mbridge/glm4moe_lite.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,8 +130,11 @@ def _convert_mtp_param(self, name: str) -> tuple[list[str]]:
if name in direct_name_mapping:
return [direct_name_mapping[name]]

assert "mtp.layers.0.transformer_layer" in name, "mtp not found"
proxy_name = name.replace("mtp.layers.0.transformer_layer", f"decoder.layers.{mtp_layer_id}")
_mtp_inner = next(
(p for p in ("transformer_layer", "mtp_model_layer") if f"mtp.layers.0.{p}" in name), None
)
assert _mtp_inner is not None, "mtp not found"
proxy_name = name.replace(f"mtp.layers.0.{_mtp_inner}", f"decoder.layers.{mtp_layer_id}")
if "self_attention" in proxy_name or "input_layernorm.weight" in proxy_name:
return self._weight_name_mapping_attention(proxy_name)
if "mlp" in proxy_name:
Expand Down
11 changes: 6 additions & 5 deletions miles_plugins/mbridge/mimo.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,13 +76,14 @@ def _convert_mtp_param(self, name: str) -> list[str]:
if name in direct_name_mapping:
return [direct_name_mapping[name]]

# Handle transformer components within MTP
# Check if this is a transformer_layer component
if "transformer_layer" in name:
# Handle transformer components within MTP. New Megatron renamed the MTP submodule
# transformer_layer -> mtp_model_layer; accept both.
_mtp_inner = next((p for p in ("transformer_layer", "mtp_model_layer") if p in name), None)
if _mtp_inner is not None:
# Create a proxy name to use with parent class methods
# Convert mtp.layers.{idx}.transformer_layer.* to decoder.layers.{idx}.*
# Convert mtp.layers.{idx}.<inner>.* to decoder.layers.{idx}.*
proxy_name = name.replace(
f"mtp.layers.{mtp_layer_idx}.transformer_layer",
f"mtp.layers.{mtp_layer_idx}.{_mtp_inner}",
f"decoder.layers.{mtp_layer_idx}",
)

Expand Down
5 changes: 3 additions & 2 deletions miles_plugins/mbridge/qwen3_5.py
Original file line number Diff line number Diff line change
Expand Up @@ -269,9 +269,10 @@ def _convert_mtp_param(self, name: str) -> list[str]:
if name in direct_name_mapping:
return [direct_name_mapping[name]]

if "transformer_layer" in name:
_mtp_inner = next((p for p in ("transformer_layer", "mtp_model_layer") if p in name), None)
if _mtp_inner is not None:
proxy_name = name.replace(
f"mtp.layers.{mtp_layer_idx}.transformer_layer",
f"mtp.layers.{mtp_layer_idx}.{_mtp_inner}",
f"decoder.layers.{mtp_layer_idx}",
)

Expand Down
5 changes: 3 additions & 2 deletions miles_plugins/mbridge/qwen3_next.py
Original file line number Diff line number Diff line change
Expand Up @@ -130,9 +130,10 @@ def _convert_mtp_param(self, name: str) -> list[str]:
if name in direct_mappings:
return [direct_mappings[name]]

if "transformer_layer" in name:
_mtp_inner = next((p for p in ("transformer_layer", "mtp_model_layer") if p in name), None)
if _mtp_inner is not None:
proxy_name = name.replace(
f"mtp.layers.{mtp_layer_idx}.transformer_layer",
f"mtp.layers.{mtp_layer_idx}.{_mtp_inner}",
f"decoder.layers.{mtp_layer_idx}",
)

Expand Down
Loading
Loading