DeepSpeed
ca7455ef - A zero-element parameter loses its shape when it is bound to a flat buffer (#8467)

Commit
15 days ago
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>
Author
Parents
Loading