[bug] Gemma4 with cfg.model.recompute_granularity = "full" is not working
- Dominant language
- Python
- Stars
- 17.9k
- Forks
- 4.5k
- Avg merge
- 4d 3h
- Merged PRs (30d)
- 272
Description
**Describe the bug**
As stated in the title `cfg.model.recompute_granularity = "full"` is not working with Gemma4.
**Traceback:**
```
[rank0]: Traceback (most recent call last):
[rank0]: File "/src/megatron/tgpt-megatron/examples/reproduce_example.py", line 73, in
[rank0]: pretrain(cfg, forward_step)
[rank0]: File "/app/Megatron-Bridge/src/megatron/bridge/utils/decorators.py", line 39, in wrapper
[rank0]: return func(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/src/megatron/bridge/training/pretrain.py", line 98, in pretrain
[rank0]: _pretrain(state=state, forward_step_func=forward_step_func, callback_manager=callback_manager)
[rank0]: File "/app/Megatron-Bridge/src/megatron/bridge/training/pretrain.py", line 142, in _pretrain
[rank0]: train(
[rank0]: File "/app/Megatron-Bridge/src/megatron/bridge/training/train.py", line 445, in train
[rank0]: ) = wrapped_train_step(
[rank0]: ^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/src/megatron/bridge/training/train.py", line 849, in train_step
[rank0]: losses_reduced = forward_backward_func(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/pipeline_parallel/schedules.py", line 695, in forward_backward_no_pipelining
[rank0]: output_tensor, num_tokens = forward_step(
[rank0]: ^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/pipeline_parallel/schedules.py", line 437, in forward_step
[rank0]: output_tensor, loss_func = forward_step_func(data_iterator, model)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/src/megatron/bridge/training/vlm_step.py", line 480, in forward_step
[rank0]: model_output = model(**forward_args)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py", line 1778, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py", line 1789, in _call_impl
[rank0]: return forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/distributed/data_parallel_base.py", line 22, in forward
[rank0]: return self.module(*inputs, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py", line 1778, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py", line 1789, in _call_impl
[rank0]: return forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/module.py", line 493, in forward
[rank0]: outputs = self.module(*inputs, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py", line 1778, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py", line 1789, in _call_impl
[rank0]: return forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/src/megatron/bridge/models/gemma_vl/modeling_gemma4_vl.py", line 198, in forward
[rank0]: outputs = self.language_model.forward(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/models/gpt/gpt_model.py", line 540, in forward
[rank0]: hidden_states = self.decoder(
[rank0]: ^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/transformer_block.py", line 643, in __call__
[rank0]: return super().__call__(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/module.py", line 356, in __call__
[rank0]: return super().__call__(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py", line 1778, in _wrapped_call_impl
[rank0]: return self._call_impl(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/nn/modules/module.py", line 1789, in _call_impl
[rank0]: return forward_call(*args, **kwargs)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/transformer_block.py", line 784, in forward
[rank0]: checkpointed_result = self._checkpointed_forward(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/transformer_block.py", line 555, in _checkpointed_forward
[rank0]: hidden_states, context = checkpoint_handler(custom(layer_idx, chunk_end))
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/transformer/transformer_block.py", line 535, in checkpoint_handler
[rank0]: return tensor_parallel.checkpoint(
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/app/Megatron-Bridge/3rdparty/Megatron-LM/megatron/core/tensor_parallel/random.py", line 642, in checkpoint
[rank0]: return CheckpointFunction.apply(function, distribute_saved_activations, *args)
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: File "/usr/local/lib/python3.12/dist-packages/torch/autograd/function.py", line 583, in apply
[rank0]: return super().apply(*args, **kwargs) # type: ignore[misc]
[rank0]: ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
[rank0]: TypeError: save_for_backward can only save variables, but argument 4 is of type tuple
```
**Steps/Code to reproduce bug**
```python
import torch
from megatron.bridge import AutoBridge
from megatron.bridge.training.vlm_step import forward_step
from megatron.bridge.training.pretrain import pretrain
from megatron.bridge.recipes.common import _pretrain_common
def get_config():
reproduce_issue = True
model_path = "google/gemma-4-26B-A4B"
cfg = _pretrain_common()
cfg.model = AutoBridge.from_hf_pretrained(model_path).to_megatron_provider(load_weights=False)
cfg.tokenizer.tokenizer_model = model_path
cfg.dataset.seq_length = 128
cfg.model.seq_length = cfg.dataset.seq_length
cfg.logger.log_interval = 1
cfg.checkpoint.save = None
cfg.checkpoint.load = None
cfg.checkpoint.pretrained_checkpoint = None
cfg.dataset.dataloader_type = "batch"
cfg.dataset.num_workers = 8
cfg.train.train_iters = 300
cfg.train.eval_iters = 0
# Parallelism settings
cfg.model.tensor_model_parallel_size = 1
cfg.model.pipeline_model_parallel_size = 1
cfg.model.expert_model_parallel_size = 1
cfg.model.expert_tensor_parallel_size = 1
cfg.model.pipeline_model_parallel_layout = None
cfg.model.pipeline_dtype = torch.bfloat16
cfg.model.params_dtype = torch.bfloat16
cfg.model.virtual_pipeline_model_parallel_size = None
cfg.model.context_parallel_size = 1
cfg.model.moe_router_fusion = True
cfg.model.moe_permute_fusion = True
cfg.model.moe_grouped_gemm = True
cfg.model.sequence_parallel = True
cfg.model.context_parallel_size = 1
cfg.model.transformer_impl = "transformer_engine"
cfg.model.attention_backend = "auto"
cfg.model.cross_entropy_loss_fusion = True
cfg.model.cross_entropy_fusion_impl = "te"
cfg.model.freeze_vision_model = True
cfg.model.freeze_vision_projection = True
cfg.optimizer.optimizer_cpu_offload=True
cfg.optimizer.optimizer_offload_fraction=1.0
# Trim model
cfg.model.num_layers = 2
cfg.model.num_moe_experts = 4
cfg.model.moe_router_topk = 2
cfg.model.interleaved_attn_pattern = (1, 1)
if reproduce_issue:
cfg.model.recompute_granularity = "full"
cfg.model.recompute_method = "uniform"
cfg.model.recompute_num_layers = 1
cfg.model.finalize()
return cfg
if __name__ == "__main__":
cfg = get_config()
pretrain(cfg, forward_step)
```
**To run the code:**
`torchrun --nproc_per_node 1 reproduce_example.py`
**Additional context**
Megatron-LM: `2d1fa8d372a3990b0bb1334cd686f15005ee138f`
Megatron-Bridge: `5d1b2c8da43d6d751333589b981da08d79e6fa79`
Contributor guide
Assessment
This issue has not been assessed yet.