已开启
fully_shard ms后端精度问题,1000 step fully_shard loss偏低 #180
MengXY107创建于  6月1日
MengXY107
MengXY107成员
6月1日 创建

该问题是怎么引起的?

重现步骤

import os

os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore"
import numpy as np
import mindspore as ms
from mindspore import nn, Tensor, mint
from mindspore.common.api import _no_grad
from mindspore.dataset import GeneratorDataset
import mindspore.communication.management as D
from hyper_parallel import init_device_mesh, fully_shard
from hyper_parallel.core.activation_checkpoint import checkpoint_wrapper, swap_wrapper, CheckpointPolicy, SwapManager
from hyper_parallel.core.dtensor.init_weights import init_empty_weights
from hyper_parallel.core.fully_shard.utils import MixedPrecisionPolicy
from hyper_parallel.platform import get_platform
from mindspore import communication as dist

platform = get_platform()

class SimpleMLP(nn.Cell):
    """2-layer MLP for DP/FSDP/FSTP testing."""

    def __init__(self, in_size=32, hidden_size=64, out_size=24):
        super().__init__()
        self.layer0 = nn.Dense(in_size, hidden_size, weight_init="normal", bias_init="zeros")
        self.layer1 = nn.Dense(hidden_size, out_size, weight_init="normal", bias_init="zeros")

    def construct(self, x):
        x = self.layer0(x)
        x = mint.nn.ReLU()(x)
        x = self.layer1(x)
        return x


def test_parallel_checkpoint_wrapper_001():
    rank = D.get_rank()
    ms.set_context(mode=ms.PYNATIVE_MODE)
    ms.set_deterministic(True)
    ms.set_seed(42)
    np.random.seed(42)
    dist.init()
    in_size = 32
    hidden_size = 64
    out_size = 24
    batch_size = 1
    num_steps = 1000
    standalone_net = SimpleMLP(in_size, hidden_size, out_size)
    lazy_net = SimpleMLP(in_size, hidden_size, out_size)
    params_standalone = standalone_net.parameters_dict()
    params_lazy = lazy_net.parameters_dict()
    for key in params_standalone:
        params_standalone[key].set_data(params_lazy[key])
    # def recompute_policy_fn(ctx, op, *args, **kwargs):
    #     return CheckpointPolicy.MUST_RECOMPUTE
    #
    # lazy_net = checkpoint_wrapper(lazy_net, policy_fn=recompute_policy_fn)
    mesh = init_device_mesh(device_type="npu", mesh_shape=(4,), mesh_dim_names=("dp",))
    mp_policy = MixedPrecisionPolicy(cast_forward_inputs=True)
    fully_shard(lazy_net, mesh=mesh, reshard_after_forward=True, mp_policy=mp_policy)
    np.random.seed(42)
    train_data = []
    for i in range(1000):
        input_ids = np.random.RandomState(44 + i).randn(batch_size, in_size).astype(np.float32)
        labels = np.random.RandomState(44 + i).randn(batch_size, out_size).astype(np.float32)
        train_data.append((Tensor(input_ids), Tensor(labels)))
    standalone_dataset = GeneratorDataset(train_data[:], ['data', 'label'], shuffle=False)
    distributed_dataset = GeneratorDataset(train_data[:], ['data', 'label'], shuffle=False)
    loss_fn = nn.MSELoss()

    def train_step(model, optimizer, inputs, labels):
        def forward_fn(inputs, labels):
            predictions = model(inputs)
            loss = loss_fn(predictions, labels)
            return loss

        grad_fn = ms.value_and_grad(forward_fn, None, optimizer.parameters)
        loss, grads = grad_fn(inputs, labels)
        optimizer(grads)
        return loss

    opt_standalone = nn.Adam(standalone_net.trainable_params(), learning_rate=1e-4)
    opt_distributed = nn.Adam(lazy_net.trainable_params(), learning_rate=1e-4)

    # Train standalone model
    print(f"Rank {rank}: Training standalone model for {num_steps} steps...")
    loss_list_standalone = []
    step = 0
    for data, label in standalone_dataset.create_tuple_iterator():
        loss = train_step(standalone_net, opt_standalone, data, label)
        print(f"step: {step}, standalone loss: {loss}")
        loss_list_standalone.append(float(loss.asnumpy()))
        step += 1
        if step >= num_steps:
            break

    # Train lazy model (FSDP)
    print(f"Rank {rank}: Training FSDP model for {num_steps} steps...")
    loss_list_lazy = []
    step = 0
    for data, label in distributed_dataset.create_tuple_iterator():
        loss = train_step(lazy_net, opt_distributed, data, label)
        print(f"step: {step}, fully_shard loss: {loss}")
        loss_list_lazy.append(float(loss.asnumpy()))
        step += 1
        if step >= num_steps:
            break

    # print(f"Rank {rank}: Standalone losses:  {loss_list_standalone}")
    # print(f"Rank {rank}: FSDP losses:        {loss_list_lazy}")
    # # comparebase.compare_nparray(np.array(loss_list_standalone), np.array(loss_list_lazy), rtol=1e-4, atol=1e-4)

