已合并
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
已合并
共 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 | |||
| 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 | + | ||
| 3554 | class TestMoeDistributeCombineAddRmsNorm(TestCase): | 3625 | class 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(): |
此处的逻辑增加了std::max的取值,与之前是否功能一致?