From 3210486986dbee1c37ed100266000403fcd112a1 Mon Sep 17 00:00:00 2001 From: sdjasj <1594576288@qq.com> Date: Wed, 29 Jul 2026 20:55:50 +0800 Subject: [PATCH] preserve initialized optimizer state in ZeRO --- colossalai/zero/low_level/low_level_optim.py | 36 +++++++++++++++++ .../test_zero/test_low_level/test_zero1_2.py | 40 +++++++++++++++++++ 2 files changed, 76 insertions(+) diff --git a/colossalai/zero/low_level/low_level_optim.py b/colossalai/zero/low_level/low_level_optim.py index c530ff009fb1..846747b17205 100644 --- a/colossalai/zero/low_level/low_level_optim.py +++ b/colossalai/zero/low_level/low_level_optim.py @@ -208,6 +208,13 @@ def __init__( master_param_current_rank = self._create_master_param_current_rank(group_params) self._master_param_groups_of_current_rank[group_id] = master_param_current_rank + # Some optimizers, such as Adagrad, eagerly initialize per-parameter + # state in their constructor. Move that state from the working + # parameters to the newly-created master shards before replacing the + # optimizer parameter group. + for working_param, master_param in zip(group_params, master_param_current_rank): + self._partition_initialized_state(working_param, master_param) + # need to replace the params in the `params` field in the optimizer # so that when the optimizer calls step(), it only updates the tensors # managed by this data parallel rank @@ -298,6 +305,35 @@ def _create_master_param_current_rank(self, param_list): return params_current_rank + def _partition_initialized_state(self, working_param: Tensor, master_param: Tensor) -> None: + """Move eagerly initialized optimizer state to the local master shard.""" + if working_param not in self.optim.state: + return + + working_state = self.optim.state.pop(working_param) + master_state = {} + bucket_store = self.pid_to_bucket_store[id(working_param)] + padding_size = self.get_param_padding_size(working_param) + + for key, value in working_state.items(): + if isinstance(value, Tensor) and key != "step" and value.numel() == working_param.numel(): + flat_value = value.detach().flatten() + if padding_size > 0: + flat_value = torch.nn.functional.pad(flat_value, [0, padding_size]) + state_shards = flat_value.split(flat_value.numel() // bucket_store.world_size) + state_shard = state_shards[bucket_store.local_rank].clone() + if torch.is_floating_point(state_shard): + state_shard = state_shard.to(dtype=master_param.dtype) + master_state[key] = state_shard.to(master_param.device) + elif isinstance(value, Tensor): + # Scalar state (for example Adagrad's step counter) has its own + # device requirements, so preserve the original device. + master_state[key] = value.detach().clone() + else: + master_state[key] = copy.deepcopy(value) + + self.optim.state[master_param] = master_state + ########################### # Backward Reduction Hook # ########################### diff --git a/tests/test_zero/test_low_level/test_zero1_2.py b/tests/test_zero/test_low_level/test_zero1_2.py index 103854f869c7..687fb7e637f9 100644 --- a/tests/test_zero/test_low_level/test_zero1_2.py +++ b/tests/test_zero/test_low_level/test_zero1_2.py @@ -211,6 +211,35 @@ def exam_zero_1_torch_ddp(dtype: torch.dtype, master_weights: bool, extra_dp_siz loose_close(p, z1p, dtype=dtype) +def exam_adagrad_initialized_state(): + """Adagrad initializes state before ZeRO replaces parameters with master shards.""" + rank = dist.get_rank() + world_size = dist.get_world_size() + seed_all(1024) + model = MlpModel().cuda().half() + optimizer = torch.optim.Adagrad(model.parameters(), lr=1e-3, initial_accumulator_value=0.25) + initialized_sums = {param: optimizer.state[param]["sum"].clone() for param in model.parameters()} + + optimizer = LowLevelZeroOptimizer( + optimizer, + initial_scale=8.332635365271916, + partition_grad=True, + ) + + for working_params, master_params in zip( + optimizer._working_param_groups.values(), optimizer._master_param_groups_of_current_rank.values() + ): + for working_param, master_param in zip(working_params, master_params): + expected_sum = split_ddp_grad(initialized_sums[working_param], world_size)[rank].float() + assert_close(optimizer.optim.state[master_param]["sum"], expected_sum) + assert optimizer.optim.state[master_param]["step"].device.type == "cpu" + assert working_param not in optimizer.optim.state + + output = model(torch.randn(8, 123, device="cuda", dtype=torch.float16)) + optimizer.backward(output.float().square().mean()) + optimizer.step() + + def run_dist(rank, world_size, port): colossalai.launch(rank=rank, world_size=world_size, port=port, host="localhost") @@ -218,11 +247,22 @@ def run_dist(rank, world_size, port): exam_zero_1_2() +def run_adagrad_dist(rank, world_size, port): + colossalai.launch(rank=rank, world_size=world_size, port=port, host="localhost") + exam_adagrad_initialized_state() + + @pytest.mark.dist @rerun_if_address_is_in_use() def test_zero_1_2(): spawn(run_dist, 4) +@pytest.mark.dist +@rerun_if_address_is_in_use() +def test_adagrad_initialized_state(): + spawn(run_adagrad_dist, 2) + + if __name__ == "__main__": test_zero_1_2()