diff --git a/megatron/core/models/mimo/model/base.py b/megatron/core/models/mimo/model/base.py index 334a7883d92..e52e9ab8258 100644 --- a/megatron/core/models/mimo/model/base.py +++ b/megatron/core/models/mimo/model/base.py @@ -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: diff --git a/tests/unit_tests/models/mimo/test_mimo_zero_grad_buffer.py b/tests/unit_tests/models/mimo/test_mimo_zero_grad_buffer.py new file mode 100644 index 00000000000..ba7fb8d6fa4 --- /dev/null +++ b/tests/unit_tests/models/mimo/test_mimo_zero_grad_buffer.py @@ -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()