19
csrc/mc2/dispatch_gmm_combine_decode/CMakeLists.txt
Normal file
19
csrc/mc2/dispatch_gmm_combine_decode/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
|
||||
if(NOT ENABLE_TEST AND NOT BENCHMARK)
|
||||
list(REMOVE_ITEM CURRENT_DIRS tests)
|
||||
endif()
|
||||
foreach(SUB_DIR ${CURRENT_DIRS})
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
|
||||
add_subdirectory(${SUB_DIR})
|
||||
endif()
|
||||
endforeach()
|
||||
@@ -0,0 +1,83 @@
|
||||
/*
|
||||
* Copyright (c) Huawei Technologies Co., Ltd. 2026. All rights reserved.
|
||||
*
|
||||
* 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.
|
||||
*/
|
||||
#ifndef DISPATCH_GMM_COMBINE_TORCH_ADPT_H
|
||||
#define DISPATCH_GMM_COMBINE_TORCH_ADPT_H
|
||||
|
||||
namespace vllm_ascend {
|
||||
|
||||
std::tuple<at::Tensor, at::Tensor> dispatch_gmm_combine_decode(
|
||||
const at::Tensor &x,
|
||||
const at::Tensor &expert_ids,
|
||||
const at::TensorList &gmm1_permuted_weight,
|
||||
const at::TensorList &gmm1_permuted_weight_scale,
|
||||
const at::TensorList &gmm2_weight,
|
||||
const at::TensorList &gmm2_weight_scale,
|
||||
const at::Tensor &expert_scales,
|
||||
const c10::optional<at::Tensor> &expert_smooth_scales,
|
||||
const c10::optional<at::Tensor> &x_active_mask,
|
||||
c10::string_view group_ep,
|
||||
int64_t ep_rank_size,
|
||||
int64_t ep_rank_id,
|
||||
int64_t moe_expert_num,
|
||||
int64_t shared_expert_num,
|
||||
int64_t shared_expert_rank_num,
|
||||
int64_t quant_mode,
|
||||
int64_t global_bs)
|
||||
{
|
||||
auto x_shape = x.sizes();
|
||||
int bs = x_shape[0];
|
||||
int h = x_shape[1];
|
||||
|
||||
at::Tensor output = at::empty({bs, h}, x.options());
|
||||
|
||||
bool is_shared_expert = (ep_rank_id < shared_expert_rank_num);
|
||||
int64_t num_local_experts = is_shared_expert ? 1 : moe_expert_num / (ep_rank_size - shared_expert_rank_num);
|
||||
auto opts = expert_ids.options().dtype(at::kLong);
|
||||
at::Tensor expert_token_nums = at::empty({num_local_experts}, opts);
|
||||
|
||||
vector<char> group_ep_chrs(group_ep.begin(), group_ep.end());
|
||||
group_ep_chrs.push_back('\0');
|
||||
char *group_ep_ptr = &group_ep_chrs[0];
|
||||
EXEC_NPU_CMD(
|
||||
// op api
|
||||
aclnnDispatchGmmCombineDecode,
|
||||
// input tensors
|
||||
x,
|
||||
expert_ids,
|
||||
gmm1_permuted_weight,
|
||||
gmm1_permuted_weight_scale,
|
||||
gmm2_weight,
|
||||
gmm2_weight_scale,
|
||||
expert_scales,
|
||||
expert_smooth_scales,
|
||||
x_active_mask,
|
||||
//input attrs
|
||||
group_ep_ptr,
|
||||
ep_rank_size,
|
||||
ep_rank_id,
|
||||
moe_expert_num,
|
||||
shared_expert_num,
|
||||
shared_expert_rank_num,
|
||||
quant_mode,
|
||||
global_bs,
|
||||
// output tensors
|
||||
output,
|
||||
expert_token_nums);
|
||||
return {output, expert_token_nums};
|
||||
}
|
||||
|
||||
}
|
||||
#endif
|
||||
32
csrc/mc2/dispatch_gmm_combine_decode/op_host/CMakeLists.txt
Normal file
32
csrc/mc2/dispatch_gmm_combine_decode/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,32 @@
|
||||
set(_DISPATCH_GMM_INC_OPTS)
|
||||
if (EXISTS ${CMAKE_SOURCE_DIR}/third_party/catlass/include)
|
||||
list(APPEND _DISPATCH_GMM_INC_OPTS -I${CMAKE_SOURCE_DIR}/third_party/catlass/include)
|
||||
else()
|
||||
message(FATAL_ERROR "dependency catlass is missing, you can fetch it by running 'git submodule update --init --recursive'")
|
||||
endif()
|
||||
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnnInner PRIVATE
|
||||
dispatch_gmm_combine_decode_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME DispatchGmmCombineDecode
|
||||
OPTIONS
|
||||
--cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-DASCENDC_DUMP=0
|
||||
-DCATLASS_ARCH=2201
|
||||
${_DISPATCH_GMM_INC_OPTS}
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE dispatch_gmm_combine_decode ACLNNTYPE aclnn_inner)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_SOURCE_DIR}/utils/inc/kernel
|
||||
)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
class DispatchGmmCombineDecode : public OpDef
|
||||
{
|
||||
public:
|
||||
explicit DispatchGmmCombineDecode(const char *name) : OpDef(name)
|
||||
{
|
||||
this->Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
|
||||
ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BF16,
|
||||
ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("expert_ids")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32,
|
||||
ge::DT_INT32, ge::DT_INT32, ge::DT_INT32, ge::DT_INT32})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("gmm1_permuted_weight")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8,
|
||||
ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_BF16, ge::DT_FLOAT16})
|
||||
.Format({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("gmm1_permuted_weight_scale")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_BF16,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("gmm2_weight")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8,
|
||||
ge::DT_BF16, ge::DT_FLOAT16, ge::DT_INT8, ge::DT_INT8,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_INT8, ge::DT_INT8,
|
||||
ge::DT_INT8, ge::DT_INT8, ge::DT_BF16, ge::DT_FLOAT16})
|
||||
.Format({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ,
|
||||
ge::FORMAT_FRACTAL_NZ, ge::FORMAT_FRACTAL_NZ, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("gmm2_weight_scale")
|
||||
.ParamType(DYNAMIC)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_BF16,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT16,
|
||||
ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_BF16,
|
||||
ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT, ge::DT_FLOAT16,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("expert_scales")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("expert_smooth_scales")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT,
|
||||
ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Input("x_active_mask")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL,
|
||||
ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL,
|
||||
ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL,
|
||||
ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL,
|
||||
ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL, ge::DT_BOOL})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("output")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_BF16, ge::DT_BF16, ge::DT_BF16,
|
||||
ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_BF16,
|
||||
ge::DT_BF16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT16,
|
||||
ge::DT_FLOAT16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Output("expert_token_nums")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64,
|
||||
ge::DT_INT64, ge::DT_INT64, ge::DT_INT64, ge::DT_INT64})
|
||||
.Format({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
|
||||
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND,
|
||||
ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
|
||||
this->Attr("group_ep").String();
|
||||
this->Attr("ep_rank_size").Int();
|
||||
this->Attr("ep_rank_id").Int();
|
||||
this->Attr("moe_expert_num").Int();
|
||||
this->Attr("share_expert_num").Int();
|
||||
this->Attr("share_expert_rank_num").Int();
|
||||
this->Attr("quant_mode").Int();
|
||||
this->Attr("global_bs").Int();
|
||||
|
||||
this->MC2().HcclGroup({"group_ep"});
|
||||
this->AICore().AddConfig("ascend910_93");
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(DispatchGmmCombineDecode);
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,95 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
#include <cstdint>
|
||||
#include "log/ops_log.h"
|
||||
#include "error/ops_error.h"
|
||||
#include "graph/utils/type_utils.h"
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ge {
|
||||
constexpr uint32_t EXPAND_X_INDEX = 0;
|
||||
constexpr uint32_t EXPERT_IDS_INDEX = 1;
|
||||
constexpr uint32_t OUTPUT_X_INDEX = 0;
|
||||
constexpr uint32_t OUTPUT_EXPERT_TOKEN_NUMS = 1;
|
||||
|
||||
constexpr uint32_t ATTR_GROUP_EP_INDEX = 0;
|
||||
constexpr uint32_t ATTR_EP_RANK_SIZE_INDEX = 1;
|
||||
constexpr uint32_t ATTR_EP_RANK_ID_INDEX = 2;
|
||||
constexpr uint32_t ATTR_MOE_EXPERT_NUM_INDEX = 3;
|
||||
constexpr uint32_t ATTR_SHARE_EXPERT_NUM_INDEX = 4;
|
||||
constexpr uint32_t ATTR_SHARE_EXPERT_RANK_NUM_INDEX = 5;
|
||||
constexpr uint32_t ATTR_QUANT_MODE_INDEX = 6;
|
||||
constexpr uint32_t ATTR_GLOBAL_BS_INDEX = 7;
|
||||
|
||||
static ge::graphStatus InferShape(gert::InferShapeContext *context)
|
||||
{
|
||||
const char *nodeName = context->GetNodeName();
|
||||
// infer output shape
|
||||
const gert::Shape *expandXShape = context->GetInputShape(EXPAND_X_INDEX);
|
||||
const gert::Shape *expertIdsShape = context->GetInputShape(EXPERT_IDS_INDEX);
|
||||
gert::Shape *expandXOutShape = context->GetOutputShape(OUTPUT_X_INDEX);
|
||||
gert::Shape *expertTokenNumsShape = context->GetOutputShape(OUTPUT_EXPERT_TOKEN_NUMS);
|
||||
if (expandXShape == nullptr || expertIdsShape == nullptr || expandXOutShape == nullptr ||
|
||||
expertTokenNumsShape == nullptr) {
|
||||
return GRAPH_FAILED;
|
||||
}
|
||||
if (expandXShape->GetDimNum() < 2 || expertIdsShape->GetDimNum() < 1) {
|
||||
return GRAPH_FAILED;
|
||||
}
|
||||
|
||||
int bs = expertIdsShape->GetDim(0);
|
||||
int h = expandXShape->GetDim(1);
|
||||
|
||||
expandXOutShape->SetDimNum(expandXShape->GetDimNum());
|
||||
expandXOutShape->SetDim(0, bs);
|
||||
expandXOutShape->SetDim(1, h);
|
||||
|
||||
// infer recvCount shape
|
||||
auto attrs = context->GetAttrs();
|
||||
OPS_ERR_IF(attrs == nullptr, OPS_LOG_E(nodeName, "attrs is nullptr."), return ge::GRAPH_FAILED);
|
||||
|
||||
auto epRankSizePtr = attrs->GetAttrPointer<int64_t>(ATTR_EP_RANK_SIZE_INDEX);
|
||||
auto epRankIdPtr = attrs->GetAttrPointer<int64_t>(ATTR_EP_RANK_ID_INDEX);
|
||||
auto moeExpertNumPtr = attrs->GetAttrPointer<int64_t>(ATTR_MOE_EXPERT_NUM_INDEX);
|
||||
auto sharedExpertRankNumPtr = attrs->GetAttrPointer<int64_t>(ATTR_SHARE_EXPERT_RANK_NUM_INDEX);
|
||||
|
||||
OPS_ERR_IF(epRankIdPtr == nullptr, OPS_LOG_E(nodeName, "epRankIdPtr is nullptr."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(moeExpertNumPtr == nullptr, OPS_LOG_E(nodeName, "moeExpertNumPtr is nullptr."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(epRankSizePtr == nullptr, OPS_LOG_E(nodeName, "epRankSizePtr is nullptr."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(sharedExpertRankNumPtr == nullptr, OPS_LOG_E(nodeName, "sharedExpertRankNumPtr is nullptr."),
|
||||
return ge::GRAPH_FAILED);
|
||||
uint32_t epRankSize = static_cast<uint32_t>(*epRankSizePtr);
|
||||
uint32_t moeExpertNum = static_cast<uint32_t>(*moeExpertNumPtr);
|
||||
uint32_t epRankId = static_cast<uint32_t>(*epRankIdPtr);
|
||||
uint32_t sharedExpertRankNum = static_cast<uint32_t>(*sharedExpertRankNumPtr);
|
||||
|
||||
expertTokenNumsShape->SetDimNum(1);
|
||||
bool isShareExpert = (epRankId < sharedExpertRankNum);
|
||||
if (isShareExpert) {
|
||||
expertTokenNumsShape->SetDim(0, 1);
|
||||
} else {
|
||||
expertTokenNumsShape->SetDim(0, moeExpertNum / (epRankSize - sharedExpertRankNum));
|
||||
}
|
||||
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus InferDataType(gert::InferDataTypeContext *context)
|
||||
{
|
||||
const auto expandXDataType = context->GetInputDataType(EXPAND_X_INDEX);
|
||||
context->SetOutputDataType(OUTPUT_X_INDEX, expandXDataType);
|
||||
context->SetOutputDataType(OUTPUT_EXPERT_TOKEN_NUMS, ge::DT_INT64);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP(DispatchGmmCombineDecode).InferShape(InferShape).InferDataType(InferDataType);
|
||||
} // namespace ge
|
||||
@@ -0,0 +1,520 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#include <cstdio>
|
||||
#include <cstdint>
|
||||
#include <string>
|
||||
|
||||
#include "log/ops_log.h"
|
||||
#include "error/ops_error.h"
|
||||
#include "graph/utils/type_utils.h"
|
||||
#include "register/op_def_registry.h"
|
||||
#include "../op_kernel/dispatch_gmm_combine_decode_tiling.h"
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "tiling/hccl/hccl_tiling.h"
|
||||
|
||||
using namespace ge;
|
||||
namespace {
|
||||
constexpr uint32_t OP_TYPE_ALL_TO_ALL = 8;
|
||||
constexpr uint32_t SYSTEM_NEED_WORKSPACE = 16 * 1024 * 1024;
|
||||
constexpr uint32_t GM_ALIGN_SIZE = 512;
|
||||
constexpr uint32_t TOKEN_DTYPE_BYTE_SIZE = 2;
|
||||
constexpr uint32_t L1_TILE_BYTE_SIZE = 32 * 1024;
|
||||
constexpr uint32_t CUBE_WORKSPACE_STAGE = 4;
|
||||
constexpr uint32_t RESERVED_WORKSPACE_SIZE = 256 * 1024;
|
||||
|
||||
constexpr uint32_t INPUT_X_INDEX = 0;
|
||||
constexpr uint32_t INPUT_EXPERT_IDS_INDEX = 1;
|
||||
constexpr uint32_t INPUT_GMM1_WEIGHT_INDEX = 2;
|
||||
constexpr uint32_t INPUT_GMM1_WEIGHT_SCALE_INDEX = 3;
|
||||
constexpr uint32_t INPUT_GMM2_WEIGHT_INDEX = 4;
|
||||
constexpr uint32_t INPUT_GMM2_WEIGHT_SCALE_INDEX = 5;
|
||||
constexpr uint32_t INPUT_EXPERT_SCALE_INDEX = 6;
|
||||
constexpr uint32_t INPUT_SMOOTH_SCALE_INDEX = 7;
|
||||
constexpr uint32_t INPUT_SHARE_X_ACTIVE_MASK_INDEX = 8;
|
||||
|
||||
constexpr uint32_t ATTR_GROUP_EP_INDEX = 0;
|
||||
constexpr uint32_t ATTR_EP_RANK_SIZE_INDEX = 1;
|
||||
constexpr uint32_t ATTR_EP_RANK_ID_INDEX = 2;
|
||||
constexpr uint32_t ATTR_MOE_EXPERT_NUM_INDEX = 3;
|
||||
constexpr uint32_t ATTR_SHARE_EXPERT_NUM_INDEX = 4;
|
||||
constexpr uint32_t ATTR_SHARE_EXPERT_RANK_NUM_INDEX = 5;
|
||||
constexpr uint32_t ATTR_QUANT_MODE_INDEX = 6;
|
||||
constexpr uint32_t ATTR_GLOBAL_BS_INDEX = 7;
|
||||
|
||||
constexpr uint32_t MIN_BATCH_SIZE = 1;
|
||||
constexpr uint32_t MAX_BATCH_SIZE = 256;
|
||||
constexpr uint32_t MAX_MOE_EXERT_NUM = 512;
|
||||
constexpr uint32_t SUPPORT_TOP_K = 12;
|
||||
constexpr uint32_t ONE_DIMS = 1;
|
||||
constexpr uint32_t TWO_DIMS = 2;
|
||||
constexpr uint32_t MIN_TOKEN_LENGTH = 512;
|
||||
constexpr uint32_t MAX_TOKEN_LENGTH = 7168;
|
||||
constexpr uint32_t MIN_GMM1_HIDDEN = 1024;
|
||||
constexpr uint32_t MAX_GMM1_HIDDEN = 6144;
|
||||
constexpr uint32_t TENSOR_HIDDEN_INDEX = 1;
|
||||
constexpr uint32_t SINGLE_HIDDEN_INDEX = 2;
|
||||
constexpr uint32_t MAX_TENSOR_COUNT = 256;
|
||||
} // namespace
|
||||
|
||||
namespace optiling {
|
||||
static size_t CeilUp(size_t x, size_t y)
|
||||
{
|
||||
return (x + y - 1) / y * y;
|
||||
}
|
||||
|
||||
static uint32_t CountTensorListLen(gert::TilingContext *context, int descIndex)
|
||||
{
|
||||
int count = 0;
|
||||
for (uint32_t i = 0; i < MAX_TENSOR_COUNT; i++) {
|
||||
auto tensorElement = context->GetDynamicInputTensor(descIndex, i);
|
||||
if (tensorElement == nullptr) {
|
||||
break;
|
||||
}
|
||||
count++;
|
||||
}
|
||||
return count;
|
||||
}
|
||||
|
||||
static ge::graphStatus CheckGmm1Shape(gert::TilingContext *context, DispatchGmmCombineDecodeTilingData *tilingData)
|
||||
{
|
||||
const char *nodeName = context->GetNodeName();
|
||||
uint32_t moeExpertNumPerRank = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
uint32_t h = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
uint32_t gmm1ListLen = CountTensorListLen(context, INPUT_GMM1_WEIGHT_INDEX);
|
||||
auto gmm1FirstTensorElement = context->GetDynamicInputTensor(INPUT_GMM1_WEIGHT_INDEX, 0);
|
||||
auto gmm1FirstTensorElementShape = gmm1FirstTensorElement->GetOriginShape();
|
||||
uint32_t elementDims = gmm1FirstTensorElementShape.GetDimNum();
|
||||
ge::DataType gmm1DataType = gmm1FirstTensorElement->GetDataType();
|
||||
if (gmm1DataType == ge::DT_BF16 || gmm1DataType == ge::DT_FLOAT16) {
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isBf16Fp16W = true;
|
||||
} else {
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isBf16Fp16W = false;
|
||||
}
|
||||
auto gmm1WeightDesc = context->GetDynamicInputDesc(INPUT_GMM1_WEIGHT_INDEX, 0);
|
||||
if (GetPrimaryFormat(gmm1WeightDesc->GetStorageFormat()) == ge::FORMAT_ND) {
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isNDFormat = true;
|
||||
}
|
||||
|
||||
OPS_ERR_IF(elementDims != 2 && elementDims != 3, OPS_LOG_E(nodeName, "gmm1Weight shape is invalid."),
|
||||
return ge::GRAPH_FAILED);
|
||||
if (gmm1ListLen > 1) { // List
|
||||
OPS_ERR_IF(h != gmm1FirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm1Weight input length does not equals to token hidden size."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(gmm1ListLen != moeExpertNumPerRank,
|
||||
OPS_LOG_E(nodeName, "gmm1Weight does not match local expert number perRank."),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen =
|
||||
gmm1FirstTensorElementShape.GetDim(TENSOR_HIDDEN_INDEX);
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isTensorList = true;
|
||||
} else { // Single
|
||||
if (elementDims == 2) { // one localExpert perRank
|
||||
OPS_ERR_IF(h != gmm1FirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm1Weight input length does not equals to token hidden size."),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen =
|
||||
gmm1FirstTensorElementShape.GetDim(SINGLE_HIDDEN_INDEX - 1);
|
||||
} else { // multi localExperts perRank
|
||||
OPS_ERR_IF(moeExpertNumPerRank != gmm1FirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm1Weight does not match local expert number per rank."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(h != gmm1FirstTensorElementShape.GetDim(1),
|
||||
OPS_LOG_E(nodeName, "gmm1Weight input length does not equals to token hidden size."),
|
||||
return ge::GRAPH_FAILED);
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen =
|
||||
gmm1FirstTensorElementShape.GetDim(SINGLE_HIDDEN_INDEX);
|
||||
}
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isTensorList = false;
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus CheckGmm1ScaleShape(gert::TilingContext *context,
|
||||
DispatchGmmCombineDecodeTilingData *tilingData)
|
||||
{
|
||||
if (tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isBf16Fp16W) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
const char *nodeName = context->GetNodeName();
|
||||
uint32_t moeExpertNumPerRank = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
uint32_t n = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen;
|
||||
|
||||
|
||||
uint32_t gmm1ScaleListLen = CountTensorListLen(context, INPUT_GMM1_WEIGHT_SCALE_INDEX);
|
||||
auto gmm1ScaleFirstTensorElement = context->GetDynamicInputTensor(INPUT_GMM1_WEIGHT_SCALE_INDEX, 0);
|
||||
auto gmm1ScaleFirstTensorElementShape = gmm1ScaleFirstTensorElement->GetOriginShape();
|
||||
uint32_t elementDims = gmm1ScaleFirstTensorElementShape.GetDimNum();
|
||||
OPS_ERR_IF(elementDims != 1 && elementDims != 2, OPS_LOG_E(nodeName, "gmm1WeightScale shape is invalid."),
|
||||
return ge::GRAPH_FAILED);
|
||||
if (gmm1ScaleListLen > 1) { // List
|
||||
OPS_ERR_IF(n != gmm1ScaleFirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm1Scale length does not equals to gmm1 hidden size."), return ge::GRAPH_FAILED);
|
||||
} else { // Single
|
||||
if (elementDims == 1) { // one localExpert perRank
|
||||
OPS_ERR_IF(n != gmm1ScaleFirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm1Scale length does not equals to gmm1 hidden size."), return ge::GRAPH_FAILED);
|
||||
} else { // multi localExperts perRank
|
||||
OPS_ERR_IF(moeExpertNumPerRank != gmm1ScaleFirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm1Scale does not match local expert number perRank."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(n != gmm1ScaleFirstTensorElementShape.GetDim(1),
|
||||
OPS_LOG_E(nodeName, "gmm1Scale length does not equals to gmm1 hidden size."), return ge::GRAPH_FAILED);
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus CheckGmm2Shape(gert::TilingContext *context, DispatchGmmCombineDecodeTilingData *tilingData)
|
||||
{
|
||||
const char *nodeName = context->GetNodeName();
|
||||
uint32_t moeExpertNumPerRank = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
uint32_t h = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
uint32_t n = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen;
|
||||
|
||||
uint32_t gmm2ListLen = CountTensorListLen(context, INPUT_GMM2_WEIGHT_INDEX);
|
||||
auto gmm2FirstTensorElement = context->GetDynamicInputTensor(INPUT_GMM2_WEIGHT_INDEX, 0);
|
||||
auto gmm2FirstTensorElementShape = gmm2FirstTensorElement->GetOriginShape();
|
||||
uint32_t elementDims = gmm2FirstTensorElementShape.GetDimNum();
|
||||
OPS_ERR_IF(elementDims != 2 && elementDims != 3, OPS_LOG_E(nodeName, "gmm2Weight shape is invalid."),
|
||||
return ge::GRAPH_FAILED);
|
||||
auto gmm2WeightDesc = context->GetDynamicInputDesc(INPUT_GMM2_WEIGHT_INDEX, 0);
|
||||
if (GetPrimaryFormat(gmm2WeightDesc->GetStorageFormat()) == ge::FORMAT_ND) {
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isNDFormat = true;
|
||||
}
|
||||
if (gmm2ListLen > 1) { // List
|
||||
OPS_ERR_IF(gmm2ListLen != moeExpertNumPerRank,
|
||||
OPS_LOG_E(nodeName, "gmm2 does not match local expert number perRank."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(n / 2 != gmm2FirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm2 does not equals to token hidden size."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(h != gmm2FirstTensorElementShape.GetDim(1),
|
||||
OPS_LOG_E(nodeName, "gmm2 does not match half of gmm1 hidden size."), return ge::GRAPH_FAILED);
|
||||
} else { // Single
|
||||
if (elementDims == 2) { // one localExpert perRank
|
||||
OPS_ERR_IF(n / 2 != gmm2FirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm2Weight does not equals to token hidden size."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(h != gmm2FirstTensorElementShape.GetDim(1),
|
||||
OPS_LOG_E(nodeName, "gmm2Weight does not match half of gmm1 hidden size."), return ge::GRAPH_FAILED);
|
||||
} else { // multi localExperts perRank
|
||||
OPS_ERR_IF(moeExpertNumPerRank != gmm2FirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm2Weight does not match local expert num perRank."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(n / 2 != gmm2FirstTensorElementShape.GetDim(1),
|
||||
OPS_LOG_E(nodeName, "gmm2Weight does not equals to token hidden size."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(h != gmm2FirstTensorElementShape.GetDim(2),
|
||||
OPS_LOG_E(nodeName, "gmm2Weight does not match half of gmm1 hidden size."), return ge::GRAPH_FAILED);
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus CheckGmm2ScaleShape(gert::TilingContext *context,
|
||||
DispatchGmmCombineDecodeTilingData *tilingData)
|
||||
{
|
||||
if (tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isBf16Fp16W) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
const char *nodeName = context->GetNodeName();
|
||||
uint32_t moeExpertNumPerRank = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
uint32_t h = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
|
||||
uint32_t gmm2ScaleListLen = CountTensorListLen(context, INPUT_GMM2_WEIGHT_SCALE_INDEX);
|
||||
auto gmm2ScaleFirstTensorElement = context->GetDynamicInputTensor(INPUT_GMM2_WEIGHT_SCALE_INDEX, 0);
|
||||
auto gmm2ScaleFirstTensorElementShape = gmm2ScaleFirstTensorElement->GetOriginShape();
|
||||
uint32_t elementDims = gmm2ScaleFirstTensorElementShape.GetDimNum();
|
||||
OPS_ERR_IF(elementDims != 1 && elementDims != 2, OPS_LOG_E(nodeName, "gmm2WeightScale shape is invalid."),
|
||||
return ge::GRAPH_FAILED);
|
||||
if (gmm2ScaleListLen > 1) { // List
|
||||
OPS_ERR_IF(h != gmm2ScaleFirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm2Scale does not match token hidden size."), return ge::GRAPH_FAILED);
|
||||
} else { // Single
|
||||
if (elementDims == 1) { // one localExpert perRank
|
||||
OPS_ERR_IF(h != gmm2ScaleFirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm2Scale does not match token hidden size."), return ge::GRAPH_FAILED);
|
||||
} else { // multi localExperts perRank
|
||||
OPS_ERR_IF(moeExpertNumPerRank != gmm2ScaleFirstTensorElementShape.GetDim(0),
|
||||
OPS_LOG_E(nodeName, "gmm2Scale does not match local expert number perRank."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(h != gmm2ScaleFirstTensorElementShape.GetDim(1),
|
||||
OPS_LOG_E(nodeName, "gmm2Scale does not match token hidden size."), return ge::GRAPH_FAILED);
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus CheckWeightTensorList(gert::TilingContext *context,
|
||||
DispatchGmmCombineDecodeTilingData *tilingData)
|
||||
{
|
||||
if (CheckGmm1Shape(context, tilingData) == ge::GRAPH_SUCCESS &&
|
||||
CheckGmm1ScaleShape(context, tilingData) == ge::GRAPH_SUCCESS &&
|
||||
CheckGmm2Shape(context, tilingData) == ge::GRAPH_SUCCESS &&
|
||||
CheckGmm2ScaleShape(context, tilingData) == ge::GRAPH_SUCCESS) {
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
ge::graphStatus CheckXActiveMaskShape(gert::TilingContext *context, const char *nodeName,
|
||||
DispatchGmmCombineDecodeTilingData &tilingData)
|
||||
{
|
||||
uint32_t epRankId = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.epRankId;
|
||||
uint32_t moeExpertNum = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNum;
|
||||
uint32_t sharedExpertRankNum = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertRankNum;
|
||||
uint32_t moeExpertNumPerRank = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
uint32_t batchSize = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.bs;
|
||||
uint32_t h = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
uint64_t gmm1WeightDim2 = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen;
|
||||
uint32_t localExpertNum = epRankId < sharedExpertRankNum ? 1 : moeExpertNumPerRank;
|
||||
const gert::StorageShape* xActiveMaskStorageShape = context->GetOptionalInputShape(
|
||||
INPUT_SHARE_X_ACTIVE_MASK_INDEX);
|
||||
if (xActiveMaskStorageShape != nullptr) {
|
||||
OPS_ERR_IF(xActiveMaskStorageShape->GetStorageShape().GetDimNum() != ONE_DIMS,
|
||||
OPS_LOG_E(nodeName, " xActiveMask scale shape dims must be 1, but current dim num is %lu.",
|
||||
xActiveMaskStorageShape->GetStorageShape().GetDimNum()),
|
||||
return ge::GRAPH_FAILED);
|
||||
const int64_t xActiveMaskDim0 = xActiveMaskStorageShape->GetStorageShape().GetDim(0);
|
||||
OPS_ERR_IF(xActiveMaskDim0 != batchSize, OPS_LOG_E(nodeName,
|
||||
"xActiveMask Dim0 must be batchSize(%u), but current dim is %lu.", batchSize, xActiveMaskDim0),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus CheckData(const char *nodeName, DispatchGmmCombineDecodeTilingData &tilingData)
|
||||
{
|
||||
uint32_t batchSize = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.bs;
|
||||
OPS_ERR_IF(batchSize < MIN_BATCH_SIZE, OPS_LOG_E(nodeName, "batchSize(bs) must >= %d.", MIN_BATCH_SIZE),
|
||||
return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(batchSize > MAX_BATCH_SIZE, OPS_LOG_E(nodeName, "batchSize(bs) must <= %d.", MAX_BATCH_SIZE),
|
||||
return ge::GRAPH_FAILED);
|
||||
uint32_t tokenLength = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
OPS_ERR_IF(
|
||||
tokenLength < MIN_TOKEN_LENGTH || tokenLength > MAX_TOKEN_LENGTH,
|
||||
OPS_LOG_E(nodeName, "tokenLength(h) is invalid. Only support [%u, %u].", MIN_TOKEN_LENGTH, MAX_TOKEN_LENGTH),
|
||||
return ge::GRAPH_FAILED);
|
||||
uint32_t gmm1HLen = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen;
|
||||
OPS_ERR_IF(
|
||||
gmm1HLen < MIN_GMM1_HIDDEN || gmm1HLen > MAX_GMM1_HIDDEN,
|
||||
OPS_LOG_E(nodeName, "gmm1 hidden size is invalid. Only support [%u, %u].", MIN_GMM1_HIDDEN, MAX_GMM1_HIDDEN),
|
||||
return ge::GRAPH_FAILED);
|
||||
uint32_t topK = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.k;
|
||||
OPS_ERR_IF(topK > SUPPORT_TOP_K, OPS_LOG_E(nodeName, "topK(k) must <= %d.", SUPPORT_TOP_K),
|
||||
return ge::GRAPH_FAILED);
|
||||
uint32_t globalBatchSize = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.globalBs;
|
||||
uint32_t epRankSize = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.epRankSize;
|
||||
if (globalBatchSize == 0) {
|
||||
globalBatchSize = epRankSize * batchSize;
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.globalBs = globalBatchSize;
|
||||
} else {
|
||||
OPS_ERR_IF(globalBatchSize < 0, OPS_LOG_E(nodeName, "globalBatchSize must >= 0."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(globalBatchSize % epRankSize > 0,
|
||||
OPS_LOG_E(nodeName, "globalBatchSize must be divisible by epRankSize."),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
uint32_t moeExpertNumPerRank = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
uint32_t recvAivNum = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.aivNum / 2;
|
||||
OPS_ERR_IF(
|
||||
moeExpertNumPerRank > recvAivNum,
|
||||
OPS_LOG_E(nodeName, "moeExpertNumPerRank must <= (aivNum/2)(%u), but got %u", recvAivNum, moeExpertNumPerRank),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus GetAttrAndSetTilingData(gert::TilingContext *context, const char *nodeName,
|
||||
DispatchGmmCombineDecodeTilingData &tilingData, std::string &groupEp)
|
||||
{
|
||||
auto attrs = context->GetAttrs();
|
||||
OPS_ERR_IF(attrs == nullptr, OPS_LOG_E(nodeName, "attrs is nullptr."), return ge::GRAPH_FAILED);
|
||||
|
||||
auto groupEpPtr = attrs->GetAttrPointer<char>(static_cast<int>(ATTR_GROUP_EP_INDEX));
|
||||
auto epRankSizePtr = attrs->GetAttrPointer<int64_t>(ATTR_EP_RANK_SIZE_INDEX);
|
||||
auto epRankIdPtr = attrs->GetAttrPointer<int64_t>(ATTR_EP_RANK_ID_INDEX);
|
||||
auto moeExpertNumPtr = attrs->GetAttrPointer<int64_t>(ATTR_MOE_EXPERT_NUM_INDEX);
|
||||
auto sharedExpertNumPtr = attrs->GetAttrPointer<int64_t>(ATTR_SHARE_EXPERT_NUM_INDEX);
|
||||
auto sharedExpertRankNumPtr = attrs->GetAttrPointer<int64_t>(ATTR_SHARE_EXPERT_RANK_NUM_INDEX);
|
||||
auto quantModePtr = attrs->GetAttrPointer<int64_t>(ATTR_QUANT_MODE_INDEX);
|
||||
auto globalBsPtr = attrs->GetAttrPointer<int64_t>(ATTR_GLOBAL_BS_INDEX);
|
||||
|
||||
uint32_t epRankSize = static_cast<uint32_t>(*epRankSizePtr);
|
||||
uint32_t epRankId = static_cast<uint32_t>(*epRankIdPtr);
|
||||
uint32_t moeExpertNum = static_cast<uint32_t>(*moeExpertNumPtr);
|
||||
uint32_t sharedExpertNum = static_cast<uint32_t>(*sharedExpertNumPtr);
|
||||
uint32_t sharedExpertRankNum = static_cast<uint32_t>(*sharedExpertRankNumPtr);
|
||||
uint32_t moeExpertNumPerRank = moeExpertNum / (epRankSize - sharedExpertRankNum);
|
||||
|
||||
OPS_ERR_IF(epRankId < 0, OPS_LOG_E(nodeName, "epRankId must >= 0."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(epRankId >= epRankSize, OPS_LOG_E(nodeName, "epRankId must < epRankSize."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(moeExpertNum > MAX_MOE_EXERT_NUM, OPS_LOG_E(nodeName, "moeExpertNum must <= %d.", MAX_MOE_EXERT_NUM),
|
||||
return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(moeExpertNum <= 0, OPS_LOG_E(nodeName, "moeExpertNum must > 0."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(sharedExpertNum != 1, OPS_LOG_E(nodeName, "sharedExpertNum must be 1."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(moeExpertNum % (epRankSize - sharedExpertRankNum) != 0,
|
||||
OPS_LOG_E(nodeName, "moeExpertNum must be divisible by (epRankSize - sharedExpertRankNum)."),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
groupEp = std::string(groupEpPtr);
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.epRankSize = epRankSize;
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.epRankId = epRankId;
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNum = moeExpertNum;
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertNum = sharedExpertNum;
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertRankNum = sharedExpertRankNum;
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.quantMode = static_cast<uint32_t>(*quantModePtr);
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.globalBs = static_cast<uint32_t>(*globalBsPtr);
|
||||
tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank = moeExpertNumPerRank;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static void SetHcommCfg(const gert::TilingContext *context, DispatchGmmCombineDecodeTilingData *tiling, const std::string groupEp)
|
||||
{
|
||||
const char *nodeName = context->GetNodeName();
|
||||
OPS_LOG_D(nodeName, "DispatchGmmCombineDecode groupEp = %s", groupEp.c_str());
|
||||
uint32_t opType = OP_TYPE_ALL_TO_ALL;
|
||||
std::string algConfigAllToAllStr = "AlltoAll=level0:fullmesh;level1:pairwise";
|
||||
std::string algConfigAllGatherStr = "AllGather=level0:ring";
|
||||
|
||||
AscendC::Mc2CcTilingConfig mc2CcTilingConfig(groupEp, opType, algConfigAllToAllStr);
|
||||
mc2CcTilingConfig.GetTiling(tiling->mc2InitTiling);
|
||||
mc2CcTilingConfig.GetTiling(tiling->mc2CcTiling);
|
||||
}
|
||||
|
||||
static ge::graphStatus SetWorkSpace(gert::TilingContext *context, const char *nodeName,
|
||||
DispatchGmmCombineDecodeTilingData &tilingData)
|
||||
{
|
||||
size_t *workSpaces = context->GetWorkspaceSizes(1);
|
||||
OPS_ERR_IF(workSpaces == nullptr, OPS_LOG_E(nodeName, "workSpaces is nullptr."), return ge::GRAPH_FAILED);
|
||||
size_t maxTokenNum;
|
||||
uint32_t epRankSize = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.epRankSize;
|
||||
uint32_t epRankId = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.epRankId;
|
||||
uint32_t sharedExpertRankNum = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertRankNum;
|
||||
uint32_t batchSize = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.bs;
|
||||
uint32_t globalBs = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.globalBs;
|
||||
uint32_t maxBatchSize = globalBs / epRankSize;
|
||||
uint32_t topK = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.k;
|
||||
uint32_t moeExpertNumPerRank = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
uint32_t h = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
uint32_t aicNum = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.aicNum;
|
||||
uint64_t gmm1HLen = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen;
|
||||
uint64_t gmm2HLen = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen / 2;
|
||||
if (epRankId < sharedExpertRankNum) {
|
||||
maxTokenNum = maxBatchSize * epRankSize / sharedExpertRankNum;
|
||||
} else {
|
||||
maxTokenNum = maxBatchSize * epRankSize * std::min(topK, moeExpertNumPerRank);
|
||||
}
|
||||
uint32_t wTypeSize = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.isBf16Fp16W ? TOKEN_DTYPE_BYTE_SIZE : sizeof(int8_t);
|
||||
|
||||
// hbm input = x: float16 or bf16
|
||||
// buf1 dispatch (Only AIV) => x1: float16 or bf16
|
||||
// buf2 gmm1 (Only AIC) => y1: float
|
||||
// sync
|
||||
// buf3 swiglu (Only AIV) => x2: float16 or bf16
|
||||
// sync ?
|
||||
// buf4 gmm2 (AIC & AIV) => y2: float16 or bf16
|
||||
// hbm combine (Only AIV) => output: float16 or bf16
|
||||
|
||||
size_t x1TokenSize = maxTokenNum * h * wTypeSize; // x1: float16 or bf16
|
||||
size_t x2TokenSize = maxTokenNum * gmm2HLen * wTypeSize; // x2: float16 or bf16
|
||||
size_t maxTokenSize = x1TokenSize < x2TokenSize ? x2TokenSize : x1TokenSize;
|
||||
maxTokenSize = CeilUp(maxTokenSize, GM_ALIGN_SIZE);
|
||||
size_t tokenScaleSize = tilingData.disGmmDeqSwigluQuantGmmDeqComInfo.isBf16Fp16W ? 0 : CeilUp(maxTokenNum * sizeof(float), GM_ALIGN_SIZE);
|
||||
size_t CVSwapBufferSize =
|
||||
CeilUp(aicNum * L1_TILE_BYTE_SIZE * CUBE_WORKSPACE_STAGE * sizeof(int32_t), GM_ALIGN_SIZE);
|
||||
size_t swigluOutSize = maxTokenNum * gmm1HLen * sizeof(float); // y1: float
|
||||
size_t gmm2DepOutSize = maxTokenNum * h * TOKEN_DTYPE_BYTE_SIZE; // y2: float
|
||||
size_t maxSwigluGmm2Size = swigluOutSize < gmm2DepOutSize ? gmm2DepOutSize : swigluOutSize;
|
||||
maxSwigluGmm2Size = CeilUp(maxSwigluGmm2Size, GM_ALIGN_SIZE);
|
||||
size_t groupListSize = CeilUp(moeExpertNumPerRank * sizeof(int64_t), GM_ALIGN_SIZE);
|
||||
size_t expandIdxSize = CeilUp(batchSize * topK * sizeof(int32_t), GM_ALIGN_SIZE);
|
||||
size_t epSendCountSize = CeilUp(epRankSize * moeExpertNumPerRank * sizeof(int32_t), GM_ALIGN_SIZE);
|
||||
size_t resveredSize = CeilUp(RESERVED_WORKSPACE_SIZE, GM_ALIGN_SIZE);
|
||||
size_t usrSize = maxTokenSize + tokenScaleSize + CVSwapBufferSize + maxSwigluGmm2Size + groupListSize + expandIdxSize +
|
||||
epSendCountSize + resveredSize;
|
||||
|
||||
workSpaces[0] = SYSTEM_NEED_WORKSPACE + usrSize;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus DispatchGmmCombineDecodeTilingFuncImpl(gert::TilingContext *context)
|
||||
{
|
||||
const char *nodeName = context->GetNodeName();
|
||||
DispatchGmmCombineDecodeTilingData *tilingData = context->GetTilingData<DispatchGmmCombineDecodeTilingData>();
|
||||
OPS_ERR_IF(tilingData == nullptr, OPS_LOG_E(nodeName, "tilingData is nullptr."), return ge::GRAPH_FAILED);
|
||||
std::string groupEp = "";
|
||||
|
||||
const gert::StorageShape *xStorageShape = context->GetInputShape(INPUT_X_INDEX);
|
||||
OPS_ERR_IF(xStorageShape == nullptr, OPS_LOG_E(nodeName, "x shape is null."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(xStorageShape->GetStorageShape().GetDimNum() != TWO_DIMS,
|
||||
OPS_LOG_E(nodeName, "x shape dims must be 2, but current dim num is %lu.",
|
||||
xStorageShape->GetStorageShape().GetDimNum()),
|
||||
return ge::GRAPH_FAILED);
|
||||
const int64_t batchSize = xStorageShape->GetStorageShape().GetDim(0);
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.bs = batchSize;
|
||||
const int64_t hiddenSize = xStorageShape->GetStorageShape().GetDim(1);
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.h = hiddenSize;
|
||||
|
||||
const gert::StorageShape *expertIdsStorageShape = context->GetInputShape(INPUT_EXPERT_IDS_INDEX);
|
||||
OPS_ERR_IF(expertIdsStorageShape == nullptr, OPS_LOG_E(nodeName, "expertIds shape is null."),
|
||||
return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(expertIdsStorageShape->GetStorageShape().GetDimNum() != TWO_DIMS,
|
||||
OPS_LOG_E(nodeName, "expertIds shape dims must be 2, but current dim num is %lu.",
|
||||
expertIdsStorageShape->GetStorageShape().GetDimNum()),
|
||||
return ge::GRAPH_FAILED);
|
||||
const int64_t topK = expertIdsStorageShape->GetStorageShape().GetDim(1);
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.k = topK;
|
||||
OPS_ERR_IF(GetAttrAndSetTilingData(context, nodeName, *tilingData, groupEp) != ge::GRAPH_SUCCESS,
|
||||
OPS_LOG_E(nodeName, "Get attr and set tiling data failed."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(CheckWeightTensorList(context, tilingData) != ge::GRAPH_SUCCESS,
|
||||
OPS_LOG_E(nodeName, "CheckWeightTensorList failed."), return ge::GRAPH_FAILED);
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
|
||||
uint32_t aicNum = ascendcPlatform.GetCoreNumAic();
|
||||
uint32_t aivNum = ascendcPlatform.GetCoreNumAiv();
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.aicNum = aicNum;
|
||||
tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.aivNum = aivNum;
|
||||
OPS_ERR_IF(CheckData(nodeName, *tilingData) != ge::GRAPH_SUCCESS,
|
||||
OPS_LOG_E(nodeName, "CheckData failed."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(CheckXActiveMaskShape(context, nodeName, *tilingData) != ge::GRAPH_SUCCESS,
|
||||
OPS_LOG_E(nodeName, "CheckXActiveMaskShape failed."), return ge::GRAPH_FAILED);
|
||||
OPS_ERR_IF(SetWorkSpace(context, nodeName, *tilingData) != ge::GRAPH_SUCCESS,
|
||||
OPS_LOG_E(nodeName, "Tiling set workspace failed."), return ge::GRAPH_FAILED);
|
||||
SetHcommCfg(context, tilingData, groupEp);
|
||||
const gert::StorageShape* xActiveMaskStorageShape = context->GetOptionalInputShape(
|
||||
INPUT_SHARE_X_ACTIVE_MASK_INDEX);
|
||||
bool xActiveMaskEnable = (xActiveMaskStorageShape != nullptr);
|
||||
uint64_t tilingKey = 0;
|
||||
if (xActiveMaskEnable) {
|
||||
tilingKey |= EXEC_FLAG_X_ACTIVE_MASK;
|
||||
}
|
||||
if (tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank != 1) {
|
||||
tilingKey |= EXEC_FLAG_DEEP_FUSE;
|
||||
}
|
||||
if (tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isTensorList) {
|
||||
tilingKey |= EXEC_FLAG_TENSOR_LIST;
|
||||
}
|
||||
if (tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.isNDFormat) {
|
||||
tilingKey |= EXEC_FLAG_ND_FORMAT;
|
||||
}
|
||||
|
||||
context->SetTilingKey(tilingKey);
|
||||
context->SetBlockDim(aicNum);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus DispatchGmmCombineDecodeTilingFunc(gert::TilingContext *context)
|
||||
{
|
||||
ge::graphStatus ret = DispatchGmmCombineDecodeTilingFuncImpl(context);
|
||||
return ret;
|
||||
}
|
||||
|
||||
struct DispatchGmmCombineDecodeCompileInfo {};
|
||||
ge::graphStatus TilingParseForDispatchGmmCombineDecode(gert::TilingParseContext *context)
|
||||
{
|
||||
(void)context;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(DispatchGmmCombineDecode)
|
||||
.Tiling(DispatchGmmCombineDecodeTilingFunc)
|
||||
.TilingParse<DispatchGmmCombineDecodeCompileInfo>(TilingParseForDispatchGmmCombineDecode);
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,103 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#include <string.h>
|
||||
#include "graph/types.h"
|
||||
#include "aclnn/opdev/platform.h"
|
||||
#include "aclnn_dispatch_gmm_combine_decode.h"
|
||||
|
||||
enum NnopbaseHcclServerType {
|
||||
NNOPBASE_HCCL_SERVER_TYPE_AICPU = 0,
|
||||
NNOPBASE_HCCL_SERVER_TYPE_MTE,
|
||||
NNOPBASE_HCCL_SERVER_TYPE_END
|
||||
};
|
||||
extern "C" void __attribute__((weak)) NnopbaseSetHcclServerType(void *executor, NnopbaseHcclServerType sType);
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
extern aclnnStatus aclnnInnerDispatchGmmCombineDecodeGetWorkspaceSize(
|
||||
const aclTensor *x,
|
||||
const aclTensor *expertIds,
|
||||
const aclTensorList *gmm1PermutedWeight,
|
||||
const aclTensorList *gmm1PermutedWeightScale,
|
||||
const aclTensorList *gmm2Weight,
|
||||
const aclTensorList *gmm2WeightScale,
|
||||
const aclTensor *expertScales,
|
||||
const aclTensor *expertSmoothScalesOptional,
|
||||
const aclTensor *xActiveMaskOptional,
|
||||
char *groupEp,
|
||||
int64_t epRankSize,
|
||||
int64_t epRankId,
|
||||
int64_t moeExpertNum,
|
||||
int64_t shareExpertNum,
|
||||
int64_t shareExpertRankNum,
|
||||
int64_t quantMode,
|
||||
int64_t globalBs,
|
||||
const aclTensor *output,
|
||||
const aclTensor *expertTokenNums,
|
||||
uint64_t *workspaceSize,
|
||||
aclOpExecutor **executor);
|
||||
extern aclnnStatus aclnnInnerDispatchGmmCombineDecode(
|
||||
void *workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor *executor,
|
||||
aclrtStream stream);
|
||||
|
||||
aclnnStatus aclnnDispatchGmmCombineDecodeGetWorkspaceSize(
|
||||
const aclTensor *x,
|
||||
const aclTensor *expertIds,
|
||||
const aclTensorList *gmm1PermutedWeight,
|
||||
const aclTensorList *gmm1PermutedWeightScale,
|
||||
const aclTensorList *gmm2Weight,
|
||||
const aclTensorList *gmm2WeightScale,
|
||||
const aclTensor *expertScales,
|
||||
const aclTensor *expertSmoothScalesOptional,
|
||||
const aclTensor *xActiveMaskOptional,
|
||||
char *groupEp,
|
||||
int64_t epRankSize,
|
||||
int64_t epRankId,
|
||||
int64_t moeExpertNum,
|
||||
int64_t shareExpertNum,
|
||||
int64_t shareExpertRankNum,
|
||||
int64_t quantMode,
|
||||
int64_t globalBs,
|
||||
const aclTensor *output,
|
||||
const aclTensor *expertTokenNums,
|
||||
uint64_t *workspaceSize,
|
||||
aclOpExecutor **executor)
|
||||
{
|
||||
return aclnnInnerDispatchGmmCombineDecodeGetWorkspaceSize(x, expertIds, gmm1PermutedWeight, gmm1PermutedWeightScale,
|
||||
gmm2Weight, gmm2WeightScale, expertScales, expertSmoothScalesOptional, xActiveMaskOptional, groupEp, epRankSize,
|
||||
epRankId, moeExpertNum, shareExpertNum, shareExpertRankNum, quantMode, globalBs,
|
||||
output, expertTokenNums, workspaceSize, executor);
|
||||
}
|
||||
|
||||
aclnnStatus aclnnDispatchGmmCombineDecode(
|
||||
void *workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor *executor,
|
||||
aclrtStream stream)
|
||||
{
|
||||
if (NnopbaseSetHcclServerType) {
|
||||
if (op::GetCurrentPlatformInfo().GetSocVersion() == op::SocVersion::ASCEND910B) {
|
||||
NnopbaseSetHcclServerType(executor, NNOPBASE_HCCL_SERVER_TYPE_AICPU);
|
||||
} else {
|
||||
NnopbaseSetHcclServerType(executor, NNOPBASE_HCCL_SERVER_TYPE_MTE);
|
||||
}
|
||||
}
|
||||
return aclnnInnerDispatchGmmCombineDecode(workspace, workspaceSize, executor, stream);
|
||||
}
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef DISPATCH_GMM_COMBINE_DECODE
|
||||
#define DISPATCH_GMM_COMBINE_DECODE
|
||||
|
||||
#include "aclnn/acl_meta.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
__attribute__((visibility("default"))) aclnnStatus aclnnDispatchGmmCombineDecodeGetWorkspaceSize(
|
||||
const aclTensor *x,
|
||||
const aclTensor *expertIds,
|
||||
const aclTensorList *gmm1PermutedWeight,
|
||||
const aclTensorList *gmm1PermutedWeightScale,
|
||||
const aclTensorList *gmm2Weight,
|
||||
const aclTensorList *gmm2WeightScale,
|
||||
const aclTensor *expertScales,
|
||||
const aclTensor *expertSmoothScalesOptional,
|
||||
const aclTensor *xActiveMaskOptional,
|
||||
char *groupEp,
|
||||
int64_t epRankSize,
|
||||
int64_t epRankId,
|
||||
int64_t moeExpertNum,
|
||||
int64_t shareExpertNum,
|
||||
int64_t shareExpertRankNum,
|
||||
int64_t quantMode,
|
||||
int64_t globalBs,
|
||||
const aclTensor *output,
|
||||
const aclTensor *expertTokenNums,
|
||||
uint64_t *workspaceSize,
|
||||
aclOpExecutor **executor);
|
||||
|
||||
__attribute__((visibility("default"))) aclnnStatus aclnnDispatchGmmCombineDecode(
|
||||
void *workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor *executor,
|
||||
aclrtStream stream);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,53 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#include "dispatch_gmm_combine_decode.h"
|
||||
#include "dispatch_gmm_combine_decode_bf16_fp16.h"
|
||||
#include <kernel_operator.h>
|
||||
#include "lib/matmul_intf.h"
|
||||
|
||||
extern "C" __global__ __aicore__ void dispatch_gmm_combine_decode(
|
||||
// input
|
||||
GM_ADDR x, GM_ADDR expert_ids, GM_ADDR gmm1_permuted_weight, GM_ADDR gmm1_permuted_weight_scale,
|
||||
GM_ADDR gmm2_weight, GM_ADDR gmm2_weight_scale, GM_ADDR expert_scales, GM_ADDR expert_smooth_scales,
|
||||
GM_ADDR x_active_mask,
|
||||
// output
|
||||
GM_ADDR output, GM_ADDR expertTokenNums,
|
||||
// system
|
||||
GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
icache_preload(8);
|
||||
REGISTER_TILING_DEFAULT(DispatchGmmCombineDecodeTilingData);
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2); // 1C2V
|
||||
GET_TILING_DATA(tiling_data, tiling);
|
||||
|
||||
#if (ORIG_DTYPE_GMM1_PERMUTED_WEIGHT == DT_INT8)
|
||||
if constexpr (TILING_KEY_IS(0) || TILING_KEY_IS(1) || TILING_KEY_IS(2) || TILING_KEY_IS(3) ||
|
||||
TILING_KEY_IS(4) || TILING_KEY_IS(5) || TILING_KEY_IS(6) || TILING_KEY_IS(7) ||
|
||||
TILING_KEY_IS(8) || TILING_KEY_IS(9) || TILING_KEY_IS(10) || TILING_KEY_IS(11) ||
|
||||
TILING_KEY_IS(12) || TILING_KEY_IS(13) || TILING_KEY_IS(14) || TILING_KEY_IS(15)) {
|
||||
DispatchGmmCombineDecodeImpl::DispatchGmmCombineDecode<
|
||||
DTYPE_X, DTYPE_GMM1_PERMUTED_WEIGHT_SCALE, DTYPE_GMM2_WEIGHT_SCALE, int8_t, int32_t, false, TILING_KEY_VAR> op;
|
||||
op.Init(x, expert_ids, gmm1_permuted_weight, gmm1_permuted_weight_scale, gmm2_weight, gmm2_weight_scale,
|
||||
expert_scales, expert_smooth_scales, x_active_mask, output, expertTokenNums, workspace, nullptr, &tiling_data);
|
||||
op.Process();
|
||||
}
|
||||
#elif (ORIG_DTYPE_GMM1_PERMUTED_WEIGHT == DT_BF16 || ORIG_DTYPE_GMM1_PERMUTED_WEIGHT == DT_FLOAT16)
|
||||
if constexpr (TILING_KEY_IS(0) || TILING_KEY_IS(1) || TILING_KEY_IS(2) || TILING_KEY_IS(3) ||
|
||||
TILING_KEY_IS(4) || TILING_KEY_IS(5) || TILING_KEY_IS(6) || TILING_KEY_IS(7) ||
|
||||
TILING_KEY_IS(8) || TILING_KEY_IS(9) || TILING_KEY_IS(10) || TILING_KEY_IS(11) ||
|
||||
TILING_KEY_IS(12) || TILING_KEY_IS(13) || TILING_KEY_IS(14) || TILING_KEY_IS(15)) {
|
||||
DispatchGmmCombineDecodeBf16Fp16Impl::DispatchGmmCombineDecodeBf16Fp16<
|
||||
DTYPE_GMM1_PERMUTED_WEIGHT, DTYPE_GMM1_PERMUTED_WEIGHT_SCALE, DTYPE_GMM2_WEIGHT_SCALE, DTYPE_GMM1_PERMUTED_WEIGHT, int32_t, false, TILING_KEY_VAR> op;
|
||||
op.Init(x, expert_ids, gmm1_permuted_weight, gmm1_permuted_weight_scale, gmm2_weight, gmm2_weight_scale,
|
||||
expert_scales, expert_smooth_scales, x_active_mask, output, expertTokenNums, workspace, nullptr, &tiling_data);
|
||||
op.Process();
|
||||
}
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,456 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef DISPATCH_GMM_COMBINE_DECODE_H
|
||||
#define DISPATCH_GMM_COMBINE_DECODE_H
|
||||
|
||||
#include "lib/matmul_intf.h"
|
||||
#include <kernel_operator.h>
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/arch.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/epilogue/tile/tile_broadcast_mul.hpp"
|
||||
#include "catlass/epilogue/tile/tile_broadcast_one_blk.hpp"
|
||||
#include "catlass/epilogue/tile/tile_swizzle.hpp"
|
||||
#include "catlass/gemm/block/block_swizzle.hpp"
|
||||
#include "dispatch_gmm_combine_decode/gemm/kernel/grouped_matmul_slice_m_per_token_dequant_multistage_workspace.h"
|
||||
#include "catlass/gemm/gemm_type.hpp"
|
||||
#include "dispatch_gmm_combine_decode/epilogue/dispatch_policy.h"
|
||||
#include "dispatch_gmm_combine_decode/gemm/dispatch_policy.h"
|
||||
#include "dispatch_gmm_combine_decode/epilogue/block/block_epilogue.h"
|
||||
#include "dispatch_gmm_combine_decode/gemm/block/block_mmad.h"
|
||||
#include "dispatch_gmm_combine_decode/gemm/kernel/grouped_matmul_slice_m_per_token_dequant_swiglu_quant_multistage_workspace.h"
|
||||
|
||||
#include "dispatch_gmm_combine_decode/raw_distributed/cam_moe_distribute_dispatch.h"
|
||||
|
||||
#include "dispatch_gmm_combine_decode_tiling.h"
|
||||
#include "dispatch_gmm_combine_decode_base.h"
|
||||
|
||||
using namespace Catlass;
|
||||
namespace DispatchGmmCombineDecodeImpl {
|
||||
using MmadAtlasA2Custom =
|
||||
Gemm::MmadAtlasA2PreloadAsyncWithCallback<CUSTOM_PRELOAD_STAGES, CUSTOM_L1_STAGES, CUSTOM_L0A_STAGES,
|
||||
CUSTOM_L0B_STAGES, CUSTOM_L0C_STAGES, CUSTOM_ENABLE_UNIT_FLAG,
|
||||
CUSTOM_ENABLE_SHUFFLE_K>;
|
||||
|
||||
using Gmm1L1TileShape = GemmShape<GMM1_L1M, GMM1_L1N, GMM1_L1K>;
|
||||
using Gmm1L0TileShape = GemmShape<GMM1_L1M, GMM1_L1N, GMM1_L0K>;
|
||||
using Gmm1EpilogueTileShape = MatrixShape<GMM1_EPIM, Gmm1L1TileShape::N>;
|
||||
using Gmm1BlockScheduler = typename Gemm::Block::GemmIdentityBlockSwizzle<GMM1_SWIZZLE_OFFSET, GMM1_SWIZZLE_DIRECTION>;
|
||||
|
||||
using Gmm2L1TileShape = GemmShape<GMM2_L1M, GMM2_L1N, GMM2_L1K>;
|
||||
using Gmm2L0TileShape = GemmShape<Gmm2L1TileShape::M, Gmm2L1TileShape::N, GMM2_L0K>;
|
||||
using Gmm2EpilogueTileShape = MatrixShape<GMM2_EPIM, Gmm2L1TileShape::N>;
|
||||
using Gmm2BlockScheduler = typename Gemm::Block::GemmIdentityBlockSwizzle<GMM2_SWIZZLE_OFFSET, GMM2_SWIZZLE_DIRECTION>;
|
||||
using Gmm2DispatchPolicy =
|
||||
Gemm::MmadAtlasA2PreloadAsyncWithCallbackResidentA<CUSTOM_PRELOAD_STAGES, GMM2_L1A_STAGES, GMM2_L1B_STAGES,
|
||||
GMM2_L0A_STAGES, GMM2_L0B_STAGES, CUSTOM_L0C_STAGES,
|
||||
CUSTOM_ENABLE_UNIT_FLAG, CUSTOM_ENABLE_SHUFFLE_K>;
|
||||
|
||||
template <TemplateMC2TypeClass, class L1TileShape_, class L0TileShape_, class EpilogueTileShape_,
|
||||
class BlockScheduler_, class DispatchPolicy_ = MmadAtlasA2Custom>
|
||||
CATLASS_DEVICE void GmmDeqSwigluQuant(GemmCoord problemShape, uint32_t groupCount, GM_ADDR gmGroupList, GM_ADDR gmA,
|
||||
layout::RowMajor layoutA, GM_ADDR gmB,
|
||||
typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type layoutB,
|
||||
GM_ADDR gmScale,
|
||||
layout::VectorLayout layoutScale, GM_ADDR gmPerTokenScale,
|
||||
layout::VectorLayout layoutPerTokenScale, GM_ADDR gmD, layout::RowMajor layoutD,
|
||||
GM_ADDR gmDequantScale, layout::VectorLayout layoutDequantScale, GM_ADDR gmWorkspace,
|
||||
GM_ADDR gmX, GM_ADDR debugGm, GM_ADDR gmexpertIds, GM_ADDR gmExpandIdx,
|
||||
GM_ADDR gmEpSendCount, GM_ADDR xActiveMask, GM_ADDR gmResvered, GM_ADDR gmExpertTokenNums,
|
||||
uint32_t epRankSize, uint32_t epRankId, uint32_t moeExpertNum,
|
||||
uint32_t moeExpertNumPerRank, uint32_t sharedExpertNum, uint32_t sharedExpertRankNum,
|
||||
uint32_t quantMode, uint32_t globalBs, uint32_t bs, uint32_t topK, uint32_t tokenLen)
|
||||
{
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
using DispatchPolicy = DispatchPolicy_;
|
||||
using L1TileShape = L1TileShape_;
|
||||
using L0TileShape = L0TileShape_;
|
||||
|
||||
using AType = Gemm::GemmType<int8_t, layout::RowMajor>;
|
||||
using LayoutB = typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type;
|
||||
using BType = Gemm::GemmType<int8_t, LayoutB>;
|
||||
using CType = Gemm::GemmType<int32_t, layout::RowMajor>;
|
||||
|
||||
using BlockMmad = Gemm::Block::BlockMmad<DispatchPolicy, L1TileShape, L0TileShape, AType, BType, CType>;
|
||||
|
||||
constexpr uint32_t ubStages = 1;
|
||||
using EpilogueDispatchPolicy = Epilogue::EpilogueAtlasA2PerTokenDequantSwiglu<ubStages, 0>;
|
||||
using ScaleType = Gemm::GemmType<W1ScaleType, layout::VectorLayout>;
|
||||
using PerTokenScaleType = Gemm::GemmType<float, layout::VectorLayout>;
|
||||
using DType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
|
||||
using RowBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
using BroadcastOneBlkType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
using OneBlkColumnBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
|
||||
using EpilogueTileShape = EpilogueTileShape_;
|
||||
using TileRowBroadcastMul = Epilogue::Tile::TileRowBroadcastMul<ArchTag, RowBroadcastMulType, EpilogueTileShape>;
|
||||
using TileBroadcastOneBlk =
|
||||
Epilogue::Tile::TileBroadcastOneBlk<ArchTag, BroadcastOneBlkType, EpilogueTileShape::ROW>;
|
||||
using TileOneBlkColumnBroadcastMul =
|
||||
Epilogue::Tile::TileOneBlkColumnBroadcastMul<ArchTag, OneBlkColumnBroadcastMulType, EpilogueTileShape>;
|
||||
using TileCopy = Epilogue::Tile::TileCopy<ArchTag, CType, ScaleType, PerTokenScaleType, DType>;
|
||||
using TileScheduler = Epilogue::Tile::EpilogueHorizontalTileSwizzle;
|
||||
|
||||
using BlockEpilogue = Epilogue::Block::BlockEpilogue<EpilogueDispatchPolicy, CType, ScaleType, PerTokenScaleType,
|
||||
DType, TileRowBroadcastMul, TileBroadcastOneBlk,
|
||||
TileOneBlkColumnBroadcastMul, TileCopy, TileScheduler>;
|
||||
|
||||
using BlockScheduler = BlockScheduler_;
|
||||
|
||||
// kernel level
|
||||
using ElementGroupList = int64_t;
|
||||
|
||||
using GemmKernel = typename std::conditional<
|
||||
(EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) != 0,
|
||||
Gemm::Kernel::GroupedMatmulSliceMPerTokenDequantSwigluQuantMultiStageWorkspace<
|
||||
TemplateMC2TypeFunc, BlockMmad, BlockEpilogue, BlockScheduler, WORKSPACE_STAGES, ElementGroupList>,
|
||||
Gemm::Kernel::GroupedMatmulSliceMPerTokenDequantSwigluQuantMultiStageWorkspaceWithShallowDispatch<
|
||||
TemplateMC2TypeFunc, BlockMmad, BlockEpilogue, BlockScheduler, WORKSPACE_STAGES, ElementGroupList>>::type;
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
typename GemmKernel::Params params{problemShape,
|
||||
groupCount,
|
||||
gmGroupList,
|
||||
gmA,
|
||||
layoutA,
|
||||
gmB,
|
||||
layoutB,
|
||||
gmScale,
|
||||
layoutScale,
|
||||
gmPerTokenScale,
|
||||
layoutPerTokenScale,
|
||||
gmD,
|
||||
layoutD,
|
||||
gmDequantScale,
|
||||
layoutDequantScale,
|
||||
gmWorkspace,
|
||||
gmX,
|
||||
debugGm,
|
||||
gmexpertIds,
|
||||
gmExpandIdx,
|
||||
gmEpSendCount,
|
||||
xActiveMask,
|
||||
gmResvered,
|
||||
gmExpertTokenNums,
|
||||
epRankSize,
|
||||
epRankId,
|
||||
moeExpertNum,
|
||||
moeExpertNumPerRank,
|
||||
sharedExpertNum,
|
||||
sharedExpertRankNum,
|
||||
quantMode,
|
||||
globalBs,
|
||||
bs,
|
||||
topK,
|
||||
tokenLen};
|
||||
// call a kernel
|
||||
GemmKernel gemm;
|
||||
gemm(params);
|
||||
} else {
|
||||
typename GemmKernel::Params params{problemShape,
|
||||
groupCount,
|
||||
gmGroupList,
|
||||
gmA,
|
||||
layoutA,
|
||||
gmB,
|
||||
layoutB,
|
||||
gmScale,
|
||||
layoutScale,
|
||||
gmPerTokenScale,
|
||||
layoutPerTokenScale,
|
||||
gmD,
|
||||
layoutD,
|
||||
gmDequantScale,
|
||||
layoutDequantScale,
|
||||
gmWorkspace};
|
||||
// call a kernel
|
||||
GemmKernel gemm;
|
||||
gemm(params);
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass, class L1TileShape_, class L0TileShape_, class EpilogueTileShape_, class BlockScheduler_,
|
||||
class DispatchPolicy_ = MmadAtlasA2Custom>
|
||||
CATLASS_DEVICE void GmmDeq(GemmCoord problemShape, uint32_t groupCount, GM_ADDR gmGroupList, GM_ADDR gmA,
|
||||
layout::RowMajor layoutA, GM_ADDR gmB,
|
||||
typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type layoutB,
|
||||
GM_ADDR gmScale,
|
||||
layout::VectorLayout layoutScale, GM_ADDR gmPerTokenScale,
|
||||
layout::VectorLayout layoutPerTokenScale, GM_ADDR gmD, layout::RowMajor layoutD,
|
||||
GM_ADDR gmWorkspace, void *combiner)
|
||||
{
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
using DispatchPolicy = DispatchPolicy_;
|
||||
using L1TileShape = L1TileShape_;
|
||||
using L0TileShape = L0TileShape_;
|
||||
|
||||
using AType = Gemm::GemmType<int8_t, layout::RowMajor>;
|
||||
using LayoutB = typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type;
|
||||
using BType = Gemm::GemmType<int8_t, LayoutB>;
|
||||
using CType = Gemm::GemmType<int32_t, layout::RowMajor>;
|
||||
|
||||
using BlockMmad = Gemm::Block::BlockMmad<DispatchPolicy, L1TileShape, L0TileShape, AType, BType, CType>;
|
||||
|
||||
constexpr uint32_t ubStages = 1;
|
||||
using EpilogueDispatchPolicy = Epilogue::EpilogueAtlasA2PerTokenDequantCombine<ubStages, EXEC_FLAG>;
|
||||
using ScaleType = Gemm::GemmType<W2ScaleType, layout::VectorLayout>;
|
||||
using PerTokenScaleType = Gemm::GemmType<float, layout::VectorLayout>;
|
||||
using DType = Gemm::GemmType<ExpandXType, layout::RowMajor>;
|
||||
|
||||
using RowBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
using BroadcastOneBlkType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
using OneBlkColumnBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
|
||||
using EpilogueTileShape = EpilogueTileShape_;
|
||||
using TileRowBroadcastMul = Epilogue::Tile::TileRowBroadcastMul<ArchTag, RowBroadcastMulType, EpilogueTileShape>;
|
||||
using TileBroadcastOneBlk =
|
||||
Epilogue::Tile::TileBroadcastOneBlk<ArchTag, BroadcastOneBlkType, EpilogueTileShape::ROW>;
|
||||
using TileOneBlkColumnBroadcastMul =
|
||||
Epilogue::Tile::TileOneBlkColumnBroadcastMul<ArchTag, OneBlkColumnBroadcastMulType, EpilogueTileShape>;
|
||||
using TileCopy = Epilogue::Tile::TileCopy<ArchTag, CType, ScaleType, PerTokenScaleType, DType>;
|
||||
using TileScheduler = Epilogue::Tile::EpilogueHorizontalTileSwizzle;
|
||||
|
||||
using BlockEpilogue = Epilogue::Block::BlockEpilogue<EpilogueDispatchPolicy, CType, ScaleType, PerTokenScaleType,
|
||||
DType, TileRowBroadcastMul, TileBroadcastOneBlk,
|
||||
TileOneBlkColumnBroadcastMul, TileCopy, TileScheduler>;
|
||||
|
||||
using BlockScheduler = BlockScheduler_;
|
||||
|
||||
// kernel level
|
||||
using ElementGroupList = int64_t;
|
||||
using GemmKernel = Gemm::Kernel::GroupedMatmulSliceMPerTokenDequantMultiStageWorkspace<
|
||||
TemplateMC2TypeFunc, BlockMmad, BlockEpilogue, BlockScheduler, WORKSPACE_STAGES, ElementGroupList>;
|
||||
|
||||
typename GemmKernel::Params params{
|
||||
problemShape, groupCount, gmGroupList, gmA, layoutA, gmB, layoutB, gmScale,
|
||||
layoutScale, gmPerTokenScale, layoutPerTokenScale, gmD, layoutD, gmWorkspace, combiner};
|
||||
|
||||
// call a kernel
|
||||
GemmKernel gemm;
|
||||
gemm(params);
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
class DispatchGmmCombineDecode
|
||||
{
|
||||
public:
|
||||
__aicore__ inline DispatchGmmCombineDecode(){};
|
||||
__aicore__ inline void Init(
|
||||
// input
|
||||
GM_ADDR x, GM_ADDR expert_ids, GM_ADDR gmm1_permuted_weight, GM_ADDR gmm1_permuted_weight_scale,
|
||||
GM_ADDR gmm2_weight, GM_ADDR gmm2_weight_scale, GM_ADDR expert_scales, GM_ADDR expert_smooth_scales, GM_ADDR x_active_mask,
|
||||
// output
|
||||
GM_ADDR output, GM_ADDR expertTokenNums,
|
||||
// system
|
||||
GM_ADDR workspaceGM, AscendC::TPipe *pipe, const DispatchGmmCombineDecodeTilingData *tilingData);
|
||||
__aicore__ inline void Process();
|
||||
|
||||
private:
|
||||
GM_ADDR gmX_;
|
||||
GM_ADDR gmexpertIds_;
|
||||
GM_ADDR gmPermuteWeight1_;
|
||||
GM_ADDR gmPermuteScale1_;
|
||||
GM_ADDR gmWeight2_;
|
||||
GM_ADDR gmScale2_;
|
||||
GM_ADDR gmOutput_;
|
||||
GM_ADDR gmExpertTokenNums_;
|
||||
GM_ADDR workspaceGM_;
|
||||
GM_ADDR gmSmoothScales_;
|
||||
GM_ADDR gmexpertScales_;
|
||||
GM_ADDR xActiveMask_;
|
||||
|
||||
uint32_t maxTokenNum_{0};
|
||||
uint32_t gmm1OutputDim_{0};
|
||||
uint32_t tokenHiddenSize_{0};
|
||||
uint32_t groupCount_{0};
|
||||
uint32_t gmm2OutputDim_{0};
|
||||
uint32_t gmm2InputDim_{0};
|
||||
uint32_t globalRankId_{0};
|
||||
uint32_t winSizePerRank_{0};
|
||||
uint32_t blockDim_{0};
|
||||
uint32_t epRankSize_{0};
|
||||
uint32_t epRankId_{0};
|
||||
uint32_t moeExpertNum_{0};
|
||||
uint32_t moeExpertNumPerRank_{0};
|
||||
uint32_t sharedExpertNum_{0};
|
||||
uint32_t sharedExpertRankNum_{0};
|
||||
uint32_t quantMode_{0};
|
||||
uint32_t globalBs_{0};
|
||||
uint32_t bs_{0};
|
||||
uint32_t maxBs_{0};
|
||||
uint32_t topK_{0};
|
||||
|
||||
AscendC::TPipe *tpipe_{nullptr};
|
||||
__gm__ HcclOpResParam *winContext_{nullptr};
|
||||
const DispatchGmmCombineDecodeTilingData *tilingData_;
|
||||
};
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void DispatchGmmCombineDecode<TemplateMC2TypeFunc>::Init(
|
||||
// input
|
||||
GM_ADDR x, GM_ADDR expert_ids, GM_ADDR gmm1_permuted_weight, GM_ADDR gmm1_permuted_weight_scale,
|
||||
GM_ADDR gmm2_weight, GM_ADDR gmm2_weight_scale, GM_ADDR expert_scales, GM_ADDR expert_smooth_scales,
|
||||
GM_ADDR x_active_mask,
|
||||
// output
|
||||
GM_ADDR output, GM_ADDR expertTokenNums,
|
||||
// system
|
||||
GM_ADDR workspaceGM, AscendC::TPipe *pipe, const DispatchGmmCombineDecodeTilingData *tilingData)
|
||||
{
|
||||
tpipe_ = pipe;
|
||||
blockDim_ = AscendC::GetBlockNum();
|
||||
winContext_ = (__gm__ HcclOpResParam *)AscendC::GetHcclContext<AscendC::HCCL_GROUP_ID_0>();
|
||||
|
||||
gmSmoothScales_ = expert_smooth_scales; // not used now
|
||||
gmX_ = x; // input token
|
||||
gmexpertIds_ = expert_ids;
|
||||
gmPermuteWeight1_ = gmm1_permuted_weight;
|
||||
gmPermuteScale1_ = gmm1_permuted_weight_scale;
|
||||
gmWeight2_ = gmm2_weight;
|
||||
gmScale2_ = gmm2_weight_scale;
|
||||
gmOutput_ = output;
|
||||
gmExpertTokenNums_ = expertTokenNums;
|
||||
workspaceGM_ = workspaceGM;
|
||||
gmexpertScales_ = expert_scales;
|
||||
xActiveMask_ = x_active_mask;
|
||||
tilingData_ = tilingData;
|
||||
epRankSize_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.epRankSize;
|
||||
epRankId_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.epRankId;
|
||||
moeExpertNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNum;
|
||||
moeExpertNumPerRank_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
sharedExpertNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertNum;
|
||||
sharedExpertRankNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertRankNum;
|
||||
quantMode_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.quantMode;
|
||||
globalBs_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.globalBs;
|
||||
bs_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.bs;
|
||||
topK_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.k;
|
||||
maxBs_ = globalBs_ / epRankSize_;
|
||||
|
||||
bool isShareExpert = (epRankId_ < sharedExpertRankNum_);
|
||||
if (isShareExpert) {
|
||||
maxTokenNum_ = maxBs_ * epRankSize_ / sharedExpertRankNum_;
|
||||
} else {
|
||||
maxTokenNum_ = maxBs_ * epRankSize_ * (topK_ < moeExpertNumPerRank_ ? topK_ : moeExpertNumPerRank_);
|
||||
}
|
||||
|
||||
gmm1OutputDim_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen;
|
||||
tokenHiddenSize_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
groupCount_ = isShareExpert ? 1 : tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
gmm2OutputDim_ = tokenHiddenSize_;
|
||||
gmm2InputDim_ = gmm1OutputDim_ / 2;
|
||||
}
|
||||
|
||||
template<uint32_t EXEC_FLAG>
|
||||
__aicore__ inline auto CreateWeightLayout(uint32_t k, uint32_t n) {
|
||||
if constexpr ((EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0) {
|
||||
MatrixCoord mc{k, n};
|
||||
return layout::RowMajor::template MakeLayoutInUb<int8_t>(mc);
|
||||
} else {
|
||||
return layout::zN::template MakeLayout<int8_t>(k, n);
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void DispatchGmmCombineDecode<TemplateMC2TypeFunc>::Process()
|
||||
{
|
||||
GemmCoord gmm1ProblemShape{maxTokenNum_, gmm1OutputDim_, tokenHiddenSize_};
|
||||
GemmCoord gmm2ProblemShape{maxTokenNum_, gmm2OutputDim_, gmm2InputDim_};
|
||||
|
||||
layout::RowMajor layoutX1{maxTokenNum_, tokenHiddenSize_};
|
||||
auto layoutWeight1 = CreateWeightLayout<EXEC_FLAG>(tokenHiddenSize_, gmm1OutputDim_);
|
||||
layout::VectorLayout layoutW1Scale{gmm1OutputDim_};
|
||||
layout::VectorLayout layoutX1Scale{maxTokenNum_};
|
||||
layout::RowMajor layoutX2{maxTokenNum_, gmm2InputDim_};
|
||||
auto layoutWeight2 = CreateWeightLayout<EXEC_FLAG>(gmm2InputDim_, gmm2OutputDim_);
|
||||
layout::VectorLayout layoutW2Scale{gmm2OutputDim_};
|
||||
layout::VectorLayout layoutX2Scale{maxTokenNum_};
|
||||
layout::RowMajor layoutOutput{maxTokenNum_, gmm2OutputDim_};
|
||||
|
||||
size_t workspaceOffset = 0;
|
||||
constexpr int32_t resveredWorkSpaceSize = 256 * 1024;
|
||||
int64_t x1TokenSize = maxTokenNum_ * tokenHiddenSize_ * sizeof(int8_t);
|
||||
int64_t x2TokenSize = maxTokenNum_ * gmm2InputDim_ * sizeof(int8_t);
|
||||
int64_t maxTokenSize = x1TokenSize < x2TokenSize ? x2TokenSize : x1TokenSize;
|
||||
int64_t tokenScaleSize = maxTokenNum_ * sizeof(float);
|
||||
GM_ADDR gmX1 = workspaceGM_ + workspaceOffset;
|
||||
GM_ADDR gmX2 = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(maxTokenSize);
|
||||
GM_ADDR gmX1Scale = workspaceGM_ + workspaceOffset;
|
||||
GM_ADDR gmX2Scale = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(tokenScaleSize);
|
||||
GM_ADDR gmWorkspace = workspaceGM_ + workspaceOffset;
|
||||
GM_ADDR gmCVSwap = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(static_cast<size_t>(blockDim_) * (GMM1_L1M * GMM1_L1N) *
|
||||
WORKSPACE_STAGES * sizeof(int32_t));
|
||||
int64_t swigluOutSize = maxTokenNum_ * gmm1OutputDim_ * sizeof(float);
|
||||
int64_t gmm2OutSize = maxTokenNum_ * tokenHiddenSize_ * sizeof(ExpandXType);
|
||||
int64_t maxSwigluGmm2Size = swigluOutSize < gmm2OutSize ? gmm2OutSize : swigluOutSize;
|
||||
GM_ADDR gmSwigluOut = workspaceGM_ + workspaceOffset;
|
||||
GM_ADDR gmGmm2DepOut = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(maxSwigluGmm2Size);
|
||||
GM_ADDR gmGroupList = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(static_cast<size_t>(groupCount_) * sizeof(int64_t));
|
||||
GM_ADDR gmExpandIdx = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(static_cast<size_t>(bs_) * topK_ * sizeof(int32_t));
|
||||
GM_ADDR gmEpSendCount = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(static_cast<size_t>(epRankSize_) * groupCount_ * sizeof(int32_t));
|
||||
GM_ADDR gmResvered = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(resveredWorkSpaceSize);
|
||||
|
||||
if constexpr ((EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) == 0) {
|
||||
if constexpr (g_coreType == AscendC::AIV) {
|
||||
AscendC::TPipe tpipe;
|
||||
MoeDistributeDispatchImpl::CamMoeDistributeDispatch<ExpandXType, int8_t, false, true, false, false, EXEC_FLAG>
|
||||
dispatcher;
|
||||
dispatcher.Init(gmX_, gmexpertIds_, gmSmoothScales_, xActiveMask_, gmX1, gmX1Scale, gmExpandIdx, gmGroupList,
|
||||
gmEpSendCount, gmExpertTokenNums_, nullptr, gmWorkspace, &tpipe, tilingData_);
|
||||
dispatcher.Process();
|
||||
tpipe.Destroy();
|
||||
icache_preload(8);
|
||||
}
|
||||
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
Arch::CrossCoreFlag gmm1AivFinished{0};
|
||||
if constexpr (g_coreType == AscendC::AIV) {
|
||||
Arch::CrossCoreBarrier<0x0, PIPE_MTE3>();
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(gmm1AivFinished);
|
||||
} else {
|
||||
Arch::CrossCoreWaitFlag(gmm1AivFinished);
|
||||
}
|
||||
}
|
||||
GmmDeqSwigluQuant<TemplateMC2TypeFunc, Gmm1L1TileShape, Gmm1L0TileShape, Gmm1EpilogueTileShape,
|
||||
Gmm1BlockScheduler>(
|
||||
gmm1ProblemShape, groupCount_, gmGroupList, gmX1, layoutX1, gmPermuteWeight1_, layoutWeight1,
|
||||
gmPermuteScale1_, layoutW1Scale, gmX1Scale, layoutX1Scale, gmX2, layoutX2, gmX2Scale,
|
||||
layoutX2Scale, gmWorkspace, gmX_, gmSmoothScales_, gmexpertIds_, gmExpandIdx, gmEpSendCount, xActiveMask_, gmResvered,
|
||||
gmExpertTokenNums_, epRankSize_, epRankId_, moeExpertNum_, moeExpertNumPerRank_, sharedExpertNum_,
|
||||
sharedExpertRankNum_, quantMode_, globalBs_, bs_, topK_, tokenHiddenSize_);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
Arch::CrossCoreFlag gmm1AivFinished{0};
|
||||
if constexpr (g_coreType == AscendC::AIV) {
|
||||
Arch::CrossCoreBarrier<0x0, PIPE_MTE3>();
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(gmm1AivFinished);
|
||||
} else {
|
||||
Arch::CrossCoreWaitFlag(gmm1AivFinished);
|
||||
}
|
||||
|
||||
MoeDistributeCombineImpl::CamMoeDistributeCombine<TemplateMC2TypeFunc> combiner;
|
||||
if (g_coreType == AscendC::AIV) {
|
||||
combiner.Init(gmGmm2DepOut, gmexpertIds_, gmExpandIdx, gmEpSendCount, nullptr, gmexpertScales_, xActiveMask_, gmOutput_,
|
||||
workspaceGM_, nullptr, tilingData_);
|
||||
}
|
||||
GmmDeq<TemplateMC2TypeFunc, Gmm2L1TileShape, Gmm2L0TileShape, Gmm2EpilogueTileShape, Gmm2BlockScheduler,
|
||||
Gmm2DispatchPolicy>(gmm2ProblemShape, groupCount_, gmGroupList, gmX2, layoutX2, gmWeight2_, layoutWeight2,
|
||||
gmScale2_, layoutW2Scale, gmX2Scale, layoutX2Scale, gmGmm2DepOut,
|
||||
layoutOutput, gmWorkspace, &combiner);
|
||||
}
|
||||
} // namespace DispatchGmmCombineDecodeImpl
|
||||
#endif // DISPATCH_GMM_COMBINE_DECODE_H
|
||||
@@ -0,0 +1,17 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/epilogue/block/block_epilogue.hpp"
|
||||
|
||||
#include "block_epilogue_per_token_dequant_swiglu.h"
|
||||
#include "block_epilogue_per_token_dequant.hpp"
|
||||
|
||||
#include "block_epilogue_swiglu_bf16_fp16.h"
|
||||
#include "block_epilogue_bf16_fp16.hpp"
|
||||
@@ -0,0 +1,337 @@
|
||||
/*
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef ACT_EPILOGUE_BLOCK_EPILOGUE_BF16_FP16_HPP
|
||||
#define ACT_EPILOGUE_BLOCK_EPILOGUE_BF16_FP16_HPP
|
||||
|
||||
#include "../../raw_distributed/cam_moe_distribute_combine.h"
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
#include "catlass/epilogue/dispatch_policy.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <uint32_t UB_STAGES_, uint32_t EXEC_FLAG_,
|
||||
class CType_, class ScaleType_, class LayoutScale_, class LayoutPerTokenScale_, class DType_,
|
||||
class TileRowBroadcastMul_, class TileBroadcastOneBlk_, class TileOneBlkColumnBroadcastMul_,
|
||||
class TileCopy_, class EpilogueTileSwizzle_>
|
||||
class BlockEpilogue<EpilogueAtlasA2Combine<UB_STAGES_, EXEC_FLAG_>,
|
||||
CType_, Gemm::GemmType<ScaleType_, LayoutScale_>, Gemm::GemmType<float, LayoutPerTokenScale_>, DType_,
|
||||
TileRowBroadcastMul_, TileBroadcastOneBlk_, TileOneBlkColumnBroadcastMul_,
|
||||
TileCopy_, EpilogueTileSwizzle_>
|
||||
{
|
||||
public:
|
||||
using DispatchPolicy = EpilogueAtlasA2Combine<UB_STAGES_, EXEC_FLAG_>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
static constexpr uint32_t EXEC_FLAG = EXEC_FLAG_;
|
||||
|
||||
// Data infos
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using ElementRawScale = ScaleType_;
|
||||
using ElementFp32Scale = float;
|
||||
using LayoutScale = LayoutScale_;
|
||||
using ElementPerTokenScale = float;
|
||||
using LayoutPerTokenScale = LayoutPerTokenScale_;
|
||||
using ElementD = typename DType_::Element;
|
||||
using LayoutD = typename DType_::Layout;
|
||||
|
||||
// Check data infos
|
||||
static_assert(std::is_same_v<ElementC, float> &&
|
||||
(std::is_same_v<ElementD, half> || std::is_same_v<ElementD, bfloat16_t>),
|
||||
"The element type template parameters of BlockEpilogue are wrong");
|
||||
static_assert(std::is_same_v<LayoutC, layout::RowMajor> && std::is_same_v<LayoutScale, layout::VectorLayout> &&
|
||||
std::is_same_v<LayoutPerTokenScale, layout::VectorLayout> &&
|
||||
std::is_same_v<LayoutD, layout::RowMajor>,
|
||||
"The layout template parameters of BlockEpilogue are wrong");
|
||||
|
||||
// Tile compute ops
|
||||
using TileRowBroadcastMul = TileRowBroadcastMul_;
|
||||
using TileBroadcastOneBlk = TileBroadcastOneBlk_;
|
||||
using TileOneBlkColumnBroadcastMul = TileOneBlkColumnBroadcastMul_;
|
||||
|
||||
// Tile copy
|
||||
using CopyGmToUbC = typename TileCopy_::CopyGmToUbC;
|
||||
using CopyGmToUbScale = typename TileCopy_::CopyGmToUbX;
|
||||
using CopyGmToUbPerTokenScale = typename TileCopy_::CopyGmToUbY;
|
||||
using CopyUbToGmD = typename TileCopy_::CopyUbToGmD;
|
||||
|
||||
using EpilogueTileSwizzle = EpilogueTileSwizzle_;
|
||||
|
||||
using TileShape = typename TileRowBroadcastMul::TileShape;
|
||||
|
||||
static_assert(TileShape::ROW == TileBroadcastOneBlk::COMPUTE_LENGTH &&
|
||||
std::is_same_v<TileShape, typename TileOneBlkColumnBroadcastMul::TileShape>,
|
||||
"TileShape must be consistent for all tile compute ops");
|
||||
|
||||
static_assert((UB_STAGES * (TileShape::COUNT * sizeof(ElementC) + TileShape::COUNT * sizeof(ElementD)) +
|
||||
TileShape::ROW * BYTE_PER_BLK) <= ArchTag::UB_SIZE,
|
||||
"TileShape is too large to fit in UB");
|
||||
|
||||
struct Params {
|
||||
__gm__ ElementRawScale *ptrScale{nullptr};
|
||||
LayoutScale layoutScale{};
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale{nullptr};
|
||||
LayoutPerTokenScale layoutPerTokenScale{};
|
||||
__gm__ ElementD *ptrD{nullptr};
|
||||
LayoutD layoutD{};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params() {};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params(__gm__ ElementRawScale *ptrScale_, LayoutScale const &layoutScale_,
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale_, LayoutPerTokenScale const &layoutPerTokenScale_,
|
||||
__gm__ ElementD *ptrD_, LayoutD const &layoutD_)
|
||||
: ptrScale(ptrScale_),
|
||||
layoutScale(layoutScale_),
|
||||
ptrPerTokenScale(ptrPerTokenScale_),
|
||||
layoutPerTokenScale(layoutPerTokenScale_),
|
||||
ptrD(ptrD_),
|
||||
layoutD(layoutD_)
|
||||
{}
|
||||
};
|
||||
|
||||
CATLASS_DEVICE void AlignUbOffset()
|
||||
{
|
||||
size_t ubMask = ubOffset & (MoeDistributeCombineImpl::UB_ALIGN - 1);
|
||||
if (ubMask != 0) {
|
||||
ubOffset += MoeDistributeCombineImpl::UB_ALIGN - ubMask;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> &resource, MoeDistributeCombineImpl::CombineCalcInfo &calcInfo,
|
||||
Params const ¶ms = Params{})
|
||||
: resource(resource), calcInfo(calcInfo), params(params)
|
||||
{
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(ElementC);
|
||||
ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(ElementD);
|
||||
|
||||
eventUbCVMTE2List[i] = eventVMTE2++;
|
||||
eventUbCMTE2VList[i] = eventMTE2V++;
|
||||
eventUbDMTE3VList[i] = eventMTE3V++;
|
||||
eventUbDVMTE3List[i] = eventVMTE3++;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
AlignUbOffset();
|
||||
epSendCountLocal_ = resource.ubBuf.template GetBufferByByte<int32_t>(ubOffset);
|
||||
ubOffset += calcInfo.moeSendNum_ * sizeof(int32_t);
|
||||
AlignUbOffset();
|
||||
AscendC::GlobalTensor<int32_t> epSendCountGM;
|
||||
epSendCountGM.SetGlobalBuffer((__gm__ int32_t *)calcInfo.epSendCount_);
|
||||
uint32_t epSendCountSize = calcInfo.isShardExpert_ ? calcInfo.epWorldSize_ : calcInfo.moeSendNum_;
|
||||
AscendC::DataCopyExtParams epSendCntParams = {1U, static_cast<uint32_t>(epSendCountSize * sizeof(uint32_t)),
|
||||
0U, 0U, 0U};
|
||||
AscendC::DataCopyPadExtParams<int32_t> copyPadParams{false, 0U, 0U, 0U};
|
||||
AscendC::DataCopyPad(epSendCountLocal_, epSendCountGM, epSendCntParams, copyPadParams);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_S>(eventMTE2S);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_S>(eventMTE2S);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void UpdateParams(Params const ¶ms_)
|
||||
{
|
||||
params = params_;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE GM_ADDR GetWinAddrByRankId(const int32_t rankId, const uint8_t expertLocalId = 0U)
|
||||
{
|
||||
return (GM_ADDR)((calcInfo.epRankId_ == rankId)
|
||||
? calcInfo.epWinContext_->localWindowsIn
|
||||
: ((HcclRankRelationResV2 *)(calcInfo.epWinContext_->remoteRes[rankId].nextDevicePtr))
|
||||
->windowsIn) +
|
||||
calcInfo.winDataSizeOffset_ + expertLocalId * calcInfo.expertPerSizeOnWin_ + rankId * OPT_RANK_OFFSET;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE void SetCombineSendEpRank(uint32_t epRank, uint32_t &remoteEpRank, uint32_t &localEpRank)
|
||||
{
|
||||
if ((calcInfo.isShardExpert_) && (epRank < calcInfo.sharedExpertRankNum_)) {
|
||||
remoteEpRank = calcInfo.epRankId_;
|
||||
localEpRank = epRank;
|
||||
} else {
|
||||
remoteEpRank = epRank;
|
||||
localEpRank = calcInfo.epRankId_;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE void DoCombineSend(AscendC::LocalTensor<ElementD> &ubD, layout::RowMajor &layoutGmTileD,
|
||||
LayoutD &layoutUbD, int64_t groupOffsetD, uint32_t expertIdx, uint32_t tileOffsetD)
|
||||
{
|
||||
const uint32_t copyTokenLen = layoutGmTileD.shape(1) * sizeof(ElementD);
|
||||
const uint32_t copyTokenSrcStride =
|
||||
(layoutUbD.stride(0) - layoutUbD.shape(1)) / (BYTE_PER_C0 / sizeof(ElementD));
|
||||
const uint32_t copyTokenDstStride = (layoutGmTileD.stride(0) - layoutGmTileD.shape(1)) * sizeof(ElementD);
|
||||
|
||||
int64_t offsetD = groupOffsetD + tileOffsetD;
|
||||
uint32_t startToken = offsetD / calcInfo.axisH_;
|
||||
uint32_t tokenOffset = offsetD - startToken * calcInfo.axisH_;
|
||||
uint32_t itToken = startToken;
|
||||
uint32_t endToken = startToken + layoutGmTileD.shape(0);
|
||||
constexpr uint32_t epRankStart = 0;
|
||||
uint32_t sendCount =
|
||||
expertIdx == 0 && epRankStart == 0 ? 0 : epSendCountLocal_.GetValue(expertOffset + epRankStart - 1);
|
||||
for (uint32_t epRank = epRankStart; epRank < calcInfo.epWorldSize_ && itToken < endToken; ++epRank) {
|
||||
uint32_t prevSendCount = sendCount;
|
||||
sendCount = epSendCountLocal_.GetValue(expertOffset + epRank);
|
||||
if (prevSendCount <= itToken && itToken < sendCount) {
|
||||
uint32_t copyTokenCount = (sendCount < endToken ? sendCount : endToken) - itToken;
|
||||
AscendC::DataCopyExtParams dataCopyParams(copyTokenCount, copyTokenLen, copyTokenSrcStride,
|
||||
copyTokenDstStride, 0);
|
||||
uint32_t remoteEpRank;
|
||||
uint32_t localEpRank;
|
||||
SetCombineSendEpRank(epRank, remoteEpRank, localEpRank);
|
||||
GM_ADDR rankGM = GetWinAddrByRankId(remoteEpRank, expertIdx) +
|
||||
localEpRank * calcInfo.moeExpertPerRankNum_ * calcInfo.expertPerSizeOnWin_;
|
||||
AscendC::GlobalTensor<ElementD> rankWindow;
|
||||
rankWindow.SetGlobalBuffer((__gm__ ElementD *)rankGM);
|
||||
AscendC::DataCopyPad(rankWindow[(itToken - prevSendCount) * calcInfo.axisH_ + tokenOffset],
|
||||
ubD[(itToken - startToken) * layoutUbD.stride(0)], dataCopyParams);
|
||||
itToken += copyTokenCount;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(int64_t groupOffsetD, uint32_t expertIdx, GemmCoord const &blockShapeMNK,
|
||||
GemmCoord const &blockCoordMNK, GemmCoord const &actualBlockShapeMNK,
|
||||
AscendC::GlobalTensor<ElementC> const &gmBlockC, LayoutC const &layoutBlockC,
|
||||
Callback &&callback = Callback{})
|
||||
{
|
||||
if (actualBlockShapeMNK.k() == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
expertOffset = expertIdx * calcInfo.epWorldSize_;
|
||||
}
|
||||
|
||||
callback();
|
||||
// Calculate the offset of the current block
|
||||
MatrixCoord blockShape = blockShapeMNK.GetCoordMN();
|
||||
MatrixCoord blockCoord = blockCoordMNK.GetCoordMN();
|
||||
MatrixCoord actualBlockShape = actualBlockShapeMNK.GetCoordMN();
|
||||
MatrixCoord blockOffset = blockCoord * blockShape;
|
||||
|
||||
AscendC::GlobalTensor<ElementRawScale> gmScale;
|
||||
gmScale.SetGlobalBuffer(params.ptrScale);
|
||||
AscendC::GlobalTensor<ElementPerTokenScale> gmPerTokenScale;
|
||||
gmPerTokenScale.SetGlobalBuffer(params.ptrPerTokenScale);
|
||||
AscendC::GlobalTensor<ElementD> gmD;
|
||||
gmD.SetGlobalBuffer(params.ptrD);
|
||||
|
||||
auto ubTileStride = MakeCoord(static_cast<int64_t>(TileShape::COLUMN), 1L);
|
||||
auto tileShape = TileShape::ToCoord();
|
||||
EpilogueTileSwizzle epilogueTileSwizzle(actualBlockShape, tileShape);
|
||||
uint32_t tileLoops = epilogueTileSwizzle.GetLoops();
|
||||
uint32_t subblockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subblockNum = AscendC::GetSubBlockNum();
|
||||
for (uint32_t loopIdx = subblockIdx; loopIdx < tileLoops; loopIdx += subblockNum) {
|
||||
auto tileCoord = epilogueTileSwizzle.GetTileCoord(loopIdx);
|
||||
auto actualTileShape = epilogueTileSwizzle.GetActualTileShape(tileCoord);
|
||||
auto tileOffsetInBlock = tileCoord * tileShape;
|
||||
auto tileOffset = blockOffset + tileOffsetInBlock;
|
||||
|
||||
auto gmTileC = gmBlockC[layoutBlockC.GetOffset(tileOffsetInBlock)];
|
||||
auto layoutGmTileC = layoutBlockC.GetTileLayout(actualTileShape);
|
||||
|
||||
auto &ubC = ubCList[ubListId];
|
||||
LayoutC layoutUbC{actualTileShape, ubTileStride};
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
copyGmToUbC(ubC, gmTileC, layoutUbC, layoutGmTileC);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
|
||||
auto &ubD = ubDList[ubListId];
|
||||
LayoutD layoutUbD{actualTileShape, ubTileStride};
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
AscendC::Cast(ubD, ubC, AscendC::RoundMode::CAST_RINT, TileShape::COUNT);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
|
||||
auto tileOffsetD = params.layoutD.GetOffset(tileOffset);
|
||||
auto layoutGmTileD = params.layoutD.GetTileLayout(actualTileShape);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
DoCombineSend(ubD, layoutGmTileD, layoutUbD, groupOffsetD, expertIdx, tileOffsetD);
|
||||
} else {
|
||||
auto gmTileD = gmD[tileOffsetD];
|
||||
copyUbToGmD(gmTileD, ubD, layoutGmTileD, layoutUbD);
|
||||
}
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
|
||||
ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
Params params;
|
||||
Arch::Resource<ArchTag> &resource;
|
||||
MoeDistributeCombineImpl::CombineCalcInfo calcInfo;
|
||||
|
||||
AscendC::LocalTensor<ElementC> ubCList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementD> ubDList[UB_STAGES];
|
||||
|
||||
int32_t eventUbCVMTE2List[UB_STAGES];
|
||||
int32_t eventUbCMTE2VList[UB_STAGES];
|
||||
int32_t eventUbScaleVMTE2List[UB_STAGES];
|
||||
int32_t eventUbScaleMTE2VList[UB_STAGES];
|
||||
int32_t eventUbPerTokenScaleVMTE2List[UB_STAGES];
|
||||
int32_t eventUbPerTokenScaleMTE2VList[UB_STAGES];
|
||||
int32_t eventUbDMTE3VList[UB_STAGES];
|
||||
int32_t eventUbDVMTE3List[UB_STAGES];
|
||||
|
||||
AscendC::LocalTensor<int32_t> epSendCountLocal_;
|
||||
|
||||
size_t ubOffset{0};
|
||||
int32_t eventVMTE2{0};
|
||||
int32_t eventMTE2V{0};
|
||||
int32_t eventMTE3V{0};
|
||||
int32_t eventVMTE3{0};
|
||||
int32_t eventVS{0};
|
||||
int32_t eventMTE2S{0};
|
||||
|
||||
uint32_t expertOffset;
|
||||
|
||||
uint32_t ubListId{0};
|
||||
|
||||
CopyGmToUbC copyGmToUbC;
|
||||
CopyUbToGmD copyUbToGmD;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue::Block
|
||||
|
||||
#endif // ACT_EPILOGUE_BLOCK_EPILOGUE_BF16_FP16_HPP
|
||||
@@ -0,0 +1,429 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef ACT_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_DEQUANT_HPP
|
||||
#define ACT_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_DEQUANT_HPP
|
||||
|
||||
#include "../../raw_distributed/cam_moe_distribute_combine.h"
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
#include "catlass/epilogue/dispatch_policy.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <uint32_t UB_STAGES_, uint32_t EXEC_FLAG_,
|
||||
class CType_, class ScaleType_, class LayoutScale_, class LayoutPerTokenScale_, class DType_,
|
||||
class TileRowBroadcastMul_, class TileBroadcastOneBlk_, class TileOneBlkColumnBroadcastMul_,
|
||||
class TileCopy_, class EpilogueTileSwizzle_>
|
||||
class BlockEpilogue<EpilogueAtlasA2PerTokenDequantCombine<UB_STAGES_, EXEC_FLAG_>,
|
||||
CType_, Gemm::GemmType<ScaleType_, LayoutScale_>, Gemm::GemmType<float, LayoutPerTokenScale_>, DType_,
|
||||
TileRowBroadcastMul_, TileBroadcastOneBlk_, TileOneBlkColumnBroadcastMul_,
|
||||
TileCopy_, EpilogueTileSwizzle_>
|
||||
{
|
||||
public:
|
||||
using DispatchPolicy = EpilogueAtlasA2PerTokenDequantCombine<UB_STAGES_, EXEC_FLAG_>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
static constexpr uint32_t EXEC_FLAG = EXEC_FLAG_;
|
||||
|
||||
// Data infos
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using ElementRawScale = ScaleType_;
|
||||
using ElementFp32Scale = float;
|
||||
using LayoutScale = LayoutScale_;
|
||||
using ElementPerTokenScale = float;
|
||||
using LayoutPerTokenScale = LayoutPerTokenScale_;
|
||||
using ElementD = typename DType_::Element;
|
||||
using LayoutD = typename DType_::Layout;
|
||||
|
||||
// Check data infos
|
||||
static_assert(std::is_same_v<ElementC, int32_t> &&
|
||||
(std::is_same_v<ElementD, half> || std::is_same_v<ElementD, bfloat16_t>),
|
||||
"The element type template parameters of BlockEpilogue are wrong");
|
||||
static_assert(std::is_same_v<LayoutC, layout::RowMajor> && std::is_same_v<LayoutScale, layout::VectorLayout> &&
|
||||
std::is_same_v<LayoutPerTokenScale, layout::VectorLayout> &&
|
||||
std::is_same_v<LayoutD, layout::RowMajor>,
|
||||
"The layout template parameters of BlockEpilogue are wrong");
|
||||
|
||||
// Tile compute ops
|
||||
using TileRowBroadcastMul = TileRowBroadcastMul_;
|
||||
using TileBroadcastOneBlk = TileBroadcastOneBlk_;
|
||||
using TileOneBlkColumnBroadcastMul = TileOneBlkColumnBroadcastMul_;
|
||||
|
||||
// Tile copy
|
||||
using CopyGmToUbC = typename TileCopy_::CopyGmToUbC;
|
||||
using CopyGmToUbScale = typename TileCopy_::CopyGmToUbX;
|
||||
using CopyGmToUbPerTokenScale = typename TileCopy_::CopyGmToUbY;
|
||||
using CopyUbToGmD = typename TileCopy_::CopyUbToGmD;
|
||||
|
||||
using EpilogueTileSwizzle = EpilogueTileSwizzle_;
|
||||
|
||||
using TileShape = typename TileRowBroadcastMul::TileShape;
|
||||
|
||||
static_assert(TileShape::ROW == TileBroadcastOneBlk::COMPUTE_LENGTH &&
|
||||
std::is_same_v<TileShape, typename TileOneBlkColumnBroadcastMul::TileShape>,
|
||||
"TileShape must be consistent for all tile compute ops");
|
||||
|
||||
static_assert((UB_STAGES * (TileShape::COUNT * sizeof(ElementC) +
|
||||
(std::is_same_v<ElementRawScale, ElementFp32Scale> ?
|
||||
0 : TileShape::COLUMN * sizeof(ElementRawScale)) +
|
||||
TileShape::COLUMN * sizeof(ElementFp32Scale) +
|
||||
TileShape::ROW * sizeof(ElementPerTokenScale) + TileShape::COUNT * sizeof(ElementD)) +
|
||||
(TileShape::COUNT + TileShape::COUNT) * sizeof(float) + TileShape::ROW * BYTE_PER_BLK) <=
|
||||
ArchTag::UB_SIZE,
|
||||
"TileShape is too large to fit in UB");
|
||||
struct Params {
|
||||
__gm__ ElementRawScale *ptrScale{nullptr};
|
||||
LayoutScale layoutScale{};
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale{nullptr};
|
||||
LayoutPerTokenScale layoutPerTokenScale{};
|
||||
__gm__ ElementD *ptrD{nullptr};
|
||||
LayoutD layoutD{};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params() {};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params(__gm__ ElementRawScale *ptrScale_, LayoutScale const &layoutScale_,
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale_, LayoutPerTokenScale const &layoutPerTokenScale_,
|
||||
__gm__ ElementD *ptrD_, LayoutD const &layoutD_)
|
||||
: ptrScale(ptrScale_),
|
||||
layoutScale(layoutScale_),
|
||||
ptrPerTokenScale(ptrPerTokenScale_),
|
||||
layoutPerTokenScale(layoutPerTokenScale_),
|
||||
ptrD(ptrD_),
|
||||
layoutD(layoutD_)
|
||||
{}
|
||||
};
|
||||
|
||||
CATLASS_DEVICE void AlignUbOffset()
|
||||
{
|
||||
size_t ubMask = ubOffset & (MoeDistributeCombineImpl::UB_ALIGN - 1);
|
||||
if (ubMask != 0) {
|
||||
ubOffset += MoeDistributeCombineImpl::UB_ALIGN - ubMask;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> &resource, MoeDistributeCombineImpl::CombineCalcInfo &calcInfo,
|
||||
Params const ¶ms = Params{})
|
||||
: resource(resource), calcInfo(calcInfo), params(params)
|
||||
{
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(ElementC);
|
||||
if constexpr (!std::is_same_v<ElementRawScale, ElementFp32Scale>) {
|
||||
ubRawScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementRawScale>(ubOffset);
|
||||
ubOffset += TileShape::COLUMN * sizeof(ElementRawScale);
|
||||
}
|
||||
ubFp32ScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementFp32Scale>(ubOffset);
|
||||
ubOffset += TileShape::COLUMN * sizeof(ElementFp32Scale);
|
||||
ubPerTokenScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementPerTokenScale>(ubOffset);
|
||||
ubOffset += TileShape::ROW * sizeof(ElementPerTokenScale);
|
||||
ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(ElementD);
|
||||
|
||||
eventUbCVMTE2List[i] = eventVMTE2++;
|
||||
eventUbCMTE2VList[i] = eventMTE2V++;
|
||||
eventUbScaleVMTE2List[i] = eventVMTE2++;
|
||||
eventUbScaleMTE2VList[i] = eventMTE2V++;
|
||||
eventUbPerTokenScaleVMTE2List[i] = eventVMTE2++;
|
||||
eventUbPerTokenScaleMTE2VList[i] = eventMTE2V++;
|
||||
eventUbDMTE3VList[i] = eventMTE3V++;
|
||||
eventUbDVMTE3List[i] = eventVMTE3++;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbScaleVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbPerTokenScaleVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
ubCFp32 = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(float);
|
||||
ubMul = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(float);
|
||||
ubPerTokenScaleBrcb = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += TileShape::ROW * BYTE_PER_BLK;
|
||||
ubPerTokenMul = ubCFp32;
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
AlignUbOffset();
|
||||
epSendCountLocal_ = resource.ubBuf.template GetBufferByByte<int32_t>(ubOffset);
|
||||
ubOffset += calcInfo.moeSendNum_ * sizeof(int32_t);
|
||||
AlignUbOffset();
|
||||
AscendC::GlobalTensor<int32_t> epSendCountGM;
|
||||
epSendCountGM.SetGlobalBuffer((__gm__ int32_t *)calcInfo.epSendCount_);
|
||||
uint32_t epSendCountSize = calcInfo.isShardExpert_ ? calcInfo.epWorldSize_ : calcInfo.moeSendNum_;
|
||||
AscendC::DataCopyExtParams epSendCntParams = {1U, static_cast<uint32_t>(epSendCountSize * sizeof(uint32_t)),
|
||||
0U, 0U, 0U};
|
||||
AscendC::DataCopyPadExtParams<int32_t> copyPadParams{false, 0U, 0U, 0U};
|
||||
AscendC::DataCopyPad(epSendCountLocal_, epSendCountGM, epSendCntParams, copyPadParams);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_S>(eventMTE2S);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_S>(eventMTE2S);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbScaleVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbPerTokenScaleVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void UpdateParams(Params const ¶ms_)
|
||||
{
|
||||
params = params_;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE GM_ADDR GetWinAddrByRankId(const int32_t rankId, const uint8_t expertLocalId = 0U)
|
||||
{
|
||||
return (GM_ADDR)((calcInfo.epRankId_ == rankId)
|
||||
? calcInfo.epWinContext_->localWindowsIn
|
||||
: ((HcclRankRelationResV2 *)(calcInfo.epWinContext_->remoteRes[rankId].nextDevicePtr))
|
||||
->windowsIn) +
|
||||
calcInfo.winDataSizeOffset_ + expertLocalId * calcInfo.expertPerSizeOnWin_ + rankId * OPT_RANK_OFFSET;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE void SetCombineSendEpRank(uint32_t epRank, uint32_t &remoteEpRank, uint32_t &localEpRank)
|
||||
{
|
||||
if ((calcInfo.isShardExpert_) && (epRank < calcInfo.sharedExpertRankNum_)) {
|
||||
remoteEpRank = calcInfo.epRankId_;
|
||||
localEpRank = epRank;
|
||||
} else {
|
||||
remoteEpRank = epRank;
|
||||
localEpRank = calcInfo.epRankId_;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE void DoCombineSend(AscendC::LocalTensor<ElementD> &ubD, layout::RowMajor &layoutGmTileD,
|
||||
LayoutD &layoutUbD, int64_t groupOffsetD, uint32_t expertIdx, uint32_t tileOffsetD)
|
||||
{
|
||||
const uint32_t copyTokenLen = layoutGmTileD.shape(1) * sizeof(ElementD);
|
||||
const uint32_t copyTokenSrcStride =
|
||||
(layoutUbD.stride(0) - layoutUbD.shape(1)) / (BYTE_PER_C0 / sizeof(ElementD));
|
||||
const uint32_t copyTokenDstStride = (layoutGmTileD.stride(0) - layoutGmTileD.shape(1)) * sizeof(ElementD);
|
||||
|
||||
int64_t offsetD = groupOffsetD + tileOffsetD;
|
||||
uint32_t startToken = offsetD / calcInfo.axisH_;
|
||||
uint32_t tokenOffset = offsetD - startToken * calcInfo.axisH_;
|
||||
uint32_t itToken = startToken;
|
||||
uint32_t endToken = startToken + layoutGmTileD.shape(0);
|
||||
constexpr uint32_t epRankStart = 0;
|
||||
uint32_t sendCount =
|
||||
expertIdx == 0 && epRankStart == 0 ? 0 : epSendCountLocal_.GetValue(expertOffset + epRankStart - 1);
|
||||
for (uint32_t epRank = epRankStart; epRank < calcInfo.epWorldSize_ && itToken < endToken; ++epRank) {
|
||||
uint32_t prevSendCount = sendCount;
|
||||
sendCount = epSendCountLocal_.GetValue(expertOffset + epRank);
|
||||
if (prevSendCount <= itToken && itToken < sendCount) {
|
||||
uint32_t copyTokenCount = (sendCount < endToken ? sendCount : endToken) - itToken;
|
||||
AscendC::DataCopyExtParams dataCopyParams(copyTokenCount, copyTokenLen, copyTokenSrcStride,
|
||||
copyTokenDstStride, 0);
|
||||
uint32_t remoteEpRank;
|
||||
uint32_t localEpRank;
|
||||
SetCombineSendEpRank(epRank, remoteEpRank, localEpRank);
|
||||
GM_ADDR rankGM = GetWinAddrByRankId(remoteEpRank, expertIdx) +
|
||||
localEpRank * calcInfo.moeExpertPerRankNum_ * calcInfo.expertPerSizeOnWin_;
|
||||
AscendC::GlobalTensor<ElementD> rankWindow;
|
||||
rankWindow.SetGlobalBuffer((__gm__ ElementD *)rankGM);
|
||||
AscendC::DataCopyPad(rankWindow[(itToken - prevSendCount) * calcInfo.axisH_ + tokenOffset],
|
||||
ubD[(itToken - startToken) * layoutUbD.stride(0)], dataCopyParams);
|
||||
itToken += copyTokenCount;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(int64_t groupOffsetD, uint32_t expertIdx, GemmCoord const &blockShapeMNK,
|
||||
GemmCoord const &blockCoordMNK, GemmCoord const &actualBlockShapeMNK,
|
||||
AscendC::GlobalTensor<ElementC> const &gmBlockC, LayoutC const &layoutBlockC,
|
||||
Callback &&callback = Callback{})
|
||||
{
|
||||
if (actualBlockShapeMNK.k() == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
expertOffset = expertIdx * calcInfo.epWorldSize_;
|
||||
}
|
||||
|
||||
callback();
|
||||
// Calculate the offset of the current block
|
||||
MatrixCoord blockShape = blockShapeMNK.GetCoordMN();
|
||||
MatrixCoord blockCoord = blockCoordMNK.GetCoordMN();
|
||||
MatrixCoord actualBlockShape = actualBlockShapeMNK.GetCoordMN();
|
||||
MatrixCoord blockOffset = blockCoord * blockShape;
|
||||
|
||||
AscendC::GlobalTensor<ElementRawScale> gmScale;
|
||||
gmScale.SetGlobalBuffer(params.ptrScale);
|
||||
AscendC::GlobalTensor<ElementPerTokenScale> gmPerTokenScale;
|
||||
gmPerTokenScale.SetGlobalBuffer(params.ptrPerTokenScale);
|
||||
AscendC::GlobalTensor<ElementD> gmD;
|
||||
gmD.SetGlobalBuffer(params.ptrD);
|
||||
|
||||
auto ubTileStride = MakeCoord(static_cast<int64_t>(TileShape::COLUMN), 1L);
|
||||
auto tileShape = TileShape::ToCoord();
|
||||
EpilogueTileSwizzle epilogueTileSwizzle(actualBlockShape, tileShape);
|
||||
uint32_t tileLoops = epilogueTileSwizzle.GetLoops();
|
||||
uint32_t subblockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subblockNum = AscendC::GetSubBlockNum();
|
||||
for (uint32_t loopIdx = subblockIdx; loopIdx < tileLoops; loopIdx += subblockNum) {
|
||||
auto tileCoord = epilogueTileSwizzle.GetTileCoord(loopIdx);
|
||||
auto actualTileShape = epilogueTileSwizzle.GetActualTileShape(tileCoord);
|
||||
auto tileOffsetInBlock = tileCoord * tileShape;
|
||||
auto tileOffset = blockOffset + tileOffsetInBlock;
|
||||
|
||||
auto gmTileC = gmBlockC[layoutBlockC.GetOffset(tileOffsetInBlock)];
|
||||
auto layoutGmTileC = layoutBlockC.GetTileLayout(actualTileShape);
|
||||
|
||||
auto &ubC = ubCList[ubListId];
|
||||
LayoutC layoutUbC{actualTileShape, ubTileStride};
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
copyGmToUbC(ubC, gmTileC, layoutUbC, layoutGmTileC);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
|
||||
auto scaleTileOffset = tileOffset.template GetCoordByAxis<1>();
|
||||
auto scaleTileShape = actualTileShape.template GetCoordByAxis<1>();
|
||||
|
||||
auto gmTileScale = gmScale[params.layoutScale.GetOffset(scaleTileOffset)];
|
||||
auto layoutGmTileScale = params.layoutScale.GetTileLayout(scaleTileShape);
|
||||
|
||||
auto &ubFp32Scale = ubFp32ScaleList[ubListId];
|
||||
auto layoutFp32UbScale = LayoutScale::template MakeLayoutInUb<ElementFp32Scale>(scaleTileShape);
|
||||
auto &ubRawScale = ubRawScaleList[ubListId];
|
||||
auto layoutRawUbScale = LayoutScale::template MakeLayoutInUb<ElementRawScale>(scaleTileShape);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbScaleVMTE2List[ubListId]);
|
||||
if constexpr (!std::is_same_v<ElementRawScale, ElementFp32Scale>) {
|
||||
copyGmToUbScale(ubRawScale, gmTileScale, layoutRawUbScale, layoutGmTileScale);
|
||||
} else {
|
||||
copyGmToUbScale(ubFp32Scale, gmTileScale, layoutFp32UbScale, layoutGmTileScale);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbScaleMTE2VList[ubListId]);
|
||||
|
||||
auto perTokenScaleTileOffset = tileOffset.template GetCoordByAxis<0>();
|
||||
auto perTokenScaleTileShape = actualTileShape.template GetCoordByAxis<0>();
|
||||
|
||||
auto gmTilePerTokenScale = gmPerTokenScale[params.layoutPerTokenScale.GetOffset(perTokenScaleTileOffset)];
|
||||
auto layoutGmTilePerTokenScale = params.layoutPerTokenScale.GetTileLayout(perTokenScaleTileShape);
|
||||
|
||||
auto &ubPerTokenScale = ubPerTokenScaleList[ubListId];
|
||||
auto layoutUbPerTokenScale =
|
||||
LayoutScale::template MakeLayoutInUb<ElementPerTokenScale>(perTokenScaleTileShape);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbPerTokenScaleVMTE2List[ubListId]);
|
||||
copyGmToUbPerTokenScale(ubPerTokenScale, gmTilePerTokenScale, layoutUbPerTokenScale,
|
||||
layoutGmTilePerTokenScale);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbPerTokenScaleMTE2VList[ubListId]);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
AscendC::Cast(ubCFp32, ubC, AscendC::RoundMode::CAST_RINT, TileShape::COUNT);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbScaleMTE2VList[ubListId]);
|
||||
if constexpr (!std::is_same_v<ElementRawScale, ElementFp32Scale>) {
|
||||
AscendC::Cast(ubFp32Scale, ubRawScale, AscendC::RoundMode::CAST_NONE, TileShape::COLUMN);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
tileRowBroadcastMul(ubMul, ubCFp32, ubFp32Scale);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbScaleVMTE2List[ubListId]);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbPerTokenScaleMTE2VList[ubListId]);
|
||||
tileBroadcastOneBlk(ubPerTokenScaleBrcb, ubPerTokenScale);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbPerTokenScaleVMTE2List[ubListId]);
|
||||
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
tileOneBlkColumnBroadcastMul(ubPerTokenMul, ubMul, ubPerTokenScaleBrcb);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
auto &ubD = ubDList[ubListId];
|
||||
LayoutD layoutUbD{actualTileShape, ubTileStride};
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
AscendC::Cast(ubD, ubPerTokenMul, AscendC::RoundMode::CAST_RINT, TileShape::COUNT);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
|
||||
auto tileOffsetD = params.layoutD.GetOffset(tileOffset);
|
||||
auto layoutGmTileD = params.layoutD.GetTileLayout(actualTileShape);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
DoCombineSend(ubD, layoutGmTileD, layoutUbD, groupOffsetD, expertIdx, tileOffsetD);
|
||||
} else {
|
||||
auto gmTileD = gmD[tileOffsetD];
|
||||
copyUbToGmD(gmTileD, ubD, layoutGmTileD, layoutUbD);
|
||||
}
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
|
||||
ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
Params params;
|
||||
Arch::Resource<ArchTag> &resource;
|
||||
MoeDistributeCombineImpl::CombineCalcInfo calcInfo;
|
||||
|
||||
AscendC::LocalTensor<ElementC> ubCList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementRawScale> ubRawScaleList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementFp32Scale> ubFp32ScaleList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementPerTokenScale> ubPerTokenScaleList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementD> ubDList[UB_STAGES];
|
||||
|
||||
int32_t eventUbCVMTE2List[UB_STAGES];
|
||||
int32_t eventUbCMTE2VList[UB_STAGES];
|
||||
int32_t eventUbScaleVMTE2List[UB_STAGES];
|
||||
int32_t eventUbScaleMTE2VList[UB_STAGES];
|
||||
int32_t eventUbPerTokenScaleVMTE2List[UB_STAGES];
|
||||
int32_t eventUbPerTokenScaleMTE2VList[UB_STAGES];
|
||||
int32_t eventUbDMTE3VList[UB_STAGES];
|
||||
int32_t eventUbDVMTE3List[UB_STAGES];
|
||||
|
||||
AscendC::LocalTensor<int32_t> epSendCountLocal_;
|
||||
|
||||
size_t ubOffset{0};
|
||||
int32_t eventVMTE2{0};
|
||||
int32_t eventMTE2V{0};
|
||||
int32_t eventMTE3V{0};
|
||||
int32_t eventVMTE3{0};
|
||||
int32_t eventVS{0};
|
||||
int32_t eventMTE2S{0};
|
||||
|
||||
uint32_t expertOffset;
|
||||
|
||||
uint32_t ubListId{0};
|
||||
|
||||
AscendC::LocalTensor<float> ubCFp32;
|
||||
AscendC::LocalTensor<float> ubMul;
|
||||
AscendC::LocalTensor<float> ubPerTokenScaleBrcb;
|
||||
AscendC::LocalTensor<float> ubPerTokenMul;
|
||||
|
||||
TileRowBroadcastMul tileRowBroadcastMul;
|
||||
TileBroadcastOneBlk tileBroadcastOneBlk;
|
||||
TileOneBlkColumnBroadcastMul tileOneBlkColumnBroadcastMul;
|
||||
|
||||
CopyGmToUbC copyGmToUbC;
|
||||
CopyGmToUbScale copyGmToUbScale;
|
||||
CopyGmToUbPerTokenScale copyGmToUbPerTokenScale;
|
||||
CopyUbToGmD copyUbToGmD;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue::Block
|
||||
|
||||
#endif // ACT_EPILOGUE_BLOCK_EPILOGUE_PER_TOKEN_DEQUANT_HPP
|
||||
@@ -0,0 +1,330 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/epilogue/dispatch_policy.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
|
||||
#include "../tile/tile_stride_muls.h"
|
||||
#include "../tile/tile_stride_binary.h"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <uint32_t UB_STAGES_, uint32_t EXEC_FLAG_,
|
||||
class CType_, class ScaleType_, class LayoutScale_, class LayoutPerTokenScale_,
|
||||
class DType_, class TileRowBroadcastMul_, class TileBroadcastOneBlk_, class TileOneBlkColumnBroadcastMul_,
|
||||
class TileCopy_, class EpilogueTileSwizzle_>
|
||||
class BlockEpilogue<EpilogueAtlasA2PerTokenDequantSwiglu<UB_STAGES_, EXEC_FLAG_>,
|
||||
CType_, Gemm::GemmType<ScaleType_, LayoutScale_>, Gemm::GemmType<float, LayoutPerTokenScale_>,
|
||||
DType_, TileRowBroadcastMul_, TileBroadcastOneBlk_, TileOneBlkColumnBroadcastMul_,
|
||||
TileCopy_, EpilogueTileSwizzle_>
|
||||
{
|
||||
public:
|
||||
using DispatchPolicy = EpilogueAtlasA2PerTokenDequantSwiglu<UB_STAGES_, EXEC_FLAG_>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
|
||||
// Data infos
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using ElementRawScale = ScaleType_;
|
||||
using ElementFp32Scale = float;
|
||||
using LayoutScale = LayoutScale_;
|
||||
using ElementPerTokenScale = float;
|
||||
using LayoutPerTokenScale = LayoutPerTokenScale_;
|
||||
using ElementD = typename DType_::Element;
|
||||
using LayoutD = typename DType_::Layout;
|
||||
|
||||
// Check data infos
|
||||
static_assert(std::is_same_v<ElementC, int32_t> && std::is_same_v<ElementD, float>,
|
||||
"The element type template parameters of BlockEpilogue are wrong");
|
||||
static_assert(std::is_same_v<LayoutC, layout::RowMajor> && std::is_same_v<LayoutScale, layout::VectorLayout> &&
|
||||
std::is_same_v<LayoutPerTokenScale, layout::VectorLayout> &&
|
||||
std::is_same_v<LayoutD, layout::RowMajor>,
|
||||
"The layout template parameters of BlockEpilogue are wrong");
|
||||
|
||||
// Tile compute ops
|
||||
using TileRowBroadcastMul = TileRowBroadcastMul_;
|
||||
using TileBroadcastOneBlk = TileBroadcastOneBlk_;
|
||||
using TileOneBlkColumnBroadcastMul = TileOneBlkColumnBroadcastMul_;
|
||||
|
||||
// Tile copy
|
||||
using CopyGmToUbC = typename TileCopy_::CopyGmToUbC;
|
||||
using CopyGmToUbScale = typename TileCopy_::CopyGmToUbX;
|
||||
using CopyGmToUbPerTokenScale = typename TileCopy_::CopyGmToUbY;
|
||||
using CopyUbToGmD = typename TileCopy_::CopyUbToGmD;
|
||||
|
||||
using EpilogueTileSwizzle = EpilogueTileSwizzle_;
|
||||
|
||||
using TileShape = typename TileRowBroadcastMul::TileShape;
|
||||
static_assert(TileShape::ROW * sizeof(float) % BYTE_PER_BLK == 0,
|
||||
"The per token scale granularity for word calculation must be 32 bytes aligned.");
|
||||
static_assert(TileShape::COLUMN % 2 == 0, "The n-axis needs to be divided into two parts.");
|
||||
|
||||
static_assert(TileShape::ROW == TileBroadcastOneBlk::COMPUTE_LENGTH &&
|
||||
std::is_same_v<TileShape, typename TileOneBlkColumnBroadcastMul::TileShape>,
|
||||
"TileShape must be consistent for all tile compute ops");
|
||||
|
||||
static_assert(UB_STAGES <= 2, "UB stages too large, event id is not enough.");
|
||||
|
||||
static_assert((UB_STAGES * (TileShape::COUNT * sizeof(ElementC) +
|
||||
(std::is_same_v<ElementRawScale, ElementFp32Scale> ?
|
||||
0 : TileShape::COLUMN * sizeof(ElementRawScale)) +
|
||||
TileShape::COLUMN * sizeof(ElementFp32Scale) +
|
||||
TileShape::ROW * sizeof(ElementPerTokenScale) + TileShape::COUNT * sizeof(ElementD)) +
|
||||
(TileShape::COUNT + TileShape::COUNT) * sizeof(float) + TileShape::ROW * BYTE_PER_BLK) <=
|
||||
ArchTag::UB_SIZE,
|
||||
"TileShape is too large to fit in UB");
|
||||
|
||||
struct Params {
|
||||
__gm__ ElementRawScale *ptrScale{nullptr};
|
||||
LayoutScale layoutScale{};
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale{nullptr};
|
||||
LayoutPerTokenScale layoutPerTokenScale{};
|
||||
__gm__ ElementD *ptrD{nullptr};
|
||||
LayoutD layoutD{};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params() {};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params(__gm__ ElementRawScale *ptrScale_, LayoutScale const &layoutScale_,
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale_, LayoutPerTokenScale const &layoutPerTokenScale_,
|
||||
__gm__ ElementD *ptrD_, LayoutD const &layoutD_)
|
||||
: ptrScale(ptrScale_),
|
||||
layoutScale(layoutScale_),
|
||||
ptrPerTokenScale(ptrPerTokenScale_),
|
||||
layoutPerTokenScale(layoutPerTokenScale_),
|
||||
ptrD(ptrD_),
|
||||
layoutD(layoutD_)
|
||||
{}
|
||||
};
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> const &resource, Params const ¶ms = Params{}) : params(params)
|
||||
{
|
||||
size_t ubOffset = 0;
|
||||
int32_t eventVMTE2 = 0;
|
||||
int32_t eventMTE2V = 0;
|
||||
int32_t eventMTE3V = 0;
|
||||
int32_t eventVMTE3 = 0;
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(ElementC);
|
||||
if constexpr (!std::is_same_v<ElementRawScale, ElementFp32Scale>) {
|
||||
ubRawScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementRawScale>(ubOffset);
|
||||
ubOffset += TileShape::COLUMN * sizeof(ElementRawScale);
|
||||
}
|
||||
ubFp32ScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementFp32Scale>(ubOffset);
|
||||
ubOffset += TileShape::COLUMN * sizeof(ElementFp32Scale);
|
||||
ubPerTokenScaleList[i] = resource.ubBuf.template GetBufferByByte<ElementPerTokenScale>(ubOffset);
|
||||
ubOffset += TileShape::ROW * sizeof(ElementPerTokenScale);
|
||||
ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(ElementD);
|
||||
|
||||
eventUbCVMTE2List[i] = eventVMTE2++;
|
||||
eventUbCMTE2VList[i] = eventMTE2V++;
|
||||
eventUbScaleVMTE2List[i] = eventVMTE2++;
|
||||
eventUbScaleMTE2VList[i] = eventMTE2V++;
|
||||
eventUbPerTokenScaleVMTE2List[i] = eventVMTE2++;
|
||||
eventUbPerTokenScaleMTE2VList[i] = eventMTE2V++;
|
||||
eventUbDMTE3VList[i] = eventMTE3V++;
|
||||
eventUbDVMTE3List[i] = eventVMTE3++;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbScaleVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbPerTokenScaleVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
ubTmpMxN = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(float);
|
||||
ubTmpMx32B = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
ubOffset += TileShape::ROW * BYTE_PER_BLK;
|
||||
ubDenominatorMxN = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbScaleVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbPerTokenScaleVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void UpdateParams(Params const ¶ms_)
|
||||
{
|
||||
params = params_;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(GemmCoord const &blockShapeMNK, GemmCoord const &blockCoordMNK,
|
||||
GemmCoord const &actualBlockShapeMNK, AscendC::GlobalTensor<ElementC> const &gmBlockC,
|
||||
LayoutC const &layoutBlockC, Callback &&callback = Callback{})
|
||||
{
|
||||
if (0 == actualBlockShapeMNK.k()) {
|
||||
return;
|
||||
}
|
||||
callback();
|
||||
// Calculate the offset of the current block
|
||||
MatrixCoord blockShape = blockShapeMNK.GetCoordMN();
|
||||
MatrixCoord blockCoord = blockCoordMNK.GetCoordMN();
|
||||
MatrixCoord actualBlockShape = actualBlockShapeMNK.GetCoordMN();
|
||||
MatrixCoord blockOffset = blockCoord * blockShape;
|
||||
bool isLeft = blockOffset.column() < (params.layoutD.shape(1) >> 1);
|
||||
AscendC::GlobalTensor<ElementRawScale> gmScale;
|
||||
gmScale.SetGlobalBuffer(params.ptrScale);
|
||||
AscendC::GlobalTensor<ElementPerTokenScale> gmPerTokenScale;
|
||||
gmPerTokenScale.SetGlobalBuffer(params.ptrPerTokenScale);
|
||||
AscendC::GlobalTensor<ElementD> gmD;
|
||||
gmD.SetGlobalBuffer(params.ptrD);
|
||||
|
||||
auto ubTileStride = MakeCoord(static_cast<int64_t>(TileShape::COLUMN), 1L);
|
||||
auto tileShape = TileShape::ToCoord();
|
||||
EpilogueTileSwizzle epilogueTileSwizzle(actualBlockShape, tileShape);
|
||||
uint32_t tileLoops = epilogueTileSwizzle.GetLoops();
|
||||
uint32_t subblockIdx = 0; // for 1C1V
|
||||
uint32_t subblockNum = 1; // for 1C1V
|
||||
for (uint32_t loopIdx = subblockIdx; loopIdx < tileLoops; loopIdx += subblockNum) {
|
||||
auto tileCoord = epilogueTileSwizzle.GetTileCoord(loopIdx);
|
||||
auto actualTileShape = epilogueTileSwizzle.GetActualTileShape(tileCoord);
|
||||
auto tileOffsetInBlock = tileCoord * tileShape;
|
||||
auto tileOffset = blockOffset + tileOffsetInBlock;
|
||||
|
||||
auto gmTileC = gmBlockC[layoutBlockC.GetOffset(tileOffsetInBlock)];
|
||||
auto layoutGmTileC = layoutBlockC.GetTileLayout(actualTileShape);
|
||||
|
||||
auto &ubC = ubCList[ubListId];
|
||||
LayoutC layoutUbC{actualTileShape, ubTileStride};
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
copyGmToUbC(ubC, gmTileC, layoutUbC, layoutGmTileC);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
|
||||
auto scaleTileOffset = tileOffset.template GetCoordByAxis<1>();
|
||||
auto scaleTileShape = actualTileShape.template GetCoordByAxis<1>();
|
||||
|
||||
auto gmTileScale = gmScale[params.layoutScale.GetOffset(scaleTileOffset)];
|
||||
auto layoutGmTileScale = params.layoutScale.GetTileLayout(scaleTileShape);
|
||||
|
||||
auto &ubFp32Scale = ubFp32ScaleList[ubListId];
|
||||
auto layoutFp32UbScale = LayoutScale::template MakeLayoutInUb<ElementFp32Scale>(scaleTileShape);
|
||||
auto &ubRawScale = ubRawScaleList[ubListId];
|
||||
auto layoutRawUbScale = LayoutScale::template MakeLayoutInUb<ElementRawScale>(scaleTileShape);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbScaleVMTE2List[ubListId]);
|
||||
if constexpr (!std::is_same_v<ElementRawScale, ElementFp32Scale>) {
|
||||
copyGmToUbScale(ubRawScale, gmTileScale, layoutRawUbScale, layoutGmTileScale);
|
||||
} else {
|
||||
copyGmToUbScale(ubFp32Scale, gmTileScale, layoutFp32UbScale, layoutGmTileScale);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbScaleMTE2VList[ubListId]);
|
||||
|
||||
auto perTokenScaleTileOffset = tileOffset.template GetCoordByAxis<0>();
|
||||
auto perTokenScaleTileShape = actualTileShape.template GetCoordByAxis<0>();
|
||||
|
||||
auto gmTilePerTokenScale = gmPerTokenScale[params.layoutPerTokenScale.GetOffset(perTokenScaleTileOffset)];
|
||||
auto layoutGmTilePerTokenScale = params.layoutPerTokenScale.GetTileLayout(perTokenScaleTileShape);
|
||||
|
||||
auto &ubPerTokenScale = ubPerTokenScaleList[ubListId];
|
||||
auto layoutUbPerTokenScale =
|
||||
LayoutScale::template MakeLayoutInUb<ElementPerTokenScale>(perTokenScaleTileShape);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbPerTokenScaleVMTE2List[ubListId]);
|
||||
copyGmToUbPerTokenScale(ubPerTokenScale, gmTilePerTokenScale, layoutUbPerTokenScale,
|
||||
layoutGmTilePerTokenScale);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbPerTokenScaleMTE2VList[ubListId]);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
AscendC::Cast(ubTmpMxN, ubC, AscendC::RoundMode::CAST_RINT, TileShape::COUNT);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbScaleMTE2VList[ubListId]);
|
||||
if constexpr (!std::is_same_v<ElementRawScale, ElementFp32Scale>) {
|
||||
AscendC::Cast(ubFp32Scale, ubRawScale, AscendC::RoundMode::CAST_NONE, TileShape::COLUMN);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
}
|
||||
tileRowBroadcastMul(ubTmpMxN, ubTmpMxN, ubFp32Scale);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbScaleVMTE2List[ubListId]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbPerTokenScaleMTE2VList[ubListId]);
|
||||
tileBroadcastOneBlk(ubTmpMx32B, ubPerTokenScale);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbPerTokenScaleVMTE2List[ubListId]);
|
||||
|
||||
auto &ubD = ubDList[ubListId];
|
||||
LayoutD layoutUbD{actualTileShape, ubTileStride};
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
// after dequant, the left half does x / (x + exp(-Dequant(x))), the right dose nothing
|
||||
if (isLeft) {
|
||||
tileOneBlkColumnBroadcastMul(ubTmpMxN, ubTmpMxN, ubTmpMx32B);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Muls(ubDenominatorMxN, ubTmpMxN, -1.0f, TileShape::COUNT);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Exp(ubDenominatorMxN, ubDenominatorMxN, TileShape::COUNT);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Adds(ubDenominatorMxN, ubDenominatorMxN, 1.0f, TileShape::COUNT);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
AscendC::Div(ubD, ubTmpMxN, ubDenominatorMxN, TileShape::COUNT);
|
||||
} else {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
tileOneBlkColumnBroadcastMul(ubD, ubTmpMxN, ubTmpMx32B);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
|
||||
auto gmTileD = gmD[params.layoutD.GetOffset(tileOffset)];
|
||||
auto layoutGmTileD = params.layoutD.GetTileLayout(actualTileShape);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
copyUbToGmD(gmTileD, ubD, layoutGmTileD, layoutUbD);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
Params params;
|
||||
|
||||
AscendC::LocalTensor<ElementC> ubCList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementRawScale> ubRawScaleList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementFp32Scale> ubFp32ScaleList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementPerTokenScale> ubPerTokenScaleList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementD> ubDList[UB_STAGES];
|
||||
|
||||
int32_t eventUbCVMTE2List[UB_STAGES];
|
||||
int32_t eventUbCMTE2VList[UB_STAGES];
|
||||
int32_t eventUbScaleVMTE2List[UB_STAGES];
|
||||
int32_t eventUbScaleMTE2VList[UB_STAGES];
|
||||
int32_t eventUbPerTokenScaleVMTE2List[UB_STAGES];
|
||||
int32_t eventUbPerTokenScaleMTE2VList[UB_STAGES];
|
||||
int32_t eventUbDMTE3VList[UB_STAGES];
|
||||
int32_t eventUbDVMTE3List[UB_STAGES];
|
||||
|
||||
uint32_t ubListId{0};
|
||||
|
||||
AscendC::LocalTensor<float> ubTmpMxN;
|
||||
AscendC::LocalTensor<float> ubTmpMx32B;
|
||||
AscendC::LocalTensor<float> ubDenominatorMxN;
|
||||
|
||||
TileRowBroadcastMul tileRowBroadcastMul;
|
||||
TileBroadcastOneBlk tileBroadcastOneBlk;
|
||||
TileOneBlkColumnBroadcastMul tileOneBlkColumnBroadcastMul;
|
||||
|
||||
CopyGmToUbC copyGmToUbC;
|
||||
CopyGmToUbScale copyGmToUbScale;
|
||||
CopyGmToUbPerTokenScale copyGmToUbPerTokenScale;
|
||||
CopyUbToGmD copyUbToGmD;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue::Block
|
||||
@@ -0,0 +1,259 @@
|
||||
/*
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/epilogue/dispatch_policy.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
|
||||
#include "../tile/tile_stride_muls.h"
|
||||
#include "../tile/tile_stride_binary.h"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <uint32_t UB_STAGES_, uint32_t EXEC_FLAG_,
|
||||
class CType_, class ScaleType_, class LayoutScale_, class LayoutPerTokenScale_,
|
||||
class DType_, class TileRowBroadcastMul_, class TileBroadcastOneBlk_, class TileOneBlkColumnBroadcastMul_,
|
||||
class TileCopy_, class EpilogueTileSwizzle_>
|
||||
class BlockEpilogue<EpilogueAtlasA2Swiglu<UB_STAGES_, EXEC_FLAG_>,
|
||||
CType_, Gemm::GemmType<ScaleType_, LayoutScale_>, Gemm::GemmType<float, LayoutPerTokenScale_>,
|
||||
DType_, TileRowBroadcastMul_, TileBroadcastOneBlk_, TileOneBlkColumnBroadcastMul_,
|
||||
TileCopy_, EpilogueTileSwizzle_>
|
||||
{
|
||||
public:
|
||||
using DispatchPolicy = EpilogueAtlasA2Swiglu<UB_STAGES_, EXEC_FLAG_>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
|
||||
// Data infos
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using ElementRawScale = ScaleType_;
|
||||
using ElementFp32Scale = float;
|
||||
using LayoutScale = LayoutScale_;
|
||||
using ElementPerTokenScale = float;
|
||||
using LayoutPerTokenScale = LayoutPerTokenScale_;
|
||||
using ElementD = typename DType_::Element;
|
||||
using LayoutD = typename DType_::Layout;
|
||||
|
||||
// Check data infos
|
||||
static_assert(std::is_same_v<ElementC, float> && std::is_same_v<ElementD, float>,
|
||||
"The element type template parameters of BlockEpilogue are wrong");
|
||||
static_assert(std::is_same_v<LayoutC, layout::RowMajor> && std::is_same_v<LayoutScale, layout::VectorLayout> &&
|
||||
std::is_same_v<LayoutPerTokenScale, layout::VectorLayout> &&
|
||||
std::is_same_v<LayoutD, layout::RowMajor>,
|
||||
"The layout template parameters of BlockEpilogue are wrong");
|
||||
|
||||
// Tile compute ops
|
||||
using TileRowBroadcastMul = TileRowBroadcastMul_;
|
||||
using TileBroadcastOneBlk = TileBroadcastOneBlk_;
|
||||
using TileOneBlkColumnBroadcastMul = TileOneBlkColumnBroadcastMul_;
|
||||
|
||||
// Tile copy
|
||||
using CopyGmToUbC = typename TileCopy_::CopyGmToUbC;
|
||||
using CopyGmToUbScale = typename TileCopy_::CopyGmToUbX;
|
||||
using CopyGmToUbPerTokenScale = typename TileCopy_::CopyGmToUbY;
|
||||
using CopyUbToGmD = typename TileCopy_::CopyUbToGmD;
|
||||
|
||||
using EpilogueTileSwizzle = EpilogueTileSwizzle_;
|
||||
|
||||
using TileShape = typename TileRowBroadcastMul::TileShape;
|
||||
static_assert(TileShape::ROW * sizeof(float) % BYTE_PER_BLK == 0,
|
||||
"The per token scale granularity for word calculation must be 32 bytes aligned.");
|
||||
static_assert(TileShape::COLUMN % 2 == 0, "The n-axis needs to be divided into two parts.");
|
||||
|
||||
static_assert(TileShape::ROW == TileBroadcastOneBlk::COMPUTE_LENGTH &&
|
||||
std::is_same_v<TileShape, typename TileOneBlkColumnBroadcastMul::TileShape>,
|
||||
"TileShape must be consistent for all tile compute ops");
|
||||
|
||||
static_assert(UB_STAGES <= 2, "UB stages too large, event id is not enough.");
|
||||
|
||||
static_assert((UB_STAGES * (TileShape::COUNT * sizeof(ElementC) + TileShape::COUNT * sizeof(ElementD)) +
|
||||
TileShape::ROW * BYTE_PER_BLK) <= ArchTag::UB_SIZE,
|
||||
"TileShape is too large to fit in UB");
|
||||
|
||||
struct Params {
|
||||
__gm__ ElementRawScale *ptrScale{nullptr};
|
||||
LayoutScale layoutScale{};
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale{nullptr};
|
||||
LayoutPerTokenScale layoutPerTokenScale{};
|
||||
__gm__ ElementD *ptrD{nullptr};
|
||||
LayoutD layoutD{};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params() {};
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params(__gm__ ElementRawScale *ptrScale_, LayoutScale const &layoutScale_,
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale_, LayoutPerTokenScale const &layoutPerTokenScale_,
|
||||
__gm__ ElementD *ptrD_, LayoutD const &layoutD_)
|
||||
: ptrScale(ptrScale_),
|
||||
layoutScale(layoutScale_),
|
||||
ptrPerTokenScale(ptrPerTokenScale_),
|
||||
layoutPerTokenScale(layoutPerTokenScale_),
|
||||
ptrD(ptrD_),
|
||||
layoutD(layoutD_)
|
||||
{}
|
||||
};
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> const &resource, Params const ¶ms = Params{}) : params(params)
|
||||
{
|
||||
size_t ubOffset = 0;
|
||||
int32_t eventVMTE2 = 0;
|
||||
int32_t eventMTE2V = 0;
|
||||
int32_t eventMTE3V = 0;
|
||||
int32_t eventVMTE3 = 0;
|
||||
int32_t eventMTE3MTE2 = 0;
|
||||
int32_t eventMTE2MTE3 = 0;
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
ubCList[i] = resource.ubBuf.template GetBufferByByte<ElementC>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(ElementC);
|
||||
ubDList[i] = resource.ubBuf.template GetBufferByByte<ElementD>(ubOffset);
|
||||
ubOffset += TileShape::COUNT * sizeof(ElementD);
|
||||
|
||||
eventUbCVMTE2List[i] = eventVMTE2++;
|
||||
eventUbCMTE2VList[i] = eventMTE2V++;
|
||||
eventUbDMTE3VList[i] = eventMTE3V++;
|
||||
eventUbDVMTE3List[i] = eventVMTE3++;
|
||||
eventUbMTE3MTE2List[i] = eventMTE3MTE2++;
|
||||
eventUbMTE2MTE3List[i] = eventMTE2MTE3++;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(eventUbMTE3MTE2List[i]);
|
||||
}
|
||||
ubDenominatorMxN = resource.ubBuf.template GetBufferByByte<float>(ubOffset);
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue()
|
||||
{
|
||||
for (uint32_t i = 0; i < UB_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[i]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(eventUbMTE3MTE2List[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void UpdateParams(Params const ¶ms_)
|
||||
{
|
||||
params = params_;
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(GemmCoord const &blockShapeMNK, GemmCoord const &blockCoordMNK,
|
||||
GemmCoord const &actualBlockShapeMNK, AscendC::GlobalTensor<ElementC> const &gmBlockC,
|
||||
LayoutC const &layoutBlockC, Callback &&callback = Callback{})
|
||||
{
|
||||
if (0 == actualBlockShapeMNK.k()) {
|
||||
return;
|
||||
}
|
||||
callback();
|
||||
ubListId = 0;
|
||||
// Calculate the offset of the current block
|
||||
MatrixCoord blockShape = blockShapeMNK.GetCoordMN();
|
||||
MatrixCoord blockCoord = blockCoordMNK.GetCoordMN();
|
||||
MatrixCoord actualBlockShape = actualBlockShapeMNK.GetCoordMN();
|
||||
MatrixCoord blockOffset = blockCoord * blockShape;
|
||||
bool isLeft = blockOffset.column() < (params.layoutD.shape(1) >> 1);
|
||||
AscendC::GlobalTensor<ElementRawScale> gmScale;
|
||||
gmScale.SetGlobalBuffer(params.ptrScale);
|
||||
AscendC::GlobalTensor<ElementPerTokenScale> gmPerTokenScale;
|
||||
gmPerTokenScale.SetGlobalBuffer(params.ptrPerTokenScale);
|
||||
AscendC::GlobalTensor<ElementD> gmD;
|
||||
gmD.SetGlobalBuffer(params.ptrD);
|
||||
|
||||
auto ubTileStride = MakeCoord(static_cast<int64_t>(TileShape::COLUMN), 1L);
|
||||
auto tileShape = TileShape::ToCoord();
|
||||
EpilogueTileSwizzle epilogueTileSwizzle(actualBlockShape, tileShape);
|
||||
uint32_t tileLoops = epilogueTileSwizzle.GetLoops();
|
||||
uint32_t subblockIdx = 0; // for 1C1V
|
||||
uint32_t subblockNum = 1; // for 1C1V
|
||||
for (uint32_t loopIdx = subblockIdx; loopIdx < tileLoops; loopIdx += subblockNum) {
|
||||
auto tileCoord = epilogueTileSwizzle.GetTileCoord(loopIdx);
|
||||
auto actualTileShape = epilogueTileSwizzle.GetActualTileShape(tileCoord);
|
||||
auto tileOffsetInBlock = tileCoord * tileShape;
|
||||
auto tileOffset = blockOffset + tileOffsetInBlock;
|
||||
|
||||
auto gmTileC = gmBlockC[layoutBlockC.GetOffset(tileOffsetInBlock)];
|
||||
auto layoutGmTileC = layoutBlockC.GetTileLayout(actualTileShape);
|
||||
|
||||
auto &ubC = ubCList[ubListId];
|
||||
LayoutC layoutUbC{actualTileShape, ubTileStride};
|
||||
|
||||
auto &ubD = ubDList[ubListId];
|
||||
LayoutD layoutUbD{actualTileShape, ubTileStride};
|
||||
auto gmTileD = gmD[params.layoutD.GetOffset(tileOffset)];
|
||||
auto layoutGmTileD = params.layoutD.GetTileLayout(actualTileShape);
|
||||
|
||||
if (isLeft) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
copyGmToUbC(ubC, gmTileC, layoutUbC, layoutGmTileC);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(eventUbCMTE2VList[ubListId]);
|
||||
AscendC::Muls(ubDenominatorMxN, ubC, -1.0f, TileShape::COUNT);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Exp(ubDenominatorMxN, ubDenominatorMxN, TileShape::COUNT);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Adds(ubDenominatorMxN, ubDenominatorMxN, 1.0f, TileShape::COUNT);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
AscendC::Div(ubD, ubC, ubDenominatorMxN, TileShape::COUNT);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(eventUbCVMTE2List[ubListId]);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(eventUbDVMTE3List[ubListId]);
|
||||
copyUbToGmD(gmTileD, ubD, layoutGmTileD, layoutUbD);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_V>(eventUbDMTE3VList[ubListId]);
|
||||
} else {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(eventUbMTE3MTE2List[ubListId]);
|
||||
copyGmToUbC(ubC, gmTileC, layoutUbC, layoutGmTileC);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE3>(eventUbMTE2MTE3List[ubListId]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE3>(eventUbMTE2MTE3List[ubListId]);
|
||||
copyUbToGmD(gmTileD, ubC, layoutGmTileD, layoutUbD);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(eventUbMTE3MTE2List[ubListId]);
|
||||
}
|
||||
|
||||
ubListId = (ubListId + 1 < UB_STAGES) ? (ubListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
Params params;
|
||||
|
||||
AscendC::LocalTensor<ElementC> ubCList[UB_STAGES];
|
||||
AscendC::LocalTensor<ElementD> ubDList[UB_STAGES];
|
||||
|
||||
int32_t eventUbCVMTE2List[UB_STAGES];
|
||||
int32_t eventUbCMTE2VList[UB_STAGES];
|
||||
int32_t eventUbDMTE3VList[UB_STAGES];
|
||||
int32_t eventUbDVMTE3List[UB_STAGES];
|
||||
int32_t eventUbMTE3MTE2List[UB_STAGES];
|
||||
int32_t eventUbMTE2MTE3List[UB_STAGES];
|
||||
|
||||
uint32_t ubListId{0};
|
||||
|
||||
AscendC::LocalTensor<float> ubDenominatorMxN;
|
||||
|
||||
TileRowBroadcastMul tileRowBroadcastMul;
|
||||
TileBroadcastOneBlk tileBroadcastOneBlk;
|
||||
TileOneBlkColumnBroadcastMul tileOneBlkColumnBroadcastMul;
|
||||
|
||||
CopyGmToUbC copyGmToUbC;
|
||||
CopyGmToUbScale copyGmToUbScale;
|
||||
CopyGmToUbPerTokenScale copyGmToUbPerTokenScale;
|
||||
CopyUbToGmD copyUbToGmD;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue::Block
|
||||
@@ -0,0 +1,42 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/epilogue/dispatch_policy.hpp"
|
||||
|
||||
namespace Catlass::Epilogue {
|
||||
|
||||
template <uint32_t UB_STAGES_, uint32_t EXEC_FLAG_>
|
||||
struct EpilogueAtlasA2PerTokenDequantSwiglu {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
static constexpr uint32_t EXEC_FLAG = EXEC_FLAG_;
|
||||
};
|
||||
|
||||
template <uint32_t UB_STAGES_, uint32_t EXEC_FLAG_>
|
||||
struct EpilogueAtlasA2PerTokenDequantCombine {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
static constexpr uint32_t EXEC_FLAG = EXEC_FLAG_;
|
||||
};
|
||||
|
||||
template <uint32_t UB_STAGES_, uint32_t EXEC_FLAG_>
|
||||
struct EpilogueAtlasA2Swiglu {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
static constexpr uint32_t EXEC_FLAG = EXEC_FLAG_;
|
||||
};
|
||||
|
||||
template <uint32_t UB_STAGES_, uint32_t EXEC_FLAG_>
|
||||
struct EpilogueAtlasA2Combine {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
static constexpr uint32_t UB_STAGES = UB_STAGES_;
|
||||
static constexpr uint32_t EXEC_FLAG = EXEC_FLAG_;
|
||||
};
|
||||
} // namespace Catlass::Epilogue
|
||||
@@ -0,0 +1,107 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/catlass.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Tile {
|
||||
|
||||
template <class ArchTag_, class ElementCompute_, class TileShape_, int64_t DST_STRIDE_, int64_t SRC0_STRIDE_,
|
||||
int64_t SRC1_STRIDE_>
|
||||
struct TileStrideBinary {
|
||||
using ArchTag = ArchTag_;
|
||||
using ElementCompute = ElementCompute_;
|
||||
using TileShape = TileShape_;
|
||||
static constexpr int64_t DST_STRIDE = DST_STRIDE_;
|
||||
static constexpr int64_t SRC0_STRIDE = SRC0_STRIDE_;
|
||||
static constexpr int64_t SRC1_STRIDE = SRC1_STRIDE_;
|
||||
|
||||
static constexpr uint32_t MAX_REPEAT_TIMES = 255;
|
||||
static constexpr uint32_t ELE_NUM_PER_BLK = BYTE_PER_BLK / sizeof(ElementCompute);
|
||||
|
||||
static constexpr uint32_t DST_BLK_NUM_PER_COLUMN = DST_STRIDE / ELE_NUM_PER_BLK;
|
||||
static constexpr uint32_t SRC0_BLK_NUM_PER_COLUMN = SRC0_STRIDE / ELE_NUM_PER_BLK;
|
||||
static constexpr uint32_t SRC1_BLK_NUM_PER_COLUMN = SRC1_STRIDE / ELE_NUM_PER_BLK;
|
||||
|
||||
static constexpr uint32_t ROW_NUM_PER_COMPUTE = MAX_REPEAT_TIMES;
|
||||
static constexpr uint32_t COL_NUM_PER_COMPUTE = BYTE_PER_VECTOR_FRACTAL / sizeof(ElementCompute);
|
||||
|
||||
CATLASS_DEVICE
|
||||
TileStrideBinary()
|
||||
{
|
||||
repeatParams.dstBlkStride = 1;
|
||||
repeatParams.src0BlkStride = 1;
|
||||
repeatParams.src1BlkStride = 1;
|
||||
repeatParams.dstRepStride = DST_BLK_NUM_PER_COLUMN;
|
||||
repeatParams.src0RepStride = SRC0_BLK_NUM_PER_COLUMN;
|
||||
repeatParams.src1RepStride = SRC1_BLK_NUM_PER_COLUMN;
|
||||
}
|
||||
|
||||
AscendC::BinaryRepeatParams repeatParams;
|
||||
};
|
||||
|
||||
template <class ArchTag_, class ElementCompute_, class TileShape_, int64_t DST_STRIDE_, int64_t SRC0_STRIDE_,
|
||||
int64_t SRC1_STRIDE_>
|
||||
struct TileStrideMul
|
||||
: TileStrideBinary<ArchTag_, ElementCompute_, TileShape_, DST_STRIDE_, SRC0_STRIDE_, SRC1_STRIDE_> {
|
||||
using Base = TileStrideBinary<ArchTag_, ElementCompute_, TileShape_, DST_STRIDE_, SRC0_STRIDE_, SRC1_STRIDE_>;
|
||||
|
||||
CATLASS_DEVICE
|
||||
TileStrideMul() : Base() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(AscendC::LocalTensor<typename Base::ElementCompute> const &ubDst,
|
||||
AscendC::LocalTensor<typename Base::ElementCompute> const &ubSrc0,
|
||||
AscendC::LocalTensor<typename Base::ElementCompute> const &ubSrc1)
|
||||
{
|
||||
for (uint32_t rowOffset = 0; rowOffset < Base::TileShape::ROW; rowOffset += Base::ROW_NUM_PER_COMPUTE) {
|
||||
uint32_t residueM = Base::TileShape::ROW - rowOffset;
|
||||
uint8_t repeatTimes =
|
||||
static_cast<uint8_t>((residueM > Base::ROW_NUM_PER_COMPUTE) ? Base::ROW_NUM_PER_COMPUTE : residueM);
|
||||
for (uint32_t colOffset = 0; colOffset < Base::TileShape::COLUMN; colOffset += Base::COL_NUM_PER_COMPUTE) {
|
||||
uint32_t residueN = Base::TileShape::COLUMN - colOffset;
|
||||
uint64_t mask = (residueN > Base::COL_NUM_PER_COMPUTE) ? Base::COL_NUM_PER_COMPUTE : residueN;
|
||||
AscendC::Mul(ubDst[rowOffset * Base::DST_STRIDE + colOffset],
|
||||
ubSrc0[rowOffset * Base::SRC0_STRIDE + colOffset],
|
||||
ubSrc1[rowOffset * Base::SRC1_STRIDE + colOffset], mask, repeatTimes, this->repeatParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
template <class ArchTag_, class ElementCompute_, class TileShape_, int64_t DST_STRIDE_, int64_t SRC0_STRIDE_,
|
||||
int64_t SRC1_STRIDE_>
|
||||
struct TileStrideDiv
|
||||
: TileStrideBinary<ArchTag_, ElementCompute_, TileShape_, DST_STRIDE_, SRC0_STRIDE_, SRC1_STRIDE_> {
|
||||
using Base = TileStrideBinary<ArchTag_, ElementCompute_, TileShape_, DST_STRIDE_, SRC0_STRIDE_, SRC1_STRIDE_>;
|
||||
|
||||
CATLASS_DEVICE
|
||||
TileStrideDiv() : Base() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(AscendC::LocalTensor<typename Base::ElementCompute> const &ubDst,
|
||||
AscendC::LocalTensor<typename Base::ElementCompute> const &ubSrc0,
|
||||
AscendC::LocalTensor<typename Base::ElementCompute> const &ubSrc1)
|
||||
{
|
||||
for (uint32_t rowOffset = 0; rowOffset < Base::TileShape::ROW; rowOffset += Base::ROW_NUM_PER_COMPUTE) {
|
||||
uint32_t residueM = Base::TileShape::ROW - rowOffset;
|
||||
uint8_t repeatTimes =
|
||||
static_cast<uint8_t>((residueM > Base::ROW_NUM_PER_COMPUTE) ? Base::ROW_NUM_PER_COMPUTE : residueM);
|
||||
for (uint32_t colOffset = 0; colOffset < Base::TileShape::COLUMN; colOffset += Base::COL_NUM_PER_COMPUTE) {
|
||||
uint32_t residueN = Base::TileShape::COLUMN - colOffset;
|
||||
uint64_t mask = (residueN > Base::COL_NUM_PER_COMPUTE) ? Base::COL_NUM_PER_COMPUTE : residueN;
|
||||
AscendC::Div(ubDst[rowOffset * Base::DST_STRIDE + colOffset],
|
||||
ubSrc0[rowOffset * Base::SRC0_STRIDE + colOffset],
|
||||
ubSrc1[rowOffset * Base::SRC1_STRIDE + colOffset], mask, repeatTimes, this->repeatParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue::Tile
|
||||
@@ -0,0 +1,59 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/catlass.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Tile {
|
||||
|
||||
template <class ArchTag_, class ElementCompute_, class TileShape_, class DstTileShape_, class SrcTileShape_>
|
||||
struct TileStrideMuls {
|
||||
using ArchTag = ArchTag_;
|
||||
using ElementCompute = ElementCompute_;
|
||||
using TileShape = TileShape_;
|
||||
using DstTileShape = DstTileShape_;
|
||||
using SrcTileShape = SrcTileShape_;
|
||||
|
||||
static_assert(DstTileShape::ROW == SrcTileShape::ROW && DstTileShape::ROW == TileShape::ROW, "Error");
|
||||
|
||||
CATLASS_DEVICE
|
||||
TileStrideMuls() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(AscendC::LocalTensor<ElementCompute> const &ubDst,
|
||||
AscendC::LocalTensor<ElementCompute> const &ubSrc, ElementCompute scalar)
|
||||
{
|
||||
constexpr uint32_t maxRepeatTimes = 255;
|
||||
constexpr uint32_t eleNumPerBlk = BYTE_PER_BLK / sizeof(ElementCompute);
|
||||
|
||||
constexpr uint32_t dstBlkNumPerColumn = DstTileShape::COLUMN / eleNumPerBlk;
|
||||
constexpr uint32_t srcBlkNumPerColumn = SrcTileShape::COLUMN / eleNumPerBlk;
|
||||
AscendC::UnaryRepeatParams repeatParams;
|
||||
repeatParams.dstBlkStride = 1;
|
||||
repeatParams.srcBlkStride = 1;
|
||||
repeatParams.dstRepStride = dstBlkNumPerColumn;
|
||||
repeatParams.srcRepStride = srcBlkNumPerColumn;
|
||||
|
||||
constexpr uint32_t rowNumPerCompute = maxRepeatTimes;
|
||||
constexpr uint32_t colNumPerCompute = BYTE_PER_VECTOR_FRACTAL / sizeof(ElementCompute);
|
||||
for (uint32_t rowOffset = 0; rowOffset < TileShape::ROW; rowOffset += rowNumPerCompute) {
|
||||
uint32_t residueM = TileShape::ROW - rowOffset;
|
||||
uint8_t repeatTimes = static_cast<uint8_t>((residueM > rowNumPerCompute) ? rowNumPerCompute : residueM);
|
||||
for (uint32_t colOffset = 0; colOffset < TileShape::COLUMN; colOffset += colNumPerCompute) {
|
||||
uint32_t residueN = TileShape::COLUMN - colOffset;
|
||||
uint64_t mask = (residueN > colNumPerCompute) ? colNumPerCompute : residueN;
|
||||
AscendC::Muls(ubDst[rowOffset * DstTileShape::COLUMN + colOffset],
|
||||
ubSrc[rowOffset * SrcTileShape::COLUMN + colOffset], scalar, mask, repeatTimes,
|
||||
repeatParams);
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue::Tile
|
||||
@@ -0,0 +1,13 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/gemm/block/block_mmad.hpp"
|
||||
|
||||
#include "block_mmad_preload_async_with_callback_resident_a.h"
|
||||
@@ -0,0 +1,420 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/coord.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/gemm/dispatch_policy.hpp"
|
||||
#include "catlass/gemm/helper.hpp"
|
||||
|
||||
namespace Catlass::Gemm::Block {
|
||||
|
||||
template <uint32_t PRELOAD_STAGES_, uint32_t L1A_STAGES_, uint32_t L1B_STAGES_, uint32_t L0A_STAGES_,
|
||||
uint32_t L0B_STAGES_, uint32_t L0C_STAGES_, bool ENABLE_UNIT_FLAG_, bool ENABLE_SHUFFLE_K_,
|
||||
class L1TileShape_, class L0TileShape_, class AType_, class BType_, class CType_, class BiasType_,
|
||||
class TileCopy_, class TileMmad_>
|
||||
struct BlockMmad<
|
||||
MmadAtlasA2PreloadAsyncWithCallbackResidentA<PRELOAD_STAGES_, L1A_STAGES_, L1B_STAGES_, L0A_STAGES_, L0B_STAGES_,
|
||||
L0C_STAGES_, ENABLE_UNIT_FLAG_, ENABLE_SHUFFLE_K_>,
|
||||
L1TileShape_, L0TileShape_, AType_, BType_, CType_, BiasType_, TileCopy_, TileMmad_> {
|
||||
public:
|
||||
// Type Aliases
|
||||
using DispatchPolicy =
|
||||
MmadAtlasA2PreloadAsyncWithCallbackResidentA<PRELOAD_STAGES_, L1A_STAGES_, L1B_STAGES_, L0A_STAGES_,
|
||||
L0B_STAGES_, L0C_STAGES_, ENABLE_UNIT_FLAG_, ENABLE_SHUFFLE_K_>;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
using L1TileShape = L1TileShape_;
|
||||
using L0TileShape = L0TileShape_;
|
||||
using ElementA = typename AType_::Element;
|
||||
using LayoutA = typename AType_::Layout;
|
||||
using ElementB = typename BType_::Element;
|
||||
using LayoutB = typename BType_::Layout;
|
||||
using ElementC = typename CType_::Element;
|
||||
using LayoutC = typename CType_::Layout;
|
||||
using TileMmad = TileMmad_;
|
||||
using CopyGmToL1A = typename TileCopy_::CopyGmToL1A;
|
||||
using CopyGmToL1B = typename TileCopy_::CopyGmToL1B;
|
||||
using CopyL1ToL0A = typename TileCopy_::CopyL1ToL0A;
|
||||
using CopyL1ToL0B = typename TileCopy_::CopyL1ToL0B;
|
||||
using CopyL0CToGm = typename TileCopy_::CopyL0CToGm;
|
||||
using ElementAccumulator =
|
||||
typename Gemm::helper::ElementAccumulatorSelector<ElementA, ElementB>::ElementAccumulator;
|
||||
using LayoutAInL1 = typename CopyL1ToL0A::LayoutSrc;
|
||||
using LayoutBInL1 = typename CopyL1ToL0B::LayoutSrc;
|
||||
using LayoutAInL0 = typename CopyL1ToL0A::LayoutDst;
|
||||
using LayoutBInL0 = typename CopyL1ToL0B::LayoutDst;
|
||||
using LayoutCInL0 = layout::zN;
|
||||
|
||||
using L1AAlignHelper = Gemm::helper::L1AlignHelper<ElementA, LayoutA>;
|
||||
using L1BAlignHelper = Gemm::helper::L1AlignHelper<ElementB, LayoutB>;
|
||||
|
||||
static constexpr uint32_t PRELOAD_STAGES = DispatchPolicy::PRELOAD_STAGES;
|
||||
static constexpr uint32_t L1A_STAGES = DispatchPolicy::L1A_STAGES;
|
||||
static constexpr uint32_t L1B_STAGES = DispatchPolicy::L1B_STAGES;
|
||||
static constexpr uint32_t L0A_STAGES = DispatchPolicy::L0A_STAGES;
|
||||
static constexpr uint32_t L0B_STAGES = DispatchPolicy::L0B_STAGES;
|
||||
static constexpr uint32_t L0C_STAGES = DispatchPolicy::L0C_STAGES;
|
||||
|
||||
static constexpr bool ENABLE_UNIT_FLAG = DispatchPolicy::ENABLE_UNIT_FLAG;
|
||||
static constexpr bool ENABLE_SHUFFLE_K = DispatchPolicy::ENABLE_SHUFFLE_K;
|
||||
|
||||
// L1 tile size
|
||||
static constexpr uint32_t L1A_TILE_SIZE = L1TileShape::M * L1TileShape::K * sizeof(ElementA);
|
||||
static constexpr uint32_t L1B_TILE_SIZE = L1TileShape::N * L1TileShape::K * sizeof(ElementB);
|
||||
// L0 tile size
|
||||
static constexpr uint32_t L0A_TILE_SIZE = L0TileShape::M * L0TileShape::K * sizeof(ElementA);
|
||||
static constexpr uint32_t L0B_TILE_SIZE = L0TileShape::K * L0TileShape::N * sizeof(ElementB);
|
||||
static constexpr uint32_t L0C_TILE_SIZE = L1TileShape::M * L1TileShape::N * sizeof(ElementAccumulator);
|
||||
|
||||
// Check LayoutC
|
||||
static_assert(std::is_same_v<LayoutC, layout::RowMajor>, "LayoutC only support RowMajor yet!");
|
||||
|
||||
// Check L1TileShape
|
||||
static_assert(L1A_TILE_SIZE * L1A_STAGES + L1B_TILE_SIZE * L1B_STAGES <= ArchTag::L1_SIZE,
|
||||
"L1TileShape exceeding the L1 space!");
|
||||
|
||||
// Check L0TileShape
|
||||
static_assert(L0A_TILE_SIZE * L0A_STAGES <= ArchTag::L0A_SIZE, "L0TileShape exceeding the L0A space!");
|
||||
static_assert(L0B_TILE_SIZE * L0B_STAGES <= ArchTag::L0B_SIZE, "L0TileShape exceeding the L0B space!");
|
||||
static_assert(L0C_TILE_SIZE * L0C_STAGES <= ArchTag::L0C_SIZE, "L0TileShape exceeding the L0C space!");
|
||||
|
||||
static_assert(L1TileShape::M == L0TileShape::M && L1TileShape::N == L0TileShape::N,
|
||||
"The situation where the basic blocks of L1 and L0 differ on the m and n axes is not supported yet");
|
||||
|
||||
static constexpr auto L1A_LAYOUT = LayoutAInL1::template MakeLayout<ElementA>(L1TileShape::M, L1TileShape::K);
|
||||
static constexpr auto L1B_LAYOUT = LayoutBInL1::template MakeLayout<ElementB>(L1TileShape::K, L1TileShape::N);
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockMmad(Arch::Resource<ArchTag> &resource, uint32_t l1BufAddrStart = 0)
|
||||
{
|
||||
InitL1(resource, l1BufAddrStart);
|
||||
InitL0A(resource);
|
||||
InitL0B(resource);
|
||||
InitL0C(resource);
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockMmad()
|
||||
{
|
||||
SynchronizeBlock();
|
||||
for (uint32_t i = 0; i < L1A_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[i]);
|
||||
}
|
||||
for (uint32_t i = 0; i < L1B_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[i]);
|
||||
}
|
||||
for (uint32_t i = 0; i < L0A_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[i]);
|
||||
}
|
||||
for (uint32_t i = 0; i < L0B_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[i]);
|
||||
}
|
||||
for (uint32_t i = 0; i < L0C_STAGES; ++i) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::FIX_M>(l0CEventList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(AscendC::GlobalTensor<ElementA> const &gmBlockA, LayoutA const &layoutA,
|
||||
AscendC::GlobalTensor<ElementB> const &gmBlockB, LayoutB const &layoutB,
|
||||
AscendC::GlobalTensor<ElementC> const &gmBlockC, LayoutC const &layoutC,
|
||||
GemmCoord const &actualShape, Callback const &callbackBeforeFixpipe,
|
||||
Callback const &callbackAfterFixpipe)
|
||||
{
|
||||
uint32_t kTileCount = CeilDiv<L1TileShape::K>(actualShape.k());
|
||||
bool useResidentA =
|
||||
(kTileCount == L1A_STAGES) && (!isFirstLoad) && (gmBlockA.GetPhyAddr() == lastGmBlockA.GetPhyAddr());
|
||||
isFirstLoad = false;
|
||||
lastGmBlockA = gmBlockA;
|
||||
|
||||
uint32_t mRound = RoundUp<L1AAlignHelper::M_ALIGNED>(actualShape.m());
|
||||
uint32_t nRound = RoundUp<L1BAlignHelper::N_ALIGNED>(actualShape.n());
|
||||
|
||||
uint32_t startTileIdx = 0;
|
||||
if constexpr (ENABLE_SHUFFLE_K) {
|
||||
startTileIdx = AscendC::GetBlockIdx() % kTileCount;
|
||||
}
|
||||
|
||||
for (uint32_t kLoopIdx = 0; kLoopIdx < kTileCount; ++kLoopIdx) {
|
||||
uint32_t kTileIdx = (startTileIdx + kLoopIdx < kTileCount) ? (startTileIdx + kLoopIdx)
|
||||
: (startTileIdx + kLoopIdx - kTileCount);
|
||||
|
||||
uint32_t kActual =
|
||||
(kTileIdx < kTileCount - 1) ? L1TileShape::K : (actualShape.k() - kTileIdx * L1TileShape::K);
|
||||
|
||||
// Emission load instruction from GM to L1
|
||||
MatrixCoord gmTileAOffset{0, kTileIdx * L1TileShape::K};
|
||||
MatrixCoord gmTileBOffset{kTileIdx * L1TileShape::K, 0};
|
||||
auto gmTileA = gmBlockA[layoutA.GetOffset(gmTileAOffset)];
|
||||
auto gmTileB = gmBlockB[layoutB.GetOffset(gmTileBOffset)];
|
||||
// Load first matrix A tile from GM to L1
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[l1AListId]);
|
||||
if (!useResidentA) {
|
||||
auto layoutTileA = layoutA.GetTileLayout(MakeCoord(actualShape.m(), kActual));
|
||||
copyGmToL1A(l1ATensorList[l1AListId], gmTileA, L1A_LAYOUT, layoutTileA);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(l1AEventList[l1AListId]);
|
||||
// Load first matrix B tile from GM to L1
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[l1BListId]);
|
||||
auto layoutTileB = layoutB.GetTileLayout(MakeCoord(kActual, actualShape.n()));
|
||||
copyGmToL1B(l1BTensorList[l1BListId], gmTileB, L1B_LAYOUT, layoutTileB);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE1>(l1BEventList[l1BListId]);
|
||||
|
||||
// If the number of preload instructions reaches the upper limit, perform an mmad calculation on L1 tile
|
||||
if (preloadCount == PRELOAD_STAGES) {
|
||||
L1TileMmad(l1TileMmadParamsList[l1TileMmadParamsId]);
|
||||
}
|
||||
|
||||
// Store the current load status
|
||||
uint32_t preloadL1TileMmadParamsId = (l1TileMmadParamsId + preloadCount < PRELOAD_STAGES)
|
||||
? (l1TileMmadParamsId + preloadCount)
|
||||
: (l1TileMmadParamsId + preloadCount - PRELOAD_STAGES);
|
||||
auto &l1TileMmadParams = l1TileMmadParamsList[preloadL1TileMmadParamsId];
|
||||
l1TileMmadParams.l1AListId = l1AListId;
|
||||
l1TileMmadParams.l1BListId = l1BListId;
|
||||
l1TileMmadParams.mRound = mRound;
|
||||
l1TileMmadParams.nRound = nRound;
|
||||
l1TileMmadParams.kActual = kActual;
|
||||
l1TileMmadParams.isKLoopFirst = (kLoopIdx == 0);
|
||||
l1TileMmadParams.isKLoopLast = (kLoopIdx == kTileCount - 1);
|
||||
if (kLoopIdx == kTileCount - 1) {
|
||||
l1TileMmadParams.gmBlockC = gmBlockC;
|
||||
l1TileMmadParams.layoutCInGm = layoutC.GetTileLayout(actualShape.GetCoordMN());
|
||||
l1TileMmadParams.callbackBeforeFixpipe = callbackBeforeFixpipe;
|
||||
l1TileMmadParams.callbackAfterFixpipe = callbackAfterFixpipe;
|
||||
}
|
||||
|
||||
if (preloadCount < PRELOAD_STAGES) {
|
||||
++preloadCount;
|
||||
} else {
|
||||
l1TileMmadParamsId = (l1TileMmadParamsId + 1 < PRELOAD_STAGES) ? (l1TileMmadParamsId + 1) : 0;
|
||||
}
|
||||
l1AListId = (l1AListId + 1 < L1A_STAGES) ? (l1AListId + 1) : 0;
|
||||
l1BListId = (l1BListId + 1 < L1B_STAGES) ? (l1BListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void SynchronizeBlock()
|
||||
{
|
||||
while (preloadCount > 0) {
|
||||
L1TileMmad(l1TileMmadParamsList[l1TileMmadParamsId]);
|
||||
l1TileMmadParamsId = (l1TileMmadParamsId + 1 < PRELOAD_STAGES) ? (l1TileMmadParamsId + 1) : 0;
|
||||
--preloadCount;
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
struct L1TileMmadParams {
|
||||
uint32_t l1AListId;
|
||||
uint32_t l1BListId;
|
||||
uint32_t mRound;
|
||||
uint32_t nRound;
|
||||
uint32_t kActual;
|
||||
bool isKLoopFirst;
|
||||
bool isKLoopLast;
|
||||
AscendC::GlobalTensor<ElementC> gmBlockC;
|
||||
LayoutC layoutCInGm;
|
||||
Callback callbackBeforeFixpipe;
|
||||
Callback callbackAfterFixpipe;
|
||||
|
||||
CATLASS_DEVICE
|
||||
L1TileMmadParams() = default;
|
||||
};
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitL1(Arch::Resource<ArchTag> &resource, uint32_t l1BufAddrStart)
|
||||
{
|
||||
uint32_t l1AOffset = l1BufAddrStart;
|
||||
for (uint32_t i = 0; i < L1A_STAGES; ++i) {
|
||||
l1ATensorList[i] = resource.l1Buf.template GetBufferByByte<ElementA>(l1AOffset + L1A_TILE_SIZE * i);
|
||||
l1AEventList[i] = i;
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[i]);
|
||||
}
|
||||
uint32_t l1BOffset = l1BufAddrStart + L1A_TILE_SIZE * L1A_STAGES;
|
||||
for (uint32_t i = 0; i < L1B_STAGES; ++i) {
|
||||
l1BTensorList[i] = resource.l1Buf.template GetBufferByByte<ElementB>(l1BOffset + L1B_TILE_SIZE * i);
|
||||
l1BEventList[i] = i + L1A_STAGES;
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitL0A(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
for (uint32_t i = 0; i < L0A_STAGES; ++i) {
|
||||
l0ATensorList[i] = resource.l0ABuf.template GetBufferByByte<ElementA>(L0A_TILE_SIZE * i);
|
||||
l0AEventList[i] = i;
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitL0B(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
for (uint32_t i = 0; i < L0B_STAGES; ++i) {
|
||||
l0BTensorList[i] = resource.l0BBuf.template GetBufferByByte<ElementB>(L0B_TILE_SIZE * i);
|
||||
l0BEventList[i] = i + L0A_STAGES;
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitL0C(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
for (uint32_t i = 0; i < L0C_STAGES; ++i) {
|
||||
l0CTensorList[i] = resource.l0CBuf.template GetBufferByByte<ElementAccumulator>(L0C_TILE_SIZE * i);
|
||||
l0CEventList[i] = i;
|
||||
AscendC::SetFlag<AscendC::HardEvent::FIX_M>(l0CEventList[i]);
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void L1TileMmad(L1TileMmadParams const ¶ms)
|
||||
{
|
||||
uint32_t mPartLoop = CeilDiv<L0TileShape::M>(params.mRound);
|
||||
uint32_t nPartLoop = CeilDiv<L0TileShape::N>(params.nRound);
|
||||
uint32_t kPartLoop = CeilDiv<L0TileShape::K>(params.kActual);
|
||||
auto &l1ATensor = l1ATensorList[params.l1AListId];
|
||||
auto &l1BTensor = l1BTensorList[params.l1BListId];
|
||||
|
||||
auto &l0CTensor = l0CTensorList[l0CListId];
|
||||
LayoutCInL0 layoutCInL0 = LayoutCInL0::MakeLayoutInL0C(MakeCoord(params.mRound, params.nRound));
|
||||
|
||||
if constexpr (!ENABLE_UNIT_FLAG) {
|
||||
if (params.isKLoopFirst) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::FIX_M>(l0CEventList[l0CListId]);
|
||||
}
|
||||
}
|
||||
|
||||
for (uint32_t mPartIdx = 0; mPartIdx < mPartLoop; ++mPartIdx) {
|
||||
uint32_t mPartActual =
|
||||
(mPartIdx < mPartLoop - 1) ? L0TileShape::M : (params.mRound - mPartIdx * L0TileShape::M);
|
||||
|
||||
for (uint32_t kPartIdx = 0; kPartIdx < kPartLoop; ++kPartIdx) {
|
||||
uint32_t kPartActual =
|
||||
(kPartIdx < kPartLoop - 1) ? L0TileShape::K : (params.kActual - kPartIdx * L0TileShape::K);
|
||||
|
||||
auto &l0ATile = l0ATensorList[l0AListId];
|
||||
auto layoutAInL0 = LayoutAInL0::template MakeLayout<ElementA>(mPartActual, kPartActual);
|
||||
auto l1AOffset = MakeCoord(mPartIdx, kPartIdx) * L0TileShape::ToCoordMK();
|
||||
auto l1ATile = l1ATensor[L1A_LAYOUT.GetOffset(l1AOffset)];
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[l0AListId]);
|
||||
if ((mPartIdx == 0) && (kPartIdx == 0)) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE1>(l1AEventList[params.l1AListId]);
|
||||
}
|
||||
copyL1ToL0A(l0ATile, l1ATile, layoutAInL0, L1A_LAYOUT);
|
||||
if ((mPartIdx == mPartLoop - 1) && (kPartIdx == kPartLoop - 1)) {
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1AEventList[params.l1AListId]);
|
||||
}
|
||||
|
||||
for (uint32_t nPartIdx = 0; nPartIdx < nPartLoop; ++nPartIdx) {
|
||||
uint32_t nPartActual =
|
||||
(nPartIdx < nPartLoop - 1) ? L0TileShape::N : (params.nRound - nPartIdx * L0TileShape::N);
|
||||
|
||||
auto &l0BTile = l0BTensorList[l0BListId];
|
||||
auto layoutBInL0 = LayoutBInL0::template MakeLayout<ElementB>(kPartActual, nPartActual);
|
||||
auto l1BOffset = MakeCoord(kPartIdx, nPartIdx) * L0TileShape::ToCoordKN();
|
||||
auto l1BTile = l1BTensor[L1B_LAYOUT.GetOffset(l1BOffset)];
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[l0BListId]);
|
||||
if ((kPartIdx == 0) && (nPartIdx == 0)) {
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE1>(l1BEventList[params.l1BListId]);
|
||||
}
|
||||
copyL1ToL0B(l0BTile, l1BTile, layoutBInL0, L1B_LAYOUT);
|
||||
if ((kPartIdx == kPartLoop - 1) && (nPartIdx == nPartLoop - 1)) {
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_MTE2>(l1BEventList[params.l1BListId]);
|
||||
}
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE1_M>(EVENT_ID0);
|
||||
|
||||
auto l0COffset = MakeCoord(mPartIdx, nPartIdx) * L0TileShape::ToCoordMN();
|
||||
auto l0CTile = l0CTensor[layoutCInL0.GetOffset(l0COffset)];
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE1_M>(EVENT_ID0);
|
||||
// If the current tile is the first tile on the k axis, the accumulator needs to be reset to 0
|
||||
bool initC = (params.isKLoopFirst && (kPartIdx == 0));
|
||||
// If the unit flag is enabled, the unit flag is set according to the calculation progress
|
||||
uint8_t unitFlag = 0b00;
|
||||
if constexpr (ENABLE_UNIT_FLAG) {
|
||||
if (params.isKLoopLast && (mPartIdx == mPartLoop - 1) && (kPartIdx == kPartLoop - 1) &&
|
||||
(nPartIdx == nPartLoop - 1)) {
|
||||
unitFlag = 0b11;
|
||||
} else {
|
||||
unitFlag = 0b10;
|
||||
}
|
||||
}
|
||||
tileMmad(l0CTile, l0ATile, l0BTile, mPartActual, nPartActual, kPartActual, initC, unitFlag);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0BEventList[l0BListId]);
|
||||
l0BListId = (l0BListId + 1 < L0B_STAGES) ? (l0BListId + 1) : 0;
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_MTE1>(l0AEventList[l0AListId]);
|
||||
l0AListId = (l0AListId + 1 < L0A_STAGES) ? (l0AListId + 1) : 0;
|
||||
}
|
||||
}
|
||||
|
||||
if (params.isKLoopLast) {
|
||||
auto layoutCInGm = params.layoutCInGm;
|
||||
|
||||
params.callbackBeforeFixpipe();
|
||||
|
||||
if constexpr (!ENABLE_UNIT_FLAG) {
|
||||
AscendC::SetFlag<AscendC::HardEvent::M_FIX>(l0CEventList[l0CListId]);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::M_FIX>(l0CEventList[l0CListId]);
|
||||
copyL0CToGm(params.gmBlockC, l0CTensor, layoutCInGm, layoutCInL0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::FIX_M>(l0CEventList[l0CListId]);
|
||||
} else {
|
||||
copyL0CToGm(params.gmBlockC, l0CTensor, layoutCInGm, layoutCInL0, 0b11);
|
||||
}
|
||||
l0CListId = (l0CListId + 1 < L0C_STAGES) ? (l0CListId + 1) : 0;
|
||||
|
||||
params.callbackAfterFixpipe();
|
||||
}
|
||||
}
|
||||
|
||||
AscendC::LocalTensor<ElementA> l1ATensorList[L1A_STAGES];
|
||||
AscendC::LocalTensor<ElementB> l1BTensorList[L1B_STAGES];
|
||||
int32_t l1AEventList[L1A_STAGES];
|
||||
int32_t l1BEventList[L1B_STAGES];
|
||||
uint32_t l1AListId{0};
|
||||
uint32_t l1BListId{0};
|
||||
|
||||
AscendC::LocalTensor<ElementA> l0ATensorList[L0A_STAGES];
|
||||
int32_t l0AEventList[L0A_STAGES];
|
||||
uint32_t l0AListId{0};
|
||||
|
||||
AscendC::LocalTensor<ElementB> l0BTensorList[L0B_STAGES];
|
||||
int32_t l0BEventList[L0B_STAGES];
|
||||
uint32_t l0BListId{0};
|
||||
|
||||
AscendC::LocalTensor<ElementAccumulator> l0CTensorList[L0C_STAGES_];
|
||||
int32_t l0CEventList[L0C_STAGES_];
|
||||
uint32_t l0CListId{0};
|
||||
|
||||
L1TileMmadParams l1TileMmadParamsList[PRELOAD_STAGES];
|
||||
uint32_t l1TileMmadParamsId{0};
|
||||
uint32_t preloadCount{0};
|
||||
|
||||
TileMmad tileMmad;
|
||||
CopyGmToL1A copyGmToL1A;
|
||||
CopyGmToL1B copyGmToL1B;
|
||||
CopyL1ToL0A copyL1ToL0A;
|
||||
CopyL1ToL0B copyL1ToL0B;
|
||||
CopyL0CToGm copyL0CToGm;
|
||||
|
||||
bool isFirstLoad{true};
|
||||
AscendC::GlobalTensor<ElementA> lastGmBlockA;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Gemm::Block
|
||||
@@ -0,0 +1,28 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#pragma once
|
||||
#include "catlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
namespace Catlass::Gemm {
|
||||
|
||||
template <uint32_t PRELOAD_STAGES_, uint32_t L1A_STAGES_, uint32_t L1B_STAGES_, uint32_t L0A_STAGES_,
|
||||
uint32_t L0B_STAGES_, uint32_t L0C_STAGES_, bool ENABLE_UNIT_FLAG_, bool ENABLE_SHUFFLE_K_>
|
||||
struct MmadAtlasA2PreloadAsyncWithCallbackResidentA : public MmadAtlasA2Async {
|
||||
static constexpr uint32_t PRELOAD_STAGES = PRELOAD_STAGES_; // Stages of emitting load instruction in advance
|
||||
static constexpr uint32_t L1A_STAGES = L1A_STAGES_;
|
||||
static constexpr uint32_t L1B_STAGES = L1B_STAGES_;
|
||||
static constexpr uint32_t L0A_STAGES = L0A_STAGES_;
|
||||
static constexpr uint32_t L0B_STAGES = L0B_STAGES_;
|
||||
static constexpr uint32_t L0C_STAGES = L0C_STAGES_;
|
||||
static constexpr bool ENABLE_UNIT_FLAG = ENABLE_UNIT_FLAG_;
|
||||
static constexpr bool ENABLE_SHUFFLE_K = ENABLE_SHUFFLE_K_;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Gemm
|
||||
@@ -0,0 +1,383 @@
|
||||
/*
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef ACT_GEMM_KERNEL_GROUPED_MATMUL_M_MULTISTAGE_WORKSPACE_BF16_FP16_HPP
|
||||
#define ACT_GEMM_KERNEL_GROUPED_MATMUL_M_MULTISTAGE_WORKSPACE_BF16_FP16_HPP
|
||||
|
||||
#include "ascendc/basic_api/interface/kernel_operator_list_tensor_intf.h"
|
||||
#include "../../raw_distributed/cam_moe_distribute_combine.h"
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/cross_core_sync.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/coord.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
|
||||
namespace Catlass::Gemm::Kernel {
|
||||
|
||||
template <TemplateMC2TypeClass, class BlockMmad_, class BlockEpilogue_, class BlockScheduler_,
|
||||
uint32_t WORKSPACE_STAGES_, class ElementGroupList_>
|
||||
class GroupedMatmulSliceMMultiStageWorkspace
|
||||
{
|
||||
public:
|
||||
using BlockMmad = BlockMmad_;
|
||||
using ArchTag = typename BlockMmad::ArchTag;
|
||||
using L1TileShape = typename BlockMmad::L1TileShape;
|
||||
using ElementA = typename BlockMmad::ElementA;
|
||||
using LayoutA = typename BlockMmad::LayoutA;
|
||||
using ElementB = typename BlockMmad::ElementB;
|
||||
using LayoutB = typename BlockMmad::LayoutB;
|
||||
using ElementC = typename BlockMmad::ElementC;
|
||||
using LayoutC = typename BlockMmad::LayoutC;
|
||||
using ElementAccumulator = typename BlockMmad::ElementAccumulator;
|
||||
|
||||
using BlockEpilogue = BlockEpilogue_;
|
||||
using ElementScale = typename BlockEpilogue::ElementRawScale;
|
||||
using LayoutScale = typename BlockEpilogue::LayoutScale;
|
||||
using ElementPerTokenScale = typename BlockEpilogue::ElementPerTokenScale;
|
||||
using LayoutPerTokenScale = typename BlockEpilogue::LayoutPerTokenScale;
|
||||
using ElementD = typename BlockEpilogue::ElementD;
|
||||
using LayoutD = typename BlockEpilogue::LayoutD;
|
||||
using EpilogueParams = typename BlockEpilogue::Params;
|
||||
|
||||
using BlockScheduler = BlockScheduler_;
|
||||
static constexpr uint32_t WORKSPACE_STAGES = WORKSPACE_STAGES_;
|
||||
using ElementGroupList = ElementGroupList_;
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
// Data members
|
||||
GemmCoord problemShape;
|
||||
uint32_t problemCount;
|
||||
__gm__ ElementGroupList_ *ptrGroupList;
|
||||
__gm__ ElementA *ptrA;
|
||||
LayoutA layoutA;
|
||||
__gm__ ElementB *ptrB;
|
||||
LayoutB layoutB;
|
||||
__gm__ ElementScale *ptrScale;
|
||||
LayoutScale layoutScale;
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale;
|
||||
LayoutPerTokenScale layoutPerTokenScale;
|
||||
__gm__ ElementD *ptrD;
|
||||
LayoutD layoutD;
|
||||
GM_ADDR ptrWorkspace;
|
||||
void *combiner;
|
||||
|
||||
// Methods
|
||||
CATLASS_DEVICE
|
||||
Params() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params(GemmCoord problemShape_, uint32_t problemCount_, GM_ADDR ptrGroupList_, GM_ADDR ptrA_, LayoutA layoutA_,
|
||||
GM_ADDR ptrB_, LayoutB layoutB_, GM_ADDR ptrScale_, LayoutScale layoutScale_, GM_ADDR ptrPerTokenScale_,
|
||||
LayoutPerTokenScale layoutPerTokenScale_, GM_ADDR ptrD_, LayoutD layoutD_, GM_ADDR ptrWorkspace_,
|
||||
void *combiner_)
|
||||
: problemShape(problemShape_),
|
||||
problemCount(problemCount_),
|
||||
ptrGroupList(reinterpret_cast<__gm__ ElementGroupList *>(ptrGroupList_)),
|
||||
ptrA(reinterpret_cast<__gm__ ElementA *>(ptrA_)),
|
||||
layoutA(layoutA_),
|
||||
ptrB(reinterpret_cast<__gm__ ElementB *>(ptrB_)),
|
||||
layoutB(layoutB_),
|
||||
ptrScale(reinterpret_cast<__gm__ ElementScale *>(ptrScale_)),
|
||||
layoutScale(layoutScale_),
|
||||
ptrPerTokenScale(reinterpret_cast<__gm__ ElementPerTokenScale *>(ptrPerTokenScale_)),
|
||||
layoutPerTokenScale(layoutPerTokenScale_),
|
||||
ptrD(reinterpret_cast<__gm__ ElementD *>(ptrD_)),
|
||||
layoutD(layoutD_),
|
||||
ptrWorkspace(ptrWorkspace_),
|
||||
combiner(combiner_)
|
||||
{}
|
||||
};
|
||||
|
||||
// Methods
|
||||
CATLASS_DEVICE
|
||||
GroupedMatmulSliceMMultiStageWorkspace()
|
||||
{
|
||||
Arch::FlagID flagId = 0;
|
||||
for (uint32_t stageId = 0; stageId < WORKSPACE_STAGES; ++stageId) {
|
||||
flagAicFinishStoreList[stageId] = Arch::CrossCoreFlag(flagId++);
|
||||
flagAivFinishComputeList[stageId] = Arch::CrossCoreFlag(flagId++);
|
||||
aicWaitFuncList[stageId] = {this, stageId};
|
||||
aicSetFuncList[stageId] = {this, stageId};
|
||||
}
|
||||
}
|
||||
|
||||
template <int32_t CORE_TYPE = g_coreType>
|
||||
CATLASS_DEVICE void operator()(Params const ¶ms);
|
||||
|
||||
template <>
|
||||
CATLASS_DEVICE void operator()<AscendC::AIC>(Params const ¶ms)
|
||||
{
|
||||
BlockScheduler blockScheduler;
|
||||
BlockMmad blockMmad(resource);
|
||||
|
||||
// Represent the full gm
|
||||
AscendC::GlobalTensor<ElementA> gmA;
|
||||
gmA.SetGlobalBuffer(params.ptrA);
|
||||
AscendC::GlobalTensor<ElementB> gmB;
|
||||
AscendC::ListTensorDesc gmBlistTensorDesc(reinterpret_cast<__gm__ void *>(params.ptrB));
|
||||
if constexpr (!(EXEC_FLAG & EXEC_FLAG_TENSOR_LIST)) {
|
||||
gmB.SetGlobalBuffer(reinterpret_cast<__gm__ ElementB *>(gmBlistTensorDesc.GetDataPtr<int32_t>(0)));
|
||||
}
|
||||
AscendC::GlobalTensor<ElementGroupList> groupList;
|
||||
groupList.SetGlobalBuffer(params.ptrGroupList);
|
||||
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
int64_t gmGroupOffsetA = 0;
|
||||
int64_t gmGroupOffsetB = 0;
|
||||
|
||||
AscendC::GlobalTensor<ElementC> gmC;
|
||||
gmC.SetGlobalBuffer(reinterpret_cast<__gm__ ElementC *>(params.ptrWorkspace));
|
||||
auto layoutC = layout::RowMajor{L1TileShape::M * coreNum * WORKSPACE_STAGES, L1TileShape::N};
|
||||
|
||||
uint32_t stageId = 0;
|
||||
uint32_t stageUsed = 0;
|
||||
uint32_t startCoreIdx = 0;
|
||||
for (uint32_t groupIdx = 0; groupIdx < params.problemCount; ++groupIdx) {
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_TENSOR_LIST) {
|
||||
gmB.SetGlobalBuffer(reinterpret_cast<__gm__ ElementB *>(
|
||||
gmBlistTensorDesc.GetDataPtr<int32_t>(groupIdx)));
|
||||
}
|
||||
uint32_t currentM = (groupIdx == 0) ? groupList.GetValue(groupIdx)
|
||||
: (groupList.GetValue(groupIdx) - groupList.GetValue(groupIdx - 1));
|
||||
GemmCoord inGroupProblemShape{currentM, params.problemShape.n(), params.problemShape.k()};
|
||||
|
||||
LayoutA layoutA = params.layoutA.GetTileLayout(inGroupProblemShape.GetCoordMK());
|
||||
LayoutB layoutB = params.layoutB;
|
||||
|
||||
blockScheduler.Update(inGroupProblemShape, MakeCoord(L1TileShape::M, L1TileShape::N));
|
||||
uint32_t coreLoops = blockScheduler.GetCoreLoops();
|
||||
|
||||
// Determine the starting loopIdx of the current core under the current
|
||||
// groupIdx
|
||||
uint32_t startLoopIdx = ((coreIdx < startCoreIdx) ? (coreIdx + coreNum) : coreIdx) - startCoreIdx;
|
||||
// Loop through the matmul of each groupIdx
|
||||
for (uint32_t loopIdx = startLoopIdx; loopIdx < coreLoops; loopIdx += coreNum) {
|
||||
// Compute block location
|
||||
GemmCoord blockCoord = blockScheduler.GetBlockCoord(loopIdx);
|
||||
GemmCoord actualBlockShape = blockScheduler.GetActualBlockShape(blockCoord);
|
||||
|
||||
Callback callbackBeforeFixpipe{};
|
||||
if (stageUsed == WORKSPACE_STAGES) {
|
||||
callbackBeforeFixpipe = MakeCallback(&aicWaitFuncList[stageId]);
|
||||
} else {
|
||||
++stageUsed;
|
||||
}
|
||||
Callback callbackAfterFixpipe = MakeCallback(&aicSetFuncList[stageId]);
|
||||
|
||||
// Compute initial location in logical coordinates
|
||||
MatrixCoord offsetA{blockCoord.m() * L1TileShape::M, blockCoord.k() * L1TileShape::K};
|
||||
MatrixCoord offsetB{blockCoord.k() * L1TileShape::K, blockCoord.n() * L1TileShape::N};
|
||||
MatrixCoord offsetC{(stageId * coreNum + coreIdx) * L1TileShape::M, 0};
|
||||
int64_t gmOffsetA = layoutA.GetOffset(offsetA);
|
||||
int64_t gmOffsetB = layoutB.GetOffset(offsetB);
|
||||
int64_t gmOffsetC = layoutC.GetOffset(offsetC);
|
||||
|
||||
// Compute block-scoped matrix multiply-add
|
||||
if constexpr (BlockMmad::DispatchPolicy::ASYNC) {
|
||||
blockMmad(gmA[gmGroupOffsetA + gmOffsetA], layoutA, gmB[gmGroupOffsetB + gmOffsetB], layoutB,
|
||||
gmC[gmOffsetC], layoutC, actualBlockShape, callbackBeforeFixpipe, callbackAfterFixpipe);
|
||||
} else {
|
||||
callbackBeforeFixpipe();
|
||||
blockMmad(gmA[gmGroupOffsetA + gmOffsetA], layoutA, gmB[gmGroupOffsetB + gmOffsetB], layoutB,
|
||||
gmC[gmOffsetC], layoutC, actualBlockShape);
|
||||
callbackAfterFixpipe();
|
||||
}
|
||||
|
||||
stageId = (stageId + 1 < WORKSPACE_STAGES) ? (stageId + 1) : 0;
|
||||
}
|
||||
|
||||
gmGroupOffsetA += inGroupProblemShape.m() * inGroupProblemShape.k();
|
||||
if constexpr (!(EXEC_FLAG & EXEC_FLAG_TENSOR_LIST)) {
|
||||
gmGroupOffsetB += inGroupProblemShape.k() * inGroupProblemShape.n();
|
||||
}
|
||||
startCoreIdx = (startCoreIdx + coreLoops) % coreNum;
|
||||
}
|
||||
|
||||
if constexpr (BlockMmad::DispatchPolicy::ASYNC) {
|
||||
blockMmad.SynchronizeBlock();
|
||||
}
|
||||
|
||||
while (stageUsed > 0) {
|
||||
uint32_t aivComputeStageId =
|
||||
(stageId >= stageUsed) ? (stageId - stageUsed) : (stageId + WORKSPACE_STAGES - stageUsed);
|
||||
Arch::CrossCoreWaitFlag(flagAivFinishComputeList[aivComputeStageId]);
|
||||
--stageUsed;
|
||||
}
|
||||
}
|
||||
|
||||
template <>
|
||||
CATLASS_DEVICE void operator()<AscendC::AIV>(Params const ¶ms)
|
||||
{
|
||||
auto *combiner = (MoeDistributeCombineImpl::CamMoeDistributeCombine<TemplateMC2TypeFunc> *)params.combiner;
|
||||
{
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
if (get_subblockid() == 0) {
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(MoeDistributeCombineImpl::RECV_SYNC_EVENT_ID);
|
||||
}
|
||||
}
|
||||
BlockScheduler blockScheduler;
|
||||
BlockEpilogue blockEpilogue(resource, combiner->GetCalcInfo());
|
||||
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx() / AscendC::GetSubBlockNum();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
int64_t gmGroupOffsetScale = 0;
|
||||
int64_t gmGroupOffsetPerTokenScale = 0;
|
||||
int64_t gmGroupOffsetD = 0;
|
||||
AscendC::GlobalTensor<ElementGroupList> groupList;
|
||||
groupList.SetGlobalBuffer(params.ptrGroupList);
|
||||
|
||||
AscendC::GlobalTensor<ElementC> gmC;
|
||||
gmC.SetGlobalBuffer(reinterpret_cast<__gm__ ElementC *>(params.ptrWorkspace));
|
||||
auto layoutC = layout::RowMajor{L1TileShape::M * coreNum * WORKSPACE_STAGES, L1TileShape::N};
|
||||
|
||||
uint32_t stageId = 0;
|
||||
uint32_t startCoreIdx = 0;
|
||||
AscendC::ListTensorDesc gmScaleListTensor;
|
||||
gmScaleListTensor = AscendC::ListTensorDesc(reinterpret_cast<__gm__ void *>(params.ptrScale));
|
||||
__gm__ ElementScale* gmScalePtr;
|
||||
if constexpr (!(EXEC_FLAG & EXEC_FLAG_TENSOR_LIST)) {
|
||||
gmScalePtr = reinterpret_cast<__gm__ ElementScale*>(gmScaleListTensor.GetDataPtr<int32_t>(0));
|
||||
}
|
||||
for (uint32_t groupIdx = 0; groupIdx < params.problemCount; ++groupIdx) {
|
||||
uint32_t currentM = (groupIdx == 0) ? groupList.GetValue(groupIdx)
|
||||
: (groupList.GetValue(groupIdx) - groupList.GetValue(groupIdx - 1));
|
||||
GemmCoord inGroupProblemShape{currentM, params.problemShape.n(), params.problemShape.k()};
|
||||
|
||||
LayoutScale layoutScale = params.layoutScale;
|
||||
LayoutPerTokenScale layoutPerTokenScale =
|
||||
params.layoutPerTokenScale.GetTileLayout(inGroupProblemShape.template GetCoordByAxis<0>());
|
||||
LayoutD layoutD = params.layoutD.GetTileLayout(inGroupProblemShape.GetCoordMN());
|
||||
EpilogueParams epilogueParams;
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_TENSOR_LIST) {
|
||||
gmScalePtr = reinterpret_cast<__gm__ ElementScale*>(
|
||||
gmScaleListTensor.GetDataPtr<int32_t>(groupIdx));
|
||||
epilogueParams = EpilogueParams {
|
||||
gmScalePtr, layoutScale,
|
||||
params.ptrPerTokenScale + gmGroupOffsetPerTokenScale, layoutPerTokenScale,
|
||||
params.ptrD + gmGroupOffsetD, layoutD};
|
||||
} else {
|
||||
epilogueParams = EpilogueParams{gmScalePtr + gmGroupOffsetScale,
|
||||
layoutScale,
|
||||
params.ptrPerTokenScale + gmGroupOffsetPerTokenScale,
|
||||
layoutPerTokenScale,
|
||||
params.ptrD + gmGroupOffsetD,
|
||||
layoutD};
|
||||
}
|
||||
blockScheduler.Update(inGroupProblemShape, L1TileShape::ToCoordMN());
|
||||
blockEpilogue.UpdateParams(epilogueParams);
|
||||
uint32_t coreLoops = blockScheduler.GetCoreLoops();
|
||||
|
||||
GemmCoord blockShapeMNK = L1TileShape::ToCoord();
|
||||
uint32_t startLoopIdx = ((coreIdx < startCoreIdx) ? (coreIdx + coreNum) : coreIdx) - startCoreIdx;
|
||||
for (uint32_t loopIdx = startLoopIdx; loopIdx < coreLoops; loopIdx += coreNum) {
|
||||
GemmCoord blockCoordMNK = blockScheduler.GetBlockCoord(loopIdx);
|
||||
GemmCoord actualBlockShapeMNK = blockScheduler.GetActualBlockShape(blockCoordMNK);
|
||||
|
||||
MatrixCoord offsetC{(stageId * coreNum + coreIdx) * L1TileShape::M, 0};
|
||||
int64_t gmOffsetC = layoutC.GetOffset(offsetC);
|
||||
auto gmBlockC = gmC[gmOffsetC];
|
||||
auto layoutBlockC = layoutC.GetTileLayout(actualBlockShapeMNK.GetCoordMN());
|
||||
|
||||
Arch::CrossCoreWaitFlag(flagAicFinishStoreList[stageId]);
|
||||
blockEpilogue(gmGroupOffsetD, groupIdx, blockShapeMNK, blockCoordMNK, actualBlockShapeMNK, gmBlockC,
|
||||
layoutBlockC);
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(flagAivFinishComputeList[stageId]);
|
||||
|
||||
stageId = (stageId + 1 < WORKSPACE_STAGES) ? (stageId + 1) : 0;
|
||||
}
|
||||
|
||||
if constexpr (!(EXEC_FLAG & EXEC_FLAG_TENSOR_LIST)) {
|
||||
gmGroupOffsetScale += inGroupProblemShape.n();
|
||||
}
|
||||
gmGroupOffsetPerTokenScale += inGroupProblemShape.m();
|
||||
gmGroupOffsetD += inGroupProblemShape.m() * inGroupProblemShape.n();
|
||||
|
||||
startCoreIdx = (startCoreIdx + coreLoops) % coreNum;
|
||||
}
|
||||
}
|
||||
|
||||
icache_preload(4);
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
if (get_subblockid() == 0) {
|
||||
resource.pipe.Init();
|
||||
combiner->TPipeSet(&resource.pipe);
|
||||
combiner->AllToAllSend();
|
||||
combiner->TPipeSet(nullptr);
|
||||
resource.pipe.Destroy();
|
||||
} else {
|
||||
resource.pipe.Init();
|
||||
combiner->TPipeSet(&resource.pipe);
|
||||
combiner->ReducePermute();
|
||||
combiner->TPipeSet(nullptr);
|
||||
resource.pipe.Destroy();
|
||||
}
|
||||
} else {
|
||||
resource.pipe.Init();
|
||||
combiner->TPipeSet(&resource.pipe);
|
||||
combiner->Process();
|
||||
combiner->TPipeSet(nullptr);
|
||||
resource.pipe.Destroy();
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
friend struct AicWaitFunc;
|
||||
friend struct AicSetFunc;
|
||||
|
||||
struct AicWaitFunc {
|
||||
using MatmulKernel =
|
||||
GroupedMatmulSliceMMultiStageWorkspace<TemplateMC2TypeFunc, BlockMmad, BlockEpilogue,
|
||||
BlockScheduler, WORKSPACE_STAGES, ElementGroupList>;
|
||||
|
||||
CATLASS_DEVICE
|
||||
AicWaitFunc() = default;
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()() const
|
||||
{
|
||||
Arch::CrossCoreWaitFlag(ptr->flagAivFinishComputeList[stageId]);
|
||||
}
|
||||
|
||||
MatmulKernel *ptr{nullptr};
|
||||
uint32_t stageId;
|
||||
};
|
||||
|
||||
struct AicSetFunc {
|
||||
using MatmulKernel =
|
||||
GroupedMatmulSliceMMultiStageWorkspace<TemplateMC2TypeFunc, BlockMmad, BlockEpilogue,
|
||||
BlockScheduler, WORKSPACE_STAGES, ElementGroupList>;
|
||||
|
||||
CATLASS_DEVICE
|
||||
AicSetFunc() = default;
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()() const
|
||||
{
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(ptr->flagAicFinishStoreList[stageId]);
|
||||
}
|
||||
|
||||
MatmulKernel *ptr{nullptr};
|
||||
uint32_t stageId;
|
||||
};
|
||||
|
||||
Arch::CrossCoreFlag flagAicFinishStoreList[WORKSPACE_STAGES];
|
||||
Arch::CrossCoreFlag flagAivFinishComputeList[WORKSPACE_STAGES];
|
||||
|
||||
AicWaitFunc aicWaitFuncList[WORKSPACE_STAGES];
|
||||
AicSetFunc aicSetFuncList[WORKSPACE_STAGES];
|
||||
Arch::Resource<ArchTag> resource;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Gemm::Kernel
|
||||
|
||||
#endif // ACT_GEMM_KERNEL_GROUPED_MATMUL_M_MULTISTAGE_WORKSPACE_BF16_FP16_HPP
|
||||
@@ -0,0 +1,383 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef ACT_GEMM_KERNEL_GROUPED_MATMUL_M_PER_TOKEN_DEQUANT_MULTISTAGE_WORKSPACE_HPP
|
||||
#define ACT_GEMM_KERNEL_GROUPED_MATMUL_M_PER_TOKEN_DEQUANT_MULTISTAGE_WORKSPACE_HPP
|
||||
|
||||
#include "ascendc/basic_api/interface/kernel_operator_list_tensor_intf.h"
|
||||
#include "../../raw_distributed/cam_moe_distribute_combine.h"
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/cross_core_sync.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/coord.hpp"
|
||||
#include "catlass/detail/callback.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
|
||||
namespace Catlass::Gemm::Kernel {
|
||||
|
||||
template <TemplateMC2TypeClass, class BlockMmad_, class BlockEpilogue_, class BlockScheduler_,
|
||||
uint32_t WORKSPACE_STAGES_, class ElementGroupList_>
|
||||
class GroupedMatmulSliceMPerTokenDequantMultiStageWorkspace
|
||||
{
|
||||
public:
|
||||
using BlockMmad = BlockMmad_;
|
||||
using ArchTag = typename BlockMmad::ArchTag;
|
||||
using L1TileShape = typename BlockMmad::L1TileShape;
|
||||
using ElementA = typename BlockMmad::ElementA;
|
||||
using LayoutA = typename BlockMmad::LayoutA;
|
||||
using ElementB = typename BlockMmad::ElementB;
|
||||
using LayoutB = typename BlockMmad::LayoutB;
|
||||
using ElementC = typename BlockMmad::ElementC;
|
||||
using LayoutC = typename BlockMmad::LayoutC;
|
||||
using ElementAccumulator = typename BlockMmad::ElementAccumulator;
|
||||
|
||||
using BlockEpilogue = BlockEpilogue_;
|
||||
using ElementScale = typename BlockEpilogue::ElementRawScale;
|
||||
using LayoutScale = typename BlockEpilogue::LayoutScale;
|
||||
using ElementPerTokenScale = typename BlockEpilogue::ElementPerTokenScale;
|
||||
using LayoutPerTokenScale = typename BlockEpilogue::LayoutPerTokenScale;
|
||||
using ElementD = typename BlockEpilogue::ElementD;
|
||||
using LayoutD = typename BlockEpilogue::LayoutD;
|
||||
using EpilogueParams = typename BlockEpilogue::Params;
|
||||
|
||||
using BlockScheduler = BlockScheduler_;
|
||||
static constexpr uint32_t WORKSPACE_STAGES = WORKSPACE_STAGES_;
|
||||
using ElementGroupList = ElementGroupList_;
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
// Data members
|
||||
GemmCoord problemShape;
|
||||
uint32_t problemCount;
|
||||
__gm__ ElementGroupList_ *ptrGroupList;
|
||||
__gm__ ElementA *ptrA;
|
||||
LayoutA layoutA;
|
||||
__gm__ ElementB *ptrB;
|
||||
LayoutB layoutB;
|
||||
__gm__ ElementScale *ptrScale;
|
||||
LayoutScale layoutScale;
|
||||
__gm__ ElementPerTokenScale *ptrPerTokenScale;
|
||||
LayoutPerTokenScale layoutPerTokenScale;
|
||||
__gm__ ElementD *ptrD;
|
||||
LayoutD layoutD;
|
||||
GM_ADDR ptrWorkspace;
|
||||
void *combiner;
|
||||
|
||||
// Methods
|
||||
CATLASS_DEVICE
|
||||
Params() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
Params(GemmCoord problemShape_, uint32_t problemCount_, GM_ADDR ptrGroupList_, GM_ADDR ptrA_, LayoutA layoutA_,
|
||||
GM_ADDR ptrB_, LayoutB layoutB_, GM_ADDR ptrScale_, LayoutScale layoutScale_, GM_ADDR ptrPerTokenScale_,
|
||||
LayoutPerTokenScale layoutPerTokenScale_, GM_ADDR ptrD_, LayoutD layoutD_, GM_ADDR ptrWorkspace_,
|
||||
void *combiner_)
|
||||
: problemShape(problemShape_),
|
||||
problemCount(problemCount_),
|
||||
ptrGroupList(reinterpret_cast<__gm__ ElementGroupList *>(ptrGroupList_)),
|
||||
ptrA(reinterpret_cast<__gm__ ElementA *>(ptrA_)),
|
||||
layoutA(layoutA_),
|
||||
ptrB(reinterpret_cast<__gm__ ElementB *>(ptrB_)),
|
||||
layoutB(layoutB_),
|
||||
ptrScale(reinterpret_cast<__gm__ ElementScale *>(ptrScale_)),
|
||||
layoutScale(layoutScale_),
|
||||
ptrPerTokenScale(reinterpret_cast<__gm__ ElementPerTokenScale *>(ptrPerTokenScale_)),
|
||||
layoutPerTokenScale(layoutPerTokenScale_),
|
||||
ptrD(reinterpret_cast<__gm__ ElementD *>(ptrD_)),
|
||||
layoutD(layoutD_),
|
||||
ptrWorkspace(ptrWorkspace_),
|
||||
combiner(combiner_)
|
||||
{}
|
||||
};
|
||||
|
||||
// Methods
|
||||
CATLASS_DEVICE
|
||||
GroupedMatmulSliceMPerTokenDequantMultiStageWorkspace()
|
||||
{
|
||||
Arch::FlagID flagId = 0;
|
||||
for (uint32_t stageId = 0; stageId < WORKSPACE_STAGES; ++stageId) {
|
||||
flagAicFinishStoreList[stageId] = Arch::CrossCoreFlag(flagId++);
|
||||
flagAivFinishComputeList[stageId] = Arch::CrossCoreFlag(flagId++);
|
||||
aicWaitFuncList[stageId] = {this, stageId};
|
||||
aicSetFuncList[stageId] = {this, stageId};
|
||||
}
|
||||
}
|
||||
|
||||
template <int32_t CORE_TYPE = g_coreType>
|
||||
CATLASS_DEVICE void operator()(Params const ¶ms);
|
||||
|
||||
template <>
|
||||
CATLASS_DEVICE void operator()<AscendC::AIC>(Params const ¶ms)
|
||||
{
|
||||
BlockScheduler blockScheduler;
|
||||
BlockMmad blockMmad(resource);
|
||||
|
||||
// Represent the full gm
|
||||
AscendC::GlobalTensor<ElementA> gmA;
|
||||
gmA.SetGlobalBuffer(params.ptrA);
|
||||
AscendC::GlobalTensor<ElementB> gmB;
|
||||
AscendC::ListTensorDesc gmBlistTensorDesc(reinterpret_cast<__gm__ void *>(params.ptrB));
|
||||
if constexpr (!(EXEC_FLAG & EXEC_FLAG_TENSOR_LIST)) {
|
||||
gmB.SetGlobalBuffer(reinterpret_cast<__gm__ ElementB *>(gmBlistTensorDesc.GetDataPtr<int32_t>(0)));
|
||||
}
|
||||
AscendC::GlobalTensor<ElementGroupList> groupList;
|
||||
groupList.SetGlobalBuffer(params.ptrGroupList);
|
||||
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
int64_t gmGroupOffsetA = 0;
|
||||
int64_t gmGroupOffsetB = 0;
|
||||
|
||||
AscendC::GlobalTensor<ElementC> gmC;
|
||||
gmC.SetGlobalBuffer(reinterpret_cast<__gm__ ElementC *>(params.ptrWorkspace));
|
||||
auto layoutC = layout::RowMajor{L1TileShape::M * coreNum * WORKSPACE_STAGES, L1TileShape::N};
|
||||
|
||||
uint32_t stageId = 0;
|
||||
uint32_t stageUsed = 0;
|
||||
uint32_t startCoreIdx = 0;
|
||||
for (uint32_t groupIdx = 0; groupIdx < params.problemCount; ++groupIdx) {
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_TENSOR_LIST) {
|
||||
gmB.SetGlobalBuffer(reinterpret_cast<__gm__ ElementB *>(
|
||||
gmBlistTensorDesc.GetDataPtr<int32_t>(groupIdx)));
|
||||
}
|
||||
uint32_t currentM = (groupIdx == 0) ? groupList.GetValue(groupIdx)
|
||||
: (groupList.GetValue(groupIdx) - groupList.GetValue(groupIdx - 1));
|
||||
GemmCoord inGroupProblemShape{currentM, params.problemShape.n(), params.problemShape.k()};
|
||||
|
||||
LayoutA layoutA = params.layoutA.GetTileLayout(inGroupProblemShape.GetCoordMK());
|
||||
LayoutB layoutB = params.layoutB;
|
||||
|
||||
blockScheduler.Update(inGroupProblemShape, MakeCoord(L1TileShape::M, L1TileShape::N));
|
||||
uint32_t coreLoops = blockScheduler.GetCoreLoops();
|
||||
|
||||
// Determine the starting loopIdx of the current core under the current
|
||||
// groupIdx
|
||||
uint32_t startLoopIdx = ((coreIdx < startCoreIdx) ? (coreIdx + coreNum) : coreIdx) - startCoreIdx;
|
||||
// Loop through the matmul of each groupIdx
|
||||
for (uint32_t loopIdx = startLoopIdx; loopIdx < coreLoops; loopIdx += coreNum) {
|
||||
// Compute block location
|
||||
GemmCoord blockCoord = blockScheduler.GetBlockCoord(loopIdx);
|
||||
GemmCoord actualBlockShape = blockScheduler.GetActualBlockShape(blockCoord);
|
||||
|
||||
Callback callbackBeforeFixpipe{};
|
||||
if (stageUsed == WORKSPACE_STAGES) {
|
||||
callbackBeforeFixpipe = MakeCallback(&aicWaitFuncList[stageId]);
|
||||
} else {
|
||||
++stageUsed;
|
||||
}
|
||||
Callback callbackAfterFixpipe = MakeCallback(&aicSetFuncList[stageId]);
|
||||
|
||||
// Compute initial location in logical coordinates
|
||||
MatrixCoord offsetA{blockCoord.m() * L1TileShape::M, blockCoord.k() * L1TileShape::K};
|
||||
MatrixCoord offsetB{blockCoord.k() * L1TileShape::K, blockCoord.n() * L1TileShape::N};
|
||||
MatrixCoord offsetC{(stageId * coreNum + coreIdx) * L1TileShape::M, 0};
|
||||
int64_t gmOffsetA = layoutA.GetOffset(offsetA);
|
||||
int64_t gmOffsetB = layoutB.GetOffset(offsetB);
|
||||
int64_t gmOffsetC = layoutC.GetOffset(offsetC);
|
||||
|
||||
// Compute block-scoped matrix multiply-add
|
||||
if constexpr (BlockMmad::DispatchPolicy::ASYNC) {
|
||||
blockMmad(gmA[gmGroupOffsetA + gmOffsetA], layoutA, gmB[gmGroupOffsetB + gmOffsetB], layoutB,
|
||||
gmC[gmOffsetC], layoutC, actualBlockShape, callbackBeforeFixpipe, callbackAfterFixpipe);
|
||||
} else {
|
||||
callbackBeforeFixpipe();
|
||||
blockMmad(gmA[gmGroupOffsetA + gmOffsetA], layoutA, gmB[gmGroupOffsetB + gmOffsetB], layoutB,
|
||||
gmC[gmOffsetC], layoutC, actualBlockShape);
|
||||
callbackAfterFixpipe();
|
||||
}
|
||||
|
||||
stageId = (stageId + 1 < WORKSPACE_STAGES) ? (stageId + 1) : 0;
|
||||
}
|
||||
|
||||
gmGroupOffsetA += inGroupProblemShape.m() * inGroupProblemShape.k();
|
||||
if constexpr (!(EXEC_FLAG & EXEC_FLAG_TENSOR_LIST)) {
|
||||
gmGroupOffsetB += inGroupProblemShape.k() * inGroupProblemShape.n();
|
||||
}
|
||||
startCoreIdx = (startCoreIdx + coreLoops) % coreNum;
|
||||
}
|
||||
|
||||
if constexpr (BlockMmad::DispatchPolicy::ASYNC) {
|
||||
blockMmad.SynchronizeBlock();
|
||||
}
|
||||
|
||||
while (stageUsed > 0) {
|
||||
uint32_t aivComputeStageId =
|
||||
(stageId >= stageUsed) ? (stageId - stageUsed) : (stageId + WORKSPACE_STAGES - stageUsed);
|
||||
Arch::CrossCoreWaitFlag(flagAivFinishComputeList[aivComputeStageId]);
|
||||
--stageUsed;
|
||||
}
|
||||
}
|
||||
|
||||
template <>
|
||||
CATLASS_DEVICE void operator()<AscendC::AIV>(Params const ¶ms)
|
||||
{
|
||||
auto *combiner = (MoeDistributeCombineImpl::CamMoeDistributeCombine<TemplateMC2TypeFunc> *)params.combiner;
|
||||
{
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
if (get_subblockid() == 0) {
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(MoeDistributeCombineImpl::RECV_SYNC_EVENT_ID);
|
||||
}
|
||||
}
|
||||
BlockScheduler blockScheduler;
|
||||
BlockEpilogue blockEpilogue(resource, combiner->GetCalcInfo());
|
||||
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx() / AscendC::GetSubBlockNum();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
int64_t gmGroupOffsetScale = 0;
|
||||
int64_t gmGroupOffsetPerTokenScale = 0;
|
||||
int64_t gmGroupOffsetD = 0;
|
||||
AscendC::GlobalTensor<ElementGroupList> groupList;
|
||||
groupList.SetGlobalBuffer(params.ptrGroupList);
|
||||
|
||||
AscendC::GlobalTensor<ElementC> gmC;
|
||||
gmC.SetGlobalBuffer(reinterpret_cast<__gm__ ElementC *>(params.ptrWorkspace));
|
||||
auto layoutC = layout::RowMajor{L1TileShape::M * coreNum * WORKSPACE_STAGES, L1TileShape::N};
|
||||
|
||||
uint32_t stageId = 0;
|
||||
uint32_t startCoreIdx = 0;
|
||||
AscendC::ListTensorDesc gmScaleListTensor;
|
||||
gmScaleListTensor = AscendC::ListTensorDesc(reinterpret_cast<__gm__ void *>(params.ptrScale));
|
||||
__gm__ ElementScale* gmScalePtr;
|
||||
if constexpr (!(EXEC_FLAG & EXEC_FLAG_TENSOR_LIST)) {
|
||||
gmScalePtr = reinterpret_cast<__gm__ ElementScale*>(gmScaleListTensor.GetDataPtr<int32_t>(0));
|
||||
}
|
||||
for (uint32_t groupIdx = 0; groupIdx < params.problemCount; ++groupIdx) {
|
||||
uint32_t currentM = (groupIdx == 0) ? groupList.GetValue(groupIdx)
|
||||
: (groupList.GetValue(groupIdx) - groupList.GetValue(groupIdx - 1));
|
||||
GemmCoord inGroupProblemShape{currentM, params.problemShape.n(), params.problemShape.k()};
|
||||
|
||||
LayoutScale layoutScale = params.layoutScale;
|
||||
LayoutPerTokenScale layoutPerTokenScale =
|
||||
params.layoutPerTokenScale.GetTileLayout(inGroupProblemShape.template GetCoordByAxis<0>());
|
||||
LayoutD layoutD = params.layoutD.GetTileLayout(inGroupProblemShape.GetCoordMN());
|
||||
EpilogueParams epilogueParams;
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_TENSOR_LIST) {
|
||||
gmScalePtr = reinterpret_cast<__gm__ ElementScale*>(
|
||||
gmScaleListTensor.GetDataPtr<int32_t>(groupIdx));
|
||||
epilogueParams = EpilogueParams {
|
||||
gmScalePtr, layoutScale,
|
||||
params.ptrPerTokenScale + gmGroupOffsetPerTokenScale, layoutPerTokenScale,
|
||||
params.ptrD + gmGroupOffsetD, layoutD};
|
||||
} else {
|
||||
epilogueParams = EpilogueParams{gmScalePtr + gmGroupOffsetScale,
|
||||
layoutScale,
|
||||
params.ptrPerTokenScale + gmGroupOffsetPerTokenScale,
|
||||
layoutPerTokenScale,
|
||||
params.ptrD + gmGroupOffsetD,
|
||||
layoutD};
|
||||
}
|
||||
blockScheduler.Update(inGroupProblemShape, L1TileShape::ToCoordMN());
|
||||
blockEpilogue.UpdateParams(epilogueParams);
|
||||
uint32_t coreLoops = blockScheduler.GetCoreLoops();
|
||||
|
||||
GemmCoord blockShapeMNK = L1TileShape::ToCoord();
|
||||
uint32_t startLoopIdx = ((coreIdx < startCoreIdx) ? (coreIdx + coreNum) : coreIdx) - startCoreIdx;
|
||||
for (uint32_t loopIdx = startLoopIdx; loopIdx < coreLoops; loopIdx += coreNum) {
|
||||
GemmCoord blockCoordMNK = blockScheduler.GetBlockCoord(loopIdx);
|
||||
GemmCoord actualBlockShapeMNK = blockScheduler.GetActualBlockShape(blockCoordMNK);
|
||||
|
||||
MatrixCoord offsetC{(stageId * coreNum + coreIdx) * L1TileShape::M, 0};
|
||||
int64_t gmOffsetC = layoutC.GetOffset(offsetC);
|
||||
auto gmBlockC = gmC[gmOffsetC];
|
||||
auto layoutBlockC = layoutC.GetTileLayout(actualBlockShapeMNK.GetCoordMN());
|
||||
|
||||
Arch::CrossCoreWaitFlag(flagAicFinishStoreList[stageId]);
|
||||
blockEpilogue(gmGroupOffsetD, groupIdx, blockShapeMNK, blockCoordMNK, actualBlockShapeMNK, gmBlockC,
|
||||
layoutBlockC);
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(flagAivFinishComputeList[stageId]);
|
||||
|
||||
stageId = (stageId + 1 < WORKSPACE_STAGES) ? (stageId + 1) : 0;
|
||||
}
|
||||
|
||||
if constexpr (!(EXEC_FLAG & EXEC_FLAG_TENSOR_LIST)) {
|
||||
gmGroupOffsetScale += inGroupProblemShape.n();
|
||||
}
|
||||
gmGroupOffsetPerTokenScale += inGroupProblemShape.m();
|
||||
gmGroupOffsetD += inGroupProblemShape.m() * inGroupProblemShape.n();
|
||||
|
||||
startCoreIdx = (startCoreIdx + coreLoops) % coreNum;
|
||||
}
|
||||
}
|
||||
|
||||
icache_preload(4);
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
if (get_subblockid() == 0) {
|
||||
resource.pipe.Init();
|
||||
combiner->TPipeSet(&resource.pipe);
|
||||
combiner->AllToAllSend();
|
||||
combiner->TPipeSet(nullptr);
|
||||
resource.pipe.Destroy();
|
||||
} else {
|
||||
resource.pipe.Init();
|
||||
combiner->TPipeSet(&resource.pipe);
|
||||
combiner->ReducePermute();
|
||||
combiner->TPipeSet(nullptr);
|
||||
resource.pipe.Destroy();
|
||||
}
|
||||
} else {
|
||||
resource.pipe.Init();
|
||||
combiner->TPipeSet(&resource.pipe);
|
||||
combiner->Process();
|
||||
combiner->TPipeSet(nullptr);
|
||||
resource.pipe.Destroy();
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
friend struct AicWaitFunc;
|
||||
friend struct AicSetFunc;
|
||||
|
||||
struct AicWaitFunc {
|
||||
using MatmulKernel =
|
||||
GroupedMatmulSliceMPerTokenDequantMultiStageWorkspace<TemplateMC2TypeFunc, BlockMmad, BlockEpilogue,
|
||||
BlockScheduler, WORKSPACE_STAGES, ElementGroupList>;
|
||||
|
||||
CATLASS_DEVICE
|
||||
AicWaitFunc() = default;
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()() const
|
||||
{
|
||||
Arch::CrossCoreWaitFlag(ptr->flagAivFinishComputeList[stageId]);
|
||||
}
|
||||
|
||||
MatmulKernel *ptr{nullptr};
|
||||
uint32_t stageId;
|
||||
};
|
||||
|
||||
struct AicSetFunc {
|
||||
using MatmulKernel =
|
||||
GroupedMatmulSliceMPerTokenDequantMultiStageWorkspace<TemplateMC2TypeFunc, BlockMmad, BlockEpilogue,
|
||||
BlockScheduler, WORKSPACE_STAGES, ElementGroupList>;
|
||||
|
||||
CATLASS_DEVICE
|
||||
AicSetFunc() = default;
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()() const
|
||||
{
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(ptr->flagAicFinishStoreList[stageId]);
|
||||
}
|
||||
|
||||
MatmulKernel *ptr{nullptr};
|
||||
uint32_t stageId;
|
||||
};
|
||||
|
||||
Arch::CrossCoreFlag flagAicFinishStoreList[WORKSPACE_STAGES];
|
||||
Arch::CrossCoreFlag flagAivFinishComputeList[WORKSPACE_STAGES];
|
||||
|
||||
AicWaitFunc aicWaitFuncList[WORKSPACE_STAGES];
|
||||
AicSetFunc aicSetFuncList[WORKSPACE_STAGES];
|
||||
Arch::Resource<ArchTag> resource;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Gemm::Kernel
|
||||
|
||||
#endif // ACT_GEMM_KERNEL_GROUPED_MATMUL_M_PER_TOKEN_DEQUANT_MULTISTAGE_WORKSPACE_HPP
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,848 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef CAM_MOE_DISTRIBUTE_COMBINE_H
|
||||
#define CAM_MOE_DISTRIBUTE_COMBINE_H
|
||||
#ifndef OPT_RANK_OFFSET
|
||||
#define OPT_RANK_OFFSET 512
|
||||
#endif
|
||||
|
||||
#include "kernel_operator.h"
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
#include "../../dispatch_gmm_combine_decode_base.h"
|
||||
#include "../../dispatch_gmm_combine_decode_tiling.h"
|
||||
|
||||
namespace MoeDistributeCombineImpl {
|
||||
constexpr uint8_t BUFFER_NUM = 2; // multi-buf
|
||||
constexpr uint32_t STATE_OFFSET = 512;
|
||||
constexpr uint32_t STATE_SIZE = 1024 * 1024; // 1M
|
||||
constexpr uint32_t RANK_SIZE_ON_WIN_512 = 512 * 1024;
|
||||
constexpr uint32_t RANK_SIZE_ON_WIN_256 = 256 * 1024;
|
||||
constexpr uint32_t TP_RANK_SIZE_ON_WIN = 0;
|
||||
constexpr uint32_t UB_ALIGN = 32;
|
||||
constexpr uint32_t SELF_STATE_OFFSET = 256 * 1024;
|
||||
constexpr uint8_t EP_DOMAIN = 0;
|
||||
constexpr uint8_t TP_DOMAIN = 1;
|
||||
constexpr uint64_t WIN_STATE_OFFSET = 512 * 1024;
|
||||
constexpr uint64_t STATE_WIN_OFFSET = 900 * 1024;
|
||||
constexpr uint16_t SEND_SYNC_EVENT_ID = 9;
|
||||
constexpr uint16_t RECV_SYNC_EVENT_ID = 10;
|
||||
|
||||
template <AscendC::HardEvent event>
|
||||
__aicore__ inline void SyncFunc()
|
||||
{
|
||||
int32_t eventID = static_cast<int32_t>(GetTPipePtr()->FetchEventID(event));
|
||||
AscendC::SetFlag<event>(eventID);
|
||||
AscendC::WaitFlag<event>(eventID);
|
||||
}
|
||||
|
||||
using namespace AscendC;
|
||||
|
||||
struct CombineCalcInfo {
|
||||
uint64_t expertPerSizeOnWin_;
|
||||
uint32_t epRankId_;
|
||||
uint32_t epWorldSize_;
|
||||
uint32_t moeExpertPerRankNum_;
|
||||
uint32_t sharedExpertRankNum_;
|
||||
uint32_t axisH_;
|
||||
uint32_t moeSendNum_;
|
||||
bool isShardExpert_;
|
||||
GM_ADDR epSendCount_;
|
||||
__gm__ HcclOpResParam *epWinContext_;
|
||||
uint64_t winDataSizeOffset_;
|
||||
};
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
class CamMoeDistributeCombine
|
||||
{
|
||||
public:
|
||||
__aicore__ inline CamMoeDistributeCombine(){};
|
||||
__aicore__ inline void Init(GM_ADDR expandX, GM_ADDR expertIds, GM_ADDR expandIdx, GM_ADDR epSendCount,
|
||||
GM_ADDR tpSendCount, GM_ADDR scales, GM_ADDR xActiveMask, GM_ADDR XOut, GM_ADDR workspaceGM, TPipe *pipe,
|
||||
const DispatchGmmCombineDecodeTilingData *tilingData);
|
||||
__aicore__ inline void Process();
|
||||
__aicore__ inline void AllToAllSend();
|
||||
__aicore__ inline void ReducePermute();
|
||||
|
||||
__aicore__ inline CombineCalcInfo &GetCalcInfo()
|
||||
{
|
||||
return calcInfo_;
|
||||
}
|
||||
|
||||
__aicore__ inline void TPipeSet(AscendC::TPipe *pipe)
|
||||
{
|
||||
tpipe_ = pipe;
|
||||
}
|
||||
|
||||
private:
|
||||
__aicore__ inline void InitStatusTargetSum();
|
||||
__aicore__ inline void AlltoAllBuffInit();
|
||||
__aicore__ inline void ReduceScatterTrans();
|
||||
__aicore__ inline void SetWaitTpStatusAndDisPatch();
|
||||
__aicore__ inline void CustomAdd(LocalTensor<ExpandXType> &dst, LocalTensor<ExpandXType> &src0,
|
||||
LocalTensor<ExpandXType> &src1, uint32_t dataCnt);
|
||||
__aicore__ inline void ExpertAlltoAllDispatchInnerCopyAdd(uint32_t tokenNumLoop, uint32_t srcStartTokenIdx,
|
||||
uint32_t ep, uint32_t expertIdx);
|
||||
__aicore__ inline void ExpertAlltoAllDispatchCopyAdd();
|
||||
__aicore__ inline void LocalWindowCopy();
|
||||
__aicore__ inline void BuffInit();
|
||||
__aicore__ inline void SplitCoreCal();
|
||||
__aicore__ inline void SetStatus();
|
||||
__aicore__ inline void WaitDispatch();
|
||||
__aicore__ GM_ADDR GetWinAddrByRankId(const int32_t rankId, const uint8_t domain, const uint8_t expertLocalId = 0U)
|
||||
{
|
||||
if (domain == EP_DOMAIN) {
|
||||
return (GM_ADDR)((epRankId_ == rankId)
|
||||
? epWinContext_->localWindowsIn
|
||||
: ((HcclRankRelationResV2 *)(epWinContext_->remoteRes[rankId].nextDevicePtr))
|
||||
->windowsIn) +
|
||||
winDataSizeOffset_ + expertLocalId * expertPerSizeOnWin_ + rankId * OPT_RANK_OFFSET;
|
||||
} else {
|
||||
return (GM_ADDR)((tpRankId_ == rankId)
|
||||
? tpWinContext_->localWindowsIn
|
||||
: ((HcclRankRelationResV2 *)(tpWinContext_->remoteRes[rankId].nextDevicePtr))
|
||||
->windowsIn) +
|
||||
winDataSizeOffset_ + rankId * OPT_RANK_OFFSET;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ GM_ADDR GetWinStateAddrByRankId(const int32_t rankId, const uint8_t domain)
|
||||
{
|
||||
if (domain == EP_DOMAIN) {
|
||||
return (GM_ADDR)((epRankId_ == rankId)
|
||||
? epWinContext_->localWindowsExp
|
||||
: ((HcclRankRelationResV2 *)(epWinContext_->remoteRes[rankId].nextDevicePtr))
|
||||
->windowsExp) +
|
||||
dataState_ * WIN_STATE_OFFSET;
|
||||
} else {
|
||||
return (GM_ADDR)((tpRankId_ == rankId)
|
||||
? tpWinContext_->localWindowsExp
|
||||
: ((HcclRankRelationResV2 *)(tpWinContext_->remoteRes[rankId].nextDevicePtr))
|
||||
->windowsExp) +
|
||||
dataState_ * WIN_STATE_OFFSET;
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline uint32_t MIN(uint32_t x, uint32_t y)
|
||||
{
|
||||
return (x < y) ? x : y;
|
||||
}
|
||||
|
||||
__aicore__ static void DoCombineRecv(void *ptr)
|
||||
{
|
||||
auto *combiner = (CamMoeDistributeCombine<TemplateMC2TypeFunc> *)ptr;
|
||||
combiner->ReducePermute();
|
||||
}
|
||||
|
||||
TPipe *tpipe_{nullptr};
|
||||
GlobalTensor<ExpandXType> expandXGM_;
|
||||
GlobalTensor<ExpandIdxType> expertIdsGM_;
|
||||
GlobalTensor<ExpandIdxType> expandIdxGM_;
|
||||
GlobalTensor<ExpandIdxType> epSendCountGM_;
|
||||
GlobalTensor<ExpandIdxType> tpSendCountGM_;
|
||||
GlobalTensor<float> expandScalesGM_;
|
||||
GlobalTensor<bool> xActiveMaskGM_;
|
||||
GlobalTensor<ExpandXType> expandOutGlobal_;
|
||||
GlobalTensor<ExpandXType> rankWindow_;
|
||||
GlobalTensor<int32_t> rankStates_;
|
||||
GlobalTensor<float> epStatusSpaceGlobalTensor_;
|
||||
GlobalTensor<float> tpStatusSpaceGlobalTensor_;
|
||||
GlobalTensor<ExpandXType> tpRankWindow_;
|
||||
GlobalTensor<ExpandXType> rowTmpGlobal_;
|
||||
GM_ADDR workspaceGM_;
|
||||
GM_ADDR epWindowGM_;
|
||||
GM_ADDR epStatusSpaceGm_;
|
||||
GM_ADDR tpWindowGM_;
|
||||
GM_ADDR tpStatusSpaceGm_;
|
||||
GM_ADDR stateGM_;
|
||||
|
||||
LocalTensor<ExpandXType> winTpSendCountTensor_;
|
||||
LocalTensor<ExpandXType> gmTpSendCountTensor_;
|
||||
LocalTensor<ExpandXType> outTensor_;
|
||||
LocalTensor<float> winTpSendCountFloatTensor_;
|
||||
LocalTensor<float> gmTpSendCountFloatTensor_;
|
||||
LocalTensor<ExpandIdxType> epSendCountLocal_;
|
||||
|
||||
CombineCalcInfo calcInfo_;
|
||||
uint32_t axisBS_{0};
|
||||
uint32_t axisMaxBs_{0};
|
||||
uint32_t axisBsAlignSize_{0};
|
||||
uint64_t activeMaskBsCnt_{0};
|
||||
uint32_t axisH_{0};
|
||||
uint32_t axisK_{0};
|
||||
uint32_t aivNum_{0};
|
||||
uint32_t epWorldSize_{0};
|
||||
uint32_t tpWorldSize_{0};
|
||||
uint32_t epRankId_{0};
|
||||
uint32_t tpRankId_{0};
|
||||
uint32_t coreIdx_{0}; // aiv id
|
||||
uint32_t sharedExpertRankNum_{0};
|
||||
uint32_t moeExpertNum_{0};
|
||||
uint32_t moeExpertPerRankNum_{0};
|
||||
uint32_t moeSendNum_{0}; // moeExpertPerRankNum_ * epWorldSize_
|
||||
uint32_t tpScatterNum_{0};
|
||||
uint32_t firstTpTokenEndIdx_{0};
|
||||
uint32_t firstTpTokenEndOffset_{0};
|
||||
uint32_t endTok_{0};
|
||||
__gm__ HcclOpResParam *epWinContext_{nullptr};
|
||||
__gm__ HcclOpResParam *tpWinContext_{nullptr};
|
||||
uint32_t epDataOffsetOnWin_{0};
|
||||
uint32_t tpDataOffsetOnWin_{0};
|
||||
uint32_t epStateOffsetOnWin_{0};
|
||||
uint32_t tpStateOffsetOnWin_{0};
|
||||
uint32_t axisHFloatSize_{0};
|
||||
uint32_t axisHExpandXTypeSize_{0};
|
||||
uint32_t bsKNum_{0};
|
||||
uint32_t startRankId_{0};
|
||||
uint32_t endRankId_{0};
|
||||
uint32_t sendRankNum_{0};
|
||||
uint32_t ubSize_{0};
|
||||
uint32_t dataState_{0};
|
||||
uint32_t stateOffset_{0};
|
||||
uint64_t winDataSizeOffset_{0};
|
||||
uint64_t expertPerSizeOnWin_{0};
|
||||
uint64_t totalWinSize_{0};
|
||||
TQueBind<QuePosition::VECIN, QuePosition::VECOUT, 1> moeQueue_;
|
||||
TQue<QuePosition::VECIN, 1> moeSumQueue_;
|
||||
TQueBind<QuePosition::VECIN, QuePosition::VECOUT, 1> gmTpSendCountQueue_;
|
||||
TQue<QuePosition::VECIN, 1> gmTpSendCountInQueue_;
|
||||
TQue<QuePosition::VECIN, 1> winTpSendCountInQueue_;
|
||||
TQue<QuePosition::VECOUT, 1> xOutQueue_;
|
||||
TBuf<> readStateBuf_;
|
||||
TBuf<> expertIdsBuf_;
|
||||
TBuf<> expandScalesBuf_;
|
||||
TBuf<> rowTmpFloatBuf_;
|
||||
TBuf<> sumFloatBuf_;
|
||||
TBuf<> mulBuf_;
|
||||
TBuf<> sendCountBuf_;
|
||||
TBuf<> indexCountsBuf_;
|
||||
TBuf<> winTpSendCountFloatBuf_;
|
||||
TBuf<> gmTpSendCountFloatBuf_;
|
||||
TBuf<> tokenBuf_;
|
||||
TBuf<> statusBuf_;
|
||||
TBuf<> gatherMaskOutBuf_; // gather mask output buf
|
||||
TBuf<> gatherTmpBuf_;
|
||||
TBuf<> statusSumOutBuf_;
|
||||
TBuf<> xActMaskTBuf_;
|
||||
TBuf<> xActMaskCastTBuf_;
|
||||
TBuf<> xActMaskSumTBuf_;
|
||||
float sumTarget_{0.0};
|
||||
int32_t epStateValue_;
|
||||
bool isShardExpert_{false};
|
||||
};
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::Init(
|
||||
GM_ADDR expandX, GM_ADDR expertIds, GM_ADDR expandIdx, GM_ADDR epSendCount, GM_ADDR tpSendCount, GM_ADDR scales, GM_ADDR xActiveMask,
|
||||
GM_ADDR XOut, GM_ADDR workspaceGM, TPipe *pipe, const DispatchGmmCombineDecodeTilingData *tilingData)
|
||||
{
|
||||
tpipe_ = pipe;
|
||||
coreIdx_ = GetBlockIdx();
|
||||
epRankId_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.epRankId;
|
||||
auto contextGM0 = AscendC::GetHcclContext<HCCL_GROUP_ID_0>();
|
||||
epWinContext_ = (__gm__ HcclOpResParam *)contextGM0;
|
||||
GlobalTensor<int32_t> selfDataStatusTensor;
|
||||
GM_ADDR statusDataSpaceGm = (GM_ADDR)epWinContext_->localWindowsExp;
|
||||
selfDataStatusTensor.SetGlobalBuffer((__gm__ int32_t *)(statusDataSpaceGm + STATE_WIN_OFFSET));
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(
|
||||
selfDataStatusTensor[coreIdx_ * UB_ALIGN]);
|
||||
__asm__ __volatile__("");
|
||||
dataState_ = selfDataStatusTensor(coreIdx_ * UB_ALIGN);
|
||||
if (dataState_ == 0) {
|
||||
selfDataStatusTensor(coreIdx_ * UB_ALIGN) = 1;
|
||||
} else {
|
||||
selfDataStatusTensor(coreIdx_ * UB_ALIGN) = 0;
|
||||
}
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(
|
||||
selfDataStatusTensor[coreIdx_ * UB_ALIGN]);
|
||||
__asm__ __volatile__("");
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
|
||||
workspaceGM_ = workspaceGM;
|
||||
expandXGM_.SetGlobalBuffer((__gm__ ExpandXType *)expandX);
|
||||
expertIdsGM_.SetGlobalBuffer((__gm__ ExpandIdxType *)expertIds);
|
||||
expandIdxGM_.SetGlobalBuffer((__gm__ ExpandIdxType *)expandIdx);
|
||||
epSendCountGM_.SetGlobalBuffer((__gm__ int32_t *)epSendCount);
|
||||
expandScalesGM_.SetGlobalBuffer((__gm__ float *)scales);
|
||||
xActiveMaskGM_.SetGlobalBuffer((__gm__ bool*)xActiveMask);
|
||||
expandOutGlobal_.SetGlobalBuffer((__gm__ ExpandXType *)XOut);
|
||||
axisBS_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.bs;
|
||||
activeMaskBsCnt_ = axisBS_;
|
||||
axisH_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
axisK_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.k;
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
aivNum_ = get_block_num();
|
||||
} else {
|
||||
aivNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.aivNum;
|
||||
}
|
||||
ubSize_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.totalUbSize;
|
||||
sharedExpertRankNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertRankNum;
|
||||
moeExpertNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNum;
|
||||
moeExpertPerRankNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
epWorldSize_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.epRankSize;
|
||||
axisMaxBs_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.globalBs / epWorldSize_;
|
||||
moeSendNum_ = epWorldSize_ * moeExpertPerRankNum_;
|
||||
tpWorldSize_ = 1;
|
||||
tpRankId_ = 0;
|
||||
totalWinSize_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.totalWinSize;
|
||||
stateOffset_ = (moeSendNum_ > 512) ? (STATE_OFFSET / 2) : STATE_OFFSET;
|
||||
expertPerSizeOnWin_ =
|
||||
static_cast<uint64_t>(axisMaxBs_) * static_cast<uint64_t>(axisH_) * static_cast<uint64_t>(sizeof(ExpandXType));
|
||||
winDataSizeOffset_ = static_cast<uint64_t>(dataState_) * static_cast<uint64_t>(moeSendNum_) * expertPerSizeOnWin_;
|
||||
epWindowGM_ = GetWinAddrByRankId(epRankId_, EP_DOMAIN);
|
||||
epStatusSpaceGm_ = GetWinStateAddrByRankId(epRankId_, EP_DOMAIN);
|
||||
epStatusSpaceGlobalTensor_.SetGlobalBuffer((__gm__ float *)epStatusSpaceGm_);
|
||||
epDataOffsetOnWin_ = epRankId_ * moeExpertPerRankNum_ * static_cast<uint32_t>(expertPerSizeOnWin_);
|
||||
epStateOffsetOnWin_ = epRankId_ * stateOffset_;
|
||||
isShardExpert_ = (epRankId_ < sharedExpertRankNum_);
|
||||
axisHFloatSize_ = axisH_ * sizeof(float);
|
||||
axisHExpandXTypeSize_ = axisH_ * sizeof(ExpandXType);
|
||||
bsKNum_ = axisBS_ * axisK_;
|
||||
|
||||
if constexpr (IsNeedReduceScatter) {
|
||||
tpSendCountGM_.SetGlobalBuffer((__gm__ int32_t *)tpSendCount);
|
||||
tpWindowGM_ = GetWinAddrByRankId(tpRankId_, TP_DOMAIN);
|
||||
tpStatusSpaceGm_ = GetWinStateAddrByRankId(tpRankId_, TP_DOMAIN);
|
||||
tpStatusSpaceGlobalTensor_.SetGlobalBuffer((__gm__ float *)tpStatusSpaceGm_);
|
||||
tpDataOffsetOnWin_ = tpRankId_ * TP_RANK_SIZE_ON_WIN;
|
||||
tpStateOffsetOnWin_ = tpRankId_ * stateOffset_;
|
||||
uint32_t tpScatterRankWinOffset = (tpRankId_ == 0) ? TP_RANK_SIZE_ON_WIN : 0;
|
||||
GM_ADDR rankGM = tpWindowGM_ + tpScatterRankWinOffset;
|
||||
tpRankWindow_.SetGlobalBuffer((__gm__ ExpandXType *)rankGM);
|
||||
}
|
||||
|
||||
InitStatusTargetSum();
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
coreIdx_ = get_block_idx();
|
||||
}
|
||||
SplitCoreCal();
|
||||
|
||||
calcInfo_.epRankId_ = epRankId_;
|
||||
calcInfo_.epWorldSize_ = epWorldSize_;
|
||||
calcInfo_.expertPerSizeOnWin_ = expertPerSizeOnWin_;
|
||||
calcInfo_.moeExpertPerRankNum_ = moeExpertPerRankNum_;
|
||||
calcInfo_.sharedExpertRankNum_ = sharedExpertRankNum_;
|
||||
calcInfo_.axisH_ = axisH_;
|
||||
calcInfo_.moeSendNum_ = moeSendNum_;
|
||||
calcInfo_.isShardExpert_ = isShardExpert_;
|
||||
calcInfo_.epSendCount_ = epSendCount;
|
||||
calcInfo_.epWinContext_ = epWinContext_;
|
||||
calcInfo_.winDataSizeOffset_ = winDataSizeOffset_;
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::InitStatusTargetSum()
|
||||
{
|
||||
// ep state
|
||||
GlobalTensor<int32_t> selfStatusTensor;
|
||||
selfStatusTensor.SetGlobalBuffer((__gm__ int32_t *)(epStatusSpaceGm_ + SELF_STATE_OFFSET));
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(
|
||||
selfStatusTensor[coreIdx_ * UB_ALIGN]);
|
||||
__asm__ __volatile__("");
|
||||
int32_t state = selfStatusTensor(coreIdx_ * UB_ALIGN);
|
||||
if (state == 0) {
|
||||
sumTarget_ = static_cast<float>(1.0);
|
||||
selfStatusTensor(coreIdx_ * UB_ALIGN) = 0x3F800000; // 1.0f
|
||||
epStateValue_ = 0x3F800000; // 1.0f
|
||||
} else {
|
||||
sumTarget_ = static_cast<float>(0.0);
|
||||
selfStatusTensor(coreIdx_ * UB_ALIGN) = 0;
|
||||
epStateValue_ = 0;
|
||||
}
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(
|
||||
selfStatusTensor[coreIdx_ * UB_ALIGN]);
|
||||
__asm__ __volatile__("");
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::BuffInit()
|
||||
{
|
||||
tpipe_->Reset();
|
||||
tpipe_->InitBuffer(readStateBuf_, UB_ALIGN);
|
||||
uint32_t sendNumAlign = Ceil(moeSendNum_ * sizeof(int32_t), UB_ALIGN) * UB_ALIGN;
|
||||
tpipe_->InitBuffer(sendCountBuf_, sendNumAlign);
|
||||
if constexpr (IsNeedReduceScatter) {
|
||||
tpipe_->InitBuffer(winTpSendCountInQueue_, BUFFER_NUM, axisHExpandXTypeSize_);
|
||||
tpipe_->InitBuffer(gmTpSendCountInQueue_, BUFFER_NUM, axisHExpandXTypeSize_);
|
||||
tpipe_->InitBuffer(xOutQueue_, BUFFER_NUM, axisHExpandXTypeSize_);
|
||||
if constexpr (AscendC::IsSameType<ExpandXType, bfloat16_t>::value) {
|
||||
tpipe_->InitBuffer(winTpSendCountFloatBuf_, axisHFloatSize_);
|
||||
tpipe_->InitBuffer(gmTpSendCountFloatBuf_, axisHFloatSize_);
|
||||
winTpSendCountFloatTensor_ = winTpSendCountFloatBuf_.Get<float>();
|
||||
gmTpSendCountFloatTensor_ = gmTpSendCountFloatBuf_.Get<float>();
|
||||
}
|
||||
} else {
|
||||
tpipe_->InitBuffer(gmTpSendCountQueue_, BUFFER_NUM, axisHExpandXTypeSize_);
|
||||
}
|
||||
epSendCountLocal_ = sendCountBuf_.Get<int32_t>();
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::AlltoAllBuffInit()
|
||||
{
|
||||
tpipe_->Reset();
|
||||
uint32_t bsMulTopkSizeAligned = Ceil(axisBS_ * axisK_ * sizeof(int32_t), UB_ALIGN) * UB_ALIGN;
|
||||
tpipe_->InitBuffer(readStateBuf_, UB_ALIGN);
|
||||
tpipe_->InitBuffer(statusBuf_, sendRankNum_ * UB_ALIGN);
|
||||
tpipe_->InitBuffer(expertIdsBuf_, bsMulTopkSizeAligned);
|
||||
tpipe_->InitBuffer(expandScalesBuf_, bsMulTopkSizeAligned);
|
||||
tpipe_->InitBuffer(tokenBuf_, axisH_ * sizeof(ExpandXType));
|
||||
tpipe_->InitBuffer(rowTmpFloatBuf_, axisHFloatSize_);
|
||||
tpipe_->InitBuffer(mulBuf_, axisHFloatSize_);
|
||||
tpipe_->InitBuffer(sumFloatBuf_, axisHFloatSize_);
|
||||
tpipe_->InitBuffer(indexCountsBuf_, bsMulTopkSizeAligned);
|
||||
tpipe_->InitBuffer(moeSumQueue_, BUFFER_NUM, axisHExpandXTypeSize_);
|
||||
tpipe_->InitBuffer(gatherMaskOutBuf_, epWorldSize_ * sizeof(float));
|
||||
tpipe_->InitBuffer(gatherTmpBuf_, sizeof(uint32_t));
|
||||
tpipe_->InitBuffer(statusSumOutBuf_, sizeof(float));
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_X_ACTIVE_MASK) {
|
||||
axisBsAlignSize_ = Ceil(axisBS_ * sizeof(bool), UB_ALIGN) * UB_ALIGN;
|
||||
tpipe_->InitBuffer(xActMaskTBuf_, axisBsAlignSize_);
|
||||
tpipe_->InitBuffer(xActMaskCastTBuf_, axisBsAlignSize_ * sizeof(half));
|
||||
tpipe_->InitBuffer(xActMaskSumTBuf_, axisBsAlignSize_ * sizeof(half));
|
||||
LocalTensor<bool> xActiveMaskTensor = xActMaskTBuf_.Get<bool>();
|
||||
LocalTensor<half> tempTensor = xActMaskCastTBuf_.Get<half>();
|
||||
LocalTensor<half> sumOutTensor = xActMaskSumTBuf_.Get<half>();
|
||||
DataCopyExtParams xActiveMaskParams{1U, static_cast<uint32_t>(axisBS_ * sizeof(bool)), 0U, 0U, 0U};
|
||||
DataCopyPadExtParams<bool> xActiveMaskCopyPadParams{false, 0U, 0U, 0U};
|
||||
DataCopyPad(xActiveMaskTensor, xActiveMaskGM_, xActiveMaskParams, xActiveMaskCopyPadParams);
|
||||
SyncFunc<AscendC::HardEvent::MTE2_V>();
|
||||
LocalTensor<int8_t> xActiveMaskInt8Tensor = xActiveMaskTensor.ReinterpretCast<int8_t>();
|
||||
Cast(tempTensor, xActiveMaskInt8Tensor, RoundMode::CAST_NONE, axisBS_);
|
||||
PipeBarrier<PIPE_V>();
|
||||
SumParams params{1, axisBsAlignSize_, axisBS_};
|
||||
Sum(sumOutTensor, tempTensor, params);
|
||||
SyncFunc<AscendC::HardEvent::V_S>();
|
||||
activeMaskBsCnt_ = static_cast<int32_t>(sumOutTensor.GetValue(0));
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::SplitCoreCal()
|
||||
{
|
||||
sendRankNum_ = epWorldSize_ / aivNum_;
|
||||
uint32_t remainderRankNum = epWorldSize_ % aivNum_;
|
||||
startRankId_ = sendRankNum_ * coreIdx_;
|
||||
if (coreIdx_ < remainderRankNum) {
|
||||
sendRankNum_++;
|
||||
startRankId_ += coreIdx_;
|
||||
} else {
|
||||
startRankId_ += remainderRankNum;
|
||||
}
|
||||
endRankId_ = startRankId_ + sendRankNum_;
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::ReduceScatterTrans()
|
||||
{
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(tpSendCountGM_[tpRankId_]);
|
||||
__asm__ __volatile__("");
|
||||
uint32_t offset = tpSendCountGM_.GetValue(tpRankId_) * axisH_;
|
||||
GlobalTensor<ExpandXType> dataCopyInGM = expandXGM_[offset];
|
||||
GM_ADDR rankGM = GetWinAddrByRankId(1 - tpRankId_, TP_DOMAIN) + tpDataOffsetOnWin_;
|
||||
rankWindow_.SetGlobalBuffer((__gm__ ExpandXType *)rankGM);
|
||||
uint32_t copyStartIdx = 0;
|
||||
if (startRankId_ > 0) {
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(
|
||||
epSendCountGM_[epWorldSize_ + startRankId_ - 1]);
|
||||
__asm__ __volatile__("");
|
||||
copyStartIdx = epSendCountGM_.GetValue(epWorldSize_ + startRankId_ - 1);
|
||||
}
|
||||
__asm__ __volatile__("");
|
||||
DataCacheCleanAndInvalid<int32_t, CacheLine::SINGLE_CACHE_LINE, DcciDst::CACHELINE_OUT>(
|
||||
epSendCountGM_[epWorldSize_ + endRankId_ - 1]);
|
||||
__asm__ __volatile__("");
|
||||
uint32_t copyEndIdx = epSendCountGM_.GetValue(epWorldSize_ + endRankId_ - 1);
|
||||
LocalTensor<ExpandXType> tmpUb;
|
||||
for (uint32_t tokenNumIdx = copyStartIdx; tokenNumIdx < copyEndIdx; tokenNumIdx++) {
|
||||
tmpUb = moeQueue_.AllocTensor<ExpandXType>();
|
||||
DataCopy(tmpUb, dataCopyInGM[tokenNumIdx * axisH_], axisH_);
|
||||
moeQueue_.EnQue(tmpUb);
|
||||
tmpUb = moeQueue_.DeQue<ExpandXType>();
|
||||
DataCopy(rankWindow_[tokenNumIdx * axisH_], tmpUb, axisH_);
|
||||
moeQueue_.FreeTensor<ExpandXType>(tmpUb);
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::SetWaitTpStatusAndDisPatch()
|
||||
{
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
if (startRankId_ >= epWorldSize_) {
|
||||
return;
|
||||
}
|
||||
if constexpr (IsNeedReduceScatter) {
|
||||
uint32_t tpToRankId = 1 - tpRankId_;
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
LocalTensor<float> statusFlagUb = readStateBuf_.Get<float>();
|
||||
statusFlagUb(0) = sumTarget_;
|
||||
SyncFunc<AscendC::HardEvent::S_MTE3>();
|
||||
GlobalTensor<float> tpWindowInstatusFp32Tensor_;
|
||||
stateGM_ = GetWinStateAddrByRankId(tpToRankId, TP_DOMAIN) + coreIdx_ * stateOffset_;
|
||||
tpWindowInstatusFp32Tensor_.SetGlobalBuffer((__gm__ float *)stateGM_);
|
||||
DataCopy<float>(tpWindowInstatusFp32Tensor_, statusFlagUb, 8UL);
|
||||
SyncFunc<AscendC::HardEvent::MTE3_S>();
|
||||
LocalTensor<float> statusFp32Tensor_ = readStateBuf_.Get<float>();
|
||||
float sumOfFlag = static_cast<float>(-1.0);
|
||||
uint32_t statusRankOffset = coreIdx_ * stateOffset_ / sizeof(float);
|
||||
while (sumOfFlag != sumTarget_) {
|
||||
DataCopy<float>(statusFp32Tensor_, tpStatusSpaceGlobalTensor_[statusRankOffset], 8);
|
||||
SyncFunc<AscendC::HardEvent::MTE2_S>();
|
||||
sumOfFlag = statusFp32Tensor_.GetValue(0);
|
||||
SyncFunc<AscendC::HardEvent::S_MTE2>();
|
||||
}
|
||||
}
|
||||
ExpertAlltoAllDispatchCopyAdd();
|
||||
SyncFunc<AscendC::HardEvent::MTE3_S>();
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::ExpertAlltoAllDispatchCopyAdd()
|
||||
{
|
||||
if (startRankId_ >= epWorldSize_) {
|
||||
return;
|
||||
}
|
||||
uint32_t curRankExpertNum = 0;
|
||||
DataCopyExtParams epSendCntParams;
|
||||
if (isShardExpert_) {
|
||||
curRankExpertNum = 1;
|
||||
epSendCntParams = {1U, static_cast<uint32_t>(epWorldSize_ * sizeof(uint32_t)), 0U, 0U, 0U};
|
||||
} else {
|
||||
curRankExpertNum = moeExpertPerRankNum_;
|
||||
epSendCntParams = {1U, static_cast<uint32_t>(moeSendNum_ * sizeof(uint32_t)), 0U, 0U, 0U};
|
||||
}
|
||||
DataCopyPadExtParams<int32_t> copyPadParams{false, 0U, 0U, 0U};
|
||||
DataCopyPad(epSendCountLocal_, epSendCountGM_, epSendCntParams, copyPadParams);
|
||||
SyncFunc<AscendC::HardEvent::MTE2_S>();
|
||||
uint32_t preCount = 0;
|
||||
uint32_t startTokenIdx = 0;
|
||||
uint32_t curTokenNum = 0;
|
||||
|
||||
for (uint32_t expertIdx = 0U; expertIdx < curRankExpertNum; expertIdx++) {
|
||||
uint32_t sendEpCount = endRankId_ - startRankId_;
|
||||
for (uint32_t i = 0; i < sendEpCount; ++i) {
|
||||
uint32_t ep = startRankId_ + (i + epRankId_) % sendEpCount;
|
||||
if ((ep > 0) || (expertIdx > 0U)) {
|
||||
preCount = epSendCountLocal_.GetValue(expertIdx * epWorldSize_ + ep - 1);
|
||||
} else {
|
||||
preCount = 0;
|
||||
}
|
||||
curTokenNum = epSendCountLocal_.GetValue(expertIdx * epWorldSize_ + ep) - preCount;
|
||||
if (curTokenNum == 0) {
|
||||
continue;
|
||||
}
|
||||
startTokenIdx = preCount * axisH_;
|
||||
ExpertAlltoAllDispatchInnerCopyAdd(curTokenNum, startTokenIdx, ep, expertIdx);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::ExpertAlltoAllDispatchInnerCopyAdd(
|
||||
uint32_t tokenNumLoop, uint32_t srcStartTokenIdx, uint32_t ep, uint32_t expertIdx)
|
||||
{
|
||||
GM_ADDR rankGM = GetWinAddrByRankId(ep, EP_DOMAIN, expertIdx) + epDataOffsetOnWin_;
|
||||
if ((isShardExpert_) && (ep < sharedExpertRankNum_)) {
|
||||
rankGM = GetWinAddrByRankId(epRankId_, EP_DOMAIN, expertIdx) + ep * moeExpertPerRankNum_ * expertPerSizeOnWin_;
|
||||
}
|
||||
rankWindow_.SetGlobalBuffer((__gm__ ExpandXType *)rankGM);
|
||||
uint32_t dataCnt = axisH_;
|
||||
for (uint32_t loopIdx = 0; loopIdx < tokenNumLoop; loopIdx++) {
|
||||
if constexpr (IsNeedReduceScatter) {
|
||||
gmTpSendCountTensor_ = gmTpSendCountInQueue_.AllocTensor<ExpandXType>();
|
||||
DataCopy(gmTpSendCountTensor_, expandXGM_[srcStartTokenIdx], dataCnt);
|
||||
gmTpSendCountInQueue_.EnQue(gmTpSendCountTensor_);
|
||||
|
||||
winTpSendCountTensor_ = winTpSendCountInQueue_.AllocTensor<ExpandXType>();
|
||||
DataCopy(winTpSendCountTensor_, tpRankWindow_[srcStartTokenIdx], dataCnt);
|
||||
winTpSendCountInQueue_.EnQue(winTpSendCountTensor_);
|
||||
|
||||
gmTpSendCountTensor_ = gmTpSendCountInQueue_.DeQue<ExpandXType>();
|
||||
winTpSendCountTensor_ = winTpSendCountInQueue_.DeQue<ExpandXType>();
|
||||
outTensor_ = xOutQueue_.AllocTensor<ExpandXType>();
|
||||
|
||||
CustomAdd(outTensor_, winTpSendCountTensor_, gmTpSendCountTensor_, dataCnt);
|
||||
gmTpSendCountInQueue_.FreeTensor<ExpandXType>(gmTpSendCountTensor_);
|
||||
winTpSendCountInQueue_.FreeTensor<ExpandXType>(winTpSendCountTensor_);
|
||||
xOutQueue_.EnQue(outTensor_);
|
||||
|
||||
outTensor_ = xOutQueue_.DeQue<ExpandXType>();
|
||||
DataCopy(rankWindow_[loopIdx * dataCnt], outTensor_, dataCnt);
|
||||
xOutQueue_.FreeTensor<ExpandXType>(outTensor_);
|
||||
} else {
|
||||
gmTpSendCountTensor_ = gmTpSendCountQueue_.AllocTensor<ExpandXType>();
|
||||
DataCopy(gmTpSendCountTensor_, expandXGM_[srcStartTokenIdx], dataCnt);
|
||||
ExpandXType val = expandXGM_[srcStartTokenIdx].GetValue(0);
|
||||
gmTpSendCountQueue_.EnQue(gmTpSendCountTensor_);
|
||||
gmTpSendCountTensor_ = gmTpSendCountQueue_.DeQue<ExpandXType>();
|
||||
DataCopy(rankWindow_[loopIdx * dataCnt], gmTpSendCountTensor_, dataCnt);
|
||||
gmTpSendCountQueue_.FreeTensor<ExpandXType>(gmTpSendCountTensor_);
|
||||
}
|
||||
srcStartTokenIdx += dataCnt;
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::CustomAdd(LocalTensor<ExpandXType> &dst,
|
||||
LocalTensor<ExpandXType> &src0,
|
||||
LocalTensor<ExpandXType> &src1,
|
||||
uint32_t dataCnt)
|
||||
{
|
||||
if constexpr (AscendC::IsSameType<ExpandXType, bfloat16_t>::value) {
|
||||
Cast(winTpSendCountFloatTensor_, src0, RoundMode::CAST_NONE, dataCnt);
|
||||
Cast(gmTpSendCountFloatTensor_, src1, RoundMode::CAST_NONE, dataCnt);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Add(winTpSendCountFloatTensor_, winTpSendCountFloatTensor_, gmTpSendCountFloatTensor_, dataCnt);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
Cast(dst, winTpSendCountFloatTensor_, RoundMode::CAST_ROUND, dataCnt);
|
||||
} else {
|
||||
Add(dst, src0, src1, dataCnt);
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::SetStatus()
|
||||
{
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
if (startRankId_ >= epWorldSize_) {
|
||||
return;
|
||||
}
|
||||
|
||||
LocalTensor<int32_t> statusFlagUb = readStateBuf_.Get<int32_t>();
|
||||
statusFlagUb.SetValue(0, epStateValue_);
|
||||
SyncFunc<AscendC::HardEvent::S_MTE3>();
|
||||
|
||||
for (uint32_t epIdx = startRankId_; epIdx < endRankId_; epIdx++) {
|
||||
stateGM_ = GetWinStateAddrByRankId(epIdx, EP_DOMAIN) + epStateOffsetOnWin_;
|
||||
rankStates_.SetGlobalBuffer((__gm__ int32_t *)stateGM_);
|
||||
DataCopy(rankStates_, statusFlagUb, 8);
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::WaitDispatch()
|
||||
{
|
||||
if (startRankId_ < epWorldSize_) {
|
||||
LocalTensor<float> statusTensor = statusBuf_.Get<float>();
|
||||
LocalTensor<float> gatherMaskOutTensor = gatherMaskOutBuf_.Get<float>();
|
||||
LocalTensor<uint32_t> gatherTmpTensor = gatherTmpBuf_.Get<uint32_t>();
|
||||
LocalTensor<float> statusSumOutTensor = statusSumOutBuf_.Get<float>();
|
||||
PipeBarrier<PIPE_ALL>();
|
||||
|
||||
gatherTmpTensor.SetValue(0, 1);
|
||||
uint32_t mask = 1; // gatherMask + sum
|
||||
uint64_t rsvdCnt = 0;
|
||||
DataCopyParams intriParams{static_cast<uint16_t>(sendRankNum_), 1,
|
||||
static_cast<uint16_t>((moeSendNum_ > 512) ? 7 : 15), 0}; // srcStride is 15 blocks
|
||||
float sumOfFlag = static_cast<float>(-1.0);
|
||||
float minTarget = (sumTarget_ * sendRankNum_) - (float)0.5;
|
||||
float maxTarget = (sumTarget_ * sendRankNum_) + (float)0.5;
|
||||
SumParams sumParams{1, sendRankNum_, sendRankNum_};
|
||||
SyncFunc<AscendC::HardEvent::S_V>();
|
||||
while ((sumOfFlag < minTarget) || (sumOfFlag > maxTarget)) {
|
||||
DataCopy<float>(statusTensor, epStatusSpaceGlobalTensor_[startRankId_ * stateOffset_ / sizeof(float)],
|
||||
intriParams);
|
||||
SyncFunc<AscendC::HardEvent::MTE2_V>();
|
||||
GatherMask(gatherMaskOutTensor, statusTensor, gatherTmpTensor, true, mask,
|
||||
{1, (uint16_t)sendRankNum_, 1, 0}, rsvdCnt);
|
||||
PipeBarrier<PIPE_V>();
|
||||
Sum(statusSumOutTensor, gatherMaskOutTensor, sumParams);
|
||||
SyncFunc<AscendC::HardEvent::V_S>();
|
||||
sumOfFlag = statusSumOutTensor.GetValue(0);
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(RECV_SYNC_EVENT_ID);
|
||||
AscendC::CrossCoreWaitFlag(RECV_SYNC_EVENT_ID);
|
||||
} else {
|
||||
SyncAll<true>();
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::LocalWindowCopy()
|
||||
{
|
||||
if (activeMaskBsCnt_ == 0U) {
|
||||
return;
|
||||
}
|
||||
uint32_t beginIndex = 0;
|
||||
uint32_t endIndex = 0;
|
||||
uint32_t processLen = 0;
|
||||
uint32_t tokenOffset = 0;
|
||||
if (activeMaskBsCnt_ < aivNum_) {
|
||||
uint32_t aivNumPerToken = aivNum_ / activeMaskBsCnt_; // activeMaskBsCnt_ < aivNum_
|
||||
if (coreIdx_ >= (activeMaskBsCnt_ * aivNumPerToken)) {
|
||||
return;
|
||||
}
|
||||
uint32_t tokenIndex = coreIdx_ / aivNumPerToken;
|
||||
processLen = ((axisH_ / UB_ALIGN) / aivNumPerToken) * UB_ALIGN;
|
||||
tokenOffset = processLen * (coreIdx_ % aivNumPerToken);
|
||||
if ((coreIdx_ % aivNumPerToken) == (aivNumPerToken - 1)) {
|
||||
processLen = axisH_ - ((aivNumPerToken - 1) * processLen);
|
||||
}
|
||||
beginIndex = tokenIndex;
|
||||
endIndex = beginIndex + 1U;
|
||||
} else {
|
||||
uint32_t tokenPerAivNum = activeMaskBsCnt_ / aivNum_;
|
||||
uint32_t remainderToken = activeMaskBsCnt_ % aivNum_;
|
||||
beginIndex = tokenPerAivNum * coreIdx_;
|
||||
if (coreIdx_ < remainderToken) {
|
||||
tokenPerAivNum++;
|
||||
beginIndex = tokenPerAivNum * coreIdx_;
|
||||
} else {
|
||||
beginIndex += remainderToken;
|
||||
}
|
||||
endIndex = beginIndex + tokenPerAivNum;
|
||||
processLen = axisH_;
|
||||
}
|
||||
LocalTensor<ExpandIdxType> expertIdsLocal = expertIdsBuf_.Get<ExpandIdxType>();
|
||||
LocalTensor<float> expandScalesLocal = expandScalesBuf_.Get<float>();
|
||||
|
||||
LocalTensor<float> rowTmpFloatLocal = rowTmpFloatBuf_.Get<float>();
|
||||
LocalTensor<float> mulBufLocal = mulBuf_.Get<float>();
|
||||
LocalTensor<float> sumFloatBufLocal = sumFloatBuf_.Get<float>();
|
||||
|
||||
LocalTensor<ExpandIdxType> indexCountsLocal = indexCountsBuf_.Get<ExpandIdxType>();
|
||||
const DataCopyExtParams bskParams = {1U, static_cast<uint32_t>(bsKNum_ * sizeof(uint32_t)), 0U, 0U, 0U};
|
||||
const DataCopyPadExtParams<ExpandIdxType> copyPadParams{false, 0U, 0U, 0U};
|
||||
const DataCopyPadExtParams<float> copyPadFloatParams{false, 0U, 0U, 0U};
|
||||
|
||||
DataCopyPad(indexCountsLocal, expandIdxGM_, bskParams, copyPadParams);
|
||||
DataCopyPad(expertIdsLocal, expertIdsGM_, bskParams, copyPadParams);
|
||||
DataCopyPad(expandScalesLocal, expandScalesGM_, bskParams, copyPadFloatParams);
|
||||
SyncFunc<AscendC::HardEvent::MTE2_S>();
|
||||
|
||||
for (uint32_t tokenIndex = beginIndex; tokenIndex < endIndex; tokenIndex++) {
|
||||
uint32_t index = tokenIndex * axisK_;
|
||||
SyncFunc<AscendC::HardEvent::MTE3_V>();
|
||||
Duplicate(sumFloatBufLocal, (float)0, axisH_);
|
||||
for (uint32_t i = 0; i < axisK_; i++) {
|
||||
int32_t moeExpert = expertIdsLocal.GetValue(index);
|
||||
if (moeExpert < 0) {
|
||||
index++;
|
||||
continue;
|
||||
}
|
||||
float scaleVal = expandScalesLocal.GetValue(index);
|
||||
GM_ADDR wAddr = (__gm__ uint8_t *)(epWindowGM_) +
|
||||
expertPerSizeOnWin_ * moeExpertPerRankNum_ * sharedExpertRankNum_ +
|
||||
expertPerSizeOnWin_ * moeExpert + indexCountsLocal.GetValue(index) * axisHExpandXTypeSize_ +
|
||||
tokenOffset * sizeof(ExpandXType);
|
||||
rowTmpGlobal_.SetGlobalBuffer((__gm__ ExpandXType *)wAddr);
|
||||
ExpandXType val = rowTmpGlobal_.GetValue(0);
|
||||
LocalTensor<ExpandXType> tmpUb = moeSumQueue_.AllocTensor<ExpandXType>();
|
||||
DataCopy(tmpUb, rowTmpGlobal_, processLen);
|
||||
moeSumQueue_.EnQue(tmpUb);
|
||||
tmpUb = moeSumQueue_.DeQue<ExpandXType>();
|
||||
Cast(rowTmpFloatLocal, tmpUb, AscendC::RoundMode::CAST_NONE, processLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Muls(mulBufLocal, rowTmpFloatLocal, scaleVal, processLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Add(sumFloatBufLocal, sumFloatBufLocal, mulBufLocal, processLen);
|
||||
index++;
|
||||
moeSumQueue_.FreeTensor<ExpandXType>(tmpUb);
|
||||
}
|
||||
LocalTensor<ExpandXType> rowTmpLocal = tokenBuf_.Get<ExpandXType>();
|
||||
if (sharedExpertRankNum_ > 0U) {
|
||||
uint32_t temp = (epRankId_ * activeMaskBsCnt_) / sharedExpertRankNum_;
|
||||
uint32_t moeOnShareRank = Ceil((tokenIndex + 1 + temp) * sharedExpertRankNum_, activeMaskBsCnt_) - 1 - epRankId_;
|
||||
uint32_t preCnt = (moeOnShareRank + epRankId_) * activeMaskBsCnt_ / sharedExpertRankNum_ -
|
||||
epRankId_ * activeMaskBsCnt_ / sharedExpertRankNum_;
|
||||
__gm__ ExpandXType *shareAddr =
|
||||
(__gm__ ExpandXType *)(epWindowGM_ + moeOnShareRank * expertPerSizeOnWin_ * moeExpertPerRankNum_) +
|
||||
(tokenIndex - preCnt) * axisH_ + tokenOffset;
|
||||
GlobalTensor<ExpandXType> shareTokGlobal;
|
||||
shareTokGlobal.SetGlobalBuffer((__gm__ ExpandXType *)(shareAddr));
|
||||
SyncFunc<AscendC::HardEvent::V_MTE2>();
|
||||
DataCopy(rowTmpLocal, shareTokGlobal, processLen);
|
||||
SyncFunc<AscendC::HardEvent::MTE2_V>();
|
||||
Cast(rowTmpFloatLocal, rowTmpLocal, AscendC::RoundMode::CAST_NONE, processLen);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Add(sumFloatBufLocal, sumFloatBufLocal, rowTmpFloatLocal, processLen);
|
||||
}
|
||||
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
LocalTensor<ExpandXType> sumBufLocal = tokenBuf_.Get<ExpandXType>();
|
||||
Cast(sumBufLocal, sumFloatBufLocal, AscendC::RoundMode::CAST_RINT, processLen);
|
||||
SyncFunc<AscendC::HardEvent::V_MTE3>();
|
||||
DataCopy(expandOutGlobal_[tokenIndex * axisH_ + tokenOffset], sumBufLocal, processLen);
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::Process()
|
||||
{
|
||||
SyncAll<true>();
|
||||
if constexpr (IsNeedReduceScatter) {
|
||||
tpipe_->InitBuffer(moeQueue_, BUFFER_NUM, axisHExpandXTypeSize_);
|
||||
ReduceScatterTrans();
|
||||
}
|
||||
if constexpr ((EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) == 0) {
|
||||
BuffInit();
|
||||
SetWaitTpStatusAndDisPatch();
|
||||
}
|
||||
AlltoAllBuffInit();
|
||||
SetStatus();
|
||||
WaitDispatch();
|
||||
LocalWindowCopy();
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::AllToAllSend()
|
||||
{
|
||||
if constexpr (IsNeedReduceScatter) {
|
||||
tpipe_->InitBuffer(moeQueue_, BUFFER_NUM, axisHExpandXTypeSize_);
|
||||
ReduceScatterTrans();
|
||||
}
|
||||
BuffInit();
|
||||
if constexpr ((EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) == 0) {
|
||||
SetWaitTpStatusAndDisPatch();
|
||||
AlltoAllBuffInit();
|
||||
}
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(SEND_SYNC_EVENT_ID);
|
||||
AscendC::CrossCoreWaitFlag(SEND_SYNC_EVENT_ID);
|
||||
} else {
|
||||
SyncAll<true>();
|
||||
}
|
||||
SetStatus();
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
AscendC::CrossCoreWaitFlag(RECV_SYNC_EVENT_ID);
|
||||
} else {
|
||||
SyncAll<true>();
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void CamMoeDistributeCombine<TemplateMC2TypeFunc>::ReducePermute()
|
||||
{
|
||||
AlltoAllBuffInit();
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
AscendC::CrossCoreSetFlag<0x0, PIPE_MTE3>(SEND_SYNC_EVENT_ID);
|
||||
} else {
|
||||
SyncAll<true>();
|
||||
}
|
||||
|
||||
WaitDispatch();
|
||||
LocalWindowCopy();
|
||||
|
||||
if constexpr (EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) {
|
||||
AscendC::CrossCoreWaitFlag(SEND_SYNC_EVENT_ID);
|
||||
}
|
||||
}
|
||||
} // namespace MoeDistributeCombineImpl
|
||||
|
||||
#endif // CAM_MOE_DISTRIBUTE_COMBINE_IMPL_H
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,118 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef DISPATCH_GMM_COMBINE_DECODE_BASE_H
|
||||
#define DISPATCH_GMM_COMBINE_DECODE_BASE_H
|
||||
|
||||
#include "moe_distribute_base.h"
|
||||
|
||||
#define TemplateMC2TypeClass typename ExpandXType, typename W1ScaleType, typename W2ScaleType, typename WType, typename ExpandIdxType, bool IsNeedReduceScatter, uint32_t EXEC_FLAG
|
||||
#define TemplateMC2TypeFunc ExpandXType, W1ScaleType, W2ScaleType, WType, ExpandIdxType, IsNeedReduceScatter, EXEC_FLAG
|
||||
#define TemplateDispatchTypeClass \
|
||||
typename XType, typename ExpandXOutType, bool StaticQuant, bool DynamicQuant, bool IsSmoothScaleExist, \
|
||||
bool IsNeedAllgater, uint32_t EXEC_FLAG
|
||||
#define TemplateDispatchTypeFunc XType, ExpandXOutType, StaticQuant, DynamicQuant, IsSmoothScaleExist, IsNeedAllgater, EXEC_FLAG
|
||||
|
||||
constexpr uint32_t STATE_OFFSET = 512;
|
||||
constexpr uint64_t WIN_STATE_OFFSET = 512 * 1024;
|
||||
constexpr uint64_t STATE_WIN_OFFSET = 900 * 1024;
|
||||
constexpr uint64_t GROUP_TOKEN_NUM_OFFSET = 932 * 1024;
|
||||
constexpr uint64_t SOFT_SYNC_OFFSET = 964 * 1024;
|
||||
constexpr uint32_t SELF_STATE_OFFSET = 256 * 1024;
|
||||
constexpr uint32_t SUM_TMP_TENSOR_SIZE = 1024;
|
||||
constexpr uint32_t UB_ALIGN = 32;
|
||||
constexpr uint32_t TOKEN_EXTRA_SPACE = 512;
|
||||
constexpr uint32_t INT32_COUNT_PER_BLOCK = 8;
|
||||
constexpr uint32_t SOFT_SYNC_SPACE_SIZE = 512;
|
||||
constexpr int64_t LOOP_TMP_SIZE = 4096;
|
||||
constexpr int32_t SUB_AIV_NUM = 2;
|
||||
constexpr int32_t ODD_EVEN_BASE = 2;
|
||||
constexpr int32_t BUFFER_NUM = 2;
|
||||
constexpr int32_t GATHER_SECOND_NUM = 2;
|
||||
constexpr uint32_t MAX_QUANT_ROW_ONCE = 8;
|
||||
constexpr uint32_t QUANT_SPACE_FACTOR = 176 * 1024 / 11; // up to 176KB for quant
|
||||
#ifndef OPT_RANK_OFFSET
|
||||
#define OPT_RANK_OFFSET 512
|
||||
#endif
|
||||
|
||||
#define CEIL_UP(x) ((x + UB_ALIGN - 1) / UB_ALIGN * UB_ALIGN)
|
||||
#define CEIL(x, y) (((x) + (y - 1)) / (y))
|
||||
#define UB_BLOCK_SIZE (32)
|
||||
#define GET_WIND_STATE_ADDR_BY_RANK_ID(rankId) \
|
||||
(((epRankId == rankId) \
|
||||
? ((GM_ADDR)(winContext_->localWindowsExp)) \
|
||||
: ((GM_ADDR)(((HcclRankRelationResV2 *)(winContext_->remoteRes[rankId].nextDevicePtr))->windowsExp))) + \
|
||||
dataState * WIN_STATE_OFFSET)
|
||||
#define GET_WIND_ADDR_BY_RANK_ID(rankId) \
|
||||
(((epRankId == rankId) \
|
||||
? ((GM_ADDR)(winContext_->localWindowsIn)) \
|
||||
: ((GM_ADDR)(((HcclRankRelationResV2 *)(winContext_->remoteRes[rankId].nextDevicePtr))->windowsIn))) + \
|
||||
winDataSizeOffset + rankId * OPT_RANK_OFFSET)
|
||||
#define TOKEN_FLAG_1 (0x55555555)
|
||||
#define TOKEN_FLAG_2 (0x33333333)
|
||||
#define V_TO_C_FLAG_1 (0x03030303)
|
||||
#define V_TO_C_FLAG_2 (0x05050505)
|
||||
#define CV_FLAG_INDEX 0
|
||||
#define GROUP_ID_INDEX 1
|
||||
#define PRE_COUNT_INDEX 2
|
||||
#define SELF_COUNT_INDEX 3
|
||||
#define TOTAL_COUNT_INDEX 4
|
||||
#define GROUP_TOKEN_COUNT 3 // equal to SELF_COUNT_INDEX
|
||||
#define GROUP_INFO_SIZE 32
|
||||
|
||||
__aicore__ inline static void EncreaseSyncFlag(__gm__ uint8_t *flagAddr, uint8_t idx)
|
||||
{
|
||||
// flag++, like set flag
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::GlobalTensor<uint8_t> global;
|
||||
global.SetGlobalBuffer(flagAddr + idx * SOFT_SYNC_SPACE_SIZE);
|
||||
__asm__ __volatile__("");
|
||||
AscendC::DataCacheCleanAndInvalid<uint8_t, AscendC::CacheLine::SINGLE_CACHE_LINE, AscendC::DcciDst::CACHELINE_OUT>(
|
||||
global);
|
||||
__asm__ __volatile__("");
|
||||
uint8_t value = global.GetValue(0);
|
||||
global.SetValue(0, value + 1);
|
||||
__asm__ __volatile__("");
|
||||
AscendC::DataCacheCleanAndInvalid<uint8_t, AscendC::CacheLine::SINGLE_CACHE_LINE, AscendC::DcciDst::CACHELINE_OUT>(
|
||||
global);
|
||||
__asm__ __volatile__("");
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
}
|
||||
|
||||
__aicore__ inline static void CheckSyncFlag(__gm__ uint8_t *flagAddr, uint8_t idx, uint32_t target)
|
||||
{
|
||||
// check flag, like wait flag
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::GlobalTensor<uint8_t> global;
|
||||
global.SetGlobalBuffer(flagAddr + idx * SOFT_SYNC_SPACE_SIZE);
|
||||
while (true) {
|
||||
__asm__ __volatile__("");
|
||||
AscendC::DataCacheCleanAndInvalid<uint8_t, AscendC::CacheLine::SINGLE_CACHE_LINE,
|
||||
AscendC::DcciDst::CACHELINE_OUT>(global);
|
||||
__asm__ __volatile__("");
|
||||
uint8_t value = global.GetValue(0);
|
||||
if (value >= target) {
|
||||
__asm__ __volatile__("");
|
||||
AscendC::DataCacheCleanAndInvalid<uint8_t, AscendC::CacheLine::SINGLE_CACHE_LINE,
|
||||
AscendC::DcciDst::CACHELINE_OUT>(global);
|
||||
__asm__ __volatile__("");
|
||||
break;
|
||||
}
|
||||
}
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
}
|
||||
|
||||
__aicore__ inline static void CalQuantRow(const uint32_t column, uint32_t &row)
|
||||
{
|
||||
row = QUANT_SPACE_FACTOR / column;
|
||||
row = row < MAX_QUANT_ROW_ONCE ? row : MAX_QUANT_ROW_ONCE;
|
||||
}
|
||||
|
||||
|
||||
#endif // DISPATCH_GMM_COMBINE_DECODE_BASE_H
|
||||
@@ -0,0 +1,457 @@
|
||||
/*
|
||||
* Copyright (c) 2026 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
#ifndef DISPATCH_GMM_COMBINE_DECODE_BF16_FP16_H
|
||||
#define DISPATCH_GMM_COMBINE_DECODE_BF16_FP16_H
|
||||
|
||||
#include "lib/matmul_intf.h"
|
||||
#include <kernel_operator.h>
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/arch.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/epilogue/tile/tile_broadcast_mul.hpp"
|
||||
#include "catlass/epilogue/tile/tile_broadcast_one_blk.hpp"
|
||||
#include "catlass/epilogue/tile/tile_swizzle.hpp"
|
||||
#include "catlass/gemm/block/block_swizzle.hpp"
|
||||
#include "dispatch_gmm_combine_decode/gemm/kernel/grouped_matmul_slice_m_multistage_workspace_bf16_fp16.h"
|
||||
#include "catlass/gemm/gemm_type.hpp"
|
||||
#include "dispatch_gmm_combine_decode/epilogue/dispatch_policy.h"
|
||||
#include "dispatch_gmm_combine_decode/gemm/dispatch_policy.h"
|
||||
#include "dispatch_gmm_combine_decode/epilogue/block/block_epilogue.h"
|
||||
#include "dispatch_gmm_combine_decode/gemm/block/block_mmad.h"
|
||||
#include "dispatch_gmm_combine_decode/gemm/kernel/grouped_matmul_slice_m_swiglu_multistage_workspace_bf16_fp16.h"
|
||||
|
||||
#include "dispatch_gmm_combine_decode/raw_distributed/cam_moe_distribute_dispatch.h"
|
||||
|
||||
#include "dispatch_gmm_combine_decode_tiling.h"
|
||||
#include "dispatch_gmm_combine_decode_base.h"
|
||||
|
||||
using namespace Catlass;
|
||||
|
||||
namespace DispatchGmmCombineDecodeBf16Fp16Impl {
|
||||
|
||||
using MmadAtlasA2Custom =
|
||||
Gemm::MmadAtlasA2PreloadAsyncWithCallback<CUSTOM_PRELOAD_STAGES, CUSTOM_L1_STAGES, CUSTOM_L0A_STAGES,
|
||||
CUSTOM_L0B_STAGES, CUSTOM_L0C_STAGES, CUSTOM_ENABLE_UNIT_FLAG,
|
||||
CUSTOM_ENABLE_SHUFFLE_K>;
|
||||
|
||||
using Gmm1L1TileShape = GemmShape<FP16_BF16_L1M, FP16_BF16_L1N, GMM1_L1K>;
|
||||
using Gmm1L0TileShape = GemmShape<Gmm1L1TileShape::M, Gmm1L1TileShape::N, GMM1_L0K>;
|
||||
using Gmm1EpilogueTileShape = MatrixShape<GMM1_EPIM, Gmm1L1TileShape::N>;
|
||||
using Gmm1BlockScheduler = typename Gemm::Block::GemmIdentityBlockSwizzle<GMM1_SWIZZLE_OFFSET, GMM1_SWIZZLE_DIRECTION>;
|
||||
|
||||
using Gmm2L1TileShape = GemmShape<FP16_BF16_GMM2_L1M, FP16_BF16_GMM2_L1N, GMM2_L1K>;
|
||||
using Gmm2L0TileShape = GemmShape<Gmm2L1TileShape::M, Gmm2L1TileShape::N, GMM2_L0K>;
|
||||
using Gmm2EpilogueTileShape = MatrixShape<GMM2_EPIM, Gmm2L1TileShape::N>;
|
||||
using Gmm2BlockScheduler = typename Gemm::Block::GemmIdentityBlockSwizzle<GMM2_SWIZZLE_OFFSET, GMM2_SWIZZLE_DIRECTION>;
|
||||
using Gmm2DispatchPolicy =
|
||||
Gemm::MmadAtlasA2PreloadAsyncWithCallbackResidentA<CUSTOM_PRELOAD_STAGES, GMM2_L1A_STAGES, GMM2_L1B_STAGES,
|
||||
GMM2_L0A_STAGES, GMM2_L0B_STAGES, CUSTOM_L0C_STAGES,
|
||||
CUSTOM_ENABLE_UNIT_FLAG, CUSTOM_ENABLE_SHUFFLE_K>;
|
||||
|
||||
template <TemplateMC2TypeClass, class L1TileShape_, class L0TileShape_, class EpilogueTileShape_,
|
||||
class BlockScheduler_, class DispatchPolicy_ = MmadAtlasA2Custom>
|
||||
CATLASS_DEVICE void GmmDeqSwigluQuant(GemmCoord problemShape, uint32_t groupCount, GM_ADDR gmGroupList, GM_ADDR gmA,
|
||||
layout::RowMajor layoutA, GM_ADDR gmB,
|
||||
typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type layoutB,
|
||||
GM_ADDR gmScale,
|
||||
layout::VectorLayout layoutScale, GM_ADDR gmPerTokenScale,
|
||||
layout::VectorLayout layoutPerTokenScale, GM_ADDR gmD, layout::RowMajor layoutD,
|
||||
GM_ADDR gmDequantScale, layout::VectorLayout layoutDequantScale, GM_ADDR gmWorkspace,
|
||||
GM_ADDR gmX, GM_ADDR debugGm, GM_ADDR gmexpertIds, GM_ADDR gmExpandIdx,
|
||||
GM_ADDR gmEpSendCount, GM_ADDR xActiveMask, GM_ADDR gmResvered, GM_ADDR gmExpertTokenNums,
|
||||
uint32_t epRankSize, uint32_t epRankId, uint32_t moeExpertNum,
|
||||
uint32_t moeExpertNumPerRank, uint32_t sharedExpertNum, uint32_t sharedExpertRankNum,
|
||||
uint32_t quantMode, uint32_t globalBs, uint32_t bs, uint32_t topK, uint32_t tokenLen)
|
||||
{
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
using DispatchPolicy = DispatchPolicy_;
|
||||
using L1TileShape = L1TileShape_;
|
||||
using L0TileShape = L0TileShape_;
|
||||
|
||||
using AType = Gemm::GemmType<ExpandXType, layout::RowMajor>;
|
||||
using LayoutB = typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type;
|
||||
using BType = Gemm::GemmType<WType, LayoutB>;
|
||||
using CType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
|
||||
using BlockMmad = Gemm::Block::BlockMmad<DispatchPolicy, L1TileShape, L0TileShape, AType, BType, CType>;
|
||||
|
||||
constexpr uint32_t ubStages = 1;
|
||||
using EpilogueDispatchPolicy = Epilogue::EpilogueAtlasA2Swiglu<ubStages, 0>;
|
||||
using ScaleType = Gemm::GemmType<W1ScaleType, layout::VectorLayout>;
|
||||
using PerTokenScaleType = Gemm::GemmType<float, layout::VectorLayout>;
|
||||
using DType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
|
||||
using RowBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
using BroadcastOneBlkType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
using OneBlkColumnBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
|
||||
using EpilogueTileShape = EpilogueTileShape_;
|
||||
using TileRowBroadcastMul = Epilogue::Tile::TileRowBroadcastMul<ArchTag, RowBroadcastMulType, EpilogueTileShape>;
|
||||
using TileBroadcastOneBlk =
|
||||
Epilogue::Tile::TileBroadcastOneBlk<ArchTag, BroadcastOneBlkType, EpilogueTileShape::ROW>;
|
||||
using TileOneBlkColumnBroadcastMul =
|
||||
Epilogue::Tile::TileOneBlkColumnBroadcastMul<ArchTag, OneBlkColumnBroadcastMulType, EpilogueTileShape>;
|
||||
using TileCopy = Epilogue::Tile::TileCopy<ArchTag, CType, ScaleType, PerTokenScaleType, DType>;
|
||||
using TileScheduler = Epilogue::Tile::EpilogueHorizontalTileSwizzle;
|
||||
|
||||
using BlockEpilogue = Epilogue::Block::BlockEpilogue<EpilogueDispatchPolicy, CType, ScaleType, PerTokenScaleType,
|
||||
DType, TileRowBroadcastMul, TileBroadcastOneBlk,
|
||||
TileOneBlkColumnBroadcastMul, TileCopy, TileScheduler>;
|
||||
|
||||
using BlockScheduler = BlockScheduler_;
|
||||
|
||||
// kernel level
|
||||
using ElementGroupList = int64_t;
|
||||
|
||||
using GemmKernel = typename std::conditional<
|
||||
(EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) != 0,
|
||||
Gemm::Kernel::GroupedMatmulSliceMSwigluMultiStageWorkspace<
|
||||
TemplateMC2TypeFunc, BlockMmad, BlockEpilogue, BlockScheduler, WORKSPACE_STAGES, ElementGroupList>,
|
||||
Gemm::Kernel::GroupedMatmulSliceMSwigluMultiStageWorkspaceWithShallowDispatch<
|
||||
TemplateMC2TypeFunc, BlockMmad, BlockEpilogue, BlockScheduler, WORKSPACE_STAGES, ElementGroupList>>::type;
|
||||
|
||||
if constexpr ((EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) != 0) {
|
||||
typename GemmKernel::Params params{problemShape,
|
||||
groupCount,
|
||||
gmGroupList,
|
||||
gmA,
|
||||
layoutA,
|
||||
gmB,
|
||||
layoutB,
|
||||
gmScale,
|
||||
layoutScale,
|
||||
gmPerTokenScale,
|
||||
layoutPerTokenScale,
|
||||
gmD,
|
||||
layoutD,
|
||||
gmDequantScale,
|
||||
layoutDequantScale,
|
||||
gmWorkspace,
|
||||
gmX,
|
||||
debugGm,
|
||||
gmexpertIds,
|
||||
gmExpandIdx,
|
||||
gmEpSendCount,
|
||||
xActiveMask,
|
||||
gmResvered,
|
||||
gmExpertTokenNums,
|
||||
epRankSize,
|
||||
epRankId,
|
||||
moeExpertNum,
|
||||
moeExpertNumPerRank,
|
||||
sharedExpertNum,
|
||||
sharedExpertRankNum,
|
||||
quantMode,
|
||||
globalBs,
|
||||
bs,
|
||||
topK,
|
||||
tokenLen};
|
||||
// call a kernel
|
||||
GemmKernel gemm;
|
||||
gemm(params);
|
||||
} else {
|
||||
typename GemmKernel::Params params{problemShape,
|
||||
groupCount,
|
||||
gmGroupList,
|
||||
gmA,
|
||||
layoutA,
|
||||
gmB,
|
||||
layoutB,
|
||||
gmScale,
|
||||
layoutScale,
|
||||
gmPerTokenScale,
|
||||
layoutPerTokenScale,
|
||||
gmD,
|
||||
layoutD,
|
||||
gmDequantScale,
|
||||
layoutDequantScale,
|
||||
gmWorkspace};
|
||||
// call a kernel
|
||||
GemmKernel gemm;
|
||||
gemm(params);
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass, class L1TileShape_, class L0TileShape_, class EpilogueTileShape_, class BlockScheduler_,
|
||||
class DispatchPolicy_ = MmadAtlasA2Custom>
|
||||
CATLASS_DEVICE void GmmDeq(GemmCoord problemShape, uint32_t groupCount, GM_ADDR gmGroupList, GM_ADDR gmA,
|
||||
layout::RowMajor layoutA, GM_ADDR gmB,
|
||||
typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type layoutB,
|
||||
GM_ADDR gmScale,
|
||||
layout::VectorLayout layoutScale, GM_ADDR gmPerTokenScale,
|
||||
layout::VectorLayout layoutPerTokenScale, GM_ADDR gmD, layout::RowMajor layoutD,
|
||||
GM_ADDR gmWorkspace, void *combiner)
|
||||
{
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
using DispatchPolicy = DispatchPolicy_;
|
||||
using L1TileShape = L1TileShape_;
|
||||
using L0TileShape = L0TileShape_;
|
||||
|
||||
using AType = Gemm::GemmType<ExpandXType, layout::RowMajor>;
|
||||
using LayoutB = typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type;
|
||||
using BType = Gemm::GemmType<WType, LayoutB>;
|
||||
using CType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
|
||||
using BlockMmad = Gemm::Block::BlockMmad<DispatchPolicy, L1TileShape, L0TileShape, AType, BType, CType>;
|
||||
|
||||
constexpr uint32_t ubStages = 1;
|
||||
using EpilogueDispatchPolicy = Epilogue::EpilogueAtlasA2Combine<ubStages, EXEC_FLAG>;
|
||||
using ScaleType = Gemm::GemmType<W2ScaleType, layout::VectorLayout>;
|
||||
using PerTokenScaleType = Gemm::GemmType<float, layout::VectorLayout>;
|
||||
using DType = Gemm::GemmType<ExpandXType, layout::RowMajor>;
|
||||
|
||||
using RowBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
using BroadcastOneBlkType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
using OneBlkColumnBroadcastMulType = Gemm::GemmType<float, layout::RowMajor>;
|
||||
|
||||
using EpilogueTileShape = EpilogueTileShape_;
|
||||
using TileRowBroadcastMul = Epilogue::Tile::TileRowBroadcastMul<ArchTag, RowBroadcastMulType, EpilogueTileShape>;
|
||||
using TileBroadcastOneBlk =
|
||||
Epilogue::Tile::TileBroadcastOneBlk<ArchTag, BroadcastOneBlkType, EpilogueTileShape::ROW>;
|
||||
using TileOneBlkColumnBroadcastMul =
|
||||
Epilogue::Tile::TileOneBlkColumnBroadcastMul<ArchTag, OneBlkColumnBroadcastMulType, EpilogueTileShape>;
|
||||
using TileCopy = Epilogue::Tile::TileCopy<ArchTag, CType, ScaleType, PerTokenScaleType, DType>;
|
||||
using TileScheduler = Epilogue::Tile::EpilogueHorizontalTileSwizzle;
|
||||
|
||||
using BlockEpilogue = Epilogue::Block::BlockEpilogue<EpilogueDispatchPolicy, CType, ScaleType, PerTokenScaleType,
|
||||
DType, TileRowBroadcastMul, TileBroadcastOneBlk,
|
||||
TileOneBlkColumnBroadcastMul, TileCopy, TileScheduler>;
|
||||
|
||||
using BlockScheduler = BlockScheduler_;
|
||||
|
||||
// kernel level
|
||||
using ElementGroupList = int64_t;
|
||||
using GemmKernel = Gemm::Kernel::GroupedMatmulSliceMMultiStageWorkspace<
|
||||
TemplateMC2TypeFunc, BlockMmad, BlockEpilogue, BlockScheduler, WORKSPACE_STAGES, ElementGroupList>;
|
||||
|
||||
typename GemmKernel::Params params{
|
||||
problemShape, groupCount, gmGroupList, gmA, layoutA, gmB, layoutB, gmScale,
|
||||
layoutScale, gmPerTokenScale, layoutPerTokenScale, gmD, layoutD, gmWorkspace, combiner};
|
||||
|
||||
// call a kernel
|
||||
GemmKernel gemm;
|
||||
gemm(params);
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
class DispatchGmmCombineDecodeBf16Fp16
|
||||
{
|
||||
public:
|
||||
__aicore__ inline DispatchGmmCombineDecodeBf16Fp16(){};
|
||||
__aicore__ inline void Init(
|
||||
// input
|
||||
GM_ADDR x, GM_ADDR expert_ids, GM_ADDR gmm1_permuted_weight, GM_ADDR gmm1_permuted_weight_scale,
|
||||
GM_ADDR gmm2_weight, GM_ADDR gmm2_weight_scale, GM_ADDR expert_scales, GM_ADDR expert_smooth_scales, GM_ADDR x_active_mask,
|
||||
// output
|
||||
GM_ADDR output, GM_ADDR expertTokenNums,
|
||||
// system
|
||||
GM_ADDR workspaceGM, AscendC::TPipe *pipe, const DispatchGmmCombineDecodeTilingData *tilingData);
|
||||
__aicore__ inline void Process();
|
||||
|
||||
private:
|
||||
GM_ADDR gmX_;
|
||||
GM_ADDR gmexpertIds_;
|
||||
GM_ADDR gmPermuteWeight1_;
|
||||
GM_ADDR gmPermuteScale1_;
|
||||
GM_ADDR gmWeight2_;
|
||||
GM_ADDR gmScale2_;
|
||||
GM_ADDR gmOutput_;
|
||||
GM_ADDR gmExpertTokenNums_;
|
||||
GM_ADDR workspaceGM_;
|
||||
GM_ADDR gmSmoothScales_;
|
||||
GM_ADDR gmexpertScales_;
|
||||
GM_ADDR xActiveMask_;
|
||||
|
||||
uint32_t maxTokenNum_{0};
|
||||
uint32_t gmm1OutputDim_{0};
|
||||
uint32_t tokenHiddenSize_{0};
|
||||
uint32_t groupCount_{0};
|
||||
uint32_t gmm2OutputDim_{0};
|
||||
uint32_t gmm2InputDim_{0};
|
||||
uint32_t globalRankId_{0};
|
||||
uint32_t winSizePerRank_{0};
|
||||
uint32_t blockDim_{0};
|
||||
uint32_t epRankSize_{0};
|
||||
uint32_t epRankId_{0};
|
||||
uint32_t moeExpertNum_{0};
|
||||
uint32_t moeExpertNumPerRank_{0};
|
||||
uint32_t sharedExpertNum_{0};
|
||||
uint32_t sharedExpertRankNum_{0};
|
||||
uint32_t quantMode_{0};
|
||||
uint32_t globalBs_{0};
|
||||
uint32_t bs_{0};
|
||||
uint32_t maxBs_{0};
|
||||
uint32_t topK_{0};
|
||||
|
||||
AscendC::TPipe *tpipe_{nullptr};
|
||||
__gm__ HcclOpResParam *winContext_{nullptr};
|
||||
const DispatchGmmCombineDecodeTilingData *tilingData_;
|
||||
};
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void DispatchGmmCombineDecodeBf16Fp16<TemplateMC2TypeFunc>::Init(
|
||||
// input
|
||||
GM_ADDR x, GM_ADDR expert_ids, GM_ADDR gmm1_permuted_weight, GM_ADDR gmm1_permuted_weight_scale,
|
||||
GM_ADDR gmm2_weight, GM_ADDR gmm2_weight_scale, GM_ADDR expert_scales, GM_ADDR expert_smooth_scales,
|
||||
GM_ADDR x_active_mask,
|
||||
// output
|
||||
GM_ADDR output, GM_ADDR expertTokenNums,
|
||||
// system
|
||||
GM_ADDR workspaceGM, AscendC::TPipe *pipe, const DispatchGmmCombineDecodeTilingData *tilingData)
|
||||
{
|
||||
tpipe_ = pipe;
|
||||
blockDim_ = AscendC::GetBlockNum();
|
||||
winContext_ = (__gm__ HcclOpResParam *)AscendC::GetHcclContext<AscendC::HCCL_GROUP_ID_0>();
|
||||
|
||||
gmSmoothScales_ = expert_smooth_scales; // not used now
|
||||
gmX_ = x; // input token
|
||||
gmexpertIds_ = expert_ids;
|
||||
gmPermuteWeight1_ = gmm1_permuted_weight;
|
||||
gmPermuteScale1_ = nullptr;
|
||||
gmWeight2_ = gmm2_weight;
|
||||
gmScale2_ = nullptr;
|
||||
gmOutput_ = output;
|
||||
gmExpertTokenNums_ = expertTokenNums;
|
||||
workspaceGM_ = workspaceGM;
|
||||
gmexpertScales_ = expert_scales;
|
||||
xActiveMask_ = x_active_mask;
|
||||
tilingData_ = tilingData;
|
||||
epRankSize_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.epRankSize;
|
||||
epRankId_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.epRankId;
|
||||
moeExpertNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNum;
|
||||
moeExpertNumPerRank_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
sharedExpertNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertNum;
|
||||
sharedExpertRankNum_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.sharedExpertRankNum;
|
||||
quantMode_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.quantMode;
|
||||
globalBs_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.globalBs;
|
||||
bs_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.bs;
|
||||
topK_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.k;
|
||||
maxBs_ = globalBs_ / epRankSize_;
|
||||
|
||||
bool isShareExpert = (epRankId_ < sharedExpertRankNum_);
|
||||
if (isShareExpert) {
|
||||
maxTokenNum_ = maxBs_ * epRankSize_ / sharedExpertRankNum_;
|
||||
} else {
|
||||
maxTokenNum_ = maxBs_ * epRankSize_ * (topK_ < moeExpertNumPerRank_ ? topK_ : moeExpertNumPerRank_);
|
||||
}
|
||||
|
||||
gmm1OutputDim_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.gmm1HLen;
|
||||
tokenHiddenSize_ = tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.h;
|
||||
groupCount_ = isShareExpert ? 1 : tilingData->disGmmDeqSwigluQuantGmmDeqComInfo.moeExpertNumPerRank;
|
||||
gmm2OutputDim_ = tokenHiddenSize_;
|
||||
gmm2InputDim_ = gmm1OutputDim_ / 2;
|
||||
}
|
||||
|
||||
template<uint32_t EXEC_FLAG, typename WType>
|
||||
__aicore__ inline auto CreateWeightLayout(uint32_t k, uint32_t n) {
|
||||
if constexpr ((EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0) {
|
||||
MatrixCoord mc{k, n};
|
||||
return layout::RowMajor::template MakeLayoutInUb<WType>(mc);
|
||||
} else {
|
||||
return layout::zN::template MakeLayout<WType>(k, n);
|
||||
}
|
||||
}
|
||||
|
||||
template <TemplateMC2TypeClass>
|
||||
__aicore__ inline void DispatchGmmCombineDecodeBf16Fp16<TemplateMC2TypeFunc>::Process()
|
||||
{
|
||||
using LayoutB = typename std::conditional<(EXEC_FLAG & EXEC_FLAG_ND_FORMAT) != 0, layout::RowMajor, layout::zN>::type;
|
||||
GemmCoord gmm1ProblemShape{maxTokenNum_, gmm1OutputDim_, tokenHiddenSize_};
|
||||
GemmCoord gmm2ProblemShape{maxTokenNum_, gmm2OutputDim_, gmm2InputDim_};
|
||||
|
||||
layout::RowMajor layoutX1{maxTokenNum_, tokenHiddenSize_};
|
||||
auto layoutWeight1 = CreateWeightLayout<EXEC_FLAG, WType>(tokenHiddenSize_, gmm1OutputDim_);
|
||||
layout::VectorLayout layoutW1Scale{gmm1OutputDim_};
|
||||
layout::VectorLayout layoutX1Scale{maxTokenNum_};
|
||||
layout::RowMajor layoutX2{maxTokenNum_, gmm2InputDim_};
|
||||
auto layoutWeight2 = CreateWeightLayout<EXEC_FLAG, WType>(gmm2InputDim_, gmm2OutputDim_);
|
||||
layout::VectorLayout layoutW2Scale{gmm2OutputDim_};
|
||||
layout::VectorLayout layoutX2Scale{maxTokenNum_};
|
||||
layout::RowMajor layoutOutput{maxTokenNum_, gmm2OutputDim_};
|
||||
|
||||
size_t workspaceOffset = 0;
|
||||
constexpr int32_t resveredWorkSpaceSize = 256 * 1024;
|
||||
int64_t x1TokenSize = maxTokenNum_ * tokenHiddenSize_ * sizeof(ExpandXType);
|
||||
int64_t x2TokenSize = maxTokenNum_ * gmm2InputDim_ * sizeof(ExpandXType);
|
||||
int64_t maxTokenSize = x1TokenSize < x2TokenSize ? x2TokenSize : x1TokenSize;
|
||||
GM_ADDR gmX1 = workspaceGM_ + workspaceOffset;
|
||||
GM_ADDR gmX2 = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(maxTokenSize);
|
||||
GM_ADDR gmX1Scale = nullptr;
|
||||
GM_ADDR gmX2Scale = nullptr;
|
||||
GM_ADDR gmWorkspace = workspaceGM_ + workspaceOffset;
|
||||
GM_ADDR gmCVSwap = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(static_cast<size_t>(blockDim_) * (FP16_BF16_L1M * FP16_BF16_L1N) *
|
||||
WORKSPACE_STAGES * sizeof(float));
|
||||
int64_t swigluOutSize = maxTokenNum_ * gmm1OutputDim_ * sizeof(float);
|
||||
int64_t gmm2OutSize = maxTokenNum_ * tokenHiddenSize_ * sizeof(ExpandXType);
|
||||
int64_t maxSwigluGmm2Size = swigluOutSize < gmm2OutSize ? gmm2OutSize : swigluOutSize;
|
||||
GM_ADDR gmSwigluOut = workspaceGM_ + workspaceOffset;
|
||||
GM_ADDR gmGmm2DepOut = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(maxSwigluGmm2Size);
|
||||
GM_ADDR gmGroupList = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(static_cast<size_t>(groupCount_) * sizeof(int64_t));
|
||||
GM_ADDR gmExpandIdx = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(static_cast<size_t>(bs_) * topK_ * sizeof(int32_t));
|
||||
GM_ADDR gmEpSendCount = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(static_cast<size_t>(epRankSize_) * groupCount_ * sizeof(int32_t));
|
||||
GM_ADDR gmResvered = workspaceGM_ + workspaceOffset;
|
||||
workspaceOffset += RoundUp<GM_ALIGN_BYTE>(resveredWorkSpaceSize);
|
||||
|
||||
if constexpr ((EXEC_FLAG & EXEC_FLAG_DEEP_FUSE) == 0) {
|
||||
if constexpr (g_coreType == AscendC::AIV) {
|
||||
AscendC::TPipe tpipe;
|
||||
MoeDistributeDispatchImpl::CamMoeDistributeDispatch<ExpandXType, ExpandXType, false, false, false, false, EXEC_FLAG>
|
||||
dispatcher;
|
||||
dispatcher.Init(gmX_, gmexpertIds_, gmSmoothScales_, xActiveMask_, gmX1, gmX1Scale, gmExpandIdx, gmGroupList,
|
||||
gmEpSendCount, gmExpertTokenNums_, nullptr, gmWorkspace, &tpipe, tilingData_);
|
||||
dispatcher.Process();
|
||||
tpipe.Destroy();
|
||||
icache_preload(8);
|
||||
}
|
||||
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
Arch::CrossCoreFlag gmm1AivFinished{0};
|
||||
if constexpr (g_coreType == AscendC::AIV) {
|
||||
Arch::CrossCoreBarrier<0x0, PIPE_MTE3>();
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(gmm1AivFinished);
|
||||
} else {
|
||||
Arch::CrossCoreWaitFlag(gmm1AivFinished);
|
||||
}
|
||||
}
|
||||
GmmDeqSwigluQuant<TemplateMC2TypeFunc, Gmm1L1TileShape, Gmm1L0TileShape, Gmm1EpilogueTileShape,
|
||||
Gmm1BlockScheduler>(
|
||||
gmm1ProblemShape, groupCount_, gmGroupList, gmX1, layoutX1, gmPermuteWeight1_, layoutWeight1,
|
||||
gmPermuteScale1_, layoutW1Scale, gmX1Scale, layoutX1Scale, gmX2, layoutX2, gmX2Scale,
|
||||
layoutX2Scale, gmWorkspace, gmX_, gmSmoothScales_, gmexpertIds_, gmExpandIdx, gmEpSendCount, xActiveMask_, gmResvered,
|
||||
gmExpertTokenNums_, epRankSize_, epRankId_, moeExpertNum_, moeExpertNumPerRank_, sharedExpertNum_,
|
||||
sharedExpertRankNum_, quantMode_, globalBs_, bs_, topK_, tokenHiddenSize_);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
Arch::CrossCoreFlag gmm1AivFinished{0};
|
||||
if constexpr (g_coreType == AscendC::AIV) {
|
||||
Arch::CrossCoreBarrier<0x0, PIPE_MTE3>();
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(gmm1AivFinished);
|
||||
} else {
|
||||
Arch::CrossCoreWaitFlag(gmm1AivFinished);
|
||||
}
|
||||
|
||||
MoeDistributeCombineImpl::CamMoeDistributeCombine<TemplateMC2TypeFunc> combiner;
|
||||
if (g_coreType == AscendC::AIV) {
|
||||
combiner.Init(gmGmm2DepOut, gmexpertIds_, gmExpandIdx, gmEpSendCount, nullptr, gmexpertScales_, xActiveMask_, gmOutput_,
|
||||
workspaceGM_, nullptr, tilingData_);
|
||||
}
|
||||
GmmDeq<TemplateMC2TypeFunc, Gmm2L1TileShape, Gmm2L0TileShape, Gmm2EpilogueTileShape, Gmm2BlockScheduler,
|
||||
Gmm2DispatchPolicy>(gmm2ProblemShape, groupCount_, gmGroupList, gmX2, layoutX2, gmWeight2_, layoutWeight2,
|
||||
gmScale2_, layoutW2Scale, gmX2Scale, layoutX2Scale, gmGmm2DepOut,
|
||||
layoutOutput, gmWorkspace, &combiner);
|
||||
}
|
||||
} // namespace DispatchGmmCombineDecodeBf16Fp16Impl
|
||||
#endif // DISPATCH_GMM_COMBINE_DECODE_BF16_FP16_H
|
||||
@@ -0,0 +1,84 @@
|
||||
/*
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 1.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
#ifndef DISPATCH_GMM_COMBINE_DECODE_TILING_H
|
||||
#define DISPATCH_GMM_COMBINE_DECODE_TILING_H
|
||||
|
||||
#include "kernel_tiling/kernel_tiling.h"
|
||||
|
||||
struct DispatchGmmCombineDecodeInfo {
|
||||
uint32_t epRankSize; // epRankSize
|
||||
uint32_t epRankId; // epRankId
|
||||
uint32_t moeExpertNum; // moe expert number
|
||||
uint32_t moeExpertNumPerRank; // moe expert number per rank
|
||||
uint32_t sharedExpertNum; // shared expert number
|
||||
uint32_t sharedExpertRankNum; // shared expert rank number
|
||||
uint32_t quantMode; // quant mode
|
||||
uint32_t globalBs; // globalBs = BS * worldSize
|
||||
uint32_t bs; // bs
|
||||
uint32_t k; // k
|
||||
uint32_t h; // h
|
||||
uint32_t aicNum; // aicNum
|
||||
uint32_t aivNum; // aivNum
|
||||
uint64_t totalUbSize;
|
||||
uint64_t totalWinSize;
|
||||
uint64_t gmm1HLen;
|
||||
bool isTensorList;
|
||||
bool isBf16Fp16W;
|
||||
bool isNDFormat;
|
||||
};
|
||||
|
||||
struct DispatchGmmCombineDecodeTilingData {
|
||||
Mc2InitTiling mc2InitTiling;
|
||||
Mc2CcTiling mc2CcTiling;
|
||||
DispatchGmmCombineDecodeInfo disGmmDeqSwigluQuantGmmDeqComInfo;
|
||||
};
|
||||
|
||||
constexpr uint32_t GM_ALIGN_BYTE = 512;
|
||||
constexpr uint32_t CUSTOM_PRELOAD_STAGES = 1;
|
||||
constexpr uint32_t CUSTOM_L1_STAGES = 2;
|
||||
constexpr uint32_t CUSTOM_L0A_STAGES = 2;
|
||||
constexpr uint32_t CUSTOM_L0B_STAGES = 2;
|
||||
constexpr uint32_t CUSTOM_L0C_STAGES = 1;
|
||||
constexpr bool CUSTOM_ENABLE_UNIT_FLAG = true;
|
||||
constexpr bool CUSTOM_ENABLE_SHUFFLE_K = true;
|
||||
|
||||
constexpr uint32_t FP16_BF16_L1M = 128;
|
||||
constexpr uint32_t FP16_BF16_L1N = 128;
|
||||
constexpr uint32_t GMM1_L1M = 256;
|
||||
constexpr uint32_t GMM1_L1N = 128;
|
||||
constexpr uint32_t GMM1_L1K = 512;
|
||||
constexpr uint32_t GMM1_L0K = 128;
|
||||
constexpr uint32_t GMM1_EPIM = 64;
|
||||
constexpr uint32_t GMM1_SWIZZLE_OFFSET = 3;
|
||||
constexpr uint32_t GMM1_SWIZZLE_DIRECTION = 0;
|
||||
|
||||
constexpr uint32_t FP16_BF16_GMM2_L1M = 64;
|
||||
constexpr uint32_t FP16_BF16_GMM2_L1N = 128;
|
||||
constexpr uint32_t GMM2_L1A_STAGES = 4;
|
||||
constexpr uint32_t GMM2_L1B_STAGES = 2;
|
||||
constexpr uint32_t GMM2_L0A_STAGES = 4;
|
||||
constexpr uint32_t GMM2_L0B_STAGES = 2;
|
||||
constexpr uint32_t GMM2_L1M = 128;
|
||||
constexpr uint32_t GMM2_L1N = 256;
|
||||
constexpr uint32_t GMM2_L1K = 512;
|
||||
constexpr uint32_t GMM2_L0K = 128;
|
||||
constexpr uint32_t GMM2_EPIM = 32;
|
||||
constexpr uint32_t GMM2_SWIZZLE_OFFSET = 3;
|
||||
constexpr uint32_t GMM2_SWIZZLE_DIRECTION = 0;
|
||||
|
||||
constexpr uint32_t WORKSPACE_STAGES = 4;
|
||||
|
||||
constexpr uint32_t EXEC_FLAG_DEEP_FUSE = (1U << 0);
|
||||
constexpr uint32_t EXEC_FLAG_TENSOR_LIST = (1U << 1);
|
||||
constexpr uint32_t EXEC_FLAG_X_ACTIVE_MASK = (1U << 2);
|
||||
constexpr uint32_t EXEC_FLAG_ND_FORMAT = (1U << 3);
|
||||
|
||||
#endif // DISPATCH_GMM_COMBINE_DECODE_TILING_H
|
||||
@@ -0,0 +1,365 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify.
|
||||
* Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
* This file is a part of the CANN Open Software.
|
||||
* Licensed under CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
* Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See LICENSE in the root of the software repository for the full text of the License.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file moe_distribute_base.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef MOE_DISTRIBUTE_BASE_H
|
||||
#define MOE_DISTRIBUTE_BASE_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
constexpr uint32_t LOCAL_NOTIFY_MAX_NUM = 64;
|
||||
constexpr uint32_t LOCAL_STREAM_MAX_NUM = 19U;
|
||||
constexpr uint32_t AICPU_OP_NOTIFY_MAX_NUM = 2;
|
||||
constexpr uint32_t AICPU_MAX_RANK_NUM = 128 * 1024;
|
||||
constexpr uint32_t TIME_CYCLE = 50; // 系统cycle数转换成时间的基准单位,固定为50
|
||||
|
||||
struct HcclSignalInfo {
|
||||
uint64_t resId; // 在代表event时为eventid,notify时为notifyid
|
||||
uint64_t addr;
|
||||
uint32_t devId;
|
||||
uint32_t tsId;
|
||||
uint32_t rankId;
|
||||
uint32_t flag;
|
||||
};
|
||||
|
||||
struct ListCommon {
|
||||
uint64_t nextHost;
|
||||
uint64_t preHost;
|
||||
uint64_t nextDevice;
|
||||
uint64_t preDevice;
|
||||
};
|
||||
|
||||
struct HcclStreamInfo {
|
||||
int32_t streamIds;
|
||||
uint32_t sqIds;
|
||||
uint32_t cqIds; // 记录物理cqId
|
||||
uint32_t logicCqids; // 记录逻辑cqId
|
||||
};
|
||||
|
||||
struct LocalResInfoV2 {
|
||||
uint32_t streamNum;
|
||||
uint32_t signalNum;
|
||||
HcclSignalInfo localSignals[LOCAL_NOTIFY_MAX_NUM];
|
||||
HcclStreamInfo streamInfo[LOCAL_STREAM_MAX_NUM];
|
||||
HcclStreamInfo mainStreamInfo;
|
||||
HcclSignalInfo aicpuOpNotify[AICPU_OP_NOTIFY_MAX_NUM]; // 集合通信AICPU展开资源
|
||||
ListCommon nextTagRes; // HccltagLocalResV2
|
||||
};
|
||||
|
||||
enum class rtFloatOverflowMode_t {
|
||||
RT_OVERFLOW_MODE_SATURATION = 0,
|
||||
RT_OVERFLOW_MODE_INFNAN,
|
||||
RT_OVERFLOW_MODE_UNDEF,
|
||||
};
|
||||
|
||||
struct AlgoTopoInfo {
|
||||
uint32_t userRank; // 通信域 RankID
|
||||
uint32_t userRankSize; // 通信域的Rank数量
|
||||
int32_t deviceLogicId;
|
||||
bool isSingleMeshAggregation;
|
||||
uint32_t deviceNumPerAggregation; // 每个Module中的Device数量
|
||||
uint32_t superPodNum; // 集群中总的超节点数
|
||||
uint32_t devicePhyId;
|
||||
uint32_t topoType; // TopoType
|
||||
uint32_t deviceType;
|
||||
uint32_t serverNum;
|
||||
uint32_t meshAggregationRankSize;
|
||||
uint32_t multiModuleDiffDeviceNumMode;
|
||||
uint32_t multiSuperPodDiffServerNumMode;
|
||||
uint32_t realUserRank;
|
||||
bool isDiffDeviceModule;
|
||||
bool isDiffDeviceType;
|
||||
uint32_t gcdDeviceNumPerAggregation;
|
||||
uint32_t moduleNum;
|
||||
uint32_t isUsedRdmaRankPairNum;
|
||||
uint64_t isUsedRdmaRankPair;
|
||||
uint32_t pairLinkCounterNum;
|
||||
uint64_t pairLinkCounter;
|
||||
uint32_t nicNum;
|
||||
uint64_t nicList; // niclist数组指针
|
||||
uint64_t complanRankLength; // complanRank占用的字节数
|
||||
uint64_t complanRank; // 指针
|
||||
uint64_t bridgeRankNum; // bridgeRank占用的个数
|
||||
uint64_t bridgeRank; // 指针
|
||||
uint64_t serverAndsuperPodRankLength; // serverAndsuperPodRank占用的字节数
|
||||
uint64_t serverAndsuperPodRank; // 指针
|
||||
};
|
||||
|
||||
struct HcclOpConfig {
|
||||
uint8_t deterministic; //确定性计算开关
|
||||
uint8_t retryEnable; // 是否重执行
|
||||
uint8_t highPerfEnable;
|
||||
uint8_t padding[5]; // 大小需要64By对齐,未来添加参数时减小padding
|
||||
uint8_t linkTimeOut[8]; // 发送超时时长
|
||||
uint64_t notifyWaitTime; // 超时时长,同HCCL_EXEC_TIMEOUT
|
||||
uint32_t retryHoldTime;
|
||||
uint32_t retryIntervalTime;
|
||||
bool interHccsDisable = false; //使能rdma开关
|
||||
rtFloatOverflowMode_t floatOverflowMode = rtFloatOverflowMode_t::RT_OVERFLOW_MODE_UNDEF;
|
||||
uint32_t multiQpThreshold = 512; // 多QP每个QP分担数据量最小阈值
|
||||
};
|
||||
|
||||
struct HcclMC2WorkSpace {
|
||||
uint64_t workSpace;
|
||||
uint64_t workSpaceSize;
|
||||
};
|
||||
|
||||
struct RemoteResPtr {
|
||||
uint64_t nextHostPtr;
|
||||
uint64_t nextDevicePtr;
|
||||
};
|
||||
|
||||
struct HDCommunicateParams {
|
||||
uint64_t hostAddr { 0 };
|
||||
uint64_t deviceAddr { 0 };
|
||||
uint64_t readCacheAddr { 0 };
|
||||
uint32_t devMemSize{ 0 };
|
||||
uint32_t buffLen{ 0 };
|
||||
uint32_t flag{ 0 };
|
||||
};
|
||||
|
||||
struct HcclRankRelationResV2 {
|
||||
uint32_t remoteUsrRankId;
|
||||
uint32_t remoteWorldRank;
|
||||
uint64_t windowsIn;
|
||||
uint64_t windowsOut;
|
||||
uint64_t windowsExp;
|
||||
ListCommon nextTagRes;
|
||||
};
|
||||
|
||||
struct HcclOpResParam {
|
||||
// 本地资源
|
||||
HcclMC2WorkSpace mc2WorkSpace;
|
||||
uint32_t localUsrRankId; // usrrankid
|
||||
uint32_t rankSize; // 通信域内total rank个数
|
||||
uint64_t winSize; // 每个win大小,静态图时,可能是0,如果通信域内也有动态图,则可能为非0
|
||||
uint64_t localWindowsIn; // 全F为无效值
|
||||
uint64_t localWindowsOut; // 全F为无效值
|
||||
char hcomId[128];
|
||||
// aicore识别remote window
|
||||
uint64_t winExpSize;
|
||||
uint64_t localWindowsExp;
|
||||
uint32_t rWinStart; // 为HcclRankRelationRes起始位置
|
||||
uint32_t rWinOffset; // 为HcclRemoteRes的大小
|
||||
uint64_t version;
|
||||
LocalResInfoV2 localRes;
|
||||
AlgoTopoInfo topoInfo;
|
||||
|
||||
// 外部配置参数
|
||||
HcclOpConfig config;
|
||||
uint64_t hostStateInfo;
|
||||
uint64_t aicpuStateInfo;
|
||||
uint64_t lockAddr;
|
||||
uint32_t rsv[16];
|
||||
uint32_t notifysize; // RDMA场景使用,910B/910_93为4B,其余芯片为8B
|
||||
uint32_t remoteResNum; // 有效的remoteResNum
|
||||
RemoteResPtr remoteRes[AICPU_MAX_RANK_NUM]; //数组指针,指向HcclRankRelationResV2,下标为remoteUserRankId
|
||||
|
||||
// communicate retry
|
||||
HDCommunicateParams kfcControlTransferH2DParams;
|
||||
HDCommunicateParams kfcStatusTransferD2HParams;
|
||||
uint64_t tinyMem; // for all2all
|
||||
uint64_t tinyMemSize;
|
||||
// 零拷贝场景使用
|
||||
uint64_t zeroCopyHeadPtr;
|
||||
uint64_t zeroCopyTailPtr;
|
||||
uint64_t zeroCopyRingBuffer;
|
||||
uint64_t zeroCopyIpcPtrs[16]; // 保存集合通信时每个对端的输入输出内存地址
|
||||
uint32_t zeroCopyDevicePhyId[16]; // 保存每个rank对应的物理卡Id
|
||||
|
||||
bool utraceStatusFlag;
|
||||
};
|
||||
|
||||
// Transport 内存类型
|
||||
enum class HcclAiRMAMemType : uint32_t {
|
||||
LOCAL_INPUT = 0,
|
||||
REMOTE_INPUT,
|
||||
|
||||
LOCAL_OUTPUT,
|
||||
REMOTE_OUTPUT,
|
||||
|
||||
// 可透传更多的内存,可在MAX_NUM之前追加,例如:
|
||||
// LOCAL_EXP,
|
||||
// REMOTE_EXP,
|
||||
MAX_NUM
|
||||
};
|
||||
|
||||
// Transport 内存信息
|
||||
struct HcclAiRMAMemInfo {
|
||||
uint32_t memMaxNum{0}; // 最大内存数量,等于 HcclAiRMAMemType::MAX_NUM
|
||||
uint32_t sizeOfMemDetails{0}; // sizeof(MemDetails),用于内存校验和偏移计算
|
||||
uint64_t memDetailPtr{0}; // MemDetails数组首地址, 个数: HcclAiRMAMemType::MAX_NUM
|
||||
// 可往后追加字段
|
||||
};
|
||||
|
||||
// 全部 Transport QP/Mem 信息
|
||||
struct HcclAiRMAInfo {
|
||||
uint32_t curRankId{0}; // 当前rankId
|
||||
uint32_t rankNum{0}; // rank数量
|
||||
uint32_t qpNum{0}; // 单个Transport的QP数量
|
||||
|
||||
uint32_t sizeOfAiRMAWQ{0}; // sizeof(HcclAiRMAWQ)
|
||||
uint32_t sizeOfAiRMACQ{0}; // sizeof(HcclAiRMACQ)
|
||||
uint32_t sizeOfAiRMAMem{0}; // sizeof(HcclAiRMAMemInfo)
|
||||
|
||||
// HcclAiRMAWQ二维数组首地址
|
||||
// QP个数: rankNum * qpNum
|
||||
// 计算偏移获取SQ指针:sqPtr + (dstRankId * qpNum + qpIndex) * sizeOfAiRMAWQ
|
||||
// 0 <= qpIndex < qpNum
|
||||
uint64_t sqPtr{0};
|
||||
|
||||
// HcclAiRMACQ二维数组首地址
|
||||
// QP个数: rankNum * qpNum
|
||||
// 计算偏移获取SCQ指针:scqPtr + (dstRankId * qpNum + qpIndex) * sizeOfAiRMACQ
|
||||
// 0 <= qpIndex < qpNum
|
||||
uint64_t scqPtr{0};
|
||||
|
||||
// HcclAiRMAWQ二维数组首地址
|
||||
// QP个数: rankNum * qpNum
|
||||
// 计算偏移获取RQ指针:rqPtr + (dstRankId * qpNum + qpIndex) * sizeOfAiRMAWQ
|
||||
// 0 <= qpIndex < qpNum
|
||||
uint64_t rqPtr{0};
|
||||
|
||||
// HcclAiRMACQ二维数组首地址
|
||||
// QP个数: rankNum * qpNum
|
||||
// 计算偏移获取RCQ指针: rcqPtr + (dstRankId * qpNum + qpIndex) * sizeOfAiRMACQ
|
||||
// 0 <= qpIndex < qpNum
|
||||
uint64_t rcqPtr{0};
|
||||
|
||||
// HcclAivMemInfo一维数组
|
||||
// 内存信息个数: rankNum
|
||||
// 计算偏移获取内存信息指针: memPtr + rankId * sizeOfAiRMAMem
|
||||
// srcRankId 获取自身内存信息,dstRankId 获取 Transport 内存信息
|
||||
uint64_t memPtr{0};
|
||||
// 可往后追加字段
|
||||
};
|
||||
struct CombinedCapability {
|
||||
uint64_t dataplaneModeBitmap;
|
||||
};
|
||||
|
||||
struct HcclA2CombineOpParam {
|
||||
uint64_t workSpace; // Address for communication between client and server,
|
||||
// hccl requests and clears
|
||||
uint64_t workSpaceSize; // Space for communication between client and server
|
||||
uint32_t rankId; // id of this rank
|
||||
uint32_t rankNum; // num of ranks in this comm group
|
||||
uint64_t winSize; // size of each windows memory
|
||||
uint64_t windowsIn[AscendC::HCCL_MAX_RANK_NUM]; // windows address for input, windowsIn[rankId] corresponds
|
||||
// to the local card address,
|
||||
// and others are cross-card mapping addresses.
|
||||
uint64_t windowsOut[AscendC::HCCL_MAX_RANK_NUM]; // windows address for output, windowsOut[rankId] corresponds
|
||||
// to the local card address,
|
||||
// and others are cross-card mapping addresses.
|
||||
uint8_t res[8328];
|
||||
uint8_t multiFlag;
|
||||
__gm__ AscendC::IbVerbsData *data;
|
||||
uint64_t dataSize;
|
||||
// 追加字段
|
||||
uint64_t sizeOfAiRMAInfo; // sizeof(HcclAiRMAInfo)
|
||||
uint64_t aiRMAInfo; // HcclAiRMAInfo* 单个结构体指针
|
||||
|
||||
CombinedCapability* capability; // address of the communication capability information structure on the Device
|
||||
uint64_t capabilitySize; // size of the communication capability information structure
|
||||
};
|
||||
enum class DataplaneMode : uint32_t {
|
||||
HOST = 0,
|
||||
AICPU = 1,
|
||||
AIV = 2,
|
||||
};
|
||||
|
||||
enum class DBMode : int32_t {
|
||||
INVALID_DB = -1,
|
||||
HW_DB = 0,
|
||||
SW_DB
|
||||
};
|
||||
|
||||
struct HcclAiRMAWQ {
|
||||
uint32_t wqn{0};
|
||||
uint64_t bufAddr{0};
|
||||
uint32_t wqeSize{0};
|
||||
uint32_t depth{0};
|
||||
uint64_t headAddr{0};
|
||||
uint64_t tailAddr{0};
|
||||
DBMode dbMode{DBMode::INVALID_DB}; // 0-hw/1-sw
|
||||
uint64_t dbAddr{0};
|
||||
uint32_t sl{0};
|
||||
};
|
||||
|
||||
struct HcclAiRMACQ {
|
||||
uint32_t cqn{0};
|
||||
uint64_t bufAddr{0};
|
||||
uint32_t cqeSize{0};
|
||||
uint32_t depth{0};
|
||||
uint64_t headAddr{0};
|
||||
uint64_t tailAddr{0};
|
||||
DBMode dbMode{DBMode::INVALID_DB}; // 0-hw/1-sw
|
||||
uint64_t dbAddr{0};
|
||||
};
|
||||
|
||||
struct hns_roce_rc_sq_wqe {
|
||||
uint32_t byte_4;
|
||||
uint32_t msg_len;
|
||||
uint32_t immtdata;
|
||||
uint32_t byte_16;
|
||||
uint32_t byte_20;
|
||||
uint32_t rkey;
|
||||
uint64_t remoteVA;
|
||||
};
|
||||
|
||||
|
||||
struct hns_roce_lite_wqe_data_seg {
|
||||
uint32_t len;
|
||||
uint32_t lkey;
|
||||
uint64_t localVA;
|
||||
};
|
||||
|
||||
__aicore__ inline void cacheWriteThrough(__gm__ uint8_t* sourceAddr, uint64_t length) {
|
||||
__gm__ uint8_t* start =
|
||||
(__gm__ uint8_t*)((uint64_t)sourceAddr / AscendC::CACHE_LINE_SIZE * AscendC::CACHE_LINE_SIZE);
|
||||
__gm__ uint8_t* end =
|
||||
(__gm__ uint8_t*)(((uint64_t)sourceAddr + length) / AscendC::CACHE_LINE_SIZE * AscendC::CACHE_LINE_SIZE);
|
||||
AscendC::GlobalTensor<uint8_t> global;
|
||||
global.SetGlobalBuffer(start);
|
||||
for (uint32_t i = 0; i <= end - start; i += AscendC::CACHE_LINE_SIZE) {
|
||||
AscendC::DataCacheCleanAndInvalid<uint8_t, AscendC::CacheLine::SINGLE_CACHE_LINE,
|
||||
AscendC::DcciDst::CACHELINE_OUT>(global[i]);
|
||||
}
|
||||
}
|
||||
__aicore__ inline DataplaneMode GetDataplaneMode(GM_ADDR contextGM0) {
|
||||
__gm__ HcclA2CombineOpParam *winContext_ = (__gm__ HcclA2CombineOpParam *)contextGM0;
|
||||
CombinedCapability* capability = winContext_->capability;
|
||||
uint64_t capabilitySize = winContext_->capabilitySize;
|
||||
DataplaneMode dataplaneMode = DataplaneMode::AICPU;
|
||||
if (capability == 0) {
|
||||
return dataplaneMode;
|
||||
}
|
||||
uint64_t dataplaneModeBitmap = capability->dataplaneModeBitmap;
|
||||
if ((dataplaneModeBitmap & 0x04) == 0x04) {
|
||||
dataplaneMode = DataplaneMode::AIV;
|
||||
}
|
||||
return dataplaneMode;
|
||||
}
|
||||
|
||||
__aicore__ inline int64_t GetCurrentTimestampUs()
|
||||
{
|
||||
return AscendC::GetSystemCycle() / TIME_CYCLE;
|
||||
}
|
||||
|
||||
__aicore__ inline void RecordRankCommDuration(AscendC::LocalTensor<int32_t> performanceInfoU32Tensor, uint32_t rankId, int64_t startTime)
|
||||
{
|
||||
int64_t endTime = GetCurrentTimestampUs();
|
||||
int32_t duration = static_cast<int32_t>(endTime - startTime); // int32_t可以表示2^31(us),约35min在实际场景下满足需要
|
||||
performanceInfoU32Tensor.SetValue(rankId * sizeof(int64_t) / sizeof(int32_t), duration); // 使用int32_t是因为atomicAdd不支持int64_t类型,这里只赋值到int64_t的低32位。
|
||||
}
|
||||
#endif // MOE_DISTRIBUTE_BASE_H
|
||||
Reference in New Issue
Block a user