Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions swift/trainers/mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down
Loading