Skip to content

Commit 080eb6b

Browse files
fwerkortohtana
andauthored
Fix ZeRO parameter alignment for grouped_mm (#8277)
## Summary ZeRO stage 1/2 stores model parameters as views into flattened fp16/bf16 buffers. Individual parameter views can start at non-16-byte offsets even when the flat buffer itself is aligned, which breaks alignment-sensitive kernels such as `torch._grouped_mm`. Pad parameter boundaries inside the ZeRO flat buffer so model parameter views remain 16-byte aligned without duplicating misaligned parameters. The padded layout is propagated through partition/gradient offsets, LP↔HP linkage, DeepCompile gradient buffers, checkpoint restore, and `zero_to_fp32` reconstruction. Older checkpoints without parameter-alignment padding remain loadable; their compact layout is converted when restored. ## Tests - ZeRO-1/2 BF16 regression with a deliberately misaligned parameter layout. - Verifies zero-copy aligned flat-buffer views across optimizer steps. - Verifies checkpoint reload with `load_module_only=True` and `load_optimizer_states=False`. - Synced with current `master`, including the ZeRO-1/2 DeepCompile file rename. Fixes #8276 --------- Signed-off-by: Cao Yuhang <caoyuhang@fwerkor.com> Signed-off-by: Masahiro Tanaka <tanaka.masahiro@gmail.com> Co-authored-by: Masahiro Tanaka <tanaka.masahiro@gmail.com> Co-authored-by: Masahiro Tanaka <81312776+tohtana@users.noreply.github.com>
1 parent e5887a3 commit 080eb6b

10 files changed

Lines changed: 833 additions & 112 deletions

File tree

deepspeed/checkpoint/constants.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,7 @@
1818
SINGLE_PARTITION_OF_FP32_GROUPS = "single_partition_of_fp32_groups"
1919
PARAM_GROUPS = 'param_groups'
2020
GROUP_PADDINGS = 'group_paddings'
21+
PARAM_ALIGNMENT_PADDINGS = 'param_alignment_paddings'
2122
PARTITION_COUNT = 'partition_count'
2223
ZERO_STAGE = 'zero_stage'
2324
CLIP_GRAD = 'clip_grad'

deepspeed/compile/init_z1_and_2.py

Lines changed: 41 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88

99
import torch
1010

11+
import deepspeed.comm as dist
1112
from deepspeed.accelerator import get_accelerator
1213
from .passes import zero_1_and_2_compile, zero3_compile
1314
from .backend import make_backend, launch_compile_passes, init_schedule
@@ -51,39 +52,47 @@ def _build_partition_grad_views(optimizer, group_idx):
5152
optimizer.all_grad_tensors[group_idx] = original_all_grad_tensors
5253

5354

55+
def _clone_partition_grad_views(partition_grad_views):
56+
cloned_views = []
57+
for view in partition_grad_views:
58+
cloned = view.clone().detach()
59+
if getattr(view, "_zero_padding", False):
60+
cloned._zero_padding = True
61+
cloned_views.append(cloned)
62+
return cloned_views
63+
64+
5465
def _build_flat_partition_grad_views(optimizer, group_idx):
5566
partition_size = int(optimizer.partition_size[group_idx])
5667
dtype = optimizer.gradient_accumulation_dtype
5768
device = get_accelerator().current_device_name()
5869
flat_buffer = torch.zeros(partition_size, dtype=dtype, device=device)
70+
partition_id = dist.get_rank(group=optimizer.real_dp_process_group[group_idx])
5971

6072
views = []
6173
current_size = 0
62-
for i, tensor in enumerate(optimizer.params_in_partition[group_idx]):
63-
num_elements = tensor.numel()
64-
tensor_offset = 0
65-
66-
if i == 0 and optimizer.first_offset[group_idx] > 0:
67-
tensor_offset = int(optimizer.first_offset[group_idx])
68-
num_elements -= tensor_offset
69-
70-
if num_elements > partition_size - current_size:
71-
num_elements = partition_size - current_size
72-
73-
if num_elements <= 0:
74-
continue
75-
76-
view = flat_buffer.narrow(0, current_size, int(num_elements))
77-
if tensor_offset == 0 and num_elements == tensor.numel():
78-
view = view.view(tensor.shape)
79-
views.append(view)
80-
current_size += int(num_elements)
81-
82-
if current_size >= partition_size:
83-
break
74+
for tensor in optimizer.params_in_partition[group_idx]:
75+
param_id = optimizer.get_param_id(tensor)
76+
source_offset = optimizer.grad_start_offset[group_idx][partition_id][param_id]
77+
dest_offset = optimizer.grad_partition_insertion_offset[group_idx][partition_id][param_id]
78+
num_elements = min(tensor.numel() - source_offset, partition_size - dest_offset)
79+
80+
if dest_offset > current_size:
81+
padding = flat_buffer.narrow(0, current_size, dest_offset - current_size)
82+
padding._zero_padding = True
83+
views.append(padding)
84+
85+
if num_elements > 0:
86+
view = flat_buffer.narrow(0, dest_offset, int(num_elements))
87+
if source_offset == 0 and num_elements == tensor.numel():
88+
view = view.view(tensor.shape)
89+
views.append(view)
90+
current_size = dest_offset + int(num_elements)
8491

8592
if current_size < partition_size:
86-
views.append(flat_buffer.narrow(0, current_size, partition_size - current_size))
93+
padding = flat_buffer.narrow(0, current_size, partition_size - current_size)
94+
padding._zero_padding = True
95+
views.append(padding)
8796

8897
return flat_buffer, views
8998

@@ -101,7 +110,12 @@ def init_z1_and_2(engine, backend, compile_config, compile_kwargs, schedule=None
101110
if use_z2:
102111
grad_buffer = {}
103112
for i, group in enumerate(optimizer.bit16_groups):
104-
grad_buffer[i] = [p.clone().detach() for p in _build_partition_grad_views(optimizer, i)]
113+
partition_grad_views = _build_partition_grad_views(optimizer, i)
114+
grad_buffer[i] = _clone_partition_grad_views(partition_grad_views)
115+
param_grad_buffers = [
116+
cloned for original, cloned in zip(partition_grad_views, grad_buffer[i])
117+
if not getattr(original, "_zero_padding", False)
118+
]
105119

106120
index_in_partition = 0
107121
first_in_partition = True
@@ -111,7 +125,7 @@ def init_z1_and_2(engine, backend, compile_config, compile_kwargs, schedule=None
111125
in_partition = optimizer.is_param_in_current_partition[param_id]
112126

113127
if in_partition:
114-
buf = grad_buffer[i][index_in_partition]
128+
buf = param_grad_buffers[index_in_partition]
115129
offset = optimizer.first_offset[i] if first_in_partition else 0
116130
dc.register_param(p.param_id, p.shape, p, buf, int(offset))
117131
index_in_partition += 1
@@ -157,7 +171,8 @@ def set_z1_grad_buffer(is_gradient_accumulation_boundary):
157171
flat_grad_buffer, group_grad_buffers = _build_flat_partition_grad_views(optimizer, group_idx)
158172
current_grad_buffers[group_idx] = _FlatPartitionGradBufferGroup(
159173
group_grad_buffers, flat_grad_buffer, lambda group_idx=group_idx: release_grad_buffer(group_idx))
160-
for (param_id, _, offset), grad_buffer in zip(grad_buffer_metadata[group_idx], group_grad_buffers):
174+
param_grad_buffers = [g for g in group_grad_buffers if not getattr(g, "_zero_padding", False)]
175+
for (param_id, _, offset), grad_buffer in zip(grad_buffer_metadata[group_idx], param_grad_buffers):
161176
dc.update_param_grad_buffer(param_id, grad_buffer, offset)
162177
optimizer.averaged_gradients = current_grad_buffers
163178

@@ -179,7 +194,6 @@ def release_grad_buffer(group_idx=None):
179194
schedule.append((0, [zero_1_and_2_compile.add_z1_reduce]))
180195
else:
181196
for opt in schedule:
182-
# avoid typical misconfiguration
183197
if zero3_compile.add_z3_gather_release in opt[1]:
184198
raise ValueError("The schedule contains the ZeRO-3 pass add_z3_gather_release, but ZeRO stage 1 "
185199
"or 2 is enabled. Use zero_1_and_2_compile.add_z1_reduce (stage 1) or add_z2_reduce "

deepspeed/runtime/engine.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1723,6 +1723,9 @@ def zero_allgather_partitions(self):
17231723
def zero_round_robin_gradients(self):
17241724
return self._config.zero_config.round_robin_gradients
17251725

1726+
def zero_parameter_alignment(self):
1727+
return self._config.zero_config.parameter_alignment
1728+
17261729
def zero_hpz_partition_size(self):
17271730
return self._config.zero_config.zero_hpz_partition_size
17281731

@@ -2741,6 +2744,7 @@ def _configure_zero_optimizer(self, optimizer):
27412744
ignore_unused_parameters=self.zero_ignore_unused_parameters(),
27422745
partition_grads=zero_stage == ZeroStageEnum.gradients,
27432746
round_robin_gradients=round_robin_gradients,
2747+
parameter_alignment=self.zero_parameter_alignment(),
27442748
has_moe_layers=self.has_moe_layers,
27452749
fp16_master_weights_and_gradients=self.fp16_master_weights_and_gradients(),
27462750
bf16_master_weights_and_gradients=self.bf16_master_weights_and_gradients(),

deepspeed/runtime/zero/config.py

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,7 @@
4141
"offload_optimizer": {...},
4242
"ignore_unused_parameters": [true|false],
4343
"round_robin_gradients": [true|false],
44+
"parameter_alignment": [true|false],
4445
"zero_hpz_partition_size": 1,
4546
"zero_quantized_weights": [true|false],
4647
"zero_quantized_nontrainable_weights": [true|false],
@@ -313,6 +314,14 @@ class DeepSpeedZeroConfig(DeepSpeedConfigModel):
313314
Performance benefit grows with gradient accumulation steps (more copying
314315
between optimizer steps) or GPU count (increased parallelism).
315316
"""
317+
318+
parameter_alignment: bool = False
319+
"""
320+
Pad ZeRO Stage 1 and 2 flat buffers between parameters so each parameter
321+
starts at a 16-byte-aligned address. This is disabled by default because
322+
the padding increases flat-buffer and optimizer-state memory usage.
323+
"""
324+
316325
zero_hpz_partition_size: int = Field(1, ge=0)
317326
"""
318327
Number of ranks in zero parameters partitioning secondary group

0 commit comments

Comments
 (0)