Skip to content

Commit b8d0eff

Browse files
fix(compat): adapt Bridge to Megatron-Core dev
Signed-off-by: svcnemo-autobot <svcnemo-autobot@nvidia.com>
1 parent aa319c9 commit b8d0eff

4 files changed

Lines changed: 40 additions & 10 deletions

File tree

src/megatron/bridge/models/gemma/modeling_gemma4.py

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1286,10 +1286,16 @@ def _forward_mlp(
12861286
hidden_states: Tensor,
12871287
inference_context: BaseInferenceContext | None = None,
12881288
padding_mask: Tensor | None = None,
1289+
input_ids: Tensor | None = None,
12891290
packed_seq_params: PackedSeqParams | None = None,
12901291
) -> Tensor:
1291-
"""Run HF's separate shared-expert, routed-expert, and router inputs."""
1292-
del inference_context
1292+
"""Run HF's separate shared-expert, routed-expert, and router inputs.
1293+
1294+
``input_ids`` is accepted for compatibility with Megatron-Core ``dev``'s
1295+
``TransformerLayer.forward``, which forwards it for hash-based MoE routing;
1296+
Gemma 4's separate-input MoE path does not use it.
1297+
"""
1298+
del inference_context, input_ids
12931299
residual = hidden_states.float() if self.config.fp32_residual_connection else hidden_states
12941300

12951301
moe_input = residual
@@ -1377,9 +1383,10 @@ def routing(
13771383
logits: Tensor,
13781384
padding_mask: Tensor | None = None,
13791385
input_ids: Tensor | None = None,
1386+
**kwargs: object,
13801387
) -> tuple[Tensor, Tensor | None]:
1381-
# Token identities do not affect Gemma 4 routing; retain the argument for compatibility with existing callers.
1382-
del input_ids
1388+
"""Route Gemma 4 tokens with arguments accepted by both MCore refs."""
1389+
del input_ids, kwargs
13831390
routing_probs, routing_map = super().routing(logits, padding_mask=padding_mask)
13841391
if routing_map is not None:
13851392
prob_sums = routing_probs.sum(dim=-1, keepdim=True).clamp(min=1e-20)

tests/functional_tests/test_groups/models/deepseek/test_deepseek_v4_conversion.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -144,7 +144,8 @@ def copy(src: str, dst: str) -> None:
144144

145145
indexer_prefix = f"{compressor_prefix}.indexer"
146146
copy(f"{indexer_prefix}.q_b_proj.weight", f"{ckpt_prefix}.attn.indexer.wq_b.weight")
147-
copy(f"{indexer_prefix}.weights_proj.weight", f"{ckpt_prefix}.attn.indexer.weights_proj.weight")
147+
# HF transformers nests the indexer score projection under `.scorer.`.
148+
copy(f"{indexer_prefix}.scorer.weights_proj.weight", f"{ckpt_prefix}.attn.indexer.weights_proj.weight")
148149
copy(f"{indexer_prefix}.kv_proj.weight", f"{ckpt_prefix}.attn.indexer.compressor.wkv.weight")
149150
copy(f"{indexer_prefix}.gate_proj.weight", f"{ckpt_prefix}.attn.indexer.compressor.wgate.weight")
150151
copy(f"{indexer_prefix}.position_bias", f"{ckpt_prefix}.attn.indexer.compressor.ape")

tests/unit_tests/models/gemma/test_gemma4_modeling.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1910,7 +1910,11 @@ def fake_routing(self, logits, padding_mask=None, input_ids=None):
19101910
router = object.__new__(Gemma4TopKRouter)
19111911
router.per_expert_scale = torch.tensor([1.0, 2.0, 3.0])
19121912

1913-
out_probs, out_map = Gemma4TopKRouter.routing(router, torch.zeros(2, 3))
1913+
out_probs, out_map = Gemma4TopKRouter.routing(
1914+
router,
1915+
torch.zeros(2, 3),
1916+
packed_seq_params=object(),
1917+
)
19141918

19151919
assert out_map is routing_map
19161920
torch.testing.assert_close(out_probs[0], torch.tensor([0.4, 1.2, 0.0]))

tests/unit_tests/recipes/test_glm5_perf_recipes.py

Lines changed: 22 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -20,9 +20,6 @@
2020
from types import SimpleNamespace
2121

2222
import pytest
23-
from megatron.core.models.gpt.experimental_attention_variant_module_specs import (
24-
_validate_dsa_index_share_pipeline_split,
25-
)
2623
from megatron.core.transformer.enums import LayerType
2724
from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout
2825

@@ -90,6 +87,20 @@ def _build_recipe(recipe_func: Callable[[], ConfigContainer], monkeypatch: pytes
9087
return recipe_func()
9188

9289

90+
def _dsa_source_layer_id(layer_id: int, *, skip_topk_offset: int, topk_freq: int) -> int:
91+
"""Return the zero-based source layer defined by MCore's DSA sharing contract."""
92+
# Mirrors the private MCore `_validate_dsa_index_share_pipeline_split`
93+
# helper that guarded this recipe before removal from the pinned revision.
94+
# MCore defines DSA sharing with one-based layers: layers through
95+
# max(skip_topk_offset, 1) compute their own indices, then each topk_freq
96+
# group reuses the indices computed by its first layer.
97+
layer_number = layer_id + 1
98+
sharing_offset = max(skip_topk_offset, 1)
99+
if layer_number <= sharing_offset:
100+
return layer_id
101+
return layer_number - ((layer_number - sharing_offset) % topk_freq) - 1
102+
103+
93104
@pytest.mark.parametrize("recipe_func", _RECIPES, ids=lambda recipe: recipe.__name__)
94105
def test_glm5_perf_recipes_are_flat_and_preserve_bridge_dsa_fields(
95106
recipe_func: Callable[[], ConfigContainer], monkeypatch: pytest.MonkeyPatch
@@ -183,7 +194,14 @@ def test_glm52_h100_pipeline_layout_keeps_dsa_index_sharing_within_each_vpp_chun
183194
decoder_count = stage.count(LayerType.decoder)
184195
if decoder_count:
185196
local_layer_ids = range(decoder_offset, decoder_offset + decoder_count)
186-
_validate_dsa_index_share_pipeline_split(cfg.model, local_layer_ids)
197+
local_layer_id_set = set(local_layer_ids)
198+
for layer_id in local_layer_ids:
199+
source_layer_id = _dsa_source_layer_id(
200+
layer_id,
201+
skip_topk_offset=cfg.model.dsa_indexer_skip_topk_offset,
202+
topk_freq=cfg.model.dsa_indexer_topk_freq,
203+
)
204+
assert source_layer_id in local_layer_id_set
187205
decoder_offset += decoder_count
188206

189207
assert decoder_offset == cfg.model.num_layers

0 commit comments

Comments
 (0)