报错信息

likedislike
MengXY107MengXY107成员
6月1日 修改了issue 的描述
MengXY107
MengXY107成员
6月1日 评论:
import os

os.environ["HYPER_PARALLEL_PLATFORM"] = "mindspore"
import numpy as np
import mindspore as ms
from mindspore import nn, Tensor, mint
from mindspore.common.api import _no_grad
from mindspore.dataset import GeneratorDataset
import mindspore.communication.management as D
from hyper_parallel import init_device_mesh, fully_shard, SkipDTensorDispatch
from hyper_parallel.core.activation_checkpoint import checkpoint_wrapper, swap_wrapper, CheckpointPolicy, SwapManager
from hyper_parallel.core.dtensor.init_weights import init_empty_weights
from hyper_parallel.core.fully_shard.utils import MixedPrecisionPolicy
from hyper_parallel.platform.mindspore.autograd_compat import enable_mindspore_backward_compat
from hyper_parallel.platform import get_platform
from mindspore import communication as dist

platform = get_platform()
# fully_shard delivers the reduce-scattered gradient into each sharded
# parameter's ``.grad`` via backward hooks. That path is only wired up after
# enabling the MindSpore backward-compat shim, and the optimizer must read
# ``param.grad`` (NOT ``ms.value_and_grad`` w.r.t. the sharded leaf params,
# which is disconnected from the all-gathered forward and yields zero grads).
enable_mindspore_backward_compat()

class SimpleMLP(nn.Cell):
    """2-layer MLP for DP/FSDP/FSTP testing."""

    def __init__(self, in_size=32, hidden_size=64, out_size=24):
        super().__init__()
        self.layer0 = nn.Dense(in_size, hidden_size, weight_init="normal", bias_init="zeros")
        self.layer1 = nn.Dense(hidden_size, out_size, weight_init="normal", bias_init="zeros")

    def construct(self, x):
        x = self.layer0(x)
        x = mint.nn.ReLU()(x)
        x = self.layer1(x)
        return x


