From d8f0141238a4adf05c90dfdc25cfaa8c266b754a Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Sun, 20 Sep 2026 19:34:48 +0800 Subject: [PATCH] fix: avoid duplicate gradient norm reduction in pure DDP --- src/twinkle/model/transformers/transformers.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/src/twinkle/model/transformers/transformers.py b/src/twinkle/model/transformers/transformers.py index 0087cc7f..69108972 100644 --- a/src/twinkle/model/transformers/transformers.py +++ b/src/twinkle/model/transformers/transformers.py @@ -1051,12 +1051,18 @@ def clip_grad_norm(self, max_grad_norm: float = 1.0, norm_type=2, **kwargs): ep_clip_kwargs = self.strategy.get_ep_clip_kwargs(self.model) if hasattr( self.strategy, 'get_ep_clip_kwargs') else {} + grad_norm_group = optimizer_config._dp_group + if isinstance(self.model, torch.nn.parallel.DistributedDataParallel) and ( + self.device_mesh is None or self.device_mesh.world_size == self.device_mesh.dp_world_size): + # Pure DDP gradients are replicated; only token counts need DP reduction. + grad_norm_group = None + grad_norm = normalize_and_clip_grad_norm( parameters, num_tokens=num_tokens, max_grad_norm=max_grad_norm, norm_type=norm_type, - group=optimizer_config._dp_group, + group=grad_norm_group, **ep_clip_kwargs, ) optimizer_config._last_grad_norm = grad_norm