A zero-element parameter loses its shape when it is bound to a flat buffer (#8467)
A zero-element trainable parameter lost its shape when a wrapper bound
it to a flat buffer. After that, the module's own forward failed.
`nn.Linear(8, 0, bias=False)` in a model, `deepspeed.initialize`, one
step:
```
RuntimeError: size mismatch, got input (1), mat (1x8), vec (0)
```
torch's `unflatten_dense_tensors` special-cases `numel == 0` and returns
a fresh 1-D `zeros({0})` rather than a view of the requested shape, so
`p.data = q.data` turned `(0, 8)` into `(0,)`. This happened with
`FP16_Optimizer` (fp16 or bf16 at stage 0) and with `BF16_Optimizer`
(bf16 at stage 1 with fp32 grad accumulation).
## The change
As discussed with @sfc-gh-truwase, zero-element parameters are now left
out of the optimizer groups, the same way frozen parameters are. They
have nothing to optimize.
- `is_optimized_parameter(param)` in `runtime/utils.py`: `requires_grad
and numel > 0`, using `ds_numel` under ZeRO-3.
- Every wrapper uses it to build its groups: `FP16_Optimizer`,
`FP16_UnfusedOptimizer`, `BF16_Optimizer`, ZeRO-1/2 (`stage_1_and_2.py`)
and ZeRO-3 (`_get_trainable_parameter_groups`).
- The checkpoint side uses the same predicate:
`_get_zero_frozen_param_attributes` and the frozen-fragment load in
`load_module_state_dict`. `param_shapes` is built from the optimizer
groups, so with the old `requires_grad` rule a filtered zero-element
parameter would be in neither `param_shapes` nor `frozen_param_shapes`,
and `zero_to_fp32` would drop it. It is now recorded with the frozen
parameters and rebuilt with its shape.
Grad hooks and the other bookkeeping in ZeRO-1/2 and ZeRO-3 iterate the
wrappers' own groups, so they follow the filter without further changes.
## Testing
On 2×H20, `tests/unit/runtime/zero/test_zero_numel_param_shape.py` (17
passed):
- `TestZeroNumelParameterShape`: the shape survives init, one step and a
second step for fp16 stage 0, bf16 stage 0, bf16 stage 1 with fp32
accumulation, ZeRO-1 and ZeRO-2.
- `TestZeroNumelParameterCheckpoint`: ZeRO-2 and ZeRO-3, save →
`get_fp32_state_dict_from_zero_checkpoint` → strict `load_state_dict`,
with the zero-element parameter at `(0, 8)`. With the optimizer filter
but the old frozen rule in `engine.py`, both fail with `KeyError:
'extra'`.
`tests/unit/checkpoint/test_zero_optimizer.py -k frozen`: 13 passed.
---------
Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
Co-authored-by: Olatunji Ruwase <tunji.ruwase@snowflake.com>