diff --git a/swift/trainers/mixin.py b/swift/trainers/mixin.py index f9adc8ad2f..74ce06bf7d 100644 --- a/swift/trainers/mixin.py +++ b/swift/trainers/mixin.py @@ -1108,10 +1108,38 @@ def create_loss_and_eval_metric(self, args): def create_optimizer_and_scheduler(self, num_training_steps: int): self.optimizer_callback.create_optimizer_and_scheduler(num_training_steps) + def _disable_foreach_for_deepspeed(self): + """Disable foreach for AdamW on torch<2.9 with DeepSpeed ZeRO-1/2 (no-op otherwise). + + torch<2.9 `torch._foreach_*` kernels use a signed int32 element index and write out + of bounds when a flat tensor has numel > INT32_MAX. DeepSpeed ZeRO-1/2 flattens each + param group into one FP32 flat partition per DP rank; for full-parameter SFT of large + models that partition can exceed the boundary and corrupt memory (loss NaN). Falling + back to the single-tensor path (foreach=False) avoids the bug without changing the + optimizer's group layout or checkpoint format, and costs about the same as foreach + because ZeRO feeds the base optimizer a few large flat tensors rather than many + small ones (measured ~2% on the 2 GiB flat). + """ + if version.parse(torch.__version__) >= version.parse('2.9.0'): + return + ds_config = getattr(self.args, 'deepspeed', None) or {} + # ZeRO-3 partitions parameters individually (no oversized flat tensor); only + # stage 1/2 flatten a whole group into one FP32 partition per rank. + if not isinstance(ds_config, dict) or ds_config.get('zero_optimization', {}).get('stage') not in (1, 2): + return + optimizer = self.optimizer + if optimizer is None or not isinstance(optimizer, torch.optim.AdamW): + return + optimizer.defaults['foreach'] = False + for group in optimizer.param_groups: + group['foreach'] = False + logger.info('[adamw-foreach] disabled foreach for DeepSpeed ZeRO on torch<2.9.') + def create_optimizer(self, model=None): self._optimizer_ori = self.optimizer = self.optimizer_callback.create_optimizer(model=model) if self.optimizer is not None: self.optimizer.param_groups = [pg for pg in self.optimizer.param_groups if len(pg['params']) > 0] + self._disable_foreach_for_deepspeed() return self.optimizer def create_scheduler(self, num_training_steps: int, optimizer=None):