def test_parallel_checkpoint_wrapper_001():
    rank = D.get_rank()
    ms.set_context(mode=ms.PYNATIVE_MODE)
    ms.set_deterministic(True)
    ms.set_seed(42)
    np.random.seed(42)
    dist.init()
    in_size = 32
    hidden_size = 64
    out_size = 24
    batch_size = 1
    num_steps = 1000
    standalone_net = SimpleMLP(in_size, hidden_size, out_size)
    lazy_net = SimpleMLP(in_size, hidden_size, out_size)
    params_standalone = standalone_net.parameters_dict()
    params_lazy = lazy_net.parameters_dict()
    for key in params_standalone:
        params_standalone[key].set_data(params_lazy[key])
    # def recompute_policy_fn(ctx, op, *args, **kwargs):
    #     return CheckpointPolicy.MUST_RECOMPUTE
    #
    # lazy_net = checkpoint_wrapper(lazy_net, policy_fn=recompute_policy_fn)
    mesh = init_device_mesh(device_type="npu", mesh_shape=(4,), mesh_dim_names=("dp",))
    mp_policy = MixedPrecisionPolicy(cast_forward_inputs=True)
    fully_shard(lazy_net, mesh=mesh, reshard_after_forward=True, mp_policy=mp_policy)
    np.random.seed(42)
    train_data = []
    for i in range(1000):
        input_ids = np.random.RandomState(44 + i).randn(batch_size, in_size).astype(np.float32)
        labels = np.random.RandomState(44 + i).randn(batch_size, out_size).astype(np.float32)
        train_data.append((Tensor(input_ids), Tensor(labels)))
    standalone_dataset = GeneratorDataset(train_data[:], ['data', 'label'], shuffle=False)
    distributed_dataset = GeneratorDataset(train_data[:], ['data', 'label'], shuffle=False)
    loss_fn = nn.MSELoss()

    def train_step(model, optimizer, inputs, labels, is_distributed):
        # Reset gradients before each backward. FSDP-wrapped cells expose
        # ``zero_grad``; a plain cell only has ``param.grad`` to clear.
        if is_distributed:
            model.zero_grad()
        else:
            for p in model.trainable_params():
                p.grad = None
        # Drive backward through the autograd graph so fully_shard's reduce-scatter
        # hooks fire and populate each sharded ``param.grad``.
        loss = loss_fn(model(inputs), labels)
        loss.backward()
        grads = tuple(p.grad for p in model.trainable_params())
        with SkipDTensorDispatch(), _no_grad():
            optimizer(grads)
        return loss

    opt_standalone = nn.Adam(standalone_net.trainable_params(), learning_rate=1e-4)
    opt_distributed = nn.Adam(lazy_net.trainable_params(), learning_rate=1e-4)

    # Train standalone model
    print(f"Rank {rank}: Training standalone model for {num_steps} steps...")
    loss_list_standalone = []
    step = 0
    for data, label in standalone_dataset.create_tuple_iterator():
        loss = train_step(standalone_net, opt_standalone, data, label, is_distributed=False)
        print(f"step: {step}, standalone loss: {loss}")
        loss_list_standalone.append(float(loss.asnumpy()))
        step += 1
        if step >= num_steps:
            break

    # Train lazy model (FSDP)
    print(f"Rank {rank}: Training FSDP model for {num_steps} steps...")
    loss_list_lazy = []
    step = 0
    for data, label in distributed_dataset.create_tuple_iterator():
        loss = train_step(lazy_net, opt_distributed, data, label, is_distributed=True)
        print(f"step: {step}, fully_shard loss: {loss}")
        loss_list_lazy.append(float(loss.asnumpy()))
        step += 1
        if step >= num_steps:
            break

    # Precision check: every rank trains on the same data A and the gradients are
    # reduce-scatter *mean*'d, so each FSDP rank's update is equivalent to the
    # standalone update on A. The per-step losses must therefore match the
    # standalone run within the tolerance below.
    rtol, atol = 1e-4, 1e-4
    std = np.array(loss_list_standalone)
    fsdp = np.array(loss_list_lazy)
    abs_diff = np.abs(std - fsdp)
    max_abs = float(abs_diff.max())
    worst = int(abs_diff.argmax())
    print(f"Rank {rank}: max |loss_std - loss_fsdp| = {max_abs:.3e} at step {worst} "
          f"(std={std[worst]:.6f}, fsdp={fsdp[worst]:.6f})")
    mism = np.where(~np.isclose(fsdp, std, rtol=rtol, atol=atol))[0]
    assert mism.size == 0, (
        f"Rank {rank}: FSDP vs standalone loss mismatch (rtol={rtol}, atol={atol}) at "
        f"{mism.size} step(s); first @ step {int(mism[0])}: "
        f"std={std[mism[0]]:.6f}, fsdp={fsdp[mism[0]]:.6f}, "
        f"|diff|={abs_diff[mism[0]]:.3e}"
    )
    print(f"Rank {rank}: PASS - FSDP losses match standalone within rtol={rtol}, atol={atol}")


likedislike
MengXY107
MengXY107成员
6月8日 评论:

需要使用loss.backward接口,才会把梯度挂在正确的地方。

likedislike