已合并
fix incorrect dynamic scales calculation when tp_world_size is 0 #4740
zhong-zixin创建于 4月16日
fix incorrect dynamic scales calculation when tp_world_size is 0 #4740
已合并
zhong-zixin创建于 4月16日
3 个文件变更+78-8
@@ -149,13 +149,14 @@ tensor_list npu_moe_distribute_dispatch_v2(const at::Tensor &x, const at::Tensor
149 at::Tensor dynamic_scales{nullptr};149 at::Tensor dynamic_scales{nullptr};
150 aclDataType acl_dynamic_scale_dtype = op_plugin::utils::get_dynamic_scales_dtype(x, scales, scales_dtype, quant_mode);150 aclDataType acl_dynamic_scale_dtype = op_plugin::utils::get_dynamic_scales_dtype(x, scales, scales_dtype, quant_mode);
151 auto scalar_dynamic_scale_dtype = npu_preparation::convert_to_scalar_type(acl_dynamic_scale_dtype);151 auto scalar_dynamic_scale_dtype = npu_preparation::convert_to_scalar_type(acl_dynamic_scale_dtype);
152- if (tp_world_size == 0) {152+ if (c10_npu::IsAclnnOnly()) {
153- dynamic_scales = npu_preparation::apply_tensor_without_format({a},153+ auto dynamic_scales_shape = op_plugin::utils::get_dynamic_shape(scales, quant_mode,
S
Ssongkai1114月20日

此处的逻辑增加了std::max的取值,与之前是否功能一致?

likedislike
zhong-zixin
4月20日 评论:
154+ std::max(a, a * tp_world_size), h);
155+ dynamic_scales = npu_preparation::apply_tensor_without_format(dynamic_scales_shape,
154 x.options().dtype(scalar_dynamic_scale_dtype));156 x.options().dtype(scalar_dynamic_scale_dtype));
155 } else {157 } else {
156- if (c10_npu::IsAclnnOnly()) {158+ if (tp_world_size == 0) {
157- auto dynamic_scales_shape = op_plugin::utils::get_dynamic_shape(scales, quant_mode, a, h);159+ dynamic_scales = npu_preparation::apply_tensor_without_format({a},
158- dynamic_scales = npu_preparation::apply_tensor_without_format(dynamic_scales_shape,
159 x.options().dtype(scalar_dynamic_scale_dtype));160 x.options().dtype(scalar_dynamic_scale_dtype));
160 } else {161 } else {
161 dynamic_scales = npu_preparation::apply_tensor_without_format({a * tp_world_size},162 dynamic_scales = npu_preparation::apply_tensor_without_format({a * tp_world_size},
@@ -3161,9 +3161,7 @@ def npu_moe_distribute_dispatch_v2_meta(x, expert_ids, group_ep, ep_world_size,
3161 else:3161 else:
3162 expand_x = x.new_empty(tuple([max(a, a * tp_world_size), h]), dtype=outDtype)3162 expand_x = x.new_empty(tuple([max(a, a * tp_world_size), h]), dtype=outDtype)
3163 dynamic_scales_dtype = get_dispatch_dynamic_scales_dtype(x, scales, quant_mode)3163 dynamic_scales_dtype = get_dispatch_dynamic_scales_dtype(x, scales, quant_mode)
3164- if tp_world_size == 0:3164+ if tp_world_size <= 1:
3165- dynamic_scales = x.new_empty((a), dtype=dynamic_scales_dtype)
3166- elif tp_world_size == 1:
3167 dynamic_scales_shape = get_dispatch_dynamic_shape(scales, quant_mode, a, h)3165 dynamic_scales_shape = get_dispatch_dynamic_shape(scales, quant_mode, a, h)
3168 dynamic_scales = x.new_empty(dynamic_scales_shape, dtype=dynamic_scales_dtype)3166 dynamic_scales = x.new_empty(dynamic_scales_shape, dtype=dynamic_scales_dtype)
3169 else:3167 else:
@@ -3551,6 +3551,77 @@ class TestMoeDistributeDispatch(TestCase):
3551 self.assertEqual(result[6].dtype, torch.float32)3551 self.assertEqual(result[6].dtype, torch.float32)
3552 3552 
3553 3553 
3554+class TestMoeDistributeDispatchV2(TestCase):
3555+ def _run_dispatch_v2(self, tp_world_size, quant_mode, scales=None, y_dtype=None):
3556+ with FakeTensorMode():
3557+ ep_world_size = 16
3558+ bs = 8
3559+ h = 7168
3560+ k = 8
3561+ moeExpertNum = 16
3562+ global_bs = bs * ep_world_size
3563+ 
3564+ local_moe_expert_num = moeExpertNum // ep_world_size
3565+ a = global_bs * min(local_moe_expert_num, k)
3566+ 
3567+ x = torch.randn(bs, h).to(torch.bfloat16)
3568+ expert_ids = torch.randn(bs, k).to(torch.int32)
3569+ 
3570+ result = torch_npu.npu_moe_distribute_dispatch_v2(
3571+ x, expert_ids, "group_ep", ep_world_size, 0, moeExpertNum,
3572+ scales=scales, x_active_mask=None, expert_scales=None,
3573+ group_tp="", tp_world_size=tp_world_size, tp_rank_id=0,
3574+ expert_shard_type=0, shared_expert_num=0, shared_expert_rank_num=0,
3575+ quant_mode=quant_mode, global_bs=global_bs, expert_token_nums_type=1,
3576+ y_dtype=y_dtype)
3577+ 
3578+ return result, a, h, local_moe_expert_num
3579+ 
3580+ def test_tp0_pertoken(self):
3581+ result, a, h, _ = self._run_dispatch_v2(tp_world_size=0, quant_mode=2)
3582+ self.assertEqual(result[0].shape, torch.Size([a, h]))
3583+ self.assertEqual(result[0].dtype, torch.int8)
3584+ self.assertEqual(result[1].shape, torch.Size([a]))
3585+ 
3586+ def test_tp1_pertoken(self):
3587+ result, a, h, _ = self._run_dispatch_v2(tp_world_size=1, quant_mode=2)
3588+ self.assertEqual(result[0].shape, torch.Size([a, h]))
3589+ self.assertEqual(result[0].dtype, torch.int8)
3590+ self.assertEqual(result[1].shape, torch.Size([a]))
3591+ 
3592+ def test_tp0_pergroup(self):
3593+ result, a, h, _ = self._run_dispatch_v2(
3594+ tp_world_size=0, quant_mode=3, y_dtype=torch.float8_e5m2)
3595+ expected_dim1 = math.ceil(h / 128)
3596+ self.assertEqual(result[1].shape, torch.Size([a, expected_dim1]))
3597+ 
3598+ def test_tp1_pergroup(self):
3599+ result, a, h, _ = self._run_dispatch_v2(
3600+ tp_world_size=1, quant_mode=3, y_dtype=torch.float8_e5m2)
3601+ expected_dim1 = math.ceil(h / 128)
3602+ self.assertEqual(result[1].shape, torch.Size([a, expected_dim1]))
3603+ 
3604+ def test_tp0_mx(self):
3605+ result, a, h, _ = self._run_dispatch_v2(
3606+ tp_world_size=0, quant_mode=4, y_dtype=torch.float8_e4m3fn)
3607+ expected_dim1 = (math.ceil(h / 32) + 1) // 2 * 2
3608+ self.assertEqual(result[1].shape, torch.Size([a, expected_dim1]))
3609+ self.assertEqual(result[1].dtype, torch.uint8)
3610+ 
3611+ def test_tp1_mx(self):
3612+ result, a, h, _ = self._run_dispatch_v2(
3613+ tp_world_size=1, quant_mode=4, y_dtype=torch.float8_e4m3fn)
3614+ expected_dim1 = (math.ceil(h / 32) + 1) // 2 * 2
3615+ self.assertEqual(result[1].shape, torch.Size([a, expected_dim1]))
3616+ self.assertEqual(result[1].dtype, torch.uint8)
3617+ 
3618+ def test_tp0_no_quant(self):
3619+ result, a, h, _ = self._run_dispatch_v2(tp_world_size=0, quant_mode=0)
3620+ self.assertEqual(result[0].shape, torch.Size([a, h]))
3621+ self.assertEqual(result[0].dtype, torch.bfloat16)
3622+ self.assertEqual(result[1].shape, torch.Size([a]))
3623+ 
3624+ 
3554class TestMoeDistributeCombineAddRmsNorm(TestCase):3625class TestMoeDistributeCombineAddRmsNorm(TestCase):
3555 def test_moe_distribute_combine_add_rms_norm(self):3626 def test_moe_distribute_combine_add_rms_norm(self):
3556 with FakeTensorMode():3627 with FakeTensorMode():