@@ -135,6 +135,40 @@ def test_draft_optimizer_provider_only_overrides_tagged_parameters() -> None:
135135 )
136136
137137
138+ def test_draft_weight_decay_override_keeps_norm_and_bias_at_zero_decay () -> None :
139+ """Megatron's standard wd_mult=0.0 for 1-D params survives the draft override:
140+ the scheduler multiplies get_wd(group) * wd_mult, so draft norms/biases keep
141+ zero decay even when draft.optimizer.weight_decay is set."""
142+ from megatron .core .optimizer import get_standard_config_overrides
143+
144+ bias_like = torch .nn .Parameter (torch .ones (4 ))
145+ bias_like .grad_norm_group = "draft"
146+ provider = DraftOptimizerConfigOverrideProvider (
147+ DraftOptimizerConfig (lr = 1.0e-3 , weight_decay = 0.02 )
148+ )
149+ context = OptimizerConfigOverrideProviderContext (
150+ scheduler_config = MagicMock (),
151+ optimizer_config = OptimizerConfig (lr = 2.0e-3 , min_lr = 2.0e-4 , weight_decay = 0.1 ),
152+ model = MagicMock (),
153+ )
154+ draft_overrides = provider .build_config_overrides (context )
155+ assert draft_overrides is not None
156+ standard_overrides = get_standard_config_overrides (context .optimizer_config )
157+
158+ merged = {** standard_overrides , ** draft_overrides }
159+ combined = combine_param_group_overrides (
160+ [
161+ override
162+ for key , override in merged .items ()
163+ if key .matches (bias_like , "draft.norm.bias" )
164+ ]
165+ )
166+
167+ assert combined ["start_wd" ] == 0.02
168+ assert combined ["wd_mult" ] == 0.0
169+ assert combined ["start_wd" ] * combined ["wd_mult" ] == 0.0
170+
171+
138172def test_draft_optimizer_provider_rejects_incompatible_inherited_min_lr () -> None :
139173 parameter = torch .nn .Parameter (torch .ones (1 ))
140174 parameter .grad_norm_group = "draft"
0 commit comments