Skip to content
Merged
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 3rdparty/Megatron-LM
Submodule Megatron-LM updated 371 files
71 changes: 64 additions & 7 deletions src/megatron/bridge/training/optim.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,16 +12,68 @@
# See the License for the specific language governing permissions and
# limitations under the License.

from typing import Union
from typing import Dict, Optional, Union

from megatron.core.optimizer import MegatronOptimizer, OptimizerConfig, get_megatron_optimizer
from megatron.core.optimizer import (
MegatronOptimizer,
OptimizerConfig,
ParamKey,
get_megatron_optimizer,
)
from megatron.core.optimizer.muon import get_megatron_muon_optimizer
from megatron.core.optimizer_param_scheduler import OptimizerParamScheduler
from megatron.core.optimizer_param_scheduler import OptimizerParamScheduler, ParamGroupOverride
from megatron.core.transformer.module import MegatronModule

from megatron.bridge.training.config import SchedulerConfig


def _build_config_overrides(
scheduler_config: SchedulerConfig,
model: Union[MegatronModule, list[MegatronModule]],
) -> Optional[Dict[ParamKey, ParamGroupOverride]]:
"""Build config overrides for weight decay based on scheduler configuration.

This function creates parameter-specific overrides for weight decay behavior.
By default, weight decay is skipped for bias parameters and 1D parameters.
For Qwen3-Next models, weight decay is applied to q_layernorm and k_layernorm.

Args:
scheduler_config: Scheduler configuration containing weight decay settings
model: The model or list of model chunks to collect parameter names from

Returns:
Dictionary of ParamKey to ParamGroupOverride for the optimizer
"""
config_overrides: Dict[ParamKey, ParamGroupOverride] = {}

# Collect param names that should skip weight decay
no_wd_names: list[str] = []
is_qwen3_next = scheduler_config.no_weight_decay_cond_type == "qwen3_next"

model_list = model if isinstance(model, list) else [model]
for model_chunk in model_list:
for name, param in model_chunk.named_parameters():
# Skip weight decay for bias parameters
if name.endswith(".bias"):
no_wd_names.append(name)
continue

# Skip weight decay for 1D parameters
if len(param.shape) == 1:
if is_qwen3_next:
# Qwen3-Next: apply weight decay to qk layernorm (don't add to skip list)
if "q_layernorm" in name or "k_layernorm" in name:
continue
no_wd_names.append(name)

# Create a single ParamKey with all names that should skip weight decay
if no_wd_names:
no_wd_key = ParamKey(name=tuple(no_wd_names))
config_overrides[no_wd_key] = ParamGroupOverride(wd_mult=0.0)

return config_overrides if config_overrides else None


def setup_optimizer(
optimizer_config: OptimizerConfig,
scheduler_config: SchedulerConfig,
Expand All @@ -39,16 +91,21 @@ def setup_optimizer(
Returns:
tuple containing the optimizer and scheduler
"""
# Build config overrides for weight decay based on scheduler config and model params
config_overrides = _build_config_overrides(scheduler_config, model)

if "muon" not in optimizer_config.optimizer and "soap" not in optimizer_config.optimizer:
optimizer = get_megatron_optimizer(
optimizer_config,
model,
config=optimizer_config,
model_chunks=model,
config_overrides=config_overrides,
use_gloo_process_groups=use_gloo_process_groups,
)
else:
optimizer = get_megatron_muon_optimizer(
optimizer_config,
model,
config=optimizer_config,
model_chunks=model,
config_overrides=config_overrides,
use_gloo_process_groups=use_gloo_process_groups,
layer_wise_distributed_optimizer="dist" in optimizer_config.optimizer,
)
Expand Down
2 changes: 0 additions & 2 deletions src/megatron/bridge/training/setup.py
Original file line number Diff line number Diff line change
Expand Up @@ -221,8 +221,6 @@ def modelopt_pre_wrap_hook(model):

cfg.model.timers = timers
cfg.optimizer.timers = timers
if cfg.scheduler.no_weight_decay_cond_type == "qwen3_next":
raise NotImplementedError("qwen3_next style weight decay disabled until mcore fix.")
optimizer, scheduler = setup_optimizer(
optimizer_config=cfg.optimizer,
scheduler_config=cfg.scheduler,
Expand Down
Loading
Loading