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
13 changes: 13 additions & 0 deletions megatron/core/models/mimo/model/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -278,6 +278,19 @@ def set_input_tensor(self, input_tensor):
if self.language_model is not None and hasattr(self.language_model, 'set_input_tensor'):
self.language_model.set_input_tensor(input_tensor)

def _active_submodules(self):
"""Yield this rank's present submodules."""
if self.language_model is not None:
yield self.language_model
for submodule in self.modality_submodules.values():
if submodule is not None:
yield submodule

def zero_grad_buffer(self):
"""Zero each active submodule's DDP grad buffer."""
for module in self._active_submodules():
module.zero_grad_buffer()

def get_text_embeddings(
self, input_ids: torch.Tensor, position_ids: torch.Tensor, special_token_ids: Dict[str, int]
) -> torch.Tensor:
Expand Down
33 changes: 33 additions & 0 deletions tests/unit_tests/models/mimo/test_mimo_zero_grad_buffer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.

"""CPU-only tests for MimoModel.zero_grad_buffer fan-out."""

from types import SimpleNamespace
from unittest.mock import MagicMock

from megatron.core.models.mimo.model.base import MimoModel


def _stub(language_model, modality_submodules):
"""Minimal stand-in carrying the real zero_grad_buffer / _active_submodules."""
stub = SimpleNamespace(language_model=language_model, modality_submodules=modality_submodules)
stub.zero_grad_buffer = MimoModel.zero_grad_buffer.__get__(stub)
stub._active_submodules = MimoModel._active_submodules.__get__(stub)
return stub


def test_zero_grad_buffer_fans_out_to_present_submodules():
language_model = MagicMock()
vision = MagicMock()
_stub(language_model, {"vision": vision}).zero_grad_buffer()

language_model.zero_grad_buffer.assert_called_once_with()
vision.zero_grad_buffer.assert_called_once_with()


def test_zero_grad_buffer_skips_none_submodules():
vision = MagicMock()
# Encoder-only rank: language_model is None, plus a None modality entry.
_stub(None, {"vision": vision, "audio": None}).zero_grad_buffer()

vision.zero_grad_buffer.assert_called_once_with()
Loading