已关闭
fully shard list[modules]场景多次进入hook,权重状态错误 #73
DavidFFFan创建于 4月3日关闭于 4月8日
4月3日 修改了issue 的描述
4月3日 关联问题类型 由 [] 改变为 [功能]
4月3日 issue类型由 Bug 改变为 Bug-Report
4月3日 issue类型由 Bug-Report 改变为 Bug
DavidFFFan
4月3日 评论:
4月3日 评论:
问题原因为list中的module都会注册hook,但是共享同一个state,导致正向时会多次进入hook,产生多余的通信。反向时状态紊乱,报错。


4月3日 修改了issue 的描述
4月3日 修改了issue 的描述
4月3日 issue类型由 Bug 改变为 Bug-Report
4月7日 关联了pull request:fix(fully_shard): grouped forward hooks for list root modules
4月8日 将 changzherui1 设为负责人
4月8日 issue状态由 TODO 改变为 ACCEPTED
4月8日 issue状态由 ACCEPTED 改变为 WIP
4月8日 issue状态由 WIP 改变为 DONE
4月8日 关闭了 issue
该问题是怎么引起的?
fully shard 配置list [module] 配置 reshard_after_forward 为False,在反向时报错。
重现步骤
新建
hyper-parallel/tests/mindspore/st/fully_shard/_test_fully_shard_list_backward.py# Copyright 2026 Huawei Technologies Co., Ltd # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. # ============================================================================ """ST for fully_shard list input on chained submodules (MindSpore).""" import mindspore as ms from mindspore import Tensor, mint, nn from mindspore.communication import get_rank, init from hyper_parallel import DTensor, SkipDTensorDispatch, init_device_mesh from hyper_parallel.core.fully_shard.api import fully_shard from hyper_parallel.core.fully_shard.utils import MixedPrecisionPolicy ms.set_seed(42) def get_backward_grads(net): """Collect DTensor gradients from fully_shard-managed params.""" grads = [] for idx, param in enumerate(net.trainable_params()): grad = param.grad assert isinstance(grad, DTensor), f"Parameter grad {idx} is not a DTensor" assert grad.shape == param.shape, ( f"Gradient global shape mismatch at index {idx}: " f"Expected {param.shape}, got {grad.shape}" ) assert grad.local_shape == param.local_shape, ( f"Gradient local shape mismatch at index {idx}: " f"Expected {param.local_shape}, got {grad.local_shape}" ) grads.append(grad) return tuple(grads) class TinyDensePairBlock(nn.Cell): """Minimal chained dense modules to reproduce fully_shard(list[submodules]) backward flow.""" def __init__(self, hidden_size: int, output_size: int): super().__init__() self.dense1 = nn.Dense(hidden_size, hidden_size, has_bias=False) self.dense2 = nn.Dense(hidden_size, output_size, has_bias=False) def construct(self, x): x = self.dense1(x) return self.dense2(x) class TinyWrappedDenseNet(nn.Cell): """A tiny top-level network that wraps the fully_shard list submodules.""" def __init__(self, hidden_size: int, output_size: int): super().__init__() self.pre = nn.Dense(hidden_size, hidden_size, has_bias=False) self.block = TinyDensePairBlock(hidden_size, output_size) def construct(self, x): x = self.pre(x) x = self.block(x) return x def test_list_modules_backward_fully_shard(): """ Feature: fully_shard list input on chained submodules. Description: Wrap two serial submodules as one fully_shard unit via fully_shard([dense1, dense2]) and run one backward step. Expectation: Backward completes and gradients are materialized as DTensor. """ ms.set_context(mode=ms.PYNATIVE_MODE) init() rank_id = get_rank() hidden_size = 16 output_size = 8 batch_size = 8 mesh = init_device_mesh(device_type="npu", mesh_shape=(2,), mesh_dim_names=("dp",)) mp_policy = MixedPrecisionPolicy() model = TinyWrappedDenseNet(hidden_size, output_size) fully_shard(model.pre, mesh=mesh, mp_policy=mp_policy) list_modules = [model.block.dense1, model.block.dense2] fully_shard(list_modules, mesh=mesh, mp_policy=mp_policy, reshard_after_forward=False) fully_shard(model, mesh=mesh, mp_policy=mp_policy) assert model.block.dense1.hsdp_scheduler is model.block.dense2.hsdp_scheduler optimizer = nn.Adam(model.trainable_params(), learning_rate=0.01) input_data = Tensor(ms.numpy.randn(batch_size, hidden_size).astype(ms.float32)) model.zero_grad() output = model(input_data) loss = mint.sum(output) loss.backward() grads = get_backward_grads(model) with SkipDTensorDispatch(): optimizer(grads) if rank_id == 0: print(f"rank: {rank_id}, loss: {float(loss.asnumpy())}")新建
hyper-parallel/tests/mindspore/st/fully_shard/test_fully_shard_list_backward.py# Copyright 2026 Huawei Technologies Co., Ltd # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. # ============================================================================ """Launch ST for fully_shard list backward scenario (MindSpore).""" from tests.common.mark_utils import arg_mark from tests.mindspore.st.utils import msrun_case _FILE_NAME = "_test_fully_shard_list_backward.py" @arg_mark(plat_marks=["platform_ascend910b"], level_mark="level0", card_mark="allcards", essential_mark="essential") def test_ms_list_modules_backward_fully_shard(): """ Feature: fully_shard list input on chained submodules. Description: Use fully_shard([dense1, dense2]) on a minimal Dense->Dense block and run one distributed backward step. Expectation: Run success. """ msrun_case( 2, _FILE_NAME, "test_list_modules_backward_fully_shard", 18511, worker_num=2, local_worker_num=2 )执行
pytest -sv test_fully_shard_list_backward.py报错信息