Skip to content

Commit f658f69

Browse files
committed
Document and pin zero decay for draft norm and bias params
Signed-off-by: seonjinn <sna@nvidia.com>
1 parent 21d0587 commit f658f69

2 files changed

Lines changed: 39 additions & 0 deletions

File tree

nemo_rl/models/megatron/draft/optimizer.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -60,6 +60,11 @@ def build_config_overrides(
6060
if minimum_lr is not None:
6161
draft_override["min_lr"] = minimum_lr
6262
if self.draft_optimizer.weight_decay is not None:
63+
# Megatron's standard overrides still set wd_mult=0.0 for norm/bias
64+
# and 1-D draft params, and the scheduler multiplies:
65+
# weight_decay = get_wd(group) * wd_mult. So this override changes
66+
# decay only for weight matrices; draft norms/biases stay at 0,
67+
# matching the main model's convention.
6368
draft_override["start_wd"] = self.draft_optimizer.weight_decay
6469
draft_override["end_wd"] = self.draft_optimizer.weight_decay
6570

tests/unit/models/megatron/test_draft_optimizer.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
138172
def 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

Comments
 (0)