From fce4557a3cbf8f98498b7b9dc02d7d798ad758f3 Mon Sep 17 00:00:00 2001 From: qiyuw Date: Tue, 21 Apr 2026 12:10:46 -0700 Subject: [PATCH 1/3] add missing knob reduce_scatter_with_fp32_accumulation Signed-off-by: qiyuw --- megatron/core/distributed/distributed_data_parallel.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/megatron/core/distributed/distributed_data_parallel.py b/megatron/core/distributed/distributed_data_parallel.py index 35325d70ce9..1816bafa4d6 100644 --- a/megatron/core/distributed/distributed_data_parallel.py +++ b/megatron/core/distributed/distributed_data_parallel.py @@ -222,7 +222,13 @@ def _allocate_buffers_for_parameters( # kernels. # If bucketing is explicitly disabled, then put all buckets in a buffer into a single # bucket group. - bucket_groups = partition_buckets(buffers, force_single_bucket_group=disable_bucketing) + bucket_groups = partition_buckets( + buffers, + force_single_bucket_group=disable_bucketing, + reduce_scatter_with_fp32_accumulation=( + self.ddp_config.reduce_scatter_with_fp32_accumulation + ), + ) if self.ddp_config.num_distributed_optimizer_instances > 1: assert ( From 3e40da3a8595b0ad771bfed24f49e23b0cd48b68 Mon Sep 17 00:00:00 2001 From: qiyuw Date: Thu, 23 Apr 2026 11:13:28 -0700 Subject: [PATCH 2/3] lint --- megatron/core/distributed/distributed_data_parallel.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/megatron/core/distributed/distributed_data_parallel.py b/megatron/core/distributed/distributed_data_parallel.py index fccb06e8657..a711f1405d1 100644 --- a/megatron/core/distributed/distributed_data_parallel.py +++ b/megatron/core/distributed/distributed_data_parallel.py @@ -265,13 +265,15 @@ def __init__( # If bucketing is explicitly disabled, then put all buckets in a buffer into a single # bucket group. self.bucket_groups = partition_buckets( - self.buffers, force_single_bucket_group=disable_bucketing, + self.buffers, + force_single_bucket_group=disable_bucketing, reduce_scatter_with_fp32_accumulation=( self.ddp_config.reduce_scatter_with_fp32_accumulation ), ) self.expert_parallel_bucket_groups = partition_buckets( - self.expert_parallel_buffers, force_single_bucket_group=disable_bucketing, + self.expert_parallel_buffers, + force_single_bucket_group=disable_bucketing, reduce_scatter_with_fp32_accumulation=( self.ddp_config.reduce_scatter_with_fp32_accumulation ), From 5c70ea7d57f3235dd0689ccbaa705cb7d5c8a36d Mon Sep 17 00:00:00 2001 From: qiyuw Date: Fri, 24 Apr 2026 09:49:41 -0700 Subject: [PATCH 3/3] explicitly del old fp4 rowwise data Signed-off-by: qiyuw --- megatron/core/fp4_utils.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/megatron/core/fp4_utils.py b/megatron/core/fp4_utils.py index 245a04eb39e..45e57285a8d 100644 --- a/megatron/core/fp4_utils.py +++ b/megatron/core/fp4_utils.py @@ -81,7 +81,8 @@ def modify_nvfp4_rowwise_storage(fp4_tensor: torch.Tensor, new_rowwise_data: tor ), "Rowwise NVFP4 storage must be uint8" # Preserve existing values and then swap storage new_rowwise_data.detach().copy_(old_rowwise) - setattr(fp4_tensor, "_rowwise_data", new_rowwise_data) + fp4_tensor._rowwise_data = new_rowwise_data + del old_rowwise def quantize_nvfp4_param_shard(