Skip to content
Merged
Show file tree
Hide file tree
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
4 changes: 3 additions & 1 deletion deepspeed/runtime/base_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -355,7 +355,9 @@ def scale_if_loss(self, value: Any) -> Any:
return self.external_loss_scale * value
if self.torch_autocast_gradscaler:
return self.torch_autocast_gradscaler.scale(value)
return self.loss_scaler.scale_loss(value)
# Only call loss_scaler if it exists (not present in BF16_Optimizer)
if hasattr(self, 'loss_scaler') and self.loss_scaler is not None:
return self.loss_scaler.scale_loss(value)

return value

Expand Down
5 changes: 5 additions & 0 deletions deepspeed/runtime/bf16_optimizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,11 @@ def __init__(self,
], f"BF16Optimizer: Unsupported gradient accumulation data type: {grad_acc_dtype}"
self.grad_acc_dtype = grad_acc_dtype

# BF16 doesn't use loss scaling, but these attributes are needed for API compatibility
self.custom_loss_scaler = False
self.external_loss_scale = None
self.torch_autocast_gradscaler = None

self.immediate_grad_update = bfloat16_config.immediate_grad_update

self.clip_grad = clip_grad
Expand Down
4 changes: 3 additions & 1 deletion deepspeed/runtime/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,7 +144,9 @@
BFLOAT16_OPTIMIZER_STATES_DEFAULT = False

# DDP variant of BFLOAT16
DDP_BFLOAT16 = "bf16"
# DDP variant: bf16 model with bf16 grad accumulation (uses FP16_Optimizer in bf16 mode)
# Must be different from BFLOAT16 to allow proper optimizer selection
DDP_BFLOAT16 = "ddp_bf16"

#########################################
# FP16 support
Expand Down
5 changes: 5 additions & 0 deletions tests/unit/runtime/zero/test_zero_tensor_fragment.py
Original file line number Diff line number Diff line change
Expand Up @@ -173,6 +173,11 @@ def test_bf16_optimizer_fragments(self, frozen_weights):
"bf16": {
"enabled": True
},
# Use fp32 gradient accumulation to ensure BF16_Optimizer is used
# (bf16 model + bf16 grad_accum uses FP16_Optimizer which doesn't support tensor fragment APIs)
"data_types": {
"grad_accum_dtype": "fp32"
},
"zero_optimization": {
"stage": 0,
}
Expand Down