已关闭
fully shard list[modules]场景多次进入hook,权重状态错误 #73
DavidFFFan创建于  4月3日关闭于  4月8日
DavidFFFan
DavidFFFan成员
4月3日 创建

该问题是怎么引起的?

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

报错信息

  File "/hyper-parallel/hyper_parallel/platform/mindspore/fully_shard/param.py", line 589, in reduce_scatter_grad
    self._assert_in_states(ShardedState.UNSHARDED)
  File "/hyper-parallel/hyper_parallel/platform/mindspore/fully_shard/param.py", line 437, in _assert_in_states
    raise AssertionError(
AssertionError: Expected sharded_state in (<ShardedState.UNSHARDED: 2>,), got ShardedState.SHARDED
likedislike
DavidFFFanDavidFFFan成员
4月3日 修改了issue 的描述
DavidFFFanDavidFFFan成员
4月3日 关联问题类型 由 [] 改变为 [功能]
DavidFFFanDavidFFFan成员
4月3日 issue类型由 Bug 改变为 Bug-Report
DavidFFFanDavidFFFan成员
4月3日 issue类型由 Bug-Report 改变为 Bug
DavidFFFan
DavidFFFan成员
4月3日 评论:

问题原因为list中的module都会注册hook,但是共享同一个state,导致正向时会多次进入hook,产生多余的通信。反向时状态紊乱,报错。

likedislike
DavidFFFanDavidFFFan成员
4月3日 修改了issue 的描述
DavidFFFanDavidFFFan成员
4月3日 修改了issue 的描述
DavidFFFanDavidFFFan成员
4月3日 issue类型由 Bug 改变为 Bug-Report
changzheruichangzherui成员
4月7日 关联了pull request:fix(fully_shard): grouped forward hooks for list root modules
changzheruichangzherui成员
4月8日 将 changzherui1 设为负责人
changzheruichangzherui成员
4月8日 issue状态由 TODO 改变为 ACCEPTED
changzheruichangzherui成员
4月8日 issue状态由 ACCEPTED 改变为 WIP
changzheruichangzherui成员
4月8日 issue状态由 WIP 改变为 DONE
changzheruichangzherui成员
4月8日 关闭了 issue