已开启
fully_shard ms后端精度问题,1000 step fully_shard loss偏低 #180
MengXY107创建于 6月1日
6月1日 修改了issue 的描述
MengXY107
6月1日 评论:
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}")


MengXY107
6月8日 评论:
6月8日 评论:
需要使用loss.backward接口,才会把梯度挂在正确的地方。


该问题是怎么引起的?
重现步骤
报错信息