19
csrc/moe/chunk_gated_delta_rule_fwd_h/CMakeLists.txt
Normal file
19
csrc/moe/chunk_gated_delta_rule_fwd_h/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()
|
||||
32
csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/CMakeLists.txt
Normal file
32
csrc/moe/chunk_gated_delta_rule_fwd_h/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,32 @@
|
||||
set(CURRENT_CMAKE_DIR ${CMAKE_CURRENT_SOURCE_DIR})
|
||||
set(CATLASS_INCLUDE_DIR "${CMAKE_SOURCE_DIR}/third_party/catlass/include")
|
||||
get_filename_component(CATLASS_INCLUDE_DIR_ABS ${CATLASS_INCLUDE_DIR} ABSOLUTE)
|
||||
|
||||
set(COMMON_KERNEL_UTILS_DIR "${CMAKE_SOURCE_DIR}/moe/common")
|
||||
get_filename_component(COMMON_KERNEL_UTILS_DIR_ABS ${COMMON_KERNEL_UTILS_DIR} ABSOLUTE)
|
||||
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnnExc PRIVATE
|
||||
chunk_gated_delta_rule_fwd_h_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME ChunkGatedDeltaRuleFwdH
|
||||
OPTIONS
|
||||
--cce-auto-sync=off
|
||||
-Wno-deprecated-declarations
|
||||
-I${CATLASS_INCLUDE_DIR_ABS}
|
||||
-I${COMMON_KERNEL_UTILS_DIR_ABS}
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE chunk_gated_delta_rule_fwd_h ACLNNTYPE aclnn_exclude)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_CURRENT_SOURCE_DIR}
|
||||
${CATLASS_INCLUDE_DIR_ABS}
|
||||
)
|
||||
endif()
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file chunk_gated_delta_rule_fwd_h_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
|
||||
class ChunkGatedDeltaRuleFwdH : public OpDef {
|
||||
public:
|
||||
explicit ChunkGatedDeltaRuleFwdH(const char *name) : OpDef(name)
|
||||
{
|
||||
this->Input("k")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, 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})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
|
||||
this->Input("w")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, 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})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
|
||||
this->Input("u")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, 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})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
|
||||
this->Input("g")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, 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})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
|
||||
this->Input("inital_state")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, 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})
|
||||
.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})
|
||||
.IgnoreContiguous();
|
||||
|
||||
this->Input("cu_seqlens")
|
||||
.ParamType(OPTIONAL)
|
||||
.ValueDepend(OPTIONAL)
|
||||
.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})
|
||||
.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})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
|
||||
this->Input("chunk_indices")
|
||||
.ParamType(OPTIONAL)
|
||||
.ValueDepend(OPTIONAL)
|
||||
.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})
|
||||
.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})
|
||||
.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})
|
||||
.AutoContiguous();
|
||||
|
||||
this->Output("h")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, 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})
|
||||
.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});
|
||||
|
||||
this->Output("v_new")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, ge::DT_FLOAT16, ge::DT_BF16, 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})
|
||||
.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});
|
||||
|
||||
this->Output("final_state")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_BF16, ge::DT_FLOAT16, ge::DT_FLOAT, ge::DT_FLOAT, ge::DT_BF16, ge::DT_FLOAT16, 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})
|
||||
.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});
|
||||
|
||||
this->Attr("output_final_state").AttrType(REQUIRED).Bool(false);
|
||||
this->Attr("chunk_size").AttrType(REQUIRED).Int(64);
|
||||
this->Attr("inital_state_stride0").AttrType(REQUIRED).Int(0);
|
||||
|
||||
OpAICoreConfig aicore_config;
|
||||
aicore_config.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(true)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true)
|
||||
.ExtendCfgInfo("prebuildPattern.value", "Opaque")
|
||||
.ExtendCfgInfo("coreType.value", "AiCore")
|
||||
.ExtendCfgInfo("jitCompile.flag", "static_false,dynamic_false");
|
||||
|
||||
this->AICore().AddConfig("ascend910b", aicore_config);
|
||||
this->AICore().AddConfig("ascend910_93", aicore_config);
|
||||
this->AICore().AddConfig("ascend950", aicore_config);
|
||||
this->AICore().AddConfig("ascend310p", aicore_config);
|
||||
|
||||
}
|
||||
};
|
||||
|
||||
OP_ADD(ChunkGatedDeltaRuleFwdH);
|
||||
|
||||
} // namespace ops
|
||||
@@ -0,0 +1,180 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file chunk_gated_delta_rule_fwd_h_tiling.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "chunk_gated_delta_rule_fwd_h_tiling.h"
|
||||
#include <register/op_impl_registry.h>
|
||||
#include "../tiling_base/data_copy_transpose_tiling.h"
|
||||
#include "../tiling_base/tiling_templates_registry.h"
|
||||
|
||||
namespace optiling {
|
||||
static constexpr size_t INPUT_K_IDX = 0;
|
||||
static constexpr size_t INPUT_W_IDX = 1;
|
||||
static constexpr size_t INPUT_U_IDX = 2;
|
||||
static constexpr size_t INPUT_G_IDX = 3;
|
||||
static constexpr size_t INPUT_INITIAL_STATE_IDX = 4;
|
||||
static constexpr size_t INPUT_SEQLENS_IDX = 5;
|
||||
static constexpr size_t INPUT_CHUNK_INDICES_IDX = 6;
|
||||
|
||||
static constexpr size_t ATTR_STORE_FINAL_STATE_IDX = 0;
|
||||
static constexpr size_t ATTR_CHUNK_SIZE_IDX = 1;
|
||||
static constexpr size_t ATTR_INITIAL_STATE_STRIDE_IDX = 2;
|
||||
|
||||
static constexpr size_t DIM_BATCH = 0;
|
||||
static constexpr size_t DIM_HEAD_NUM = 1;
|
||||
static constexpr size_t DIM_SEQLEN = 2;
|
||||
static constexpr size_t DIM_HEAD_DIM = 3;
|
||||
|
||||
|
||||
static void ChunkGatedDeltaRuleFwdHTilingDataPrint(gert::TilingContext *context, ChunkGatedDeltaRuleFwdHTilingData &tiling)
|
||||
{
|
||||
auto nodeName = context->GetNodeName();
|
||||
OP_LOGD(nodeName, ">>>>>>>>>>>>>>> Start to print ChunkGatedDeltaRuleFwdH tiling data <<<<<<<<<<<<<<<<");
|
||||
OP_LOGD(nodeName, "=== batch: %ld", tiling.get_batch());
|
||||
OP_LOGD(nodeName, "=== seqlen: %ld", tiling.get_seqlen());
|
||||
OP_LOGD(nodeName, "=== kNumHead: %ld", tiling.get_kNumHead());
|
||||
OP_LOGD(nodeName, "=== vNumHead: %ld", tiling.get_vNumHead());
|
||||
OP_LOGD(nodeName, "=== kHeadDim: %ld", tiling.get_kHeadDim());
|
||||
OP_LOGD(nodeName, "=== vHeadDim: %ld", tiling.get_vHeadDim());
|
||||
OP_LOGD(nodeName, "=== chunkSize: %ld", tiling.get_chunkSize());
|
||||
OP_LOGD(nodeName, "=== useInitialState: %ld", tiling.get_useInitialState());
|
||||
OP_LOGD(nodeName, "=== storeFinalState: %ld", tiling.get_storeFinalState());
|
||||
OP_LOGD(nodeName, "=== dataType: %ld", tiling.get_dataType());
|
||||
OP_LOGD(nodeName, "=== isVariedLen: %ld", tiling.get_isVariedLen());
|
||||
OP_LOGD(nodeName, "=== shapeBatch: %ld", tiling.get_shapeBatch());
|
||||
OP_LOGD(nodeName, "=== tokenBatch: %f", tiling.get_tokenBatch());
|
||||
OP_LOGD(nodeName, ">>>>>>>>>>>>>>> Print ChunkGatedDeltaRuleFwdH tiling data end <<<<<<<<<<<<<<<<");
|
||||
}
|
||||
|
||||
ge::graphStatus Tiling4ChunkGatedDeltaRuleFwdH(gert::TilingContext *context)
|
||||
{
|
||||
OP_LOGD(context->GetNodeName(), "Tiling4ChunkGatedDeltaRuleFwdH start.");
|
||||
ChunkGatedDeltaRuleFwdHTilingData tiling;
|
||||
|
||||
gert::Shape kStorageShape = context->GetOptionalInputShape(INPUT_K_IDX)->GetStorageShape();
|
||||
gert::Shape uStorageShape = context->GetOptionalInputShape(INPUT_U_IDX)->GetStorageShape();
|
||||
|
||||
int64_t seqlen = kStorageShape.GetDim(DIM_SEQLEN);
|
||||
int64_t kNumHead = kStorageShape.GetDim(DIM_HEAD_NUM);
|
||||
int64_t vNumHead = uStorageShape.GetDim(DIM_HEAD_NUM);
|
||||
int64_t kHeadDim = kStorageShape.GetDim(DIM_HEAD_DIM);
|
||||
int64_t vHeadDim = uStorageShape.GetDim(DIM_HEAD_DIM);
|
||||
int64_t batch, isVariedLen, shapeBatch, tokenBatch;
|
||||
|
||||
auto cuSeqlensTensor = context->GetOptionalInputTensor(INPUT_SEQLENS_IDX);
|
||||
if (cuSeqlensTensor == nullptr) {
|
||||
isVariedLen = false;
|
||||
shapeBatch = kStorageShape.GetDim(DIM_BATCH);
|
||||
tokenBatch = 1;
|
||||
batch = shapeBatch;
|
||||
} else {
|
||||
isVariedLen = true;
|
||||
shapeBatch = 1;
|
||||
tokenBatch = cuSeqlensTensor->GetStorageShape().GetDim(DIM_BATCH) - 1;
|
||||
batch = tokenBatch;
|
||||
}
|
||||
|
||||
auto initialStateTensor = context->GetOptionalInputTensor(INPUT_INITIAL_STATE_IDX);
|
||||
bool useInitialState = initialStateTensor != nullptr;
|
||||
int64_t stateDataType = 2;
|
||||
if (useInitialState) {
|
||||
auto stateDType = initialStateTensor->GetDataType();
|
||||
if (stateDType == ge::DT_BF16) {
|
||||
stateDataType = 1;
|
||||
} else if (stateDType == ge::DT_FLOAT16) {
|
||||
stateDataType = 0;
|
||||
}
|
||||
}
|
||||
|
||||
auto gDType = context->GetOptionalInputTensor(INPUT_G_IDX)->GetDataType();
|
||||
int64_t gDataType = 2;
|
||||
if (gDType == ge::DT_BF16) {
|
||||
gDataType = 1;
|
||||
} else if (gDType == ge::DT_FLOAT16) {
|
||||
gDataType = 0;
|
||||
}
|
||||
|
||||
auto attrPtr = context->GetAttrs();
|
||||
bool storeFinalState = *(attrPtr->GetAttrPointer<bool>(ATTR_STORE_FINAL_STATE_IDX));
|
||||
int64_t chunkSize = *(attrPtr->GetAttrPointer<int64_t>(ATTR_CHUNK_SIZE_IDX));
|
||||
int64_t initalStateStride0 = *(attrPtr->GetAttrPointer<int64_t>(ATTR_INITIAL_STATE_STRIDE_IDX));
|
||||
|
||||
auto dtype = context->GetInputTensor(0)->GetDataType();
|
||||
uint64_t dataType = dtype == ge::DT_BF16 ? 1 : 0;
|
||||
|
||||
const auto ascendcPlatform = platform_ascendc::PlatformAscendC(context->GetPlatformInfo());
|
||||
uint32_t aicCoreNum = ascendcPlatform.GetCoreNumAic();
|
||||
context->SetBlockDim(aicCoreNum);
|
||||
|
||||
constexpr size_t WORKSPACE_RSV_BYTE = 16 * 1024 * 1024;
|
||||
constexpr size_t GM_ALIGN = 512;
|
||||
constexpr int64_t PING_PONG_STAGES = 2;
|
||||
|
||||
size_t workspaceOffset = ascendcPlatform.GetLibApiWorkSpaceSize();
|
||||
workspaceOffset += WORKSPACE_RSV_BYTE;
|
||||
|
||||
tiling.set_vWorkspaceOffset(workspaceOffset);
|
||||
workspaceOffset += (aicCoreNum * chunkSize * vHeadDim * sizeof(float) * PING_PONG_STAGES + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
|
||||
|
||||
tiling.set_vUpdateWorkspaceOffset(workspaceOffset);
|
||||
workspaceOffset += (aicCoreNum * chunkSize * vHeadDim * sizeof(float) * PING_PONG_STAGES + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
|
||||
|
||||
tiling.set_hWorkspaceOffset(workspaceOffset);
|
||||
workspaceOffset += (aicCoreNum * kHeadDim * vHeadDim * sizeof(float) * PING_PONG_STAGES + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
|
||||
|
||||
tiling.set_numSeqWorkspaceOffset(workspaceOffset);
|
||||
workspaceOffset += ((tokenBatch + 1) * sizeof(int64_t) + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
|
||||
|
||||
tiling.set_numChunksWorkspaceOffset(workspaceOffset);
|
||||
workspaceOffset += ((tokenBatch + 1) * sizeof(int64_t) + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
|
||||
|
||||
workspaceOffset += WORKSPACE_RSV_BYTE;
|
||||
size_t *currentWorkspace = context->GetWorkspaceSizes(1);
|
||||
currentWorkspace[0] = (workspaceOffset - 0);
|
||||
|
||||
tiling.set_batch(batch);
|
||||
tiling.set_seqlen(seqlen);
|
||||
tiling.set_kNumHead(kNumHead);
|
||||
tiling.set_vNumHead(vNumHead);
|
||||
tiling.set_kHeadDim(kHeadDim);
|
||||
tiling.set_vHeadDim(vHeadDim);
|
||||
tiling.set_chunkSize(chunkSize);
|
||||
tiling.set_initalStateStride0(initalStateStride0);
|
||||
tiling.set_useInitialState(useInitialState);
|
||||
tiling.set_storeFinalState(storeFinalState);
|
||||
tiling.set_dataType(dataType);
|
||||
tiling.set_stateDataType(stateDataType);
|
||||
tiling.set_gDataType(gDataType);
|
||||
tiling.set_isVariedLen(isVariedLen);
|
||||
tiling.set_shapeBatch(shapeBatch);
|
||||
tiling.set_tokenBatch(tokenBatch);
|
||||
|
||||
tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity());
|
||||
context->GetRawTilingData()->SetDataSize(tiling.GetDataSize());
|
||||
|
||||
ChunkGatedDeltaRuleFwdHTilingDataPrint(context, tiling);
|
||||
OP_LOGD(context->GetNodeName(), "Tiling4ChunkGatedDeltaRuleFwdH end.");
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
ge::graphStatus TilingPrepareForChunkGatedDeltaRuleFwdH(gert::TilingParseContext *context)
|
||||
{
|
||||
(void)context;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(ChunkGatedDeltaRuleFwdH)
|
||||
.Tiling(Tiling4ChunkGatedDeltaRuleFwdH)
|
||||
.TilingParse<ChunkGatedDeltaRuleFwdHCompileInfo>(TilingPrepareForChunkGatedDeltaRuleFwdH);
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,50 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file chunk_gated_delta_rule_fwd_h_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
#include <register/tilingdata_base.h>
|
||||
#include <tiling/tiling_api.h>
|
||||
|
||||
namespace optiling {
|
||||
|
||||
BEGIN_TILING_DATA_DEF(ChunkGatedDeltaRuleFwdHTilingData)
|
||||
TILING_DATA_FIELD_DEF(int64_t, batch);
|
||||
TILING_DATA_FIELD_DEF(int64_t, seqlen);
|
||||
TILING_DATA_FIELD_DEF(int64_t, kNumHead);
|
||||
TILING_DATA_FIELD_DEF(int64_t, vNumHead);
|
||||
TILING_DATA_FIELD_DEF(int64_t, kHeadDim);
|
||||
TILING_DATA_FIELD_DEF(int64_t, vHeadDim);
|
||||
TILING_DATA_FIELD_DEF(int64_t, chunkSize);
|
||||
TILING_DATA_FIELD_DEF(int64_t, initalStateStride0);
|
||||
TILING_DATA_FIELD_DEF(bool, useInitialState);
|
||||
TILING_DATA_FIELD_DEF(bool, storeFinalState);
|
||||
TILING_DATA_FIELD_DEF(int64_t, dataType);
|
||||
TILING_DATA_FIELD_DEF(int64_t, gDataType);
|
||||
TILING_DATA_FIELD_DEF(int64_t, stateDataType);
|
||||
TILING_DATA_FIELD_DEF(int64_t, isVariedLen);
|
||||
TILING_DATA_FIELD_DEF(int64_t, shapeBatch);
|
||||
TILING_DATA_FIELD_DEF(int64_t, tokenBatch);
|
||||
TILING_DATA_FIELD_DEF(int64_t, vWorkspaceOffset);
|
||||
TILING_DATA_FIELD_DEF(int64_t, vUpdateWorkspaceOffset);
|
||||
TILING_DATA_FIELD_DEF(int64_t, hWorkspaceOffset);
|
||||
TILING_DATA_FIELD_DEF(int64_t, numSeqWorkspaceOffset);
|
||||
TILING_DATA_FIELD_DEF(int64_t, numChunksWorkspaceOffset);
|
||||
END_TILING_DATA_DEF;
|
||||
|
||||
REGISTER_TILING_DATA_CLASS(ChunkGatedDeltaRuleFwdH, ChunkGatedDeltaRuleFwdHTilingData)
|
||||
|
||||
struct ChunkGatedDeltaRuleFwdHCompileInfo {};
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,230 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Tianjin University, 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.
|
||||
*/
|
||||
#include "aclnn_chunk_gated_delta_rule_fwd_h.h"
|
||||
#include "chunk_gated_delta_rule_fwd_h.h"
|
||||
#include <dlfcn.h>
|
||||
#include <new>
|
||||
#include <iostream>
|
||||
|
||||
#include "aclnn_kernels/transdata.h"
|
||||
#include "aclnn_kernels/contiguous.h"
|
||||
#include "acl/acl.h"
|
||||
#include "aclnn/aclnn_base.h"
|
||||
#include "aclnn_kernels/common/op_error_check.h"
|
||||
#include "opdev/common_types.h"
|
||||
#include "opdev/data_type_utils.h"
|
||||
#include "opdev/format_utils.h"
|
||||
#include "opdev/op_dfx.h"
|
||||
#include "opdev/op_executor.h"
|
||||
#include "opdev/op_log.h"
|
||||
#include "opdev/platform.h"
|
||||
#include "opdev/shape_utils.h"
|
||||
#include "opdev/tensor_view_utils.h"
|
||||
#include "opdev/make_op_executor.h"
|
||||
|
||||
|
||||
using namespace op;
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
struct ChunkGatedDeltaRuleFwdHParams {
|
||||
const aclTensor *k = nullptr;
|
||||
const aclTensor *w = nullptr;
|
||||
const aclTensor *u = nullptr;
|
||||
const aclTensor *gOptional = nullptr;
|
||||
const aclTensor *gkOptional = nullptr;
|
||||
const aclTensor *initalStateOptional = nullptr;
|
||||
bool outputFinalState = false;
|
||||
int64_t chunkSize = 64;
|
||||
bool saveNewValue = true;
|
||||
const aclIntArray *cuSeqlensOptional = nullptr;
|
||||
const aclIntArray *chunkIndicesOptional = nullptr;
|
||||
bool useExp2 = false;
|
||||
bool transposeStateLayout = false;
|
||||
const aclTensor *hOut = nullptr;
|
||||
const aclTensor *vNewOut = nullptr;
|
||||
const aclTensor *finalStateOut = nullptr;
|
||||
};
|
||||
|
||||
static aclnnStatus CheckNotNull(ChunkGatedDeltaRuleFwdHParams params)
|
||||
{
|
||||
CHECK_COND(params.k != nullptr, ACLNN_ERR_PARAM_NULLPTR, "k must not be nullptr.");
|
||||
CHECK_COND(params.w != nullptr, ACLNN_ERR_PARAM_NULLPTR, "w must not be nullptr.");
|
||||
CHECK_COND(params.u != nullptr, ACLNN_ERR_PARAM_NULLPTR, "u must not be nullptr.");
|
||||
|
||||
CHECK_COND(params.hOut != nullptr, ACLNN_ERR_PARAM_NULLPTR, "hOut must not be nullptr.");
|
||||
CHECK_COND(params.vNewOut != nullptr, ACLNN_ERR_PARAM_NULLPTR, "vNewOut must not be nullptr.");
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus CheckFormat(ChunkGatedDeltaRuleFwdHParams params)
|
||||
{
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus CheckShape(ChunkGatedDeltaRuleFwdHParams params)
|
||||
{
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus CheckDtype(ChunkGatedDeltaRuleFwdHParams params)
|
||||
{
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus DataContiguous(const aclTensor *&tensor, aclOpExecutor *executor)
|
||||
{
|
||||
tensor = l0op::Contiguous(tensor, executor);
|
||||
CHECK_RET(tensor != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus ParamsDataContiguous(ChunkGatedDeltaRuleFwdHParams ¶ms, aclOpExecutor *executorPtr)
|
||||
{
|
||||
CHECK_COND(DataContiguous(params.k, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
|
||||
"Contiguous k failed.");
|
||||
CHECK_COND(DataContiguous(params.w, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
|
||||
"Contiguous w failed.");
|
||||
CHECK_COND(DataContiguous(params.u, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
|
||||
"Contiguous u failed.");
|
||||
CHECK_COND(DataContiguous(params.gOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
|
||||
"Contiguous gOptional failed.");
|
||||
if (params.initalStateOptional != nullptr) {
|
||||
CHECK_COND(DataContiguous(params.initalStateOptional, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
|
||||
"Contiguous initalStateOptional failed.");
|
||||
}
|
||||
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus CheckGOptionalNonNull(const ChunkGatedDeltaRuleFwdHParams ¶ms)
|
||||
{
|
||||
CHECK_COND(params.gOptional != nullptr, ACLNN_ERR_PARAM_INVALID,
|
||||
"g is an optional-parameter slot in the API but only a non-null aclTensor is supported; nullptr is not allowed until g=None is implemented.");
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus CheckReservedOptions(const ChunkGatedDeltaRuleFwdHParams ¶ms)
|
||||
{
|
||||
CHECK_COND(params.gkOptional == nullptr, ACLNN_ERR_PARAM_INVALID,
|
||||
"gk is reserved for ChunkGatedDeltaRuleFwdH and must be nullptr.");
|
||||
CHECK_COND(params.saveNewValue, ACLNN_ERR_PARAM_INVALID,
|
||||
"save_new_value is reserved and only true is supported.");
|
||||
CHECK_COND(!params.useExp2, ACLNN_ERR_PARAM_INVALID,
|
||||
"use_exp2 is reserved and only false is supported.");
|
||||
CHECK_COND(!params.transposeStateLayout, ACLNN_ERR_PARAM_INVALID,
|
||||
"transpose_state_layout is reserved and only false is supported.");
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
static aclnnStatus CheckParams(ChunkGatedDeltaRuleFwdHParams params)
|
||||
{
|
||||
CHECK_RET(CheckNotNull(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
|
||||
CHECK_RET(CheckGOptionalNonNull(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
|
||||
CHECK_RET(CheckReservedOptions(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
|
||||
CHECK_RET(CheckFormat(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
|
||||
CHECK_RET(CheckShape(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
|
||||
CHECK_RET(CheckDtype(params) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
aclnnStatus aclnnChunkGatedDeltaRuleFwdHGetWorkspaceSize(
|
||||
const aclTensor *k,
|
||||
const aclTensor *w,
|
||||
const aclTensor *u,
|
||||
const aclTensor *gOptional,
|
||||
const aclTensor *gkOptional,
|
||||
const aclTensor *initalStateOptional,
|
||||
bool outputFinalState,
|
||||
int64_t chunkSize,
|
||||
bool saveNewValue,
|
||||
const aclIntArray *cuSeqlensOptional,
|
||||
const aclIntArray *chunkIndicesOptional,
|
||||
bool useExp2,
|
||||
bool transposeStateLayout,
|
||||
const aclTensor *hOut,
|
||||
const aclTensor *vNewOut,
|
||||
const aclTensor *finalStateOut,
|
||||
uint64_t *workspaceSize,
|
||||
aclOpExecutor **executor)
|
||||
{
|
||||
ChunkGatedDeltaRuleFwdHParams params{k,
|
||||
w,
|
||||
u,
|
||||
gOptional,
|
||||
gkOptional,
|
||||
initalStateOptional,
|
||||
outputFinalState,
|
||||
chunkSize,
|
||||
saveNewValue,
|
||||
cuSeqlensOptional,
|
||||
chunkIndicesOptional,
|
||||
useExp2,
|
||||
transposeStateLayout,
|
||||
hOut,
|
||||
vNewOut,
|
||||
finalStateOut};
|
||||
// Standard syntax, Check parameters.
|
||||
L2_DFX_PHASE_1(aclnnChunkGatedDeltaRuleFwdH,
|
||||
DFX_IN(k, w, u, gOptional, gkOptional, initalStateOptional, cuSeqlensOptional, chunkIndicesOptional),
|
||||
DFX_OUT(hOut, vNewOut, finalStateOut));
|
||||
auto uniqueExecutor = CREATE_EXECUTOR();
|
||||
CHECK_RET(uniqueExecutor.get() != nullptr, ACLNN_ERR_INNER_CREATE_EXECUTOR);
|
||||
auto executorPtr = uniqueExecutor.get();
|
||||
auto ret = CheckParams(params);
|
||||
CHECK_RET(ret == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID);
|
||||
CHECK_COND(ParamsDataContiguous(params, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
|
||||
"ParamsDataContiguous failed.");
|
||||
|
||||
// aclGetViewStrides obtains the strides and the number of strides corresponding to aclTensor
|
||||
int64_t *initialStateStridesValuePtr = nullptr;
|
||||
int64_t initialStateStridesValue = 0;
|
||||
uint64_t initialStateStridesNum = 0;
|
||||
|
||||
if (initalStateOptional != nullptr) {
|
||||
ret = aclGetViewStrides(initalStateOptional, &initialStateStridesValuePtr, &initialStateStridesNum);
|
||||
CHECK_RET(ret == ACLNN_SUCCESS, ret);
|
||||
initialStateStridesValue = initialStateStridesValuePtr[initialStateStridesNum - 2];
|
||||
}
|
||||
|
||||
auto result = l0op::ChunkGatedDeltaRuleFwdH(params.k, params.w, params.u, params.gOptional, params.initalStateOptional, params.cuSeqlensOptional, params.chunkIndicesOptional, params.outputFinalState, params.chunkSize, initialStateStridesValue, params.hOut, params.vNewOut, params.finalStateOut, executorPtr);
|
||||
CHECK_RET(result[0] != nullptr, ACLNN_ERR_PARAM_NULLPTR);
|
||||
|
||||
// If the output tensor is non-contiguous, convert the calculated contiguous tensor to non-contiguous.
|
||||
auto viewCopyResult0 = l0op::ViewCopy(result[0], params.hOut, executorPtr);
|
||||
CHECK_RET(viewCopyResult0 != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
auto viewCopyResult1 = l0op::ViewCopy(result[1], params.vNewOut, executorPtr);
|
||||
CHECK_RET(viewCopyResult1 != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
if (outputFinalState && params.finalStateOut != nullptr) {
|
||||
auto viewCopyResult2 = l0op::ViewCopy(result[2], params.finalStateOut, executorPtr);
|
||||
CHECK_RET(viewCopyResult2 != nullptr, ACLNN_ERR_INNER_NULLPTR);
|
||||
}
|
||||
|
||||
// Standard syntax, get the size of workspace needed during computation.
|
||||
*workspaceSize = uniqueExecutor->GetWorkspaceSize();
|
||||
uniqueExecutor.ReleaseTo(executor);
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
|
||||
aclnnStatus aclnnChunkGatedDeltaRuleFwdH(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)
|
||||
{
|
||||
L2_DFX_PHASE_2(aclnnChunkGatedDeltaRuleFwdH);
|
||||
CHECK_COND(CommonOpExecutorRun(workspace, workspaceSize, executor, stream) == ACLNN_SUCCESS, ACLNN_ERR_INNER,
|
||||
"This is an error in ChunkGatedDeltaRuleFwdH launch aicore.");
|
||||
return ACLNN_SUCCESS;
|
||||
}
|
||||
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
@@ -0,0 +1,77 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Tianjin University, 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.
|
||||
*/
|
||||
#ifndef OP_API_INC_ACLNN_CHUNK_GATED_DELTA_RULE_FWD_H_H
|
||||
#define OP_API_INC_ACLNN_CHUNK_GATED_DELTA_RULE_FWD_H_H
|
||||
#include "aclnn/aclnn_base.h"
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
/* function: aclnnChunkGatedDeltaRuleFwdHGetWorkspaceSize
|
||||
* parameters (order aligned with chunk_gated_delta_rule_fwd_h Python API):
|
||||
* k : required
|
||||
* w : required
|
||||
* u : required
|
||||
* gOptional : optional, only non-null aclTensor is supported
|
||||
* gkOptional : optional, reserved (must be nullptr)
|
||||
* initalStateOptional : optional
|
||||
* outputFinalState : required
|
||||
* chunkSize : required
|
||||
* saveNewValue : reserved (must be true)
|
||||
* cuSeqlensOptional : optional
|
||||
* chunkIndicesOptional : optional
|
||||
* useExp2 : reserved (must be false)
|
||||
* transposeStateLayout : reserved (must be false)
|
||||
* hOut : required
|
||||
* vNewOut : required
|
||||
* finalStateOut : optional
|
||||
* workspaceSize : size of workspace(output).
|
||||
* executor : executor context(output).
|
||||
*/
|
||||
__attribute__((visibility("default")))
|
||||
aclnnStatus aclnnChunkGatedDeltaRuleFwdHGetWorkspaceSize(
|
||||
const aclTensor *k,
|
||||
const aclTensor *w,
|
||||
const aclTensor *u,
|
||||
const aclTensor *gOptional,
|
||||
const aclTensor *gkOptional,
|
||||
const aclTensor *initalStateOptional,
|
||||
bool outputFinalState,
|
||||
int64_t chunkSize,
|
||||
bool saveNewValue,
|
||||
const aclIntArray *cuSeqlensOptional,
|
||||
const aclIntArray *chunkIndicesOptional,
|
||||
bool useExp2,
|
||||
bool transposeStateLayout,
|
||||
const aclTensor *hOut,
|
||||
const aclTensor *vNewOut,
|
||||
const aclTensor *finalStateOut,
|
||||
uint64_t *workspaceSize,
|
||||
aclOpExecutor **executor);
|
||||
|
||||
/* function: aclnnChunkGatedDeltaRuleFwdH
|
||||
* parameters :
|
||||
* workspace : workspace memory addr(input).
|
||||
* workspaceSize : size of workspace(input).
|
||||
* executor : executor context(input).
|
||||
* stream : acl stream.
|
||||
*/
|
||||
__attribute__((visibility("default")))
|
||||
aclnnStatus aclnnChunkGatedDeltaRuleFwdH(
|
||||
void *workspace,
|
||||
uint64_t workspaceSize,
|
||||
aclOpExecutor *executor,
|
||||
aclrtStream stream);
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,70 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Tianjin University, 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.
|
||||
*/
|
||||
|
||||
#include "opdev/op_log.h"
|
||||
#include "opdev/op_dfx.h"
|
||||
#include "opdev/make_op_executor.h"
|
||||
#include "chunk_gated_delta_rule_fwd_h.h"
|
||||
|
||||
using namespace op;
|
||||
|
||||
namespace l0op {
|
||||
OP_TYPE_REGISTER(ChunkGatedDeltaRuleFwdH);
|
||||
|
||||
const std::array<const aclTensor *, 3> ChunkGatedDeltaRuleFwdH(
|
||||
const aclTensor *k,
|
||||
const aclTensor *w,
|
||||
const aclTensor *u,
|
||||
const aclTensor *g,
|
||||
const aclTensor *initalStateOptional,
|
||||
const aclIntArray *cuSeqlensOptional,
|
||||
const aclIntArray *chunkIndicesOptional,
|
||||
bool outputFinalState,
|
||||
int64_t chunkSize,
|
||||
int64_t initialStateStridesValue,
|
||||
const aclTensor *hOut,
|
||||
const aclTensor *vNewOut,
|
||||
const aclTensor *finalStateOut,
|
||||
aclOpExecutor *executor)
|
||||
{
|
||||
L0_DFX(ChunkGatedDeltaRuleFwdH, k, w, u, g, initalStateOptional, cuSeqlensOptional, chunkIndicesOptional, outputFinalState, chunkSize, initialStateStridesValue, hOut, vNewOut, finalStateOut);
|
||||
|
||||
const aclTensor *actualCuSeqlens = nullptr;
|
||||
if (cuSeqlensOptional) {
|
||||
actualCuSeqlens = executor->ConvertToTensor(cuSeqlensOptional, DataType::DT_INT64);
|
||||
const_cast<aclTensor *>(actualCuSeqlens)->SetStorageFormat(Format::FORMAT_ND);
|
||||
const_cast<aclTensor *>(actualCuSeqlens)->SetViewFormat(Format::FORMAT_ND);
|
||||
const_cast<aclTensor *>(actualCuSeqlens)->SetOriginalFormat(Format::FORMAT_ND);
|
||||
} else {
|
||||
actualCuSeqlens = nullptr;
|
||||
}
|
||||
|
||||
const aclTensor *actualChunkIndices = nullptr;
|
||||
if (chunkIndicesOptional) {
|
||||
actualChunkIndices = executor->ConvertToTensor(chunkIndicesOptional, DataType::DT_INT64);
|
||||
const_cast<aclTensor *>(actualChunkIndices)->SetStorageFormat(Format::FORMAT_ND);
|
||||
const_cast<aclTensor *>(actualChunkIndices)->SetViewFormat(Format::FORMAT_ND);
|
||||
const_cast<aclTensor *>(actualChunkIndices)->SetOriginalFormat(Format::FORMAT_ND);
|
||||
} else {
|
||||
actualChunkIndices = nullptr;
|
||||
}
|
||||
|
||||
auto ret = ADD_TO_LAUNCHER_LIST_AICORE(ChunkGatedDeltaRuleFwdH,
|
||||
OP_INPUT(k, w, u, g, initalStateOptional, actualCuSeqlens, actualChunkIndices),
|
||||
OP_OUTPUT(hOut, vNewOut, finalStateOut),
|
||||
OP_ATTR(outputFinalState, chunkSize, initialStateStridesValue));
|
||||
if (ret != ACLNN_SUCCESS) {
|
||||
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ADD_TO_LAUNCHER_LIST_AICORE failed.");
|
||||
return {nullptr, nullptr, nullptr};
|
||||
}
|
||||
return {hOut, vNewOut, finalStateOut};
|
||||
}
|
||||
|
||||
} // namespace l0op
|
||||
@@ -0,0 +1,33 @@
|
||||
/**
|
||||
* Copyright (c) 2025 Tianjin University, 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.
|
||||
*/
|
||||
#ifndef OP_API_INC_LEVEL0_OP_CHUNK_GATED_DELTA_RULE_FWD_H_H
|
||||
#define OP_API_INC_LEVEL0_OP_CHUNK_GATED_DELTA_RULE_FWD_H_H
|
||||
|
||||
#include "opdev/op_executor.h"
|
||||
|
||||
namespace l0op {
|
||||
const std::array<const aclTensor *, 3> ChunkGatedDeltaRuleFwdH(
|
||||
const aclTensor *k,
|
||||
const aclTensor *w,
|
||||
const aclTensor *u,
|
||||
const aclTensor *g,
|
||||
const aclTensor *initalStateOptional,
|
||||
const aclIntArray *cuSeqlensOptional,
|
||||
const aclIntArray *chunkIndicesOptional,
|
||||
bool outputFinalState,
|
||||
int64_t chunkSize,
|
||||
int64_t initialStateStridesValue,
|
||||
const aclTensor *hOut,
|
||||
const aclTensor *vNewOut,
|
||||
const aclTensor *finalStateOut,
|
||||
aclOpExecutor *executor);
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,39 @@
|
||||
#ifndef COMPAT_310P_H
|
||||
#define COMPAT_310P_H
|
||||
|
||||
#ifndef __CCE_KT_TEST__
|
||||
#include "kernel_operator.h"
|
||||
#endif
|
||||
|
||||
// Dummy bfloat16_t only needed on 310P (dav_m200) where the compiler
|
||||
// doesn't provide a native bf16 type. On 910B/910C the compiler's
|
||||
// __clang_cce_types.h already typedefs bfloat16_t from __bf16.
|
||||
#if defined(__CCE_AICORE__) && (__CCE_AICORE__ == 200) && !defined(__bfloat16_t_defined)
|
||||
#define __bfloat16_t_defined
|
||||
#define __COMPAT_310P_ACTIVE__
|
||||
struct bfloat16_t {
|
||||
uint16_t val;
|
||||
bfloat16_t() = default;
|
||||
bfloat16_t(float v) : val(0) { (void)v; }
|
||||
operator float() const { return 0.f; }
|
||||
};
|
||||
#endif
|
||||
|
||||
// 310P has no fixpipe unit; post-matmul stores go through MTE3
|
||||
#ifndef PIPE_FIX
|
||||
#define PIPE_FIX PIPE_MTE3
|
||||
#endif
|
||||
|
||||
// 310P renames LoadDataWithSparse → LoadDataWithSparseCal
|
||||
#ifdef __COMPAT_310P_ACTIVE__
|
||||
#define LoadDataWithSparse LoadDataWithSparseCal
|
||||
#endif
|
||||
|
||||
// 310P has no AscendC::ToFloat — dummy bfloat16_t already has operator float()
|
||||
#ifdef __COMPAT_310P_ACTIVE__
|
||||
namespace AscendC {
|
||||
inline float ToFloat(bfloat16_t v) { return (float)v; }
|
||||
}
|
||||
#endif
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,179 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "../gdn_fwd_h_epilogue_policies.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/epilogue/tile/tile_copy.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <
|
||||
class HOutputType_,
|
||||
class GInputType_,
|
||||
class HInputType_,
|
||||
class HUpdateInputType_,
|
||||
class FinalStateType_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasGDNFwdHUpdate,
|
||||
HOutputType_,
|
||||
GInputType_,
|
||||
HInputType_,
|
||||
HUpdateInputType_,
|
||||
FinalStateType_
|
||||
> {
|
||||
public:
|
||||
// Type aliases
|
||||
using DispatchPolicy = EpilogueAtlasGDNFwdHUpdate;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using HElementOutput = typename HOutputType_::Element;
|
||||
using GElementInput = typename GInputType_::Element;
|
||||
using HElementInput = typename HInputType_::Element;
|
||||
using HUpdateElementInput = typename HUpdateInputType_::Element;
|
||||
using FinalStateElement = typename FinalStateType_::Element;
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
|
||||
// Bumped layout to fit kHeadDim up to 256 with subBlockNum=2 (per-subblock M up to 128).
|
||||
// Required: calc (fp32) up to 128*128*4=64KB; h (fp16) up to 128*128*2=32KB;
|
||||
// hUpdate/hOutput/finalOutput at the same offset, max needed 64KB; glast small.
|
||||
constexpr uint32_t CALC_BUF_OFFSET = 0;
|
||||
constexpr uint32_t PING_BUF_0_OFFSET = 64 * 1024;
|
||||
constexpr uint32_t PING_BUF_1_OFFSET = 96 * 1024;
|
||||
constexpr uint32_t PING_BUF_2_OFFSET = 112 * 1024;
|
||||
constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024;
|
||||
|
||||
|
||||
calcUbTensor = resource.ubBuf.template GetBufferByByte<float>(CALC_BUF_OFFSET);
|
||||
|
||||
hUpdateUbTensor = resource.ubBuf.template GetBufferByByte<float>(PING_BUF_1_OFFSET);
|
||||
hUbTensor = resource.ubBuf.template GetBufferByByte<HElementInput>(PING_BUF_0_OFFSET);
|
||||
|
||||
hOutputUbTensor = resource.ubBuf.template GetBufferByByte<HElementOutput>(PING_BUF_1_OFFSET);
|
||||
finalOutputUbTensor = resource.ubBuf.template GetBufferByByte<FinalStateElement>(PING_BUF_1_OFFSET);
|
||||
|
||||
glastUbTensor = resource.ubBuf.template GetBufferByByte<float>(PING_G_BUF_OFFSET);
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(
|
||||
AscendC::GlobalTensor<HElementOutput> hOutput,
|
||||
AscendC::GlobalTensor<FinalStateElement> finalState,
|
||||
AscendC::GlobalTensor<GElementInput> gInput,
|
||||
AscendC::GlobalTensor<HElementInput> hInput,
|
||||
AscendC::GlobalTensor<float> hUpdateInput,
|
||||
uint32_t chunkSize,
|
||||
uint32_t kHeadDim,
|
||||
uint32_t vHeadDim,
|
||||
Arch::CrossCoreFlag cube2Done,
|
||||
bool isFinalState
|
||||
)
|
||||
{
|
||||
uint32_t mActual = kHeadDim;
|
||||
uint32_t nActual = vHeadDim;
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
uint32_t mActualPerSubBlock = CeilDiv(mActual, subBlockNum);
|
||||
uint32_t mActualThisSubBlock = (subBlockIdx == 0) ? mActualPerSubBlock : (mActual - mActualPerSubBlock);
|
||||
uint32_t mOffset = subBlockIdx * mActualPerSubBlock;
|
||||
uint32_t nOffset = 0;
|
||||
int64_t offsetH = mOffset * nActual + nOffset;
|
||||
|
||||
AscendC::ResetMask();
|
||||
|
||||
AscendC::GlobalTensor<HElementOutput> hOutputThisSubBlock = hOutput[offsetH];
|
||||
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
|
||||
AscendC::GlobalTensor<HElementInput> hInputThisSubBlock = hInput[offsetH];
|
||||
AscendC::GlobalTensor<float> hUpdateInputThisSubBlock = hUpdateInput[offsetH];
|
||||
AscendC::GlobalTensor<FinalStateElement> finalStateThisSubBlock = finalState[offsetH];
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0);
|
||||
AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0);
|
||||
AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
GElementInput gLastVal = gInputThisSubBlock.GetValue(chunkSize-1);
|
||||
float gLastFloat = 0.0f;
|
||||
if constexpr(std::is_same<GElementInput, float>::value) {
|
||||
gLastFloat = gLastVal;
|
||||
} else if constexpr(std::is_same<GElementInput, half>::value) {
|
||||
gLastFloat = (float)gLastVal;
|
||||
} else if constexpr(std::is_same<GElementInput, bfloat16_t>::value) {
|
||||
gLastFloat = AscendC::ToFloat(gLastVal);
|
||||
}
|
||||
glastUbTensor.SetValue(0, gLastFloat);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::Exp(glastUbTensor, glastUbTensor, 1);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_S>(EVENT_ID0);
|
||||
float muls = glastUbTensor.GetValue(0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::Muls(calcUbTensor, calcUbTensor, muls, mActualThisSubBlock * nActual);
|
||||
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1);
|
||||
AscendC::DataCopy(hUpdateUbTensor, hUpdateInputThisSubBlock, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1);
|
||||
AscendC::Add<float>(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, mActualThisSubBlock * nActual);
|
||||
|
||||
if (isFinalState) {
|
||||
if constexpr(!std::is_same<FinalStateElement, float>::value) {
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::DataCopy(finalStateThisSubBlock, finalOutputUbTensor, mActualThisSubBlock * nActual);
|
||||
} else {
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::DataCopy(finalStateThisSubBlock, hUpdateUbTensor, mActualThisSubBlock * nActual);
|
||||
}
|
||||
} else {
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::Cast(hOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::DataCopy(hOutputThisSubBlock, hOutputUbTensor, mActualThisSubBlock * nActual);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
AscendC::LocalTensor<float> calcUbTensor;
|
||||
|
||||
AscendC::LocalTensor<HElementInput> hUbTensor;
|
||||
AscendC::LocalTensor<float> hUpdateUbTensor;
|
||||
|
||||
AscendC::LocalTensor<HElementOutput> hOutputUbTensor;
|
||||
AscendC::LocalTensor<FinalStateElement> finalOutputUbTensor;
|
||||
|
||||
AscendC::LocalTensor<float> glastUbTensor;
|
||||
|
||||
};
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,258 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_VNEW_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_VNEW_HPP
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "../gdn_fwd_h_epilogue_policies.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/epilogue/tile/tile_copy.hpp"
|
||||
|
||||
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <
|
||||
class VOutputType_,
|
||||
class GInputType_,
|
||||
class UInputType_,
|
||||
class WSInputType_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasGDNFwdHVnew,
|
||||
VOutputType_,
|
||||
GInputType_,
|
||||
UInputType_,
|
||||
WSInputType_
|
||||
> {
|
||||
public:
|
||||
// Type aliases
|
||||
using DispatchPolicy = EpilogueAtlasGDNFwdHVnew;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using VElementOutput = typename VOutputType_::Element;
|
||||
using GElementInput = typename GInputType_::Element;
|
||||
using UElementInput = typename UInputType_::Element;
|
||||
using WSElementInput = typename WSInputType_::Element;
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
|
||||
constexpr uint32_t CALC_BUF_OFFSET = 0;
|
||||
constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024;
|
||||
constexpr uint32_t PING_BUF_1_OFFSET = 64 * 1024;
|
||||
constexpr uint32_t PONG_BUF_0_OFFSET = 96 * 1024;
|
||||
constexpr uint32_t PONG_BUF_1_OFFSET = 128 * 1024;
|
||||
constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024;
|
||||
constexpr uint32_t PONG_G_BUF_OFFSET = 161 * 1024;
|
||||
constexpr uint32_t PING_G_SUB_BUF_OFFSET = 162 * 1024;
|
||||
constexpr uint32_t PONG_G_SUB_BUF_OFFSET = 163 * 1024;
|
||||
constexpr uint32_t PING_G_INPUT_BUF_OFFSET = 164 * 1024;
|
||||
constexpr uint32_t PONG_G_INPUT_BUF_OFFSET = 165 * 1024;
|
||||
constexpr uint32_t SHARE_BUF_OFFSET = 166 * 1024;
|
||||
|
||||
calcUbTensor = resource.ubBuf.template GetBufferByByte<float>(CALC_BUF_OFFSET);
|
||||
|
||||
uUbTensor_ping = resource.ubBuf.template GetBufferByByte<UElementInput>(PING_BUF_1_OFFSET);
|
||||
uUbFloatTensor_ping = resource.ubBuf.template GetBufferByByte<float>(PING_BUF_0_OFFSET);
|
||||
wsUbTensor_ping = resource.ubBuf.template GetBufferByByte<float>(PING_BUF_1_OFFSET);
|
||||
gUbTensor_ping = resource.ubBuf.template GetBufferByByte<float>(PING_G_BUF_OFFSET);
|
||||
gLastUbTensor_ping = resource.ubBuf.template GetBufferByByte<float>(PING_G_SUB_BUF_OFFSET);
|
||||
gInputUbTensor_ping = resource.ubBuf.template GetBufferByByte<GElementInput>(PING_G_INPUT_BUF_OFFSET);
|
||||
vNewOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte<VElementOutput>(PING_BUF_1_OFFSET);
|
||||
vNewDecayUbTensor_ping = resource.ubBuf.template GetBufferByByte<VElementOutput>(PING_BUF_0_OFFSET);
|
||||
|
||||
uUbTensor_pong = resource.ubBuf.template GetBufferByByte<UElementInput>(PONG_BUF_1_OFFSET);
|
||||
uUbFloatTensor_pong = resource.ubBuf.template GetBufferByByte<float>(PONG_BUF_0_OFFSET);
|
||||
wsUbTensor_pong = resource.ubBuf.template GetBufferByByte<float>(PONG_BUF_1_OFFSET);
|
||||
gUbTensor_pong = resource.ubBuf.template GetBufferByByte<float>(PONG_G_BUF_OFFSET);
|
||||
gLastUbTensor_pong = resource.ubBuf.template GetBufferByByte<float>(PONG_G_SUB_BUF_OFFSET);
|
||||
gInputUbTensor_pong = resource.ubBuf.template GetBufferByByte<GElementInput>(PONG_G_INPUT_BUF_OFFSET);
|
||||
vNewOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte<VElementOutput>(PONG_BUF_1_OFFSET);
|
||||
vNewDecayUbTensor_pong = resource.ubBuf.template GetBufferByByte<VElementOutput>(PONG_BUF_0_OFFSET);
|
||||
|
||||
shareBuffer_ = resource.ubBuf.template GetBufferByByte<uint8_t>(SHARE_BUF_OFFSET);
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(
|
||||
AscendC::GlobalTensor<VElementOutput> vnewOutput,
|
||||
AscendC::GlobalTensor<VElementOutput> vnewdecayOutput,
|
||||
AscendC::GlobalTensor<GElementInput> gInput,
|
||||
AscendC::GlobalTensor<UElementInput> uInput,
|
||||
AscendC::GlobalTensor<float> wsInput,
|
||||
uint32_t chunkSize,
|
||||
uint32_t kHeadDim,
|
||||
uint32_t vHeadDim,
|
||||
Arch::CrossCoreFlag cube1Done
|
||||
)
|
||||
{
|
||||
uint32_t mActual = chunkSize;
|
||||
uint32_t nkActual = kHeadDim;
|
||||
uint32_t nvActual = vHeadDim;
|
||||
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
uint32_t mActualPerSubBlock = CeilDiv(mActual, subBlockNum);
|
||||
uint32_t mActualThisSubBlock = (subBlockIdx == 0) ? mActualPerSubBlock : (mActual - mActualPerSubBlock);
|
||||
uint32_t mOffset = subBlockIdx * mActualPerSubBlock;
|
||||
uint32_t nOffset = 0;
|
||||
// 当前场景内部一定连续
|
||||
// k [B, H, T, D]
|
||||
// g [B, H, T]
|
||||
// 在外部offset的基础上进一步offset
|
||||
// 当前asset kdim == vHeadDim
|
||||
int64_t offsetK = mOffset * nvActual + nOffset;
|
||||
int64_t offsetD = 0; // 因为要用最后一个数减去之前所有,所以全部读入
|
||||
|
||||
uint32_t gbrcStart, gbrcRealStart, gbrcReptime, gbrcEffStart, gbrcEffEnd;
|
||||
if(subBlockIdx==0)
|
||||
{
|
||||
gbrcStart = 0;
|
||||
gbrcRealStart = 0;
|
||||
gbrcReptime = (mActualThisSubBlock + 8 - 1) / 8;
|
||||
|
||||
}
|
||||
else
|
||||
{
|
||||
gbrcStart = mActualPerSubBlock;
|
||||
gbrcRealStart = gbrcStart & ~15;
|
||||
gbrcReptime = (mActual - gbrcRealStart + 8 - 1) / 8;
|
||||
}
|
||||
gbrcEffStart = gbrcStart-gbrcRealStart;
|
||||
gbrcEffEnd = gbrcEffStart + mActualThisSubBlock;
|
||||
|
||||
AscendC::ResetMask();
|
||||
|
||||
AscendC::GlobalTensor<VElementOutput> vnewOutputThisSubBlock = vnewOutput[offsetK];
|
||||
AscendC::GlobalTensor<VElementOutput> vnewdecayOutputThisSubBlock = vnewdecayOutput[offsetK];
|
||||
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
|
||||
AscendC::GlobalTensor<UElementInput> uInputThisSubBlock = uInput[offsetK];
|
||||
AscendC::GlobalTensor<float> wsInputThisSubBlock = wsInput[offsetK];
|
||||
|
||||
pingpongFlag = isFirst ? 0 : 4;
|
||||
AscendC::LocalTensor<UElementInput> uUbTensor = isFirst ? uUbTensor_ping : uUbTensor_pong;
|
||||
AscendC::LocalTensor<float> uUbFloatTensor = isFirst ? uUbFloatTensor_ping : uUbFloatTensor_pong;
|
||||
AscendC::LocalTensor<float> wsUbTensor = isFirst ? wsUbTensor_ping : wsUbTensor_pong;
|
||||
AscendC::LocalTensor<float> gUbTensor = isFirst ? gUbTensor_ping : gUbTensor_pong;
|
||||
AscendC::LocalTensor<float> gLastUbTensor = isFirst ? gLastUbTensor_ping : gLastUbTensor_pong;
|
||||
AscendC::LocalTensor<GElementInput> gInputUbTensor = isFirst ? gInputUbTensor_ping : gInputUbTensor_pong;
|
||||
AscendC::LocalTensor<VElementOutput> vNewOutputUbTensor = isFirst ? vNewOutputUbTensor_ping : vNewOutputUbTensor_pong;
|
||||
AscendC::LocalTensor<VElementOutput> vNewDecayUbTensor = isFirst ? vNewDecayUbTensor_ping : vNewDecayUbTensor_pong;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
|
||||
if constexpr(std::is_same<GElementInput, float>::value) {
|
||||
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
|
||||
} else {
|
||||
AscendC::DataCopy(gInputUbTensor, gInputThisSubBlock, mActual);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
|
||||
if constexpr(!std::is_same<GElementInput, float>::value) {
|
||||
AscendC::Cast(gUbTensor, gInputUbTensor, AscendC::RoundMode::CAST_NONE, mActual);
|
||||
}
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_S>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_S>(EVENT_ID2 + pingpongFlag);
|
||||
float inputVal = gUbTensor.GetValue(mActual-1);
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Duplicate<float>(gLastUbTensor, inputVal, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::Sub<float>(gUbTensor, gLastUbTensor, gUbTensor, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::Exp(gUbTensor, gUbTensor, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
uint32_t dstShape_[2] = {gbrcReptime*8, nvActual};
|
||||
uint32_t srcShape_[2] = {gbrcReptime*8, 1};
|
||||
AscendC::Broadcast<float, 2, 1>(calcUbTensor, gUbTensor[gbrcRealStart], dstShape_, srcShape_, shareBuffer_);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
|
||||
AscendC::DataCopy(uUbTensor, uInputThisSubBlock, mActualThisSubBlock * nvActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::Cast(uUbFloatTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::DataCopy(wsUbTensor, wsInputThisSubBlock, mActualThisSubBlock * nvActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
|
||||
AscendC::Sub<float>(uUbFloatTensor, uUbFloatTensor, wsUbTensor, mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Cast(vNewOutputUbTensor, uUbFloatTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::DataCopy(vnewOutputThisSubBlock, vNewOutputUbTensor, mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
|
||||
AscendC::Mul(calcUbTensor[gbrcEffStart*nvActual], uUbFloatTensor, calcUbTensor[gbrcEffStart*nvActual], mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::Cast(vNewDecayUbTensor, calcUbTensor[gbrcEffStart*nvActual], AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
AscendC::DataCopy(vnewdecayOutputThisSubBlock, vNewDecayUbTensor, mActualThisSubBlock * nvActual);
|
||||
|
||||
if (isFirst) {
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
}
|
||||
|
||||
isFirst = false;
|
||||
}
|
||||
|
||||
private:
|
||||
uint32_t pingpongFlag = 0;
|
||||
bool isFirst = true;
|
||||
|
||||
AscendC::LocalTensor<float> calcUbTensor;
|
||||
|
||||
AscendC::LocalTensor<UElementInput> uUbTensor_ping;
|
||||
AscendC::LocalTensor<float> uUbFloatTensor_ping;
|
||||
AscendC::LocalTensor<float> wsUbTensor_ping;
|
||||
AscendC::LocalTensor<float> gUbTensor_ping;
|
||||
AscendC::LocalTensor<float> gLastUbTensor_ping;
|
||||
AscendC::LocalTensor<GElementInput> gInputUbTensor_ping;
|
||||
AscendC::LocalTensor<VElementOutput> vNewOutputUbTensor_ping;
|
||||
AscendC::LocalTensor<VElementOutput> vNewDecayUbTensor_ping;
|
||||
|
||||
AscendC::LocalTensor<UElementInput> uUbTensor_pong;
|
||||
AscendC::LocalTensor<float> uUbFloatTensor_pong;
|
||||
AscendC::LocalTensor<float> wsUbTensor_pong;
|
||||
AscendC::LocalTensor<float> gUbTensor_pong;
|
||||
AscendC::LocalTensor<float> gLastUbTensor_pong;
|
||||
AscendC::LocalTensor<GElementInput> gInputUbTensor_pong;
|
||||
AscendC::LocalTensor<VElementOutput> vNewOutputUbTensor_pong;
|
||||
AscendC::LocalTensor<VElementOutput> vNewDecayUbTensor_pong;
|
||||
|
||||
AscendC::LocalTensor<uint8_t> shareBuffer_;
|
||||
|
||||
};
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,27 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP
|
||||
#define CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
|
||||
namespace Catlass::Epilogue {
|
||||
|
||||
struct EpilogueAtlasGDNFwdHVnew {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
};
|
||||
|
||||
struct EpilogueAtlasGDNFwdHUpdate {
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue
|
||||
|
||||
#endif // CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP
|
||||
@@ -0,0 +1,285 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
using namespace Catlass;
|
||||
|
||||
#ifndef CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP
|
||||
#define CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP
|
||||
|
||||
// constexpr uint32_t PING_PONG_STAGES = 1;
|
||||
constexpr uint32_t PING_PONG_STAGES = 2;
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T AlignUp(T a, T b) {
|
||||
return (b == 0) ? 0 : (a + b - 1) / b * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T Min(T a, T b) {
|
||||
return (a > b) ? b : a;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T Max(T a, T b) {
|
||||
return (a > b) ? a : b;
|
||||
}
|
||||
|
||||
namespace Catlass::Gemm::Block {
|
||||
|
||||
struct GDNFwdHOffsets {
|
||||
uint32_t hSrcOffset;
|
||||
uint32_t hDstOffset;
|
||||
uint32_t uvOffset;
|
||||
uint32_t wkOffset;
|
||||
uint32_t wOffset;
|
||||
uint32_t gOffset;
|
||||
uint32_t hWorkOffset;
|
||||
uint32_t vWorkOffset;
|
||||
uint32_t initialStateOffset;
|
||||
uint32_t finalStateOffset;
|
||||
bool isInitialState;
|
||||
bool isFinalState;
|
||||
uint32_t blockTokens;
|
||||
bool isDummyHead;
|
||||
uint32_t batchIdx;
|
||||
uint32_t headIdx;
|
||||
uint32_t chunkIdx;
|
||||
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdH {
|
||||
uint32_t batch;
|
||||
uint32_t seqlen;
|
||||
uint32_t kNumHead;
|
||||
uint32_t vNumHead;
|
||||
uint32_t kHeadDim;
|
||||
uint32_t vHeadDim;
|
||||
uint32_t chunkSize;
|
||||
uint32_t initalStateStride0;
|
||||
uint32_t vBlockSize{128};
|
||||
uint32_t isVariedLen;
|
||||
uint32_t shapeBatch;
|
||||
uint32_t tokenBatch;
|
||||
bool useInitialState;
|
||||
bool storeFinalState;
|
||||
uint32_t numSeqWorkspaceOffset;
|
||||
uint32_t numChunksWorkspaceOffset;
|
||||
|
||||
uint32_t taskIdx;
|
||||
uint32_t taskLoops;
|
||||
uint32_t cubeCoreIdx;
|
||||
uint32_t cubeCoreNum;
|
||||
uint32_t vLoops;
|
||||
uint32_t taskNum;
|
||||
uint32_t headGroups;
|
||||
uint32_t totalChunks;
|
||||
uint32_t totalTokens;
|
||||
uint32_t headInnerLoop;
|
||||
|
||||
uint32_t iterId {0};
|
||||
bool hasDummyHead;
|
||||
bool isRunning;
|
||||
bool processNewTask {true};
|
||||
bool firstLoop {true};
|
||||
bool lastLoop {false};
|
||||
GDNFwdHOffsets offsets[PING_PONG_STAGES];
|
||||
int32_t currStage{PING_PONG_STAGES - 1};
|
||||
|
||||
uint32_t vIdx;
|
||||
uint32_t batchIdx;
|
||||
uint32_t baseHeadIdx;
|
||||
uint32_t chunkIdx;
|
||||
uint32_t headInnerIdx;
|
||||
uint32_t vHeadIdx;
|
||||
uint32_t kHeadIdx;
|
||||
uint32_t shapeBatchIdx;
|
||||
uint32_t tokenBatchIdx;
|
||||
|
||||
uint32_t chunkOffset;
|
||||
uint32_t tokenOffset;
|
||||
uint32_t batchChunks;
|
||||
uint32_t batchTokens;
|
||||
|
||||
AscendC::GlobalTensor<int64_t> gmSeqlen;
|
||||
AscendC::GlobalTensor<int64_t> gmNumSeq;
|
||||
AscendC::GlobalTensor<int64_t> gmNumChunks;
|
||||
|
||||
Arch::CrossCoreFlag cube1Done{0};
|
||||
Arch::CrossCoreFlag vec1Done{1};
|
||||
Arch::CrossCoreFlag cube2Done{2};
|
||||
Arch::CrossCoreFlag vec2Done{3};
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdH() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user, uint32_t coreIdx, uint32_t coreNum) {
|
||||
__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling);
|
||||
|
||||
batch = gdnFwdHTilingData->batch;
|
||||
seqlen = gdnFwdHTilingData->seqlen;
|
||||
kNumHead = gdnFwdHTilingData->kNumHead;
|
||||
vNumHead = gdnFwdHTilingData->vNumHead;
|
||||
kHeadDim = gdnFwdHTilingData->kHeadDim;
|
||||
vHeadDim = gdnFwdHTilingData->vHeadDim;
|
||||
chunkSize = gdnFwdHTilingData->chunkSize;
|
||||
initalStateStride0 = gdnFwdHTilingData->initalStateStride0;
|
||||
isVariedLen = gdnFwdHTilingData->isVariedLen;
|
||||
shapeBatch = gdnFwdHTilingData->shapeBatch;
|
||||
tokenBatch = gdnFwdHTilingData->tokenBatch;
|
||||
useInitialState = gdnFwdHTilingData->useInitialState;
|
||||
storeFinalState = gdnFwdHTilingData->storeFinalState;
|
||||
numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset;
|
||||
numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset;
|
||||
|
||||
gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens);
|
||||
gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset));
|
||||
gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset));
|
||||
|
||||
if (isVariedLen) {
|
||||
gmNumChunks.SetValue(0, 0);
|
||||
gmNumSeq.SetValue(0, 0);
|
||||
uint32_t actualBatch = 0;
|
||||
int64_t prevSeq = 0, currSeq;
|
||||
for (uint32_t b = 1; b <= tokenBatch; b++) {
|
||||
currSeq = gmSeqlen.GetValue(b);
|
||||
int64_t batchSeqLen = currSeq - prevSeq;
|
||||
if (batchSeqLen > 0) {
|
||||
actualBatch++;
|
||||
gmNumSeq.SetValue(actualBatch, currSeq);
|
||||
int64_t batchChunk = (batchSeqLen + chunkSize - 1) / chunkSize;
|
||||
gmNumChunks.SetValue(actualBatch, gmNumChunks.GetValue(actualBatch - 1) + batchChunk);
|
||||
}
|
||||
prevSeq = currSeq;
|
||||
}
|
||||
tokenBatch = actualBatch;
|
||||
batch = actualBatch;
|
||||
totalChunks = gmNumChunks.GetValue(tokenBatch);
|
||||
totalTokens = gmNumSeq.GetValue(tokenBatch);
|
||||
} else {
|
||||
totalChunks = (seqlen + chunkSize - 1) / chunkSize;
|
||||
totalTokens = seqlen;
|
||||
}
|
||||
|
||||
cubeCoreIdx = coreIdx;
|
||||
cubeCoreNum = coreNum;
|
||||
vLoops = vHeadDim / vBlockSize;
|
||||
taskNum = vLoops * batch * vNumHead;
|
||||
headGroups = vNumHead / kNumHead;
|
||||
hasDummyHead = (taskNum % (PING_PONG_STAGES * cubeCoreNum) <= cubeCoreNum) && (taskNum % (PING_PONG_STAGES * cubeCoreNum) > 0);
|
||||
taskLoops = (taskNum + cubeCoreNum * PING_PONG_STAGES - 1) / (cubeCoreNum * PING_PONG_STAGES);
|
||||
headInnerLoop = taskNum > cubeCoreNum ? PING_PONG_STAGES : 1;
|
||||
taskIdx = cubeCoreIdx * headInnerLoop;
|
||||
isRunning = taskIdx < taskNum;
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitTask() {
|
||||
iterId++;
|
||||
currStage = (currStage + 1) % PING_PONG_STAGES;
|
||||
if (processNewTask) {
|
||||
if (taskIdx >= taskNum) {
|
||||
lastLoop = true;
|
||||
isRunning = false;
|
||||
return;
|
||||
}
|
||||
vIdx = taskIdx / (batch * vNumHead);
|
||||
batchIdx = (taskIdx - vIdx * batch * vNumHead) / vNumHead;
|
||||
baseHeadIdx = taskIdx % vNumHead;
|
||||
shapeBatchIdx = isVariedLen ? 0 : batchIdx;
|
||||
tokenBatchIdx = isVariedLen ? batchIdx : 0;
|
||||
chunkOffset = isVariedLen ? gmNumChunks.GetValue(tokenBatchIdx) : 0;
|
||||
batchChunks = isVariedLen ? (gmNumChunks.GetValue(tokenBatchIdx + 1) - chunkOffset) : totalChunks;
|
||||
tokenOffset = isVariedLen ? gmNumSeq.GetValue(tokenBatchIdx) : 0;
|
||||
batchTokens = isVariedLen ? (gmNumSeq.GetValue(tokenBatchIdx + 1) - tokenOffset) : totalTokens;
|
||||
chunkIdx = 0;
|
||||
headInnerIdx = 0;
|
||||
} else {
|
||||
chunkIdx = headInnerIdx == PING_PONG_STAGES - 1 ? chunkIdx + 1 : chunkIdx;
|
||||
headInnerIdx = (headInnerIdx + 1) % PING_PONG_STAGES;
|
||||
}
|
||||
|
||||
vHeadIdx = baseHeadIdx + headInnerIdx;
|
||||
kHeadIdx = vHeadIdx / headGroups;
|
||||
offsets[currStage].isInitialState = chunkIdx == 0;
|
||||
offsets[currStage].isFinalState = chunkIdx == (batchChunks - 1);
|
||||
offsets[currStage].initialStateOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * initalStateStride0;
|
||||
offsets[currStage].finalStateOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * vHeadDim;
|
||||
offsets[currStage].hSrcOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset + chunkIdx) * kHeadDim * vHeadDim;
|
||||
offsets[currStage].hDstOffset = offsets[currStage].hSrcOffset + kHeadDim * vHeadDim;
|
||||
offsets[currStage].uvOffset = (shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * vHeadDim;
|
||||
offsets[currStage].wkOffset = (shapeBatchIdx * kNumHead * totalTokens + kHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * kHeadDim;
|
||||
offsets[currStage].wOffset = (shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * kHeadDim;
|
||||
offsets[currStage].gOffset = shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize;
|
||||
offsets[currStage].hWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * kHeadDim * vHeadDim;
|
||||
offsets[currStage].vWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * vHeadDim;
|
||||
offsets[currStage].blockTokens = offsets[currStage].isFinalState ? (batchTokens - chunkIdx * chunkSize) : chunkSize;
|
||||
offsets[currStage].isDummyHead = headInnerLoop < PING_PONG_STAGES && headInnerIdx >= headInnerLoop;
|
||||
offsets[currStage].batchIdx = batchIdx;
|
||||
offsets[currStage].headIdx = vHeadIdx;
|
||||
offsets[currStage].chunkIdx = chunkIdx;
|
||||
|
||||
processNewTask = chunkIdx == batchChunks - 1 && headInnerIdx == PING_PONG_STAGES - 1;
|
||||
if (processNewTask) {
|
||||
uint32_t currLoopIdx = taskIdx / (PING_PONG_STAGES * cubeCoreNum);
|
||||
headInnerLoop = ((currLoopIdx + 2 == taskLoops) && hasDummyHead) ? 1 : PING_PONG_STAGES;
|
||||
taskIdx = (currLoopIdx + 1) * PING_PONG_STAGES * cubeCoreNum + headInnerLoop * cubeCoreIdx;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GDNFwdHOffsets& GetStage1Offsets() {
|
||||
return offsets[currStage];
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
bool NeedProcessStage1() {
|
||||
GDNFwdHOffsets& stage1Offsets = GetStage1Offsets();
|
||||
return !(lastLoop || stage1Offsets.isDummyHead);
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GDNFwdHOffsets& GetStage2Offsets() {
|
||||
return offsets[(currStage - 1) % PING_PONG_STAGES];
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
bool NeedProcessStage2() {
|
||||
GDNFwdHOffsets& stage2Offsets = GetStage2Offsets();
|
||||
return !(iterId == 1 || (!storeFinalState && stage2Offsets.isFinalState) || stage2Offsets.isDummyHead);
|
||||
}
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdHCube : public BlockSchedulerGdnFwdH {
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdHCube() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) {
|
||||
BlockSchedulerGdnFwdH::Init(cu_seqlens, chunk_indices, tiling, user, AscendC::GetBlockIdx(), AscendC::GetBlockNum());
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdHVec : public BlockSchedulerGdnFwdH {
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdHVec() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) {
|
||||
BlockSchedulerGdnFwdH::Init(cu_seqlens, chunk_indices, tiling, user, AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), AscendC::GetBlockNum());
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // namespace Catlass::Gemm::Block
|
||||
|
||||
#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP
|
||||
@@ -0,0 +1,511 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#define CATLASS_ARCH 2201
|
||||
#define CATLASS_UNIFIED_CORE 1
|
||||
|
||||
#include "catlass/arch/arch.hpp"
|
||||
#include "catlass/arch/cross_core_sync.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/epilogue/block/block_epilogue.hpp"
|
||||
#include "../../epilogue/block/block_epilogue_gdn_fwdh_update.hpp"
|
||||
#include "../../epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp"
|
||||
#include "catlass/gemm/block/block_mmad.hpp"
|
||||
#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp"
|
||||
#include "catlass/gemm/block/block_swizzle.hpp"
|
||||
#include "../block/block_scheduler_gdn_fwd_h.hpp"
|
||||
#include "catlass/gemm/dispatch_policy.hpp"
|
||||
#include "catlass/gemm/gemm_type.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "tla/tensor.hpp"
|
||||
#include "tla/layout.hpp"
|
||||
#include "tla/tensor.hpp"
|
||||
|
||||
using _0 = tla::Int<0>;
|
||||
using _1 = tla::Int<1>;
|
||||
using _2 = tla::Int<2>;
|
||||
using _4 = tla::Int<4>;
|
||||
using _8 = tla::Int<8>;
|
||||
using _16 = tla::Int<16>;
|
||||
using _32 = tla::Int<32>;
|
||||
using _64 = tla::Int<64>;
|
||||
using _128 = tla::Int<128>;
|
||||
using _256 = tla::Int<256>;
|
||||
using _512 = tla::Int<512>;
|
||||
using _1024 = tla::Int<1024>;
|
||||
using _2048 = tla::Int<2048>;
|
||||
using _4096 = tla::Int<4096>;
|
||||
using _8192 = tla::Int<8192>;
|
||||
using _16384 = tla::Int<16384>;
|
||||
using _32768 = tla::Int<32768>;
|
||||
using _65536 = tla::Int<65536>;
|
||||
|
||||
|
||||
|
||||
|
||||
#include "kernel_operator.h"
|
||||
using namespace Catlass;
|
||||
using namespace tla;
|
||||
|
||||
namespace Catlass::Gemm::Kernel {
|
||||
|
||||
template<
|
||||
typename INPUT_TYPE,
|
||||
typename G_TYPE,
|
||||
typename STATE_TYPE,
|
||||
typename WORKSPACE_TYPE
|
||||
>
|
||||
class GDNFwdHKernel {
|
||||
public:
|
||||
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHCube;
|
||||
using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHVec;
|
||||
|
||||
using DispatchPolicyTla = Gemm::MmadPingpongTlaMulti<ArchTag, true, false>;
|
||||
using L1TileShapeTla = Shape<_128, _128, _128>;
|
||||
using L0TileShapeTla = L1TileShapeTla;
|
||||
|
||||
using WType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using HType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using VworkType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using KType = Gemm::GemmType<INPUT_TYPE, layout::ColumnMajor>;
|
||||
using HworkType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using VType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using GType = Gemm::GemmType<G_TYPE, layout::RowMajor>;
|
||||
using UType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using FinalStateType = Gemm::GemmType<STATE_TYPE, layout::RowMajor>;
|
||||
|
||||
// cube 1
|
||||
using TileCopyWH = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using BlockMmadWH = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyWH>;
|
||||
|
||||
// cube 2
|
||||
using TileCopyKV = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::ColumnMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using BlockMmadKV = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyKV>;
|
||||
|
||||
// vec 1
|
||||
using DispatchPolicyGDNFwdHVnew = Epilogue::EpilogueAtlasGDNFwdHVnew;
|
||||
using EpilogueGDNFwdHVnew = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdHVnew, VType, GType, UType, VworkType>;
|
||||
|
||||
// vec 2
|
||||
using DispatchPolicyGDNFwdHUpdate = Epilogue::EpilogueAtlasGDNFwdHUpdate;
|
||||
using EpilogueGDNFwdHUpdate = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdHUpdate, HType, GType, HType, HworkType, FinalStateType>;
|
||||
|
||||
using GDNFwdHOffsets = Catlass::Gemm::Block::GDNFwdHOffsets;
|
||||
|
||||
using ElementK = INPUT_TYPE;
|
||||
using ElementW = INPUT_TYPE;
|
||||
using ElementU = INPUT_TYPE;
|
||||
using ElementG = G_TYPE;
|
||||
using ElementH = INPUT_TYPE;
|
||||
using ElementV = INPUT_TYPE;
|
||||
using ElementVWork = WORKSPACE_TYPE;
|
||||
using ElementHWork = WORKSPACE_TYPE;
|
||||
using ElementInitialState = STATE_TYPE;
|
||||
using ElementFinalState = STATE_TYPE;
|
||||
|
||||
using LayoutW = Catlass::layout::RowMajor;
|
||||
using LayoutH = Catlass::layout::RowMajor;
|
||||
using LayoutV = Catlass::layout::RowMajor;
|
||||
using LayoutK = Catlass::layout::ColumnMajor;
|
||||
|
||||
|
||||
uint32_t batch;
|
||||
uint32_t seqlen;
|
||||
uint32_t kNumHead;
|
||||
uint32_t vNumHead;
|
||||
uint32_t kHeadDim;
|
||||
uint32_t vHeadDim;
|
||||
uint32_t chunkSize;
|
||||
uint32_t initalStateStride0;
|
||||
bool useInitialState;
|
||||
bool storeFinalState;
|
||||
uint32_t isVariedLen;
|
||||
uint32_t shapeBatch;
|
||||
uint32_t tokenBatch;
|
||||
uint32_t vWorkspaceOffset;
|
||||
uint32_t vUpdateWorkspaceOffset;
|
||||
uint32_t hWorkspaceOffset;
|
||||
uint32_t numSeqWorkspaceOffset;
|
||||
uint32_t numChunksWorkspaceOffset;
|
||||
|
||||
AscendC::GlobalTensor<ElementK> gmK;
|
||||
AscendC::GlobalTensor<ElementW> gmW;
|
||||
AscendC::GlobalTensor<ElementU> gmU;
|
||||
AscendC::GlobalTensor<ElementG> gmG;
|
||||
AscendC::GlobalTensor<ElementInitialState> gmInitialState;
|
||||
AscendC::GlobalTensor<ElementH> gmH;
|
||||
AscendC::GlobalTensor<ElementV> gmV;
|
||||
AscendC::GlobalTensor<ElementFinalState> gmFinalState;
|
||||
AscendC::GlobalTensor<ElementVWork> gmVWorkspace;
|
||||
AscendC::GlobalTensor<ElementV> gmVUpdateWorkspace;
|
||||
AscendC::GlobalTensor<ElementHWork> gmHWorkspace;
|
||||
|
||||
AscendC::GlobalTensor<int64_t> gmSeqlen;
|
||||
AscendC::GlobalTensor<int64_t> gmNumSeq;
|
||||
AscendC::GlobalTensor<int64_t> gmNumChunks;
|
||||
|
||||
CubeScheduler cubeBlockScheduler;
|
||||
VecScheduler vecBlockScheduler;
|
||||
|
||||
Arch::Resource<ArchTag> resource;
|
||||
|
||||
|
||||
__aicore__ inline GDNFwdHKernel() {}
|
||||
|
||||
__aicore__ inline void Init(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices,
|
||||
GM_ADDR h, GM_ADDR v_new, GM_ADDR final_state, GM_ADDR tiling, GM_ADDR user) {
|
||||
|
||||
__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling);
|
||||
|
||||
batch = gdnFwdHTilingData->batch;
|
||||
seqlen = gdnFwdHTilingData->seqlen;
|
||||
kNumHead = gdnFwdHTilingData->kNumHead;
|
||||
vNumHead = gdnFwdHTilingData->vNumHead;
|
||||
kHeadDim = gdnFwdHTilingData->kHeadDim;
|
||||
vHeadDim = gdnFwdHTilingData->vHeadDim;
|
||||
chunkSize = gdnFwdHTilingData->chunkSize;
|
||||
initalStateStride0 = gdnFwdHTilingData->initalStateStride0;
|
||||
useInitialState = gdnFwdHTilingData->useInitialState;
|
||||
storeFinalState = gdnFwdHTilingData->storeFinalState;
|
||||
isVariedLen = gdnFwdHTilingData->isVariedLen;
|
||||
shapeBatch = gdnFwdHTilingData->shapeBatch;
|
||||
tokenBatch = gdnFwdHTilingData->tokenBatch;
|
||||
vWorkspaceOffset = gdnFwdHTilingData->vWorkspaceOffset;
|
||||
vUpdateWorkspaceOffset = gdnFwdHTilingData->vUpdateWorkspaceOffset;
|
||||
hWorkspaceOffset = gdnFwdHTilingData->hWorkspaceOffset;
|
||||
numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset;
|
||||
numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset;
|
||||
|
||||
gmK.SetGlobalBuffer((__gm__ ElementK *)k);
|
||||
gmW.SetGlobalBuffer((__gm__ ElementW *)w);
|
||||
gmU.SetGlobalBuffer((__gm__ ElementU *)u);
|
||||
gmG.SetGlobalBuffer((__gm__ ElementG *)g);
|
||||
gmInitialState.SetGlobalBuffer((__gm__ ElementInitialState *)inital_state);
|
||||
gmH.SetGlobalBuffer((__gm__ ElementH *)h);
|
||||
gmV.SetGlobalBuffer((__gm__ ElementV *)v_new);
|
||||
gmFinalState.SetGlobalBuffer((__gm__ ElementFinalState *)final_state);
|
||||
gmVWorkspace.SetGlobalBuffer((__gm__ ElementVWork *)(user + vWorkspaceOffset));
|
||||
gmVUpdateWorkspace.SetGlobalBuffer((__gm__ ElementV *)(user + vUpdateWorkspaceOffset));
|
||||
gmHWorkspace.SetGlobalBuffer((__gm__ ElementHWork *)(user + hWorkspaceOffset));
|
||||
|
||||
gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens);
|
||||
gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset));
|
||||
gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset));
|
||||
|
||||
cubeBlockScheduler.Init(cu_seqlens, chunk_indices, tiling, user);
|
||||
}
|
||||
|
||||
__aicore__ inline void Process() {
|
||||
ProcessUnifiedCore();
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessUnifiedCore() {
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
|
||||
BlockMmadWH blockMmadWH(resource);
|
||||
BlockMmadKV blockMmadKV(resource);
|
||||
EpilogueGDNFwdHVnew epilogueGDNFwdHVnew(resource);
|
||||
|
||||
auto wLayout = tla::MakeLayout<ElementW, LayoutW>(shapeBatch * kNumHead * cubeBlockScheduler.totalTokens, kHeadDim);
|
||||
auto hLayout = tla::MakeLayout<ElementH, LayoutH>(shapeBatch * vNumHead * cubeBlockScheduler.totalChunks * kHeadDim, vHeadDim);
|
||||
auto vLayout = tla::MakeLayout<ElementVWork, LayoutV>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
|
||||
auto kLayout = tla::MakeLayout<ElementK, LayoutK>(kHeadDim, shapeBatch * kNumHead * cubeBlockScheduler.totalTokens);
|
||||
auto vworkLayout = tla::MakeLayout<ElementV, LayoutV>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
|
||||
auto hworkLayout = tla::MakeLayout<ElementHWork, LayoutH>(coreNum * kHeadDim * PING_PONG_STAGES, vHeadDim);
|
||||
|
||||
if (useInitialState) {
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensorPing = resource.ubBuf.template GetBufferByByte<ElementInitialState>(0);
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensorPong = resource.ubBuf.template GetBufferByByte<ElementInitialState>(96 * 1024);
|
||||
AscendC::LocalTensor<ElementH> hUbTensorPing = resource.ubBuf.template GetBufferByByte<ElementH>(64 * 1024);
|
||||
AscendC::LocalTensor<ElementH> hUbTensorPong = resource.ubBuf.template GetBufferByByte<ElementH>(160 * 1024);
|
||||
uint32_t totalChunks = isVariedLen ? cubeBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize);
|
||||
uint32_t stateBlockSize = kHeadDim * vHeadDim;
|
||||
uint32_t pingpongFlag = 1;
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID1);
|
||||
for (uint32_t shapeBatchIdx = 0; shapeBatchIdx < shapeBatch; shapeBatchIdx++) {
|
||||
for (uint32_t vHeadIdx = 0; vHeadIdx < vNumHead; vHeadIdx++) {
|
||||
for (uint32_t tokenBatchIdx = 0; tokenBatchIdx < cubeBlockScheduler.tokenBatch; tokenBatchIdx++) {
|
||||
uint32_t batchIdx = isVariedLen ? tokenBatchIdx : shapeBatchIdx;
|
||||
uint32_t chunkOffset = isVariedLen ? gmNumChunks.GetValue(tokenBatchIdx) : 0;
|
||||
uint32_t initialStateOffset = (batchIdx * vNumHead + vHeadIdx) * stateBlockSize;
|
||||
uint32_t hOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * stateBlockSize;
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong;
|
||||
AscendC::LocalTensor<ElementH> hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong;
|
||||
auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0;
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(event_id);
|
||||
if constexpr(!std::is_same<ElementInitialState, ElementH>::value) {
|
||||
AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateOffset], stateBlockSize);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(event_id);
|
||||
AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_NONE, stateBlockSize);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(event_id);
|
||||
AscendC::DataCopy(gmH[hOffset], hUbTensor, stateBlockSize);
|
||||
} else {
|
||||
AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateOffset], stateBlockSize);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE3>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE3>(event_id);
|
||||
AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateBlockSize);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(event_id);
|
||||
pingpongFlag = 1 - pingpongFlag;
|
||||
}
|
||||
}
|
||||
}
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID1);
|
||||
}
|
||||
|
||||
while (cubeBlockScheduler.isRunning) {
|
||||
cubeBlockScheduler.InitTask();
|
||||
GDNFwdHOffsets& stage1Offsets = cubeBlockScheduler.GetStage1Offsets();
|
||||
|
||||
// CUBE1: v_work = w @ h[i]
|
||||
if (cubeBlockScheduler.NeedProcessStage1()) {
|
||||
auto tensorW = tla::MakeTensor(gmW[stage1Offsets.wOffset], wLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorH = tla::MakeTensor(gmH[stage1Offsets.hSrcOffset], hLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorV = tla::MakeTensor(gmVWorkspace[stage1Offsets.vWorkOffset], vLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube1Shape{stage1Offsets.blockTokens, vHeadDim, kHeadDim};
|
||||
auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k()));
|
||||
auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n()));
|
||||
auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n()));
|
||||
blockMmadWH.preSetFlags();
|
||||
blockMmadWH(tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape);
|
||||
blockMmadWH.finalWaitFlags();
|
||||
}
|
||||
|
||||
// VEC1: v_new epilogue
|
||||
if (cubeBlockScheduler.NeedProcessStage1()) {
|
||||
epilogueGDNFwdHVnew(
|
||||
gmV[stage1Offsets.uvOffset], gmVUpdateWorkspace[stage1Offsets.vWorkOffset],
|
||||
gmG[stage1Offsets.gOffset], gmU[stage1Offsets.uvOffset], gmVWorkspace[stage1Offsets.vWorkOffset],
|
||||
stage1Offsets.blockTokens, kHeadDim, vHeadDim, cubeBlockScheduler.cube1Done
|
||||
);
|
||||
}
|
||||
|
||||
if (cubeBlockScheduler.iterId > 1) {
|
||||
GDNFwdHOffsets& stage2Offsets = cubeBlockScheduler.GetStage2Offsets();
|
||||
|
||||
// CUBE2: h_work = k.T @ v_update
|
||||
// BlockMmadTla has no outer M loop; m must be split when kHeadDim > L1_TILE_M.
|
||||
if (cubeBlockScheduler.NeedProcessStage2()) {
|
||||
auto tensorK = tla::MakeTensor(gmK[stage2Offsets.wkOffset], kLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[stage2Offsets.vWorkOffset], vworkLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorHwork = tla::MakeTensor(gmHWorkspace[stage2Offsets.hWorkOffset], hworkLayout, Catlass::Arch::PositionGM{});
|
||||
constexpr uint32_t L1_TILE_M_C2 = tla::get<0>(L1TileShapeTla{});
|
||||
uint32_t mLoopC2 = (kHeadDim + L1_TILE_M_C2 - 1) / L1_TILE_M_C2;
|
||||
for (uint32_t mIdx = 0; mIdx < mLoopC2; ++mIdx) {
|
||||
uint32_t mOff = mIdx * L1_TILE_M_C2;
|
||||
uint32_t mTail = kHeadDim - mOff;
|
||||
uint32_t mActual = (mTail < L1_TILE_M_C2) ? mTail : L1_TILE_M_C2;
|
||||
GemmCoord cube2Shape{mActual, vHeadDim, stage2Offsets.blockTokens};
|
||||
auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(mOff, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k()));
|
||||
auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n()));
|
||||
auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(mOff, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n()));
|
||||
blockMmadKV.preSetFlags();
|
||||
blockMmadKV(tensorBlockK, tensorBlockVwork, tensorBlockHwork, cube2Shape);
|
||||
blockMmadKV.finalWaitFlags();
|
||||
}
|
||||
}
|
||||
|
||||
// VEC2: h update epilogue
|
||||
if (cubeBlockScheduler.NeedProcessStage2()) {
|
||||
EpilogueGDNFwdHUpdate epilogueGDNFwdHUpdate(resource);
|
||||
epilogueGDNFwdHUpdate(
|
||||
gmH[stage2Offsets.hDstOffset], gmFinalState[stage2Offsets.finalStateOffset],
|
||||
gmG[stage2Offsets.gOffset], gmH[stage2Offsets.hSrcOffset],
|
||||
gmHWorkspace[stage2Offsets.hWorkOffset],
|
||||
stage2Offsets.blockTokens, kHeadDim, vHeadDim, cubeBlockScheduler.cube2Done,
|
||||
(stage2Offsets.isFinalState && storeFinalState)
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void ProcessSplitCore() {
|
||||
|
||||
if ASCEND_IS_AIC {
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
|
||||
BlockMmadWH blockMmadWH(resource);
|
||||
BlockMmadKV blockMmadKV(resource);
|
||||
|
||||
auto wLayout = tla::MakeLayout<ElementW, LayoutW>(shapeBatch * kNumHead * cubeBlockScheduler.totalTokens, kHeadDim);
|
||||
auto hLayout = tla::MakeLayout<ElementH, LayoutH>(shapeBatch * vNumHead * cubeBlockScheduler.totalChunks * kHeadDim, vHeadDim);
|
||||
auto vLayout = tla::MakeLayout<ElementVWork, LayoutV>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
|
||||
|
||||
auto kLayout = tla::MakeLayout<ElementK, LayoutK>(kHeadDim, shapeBatch * kNumHead * cubeBlockScheduler.totalTokens);
|
||||
auto vworkLayout = tla::MakeLayout<ElementV, LayoutV>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
|
||||
auto hworkLayout = tla::MakeLayout<ElementHWork, LayoutH>(coreNum * kHeadDim * PING_PONG_STAGES, vHeadDim);
|
||||
|
||||
while (cubeBlockScheduler.isRunning) {
|
||||
cubeBlockScheduler.InitTask();
|
||||
// step 1: v_work = w @ h[i]
|
||||
GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetStage1Offsets();
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done);
|
||||
if (cubeBlockScheduler.NeedProcessStage1()) {
|
||||
int64_t cube1OffsetW = cube1Offsets.wOffset;
|
||||
int64_t cube1OffsetH = cube1Offsets.hSrcOffset;
|
||||
int64_t cube1OffsetVwork = cube1Offsets.vWorkOffset;
|
||||
auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorV = tla::MakeTensor(gmVWorkspace[cube1OffsetVwork], vLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube1Shape {cube1Offsets.blockTokens, vHeadDim, kHeadDim};
|
||||
auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k()));
|
||||
auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n()));
|
||||
auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n()));
|
||||
blockMmadWH.preSetFlags();
|
||||
blockMmadWH(tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape);
|
||||
blockMmadWH.finalWaitFlags();
|
||||
}
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done);
|
||||
|
||||
if (cubeBlockScheduler.iterId > 1) {
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
|
||||
GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetStage2Offsets();
|
||||
if (cubeBlockScheduler.NeedProcessStage2()) {
|
||||
// step 3: h[i+1] = k.T @ v_work
|
||||
// BlockMmadTla has no outer M loop; m must be split when kHeadDim > L1_TILE_M.
|
||||
int64_t cube2OffsetK = cube2Offsets.wkOffset;
|
||||
int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset;
|
||||
int64_t cube2OffsetH = cube2Offsets.hWorkOffset;
|
||||
auto tensorK = tla::MakeTensor(gmK[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vworkLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorHwork = tla::MakeTensor(gmHWorkspace[cube2OffsetH], hworkLayout, Catlass::Arch::PositionGM{});
|
||||
constexpr uint32_t L1_TILE_M_C2 = tla::get<0>(L1TileShapeTla{});
|
||||
uint32_t mLoopC2 = (kHeadDim + L1_TILE_M_C2 - 1) / L1_TILE_M_C2;
|
||||
for (uint32_t mIdx = 0; mIdx < mLoopC2; ++mIdx) {
|
||||
uint32_t mOff = mIdx * L1_TILE_M_C2;
|
||||
uint32_t mTail = kHeadDim - mOff;
|
||||
uint32_t mActual = (mTail < L1_TILE_M_C2) ? mTail : L1_TILE_M_C2;
|
||||
GemmCoord cube2Shape{mActual, vHeadDim, cube2Offsets.blockTokens};
|
||||
auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(mOff, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k()));
|
||||
auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n()));
|
||||
auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(mOff, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n()));
|
||||
blockMmadKV.preSetFlags();
|
||||
blockMmadKV(tensorBlockK, tensorBlockVwork, tensorBlockHwork, cube2Shape);
|
||||
blockMmadKV.finalWaitFlags();
|
||||
}
|
||||
}
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done);
|
||||
}
|
||||
}
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done);
|
||||
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
|
||||
EpilogueGDNFwdHVnew epilogueGDNFwdHVnew(resource);
|
||||
|
||||
if (useInitialState) {
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensorPing = resource.ubBuf.template GetBufferByByte<ElementInitialState>(0);
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensorPong = resource.ubBuf.template GetBufferByByte<ElementInitialState>(96 * 1024);
|
||||
AscendC::LocalTensor<ElementH> hUbTensorPing = resource.ubBuf.template GetBufferByByte<ElementH>(64 * 1024);
|
||||
AscendC::LocalTensor<ElementH> hUbTensorPong = resource.ubBuf.template GetBufferByByte<ElementH>(160 * 1024);
|
||||
uint32_t totalChunks = isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize);
|
||||
uint32_t stateBlockSize = kHeadDim * vHeadDim;
|
||||
uint32_t pingpongFlag = 1;
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID1);
|
||||
AscendC::DataCopyParams repeatParams = {static_cast<uint16_t>(kHeadDim), static_cast<uint16_t>(vHeadDim * sizeof(ElementInitialState) / 32),
|
||||
static_cast<uint16_t>((initalStateStride0 - vHeadDim)* sizeof(ElementInitialState) / 32), static_cast<uint16_t>(0)};
|
||||
for (uint32_t shapeBatchIdx = 0; shapeBatchIdx < shapeBatch; shapeBatchIdx++) {
|
||||
for (uint32_t vHeadIdx = 0; vHeadIdx < vNumHead; vHeadIdx++) {
|
||||
for (uint32_t tokenBatchIdx = 0; tokenBatchIdx < vecBlockScheduler.tokenBatch; tokenBatchIdx++) {
|
||||
uint32_t batchIdx = isVariedLen ? tokenBatchIdx : shapeBatchIdx;
|
||||
uint32_t chunkOffset = isVariedLen ? gmNumChunks.GetValue(tokenBatchIdx) : 0;
|
||||
uint32_t initialStateSrcOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * initalStateStride0;
|
||||
uint32_t hOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * stateBlockSize;
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong;
|
||||
AscendC::LocalTensor<ElementH> hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong;
|
||||
auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0;
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(event_id);
|
||||
if constexpr(!std::is_same<ElementInitialState, ElementH>::value) {
|
||||
AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateSrcOffset], repeatParams);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(event_id);
|
||||
AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateBlockSize);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(event_id);
|
||||
AscendC::DataCopy(gmH[hOffset], hUbTensor, stateBlockSize);
|
||||
} else {
|
||||
AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateSrcOffset], repeatParams);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE3>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE3>(event_id);
|
||||
AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateBlockSize);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(event_id);
|
||||
pingpongFlag = 1 - pingpongFlag;
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID1);
|
||||
}
|
||||
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
|
||||
while (vecBlockScheduler.isRunning) {
|
||||
vecBlockScheduler.InitTask();
|
||||
// step 2:
|
||||
GDNFwdHOffsets& vec1Offsets = vecBlockScheduler.GetStage1Offsets();
|
||||
// gmV = gmU - gmVWorkspace
|
||||
// g_buf = gmG[-1] - gmG
|
||||
// g_buf = exp(g_buf)
|
||||
// gmVWorkspace = g_buf * gmV
|
||||
if (vecBlockScheduler.NeedProcessStage1()) {
|
||||
epilogueGDNFwdHVnew(
|
||||
gmV[vec1Offsets.uvOffset], gmVUpdateWorkspace[vec1Offsets.vWorkOffset],
|
||||
gmG[vec1Offsets.gOffset], gmU[vec1Offsets.uvOffset], gmVWorkspace[vec1Offsets.vWorkOffset],
|
||||
vec1Offsets.blockTokens, kHeadDim, vHeadDim, vecBlockScheduler.cube1Done
|
||||
);
|
||||
} else {
|
||||
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube1Done);
|
||||
}
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
|
||||
|
||||
if (vecBlockScheduler.iterId > 1) {
|
||||
GDNFwdHOffsets& vec2Offsets = vecBlockScheduler.GetStage2Offsets();
|
||||
if (vecBlockScheduler.NeedProcessStage2()) {
|
||||
// step 4: h[i+1] += h_work if i < num_chunks - 1 else None
|
||||
EpilogueGDNFwdHUpdate epilogueGDNFwdHUpdate(resource);
|
||||
epilogueGDNFwdHUpdate(
|
||||
gmH[vec2Offsets.hDstOffset], gmFinalState[vec2Offsets.finalStateOffset],
|
||||
gmG[vec2Offsets.gOffset],
|
||||
gmH[vec2Offsets.hSrcOffset],
|
||||
gmHWorkspace[vec2Offsets.hWorkOffset],
|
||||
vec2Offsets.blockTokens, kHeadDim, vHeadDim, vecBlockScheduler.cube2Done,
|
||||
(vec2Offsets.isFinalState && storeFinalState)
|
||||
);
|
||||
} else {
|
||||
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done);
|
||||
}
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_UPDATE_HPP
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "../gdn_fwd_h_epilogue_policies.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/epilogue/tile/tile_copy.hpp"
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <
|
||||
class HOutputType_,
|
||||
class GInputType_,
|
||||
class HInputType_,
|
||||
class HUpdateInputType_,
|
||||
class FinalStateType_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasGDNFwdHUpdate,
|
||||
HOutputType_,
|
||||
GInputType_,
|
||||
HInputType_,
|
||||
HUpdateInputType_,
|
||||
FinalStateType_
|
||||
> {
|
||||
public:
|
||||
// Type aliases
|
||||
using DispatchPolicy = EpilogueAtlasGDNFwdHUpdate;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using HElementOutput = typename HOutputType_::Element;
|
||||
using GElementInput = typename GInputType_::Element;
|
||||
using HElementInput = typename HInputType_::Element;
|
||||
using HUpdateElementInput = typename HUpdateInputType_::Element;
|
||||
using FinalStateElement = typename FinalStateType_::Element;
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
|
||||
constexpr uint32_t CALC_BUF_OFFSET = 0;
|
||||
constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024;
|
||||
constexpr uint32_t PING_BUF_1_OFFSET = 64 * 1024;
|
||||
constexpr uint32_t PING_BUF_2_OFFSET = 80 * 1024;
|
||||
constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024;
|
||||
|
||||
|
||||
calcUbTensor = resource.ubBuf.template GetBufferByByte<float>(CALC_BUF_OFFSET);
|
||||
|
||||
hUpdateUbTensor = resource.ubBuf.template GetBufferByByte<float>(PING_BUF_1_OFFSET);
|
||||
hUbTensor = resource.ubBuf.template GetBufferByByte<HElementInput>(PING_BUF_0_OFFSET);
|
||||
|
||||
hOutputUbTensor = resource.ubBuf.template GetBufferByByte<HElementOutput>(PING_BUF_1_OFFSET);
|
||||
finalOutputUbTensor = resource.ubBuf.template GetBufferByByte<FinalStateElement>(PING_BUF_1_OFFSET);
|
||||
|
||||
glastUbTensor = resource.ubBuf.template GetBufferByByte<float>(PING_G_BUF_OFFSET);
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(
|
||||
AscendC::GlobalTensor<HElementOutput> hOutput,
|
||||
AscendC::GlobalTensor<FinalStateElement> finalState,
|
||||
AscendC::GlobalTensor<GElementInput> gInput,
|
||||
AscendC::GlobalTensor<HElementInput> hInput,
|
||||
AscendC::GlobalTensor<float> hUpdateInput,
|
||||
uint32_t chunkSize,
|
||||
uint32_t kHeadDim,
|
||||
uint32_t vHeadDim,
|
||||
Arch::CrossCoreFlag cube2Done,
|
||||
bool isFinalState
|
||||
)
|
||||
{
|
||||
uint32_t mActual = kHeadDim;
|
||||
uint32_t nActual = vHeadDim;
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
uint32_t mActualPerSubBlock = CeilDiv(mActual, subBlockNum);
|
||||
uint32_t mActualThisSubBlock = (subBlockIdx == 0) ? mActualPerSubBlock : (mActual - mActualPerSubBlock);
|
||||
uint32_t mOffset = subBlockIdx * mActualPerSubBlock;
|
||||
uint32_t nOffset = 0;
|
||||
int64_t offsetH = mOffset * nActual + nOffset;
|
||||
|
||||
AscendC::ResetMask();
|
||||
|
||||
AscendC::GlobalTensor<HElementOutput> hOutputThisSubBlock = hOutput[offsetH];
|
||||
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
|
||||
AscendC::GlobalTensor<HElementInput> hInputThisSubBlock = hInput[offsetH];
|
||||
AscendC::GlobalTensor<float> hUpdateInputThisSubBlock = hUpdateInput[offsetH];
|
||||
AscendC::GlobalTensor<FinalStateElement> finalStateThisSubBlock = finalState[offsetH];
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0);
|
||||
AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0);
|
||||
AscendC::Cast(calcUbTensor, hUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
GElementInput gLastVal = gInputThisSubBlock.GetValue(chunkSize-1);
|
||||
float gLastFloat = 0.0f;
|
||||
if constexpr(std::is_same<GElementInput, float>::value) {
|
||||
gLastFloat = gLastVal;
|
||||
} else if constexpr(std::is_same<GElementInput, half>::value) {
|
||||
gLastFloat = (float)gLastVal;
|
||||
} else if constexpr(std::is_same<GElementInput, bfloat16_t>::value) {
|
||||
gLastFloat = AscendC::ToFloat(gLastVal);
|
||||
}
|
||||
glastUbTensor.SetValue(0, gLastFloat);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::Exp(glastUbTensor, glastUbTensor, 1);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_S>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_S>(EVENT_ID0);
|
||||
float muls = glastUbTensor.GetValue(0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(EVENT_ID0);
|
||||
AscendC::Muls(calcUbTensor, calcUbTensor, muls, mActualThisSubBlock * nActual);
|
||||
|
||||
Arch::CrossCoreWaitFlag(cube2Done);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1);
|
||||
AscendC::DataCopy(hUpdateUbTensor, hUpdateInputThisSubBlock, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1);
|
||||
AscendC::Add<float>(hUpdateUbTensor, calcUbTensor, hUpdateUbTensor, mActualThisSubBlock * nActual);
|
||||
|
||||
if (isFinalState) {
|
||||
if constexpr(!std::is_same<FinalStateElement, float>::value) {
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Cast(finalOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
|
||||
AscendC::DataCopy(finalStateThisSubBlock, finalOutputUbTensor, mActualThisSubBlock * nActual);
|
||||
} else {
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
|
||||
AscendC::DataCopy(finalStateThisSubBlock, hUpdateUbTensor, mActualThisSubBlock * nActual);
|
||||
}
|
||||
} else {
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Cast(hOutputUbTensor, hUpdateUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
|
||||
AscendC::DataCopy(hOutputThisSubBlock, hOutputUbTensor, mActualThisSubBlock * nActual);
|
||||
}
|
||||
}
|
||||
|
||||
private:
|
||||
AscendC::LocalTensor<float> calcUbTensor;
|
||||
|
||||
AscendC::LocalTensor<HElementInput> hUbTensor;
|
||||
AscendC::LocalTensor<float> hUpdateUbTensor;
|
||||
|
||||
AscendC::LocalTensor<HElementOutput> hOutputUbTensor;
|
||||
AscendC::LocalTensor<FinalStateElement> finalOutputUbTensor;
|
||||
|
||||
AscendC::LocalTensor<float> glastUbTensor;
|
||||
|
||||
};
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,268 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_VNEW_HPP
|
||||
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDH_VNEW_HPP
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "../gdn_fwd_h_epilogue_policies.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "catlass/matrix_coord.hpp"
|
||||
#include "catlass/epilogue/tile/tile_copy.hpp"
|
||||
|
||||
|
||||
|
||||
namespace Catlass::Epilogue::Block {
|
||||
|
||||
template <
|
||||
class VOutputType_,
|
||||
class GInputType_,
|
||||
class UInputType_,
|
||||
class WSInputType_
|
||||
>
|
||||
class BlockEpilogue <
|
||||
EpilogueAtlasGDNFwdHVnew,
|
||||
VOutputType_,
|
||||
GInputType_,
|
||||
UInputType_,
|
||||
WSInputType_
|
||||
> {
|
||||
public:
|
||||
// Type aliases
|
||||
using DispatchPolicy = EpilogueAtlasGDNFwdHVnew;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
|
||||
using VElementOutput = typename VOutputType_::Element;
|
||||
using GElementInput = typename GInputType_::Element;
|
||||
using UElementInput = typename UInputType_::Element;
|
||||
using WSElementInput = typename WSInputType_::Element;
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockEpilogue(Arch::Resource<ArchTag> &resource)
|
||||
{
|
||||
|
||||
constexpr uint32_t CALC_BUF_OFFSET = 0;
|
||||
constexpr uint32_t PING_BUF_0_OFFSET = 32 * 1024;
|
||||
constexpr uint32_t PING_BUF_1_OFFSET = 64 * 1024;
|
||||
constexpr uint32_t PONG_BUF_0_OFFSET = 96 * 1024;
|
||||
constexpr uint32_t PONG_BUF_1_OFFSET = 128 * 1024;
|
||||
constexpr uint32_t PING_G_BUF_OFFSET = 160 * 1024;
|
||||
constexpr uint32_t PONG_G_BUF_OFFSET = 161 * 1024;
|
||||
constexpr uint32_t PING_G_SUB_BUF_OFFSET = 162 * 1024;
|
||||
constexpr uint32_t PONG_G_SUB_BUF_OFFSET = 163 * 1024;
|
||||
constexpr uint32_t PING_G_INPUT_BUF_OFFSET = 164 * 1024;
|
||||
constexpr uint32_t PONG_G_INPUT_BUF_OFFSET = 165 * 1024;
|
||||
constexpr uint32_t SHARE_BUF_OFFSET = 166 * 1024;
|
||||
|
||||
calcUbTensor = resource.ubBuf.template GetBufferByByte<float>(CALC_BUF_OFFSET);
|
||||
|
||||
uUbTensor_ping = resource.ubBuf.template GetBufferByByte<UElementInput>(PING_BUF_1_OFFSET);
|
||||
uUbFloatTensor_ping = resource.ubBuf.template GetBufferByByte<float>(PING_BUF_0_OFFSET);
|
||||
wsUbTensor_ping = resource.ubBuf.template GetBufferByByte<float>(PING_BUF_1_OFFSET);
|
||||
gUbTensor_ping = resource.ubBuf.template GetBufferByByte<float>(PING_G_BUF_OFFSET);
|
||||
gLastUbTensor_ping = resource.ubBuf.template GetBufferByByte<float>(PING_G_SUB_BUF_OFFSET);
|
||||
gInputUbTensor_ping = resource.ubBuf.template GetBufferByByte<GElementInput>(PING_G_INPUT_BUF_OFFSET);
|
||||
vNewOutputUbTensor_ping = resource.ubBuf.template GetBufferByByte<VElementOutput>(PING_BUF_1_OFFSET);
|
||||
vNewDecayUbTensor_ping = resource.ubBuf.template GetBufferByByte<VElementOutput>(PING_BUF_0_OFFSET);
|
||||
|
||||
uUbTensor_pong = resource.ubBuf.template GetBufferByByte<UElementInput>(PONG_BUF_1_OFFSET);
|
||||
uUbFloatTensor_pong = resource.ubBuf.template GetBufferByByte<float>(PONG_BUF_0_OFFSET);
|
||||
wsUbTensor_pong = resource.ubBuf.template GetBufferByByte<float>(PONG_BUF_1_OFFSET);
|
||||
gUbTensor_pong = resource.ubBuf.template GetBufferByByte<float>(PONG_G_BUF_OFFSET);
|
||||
gLastUbTensor_pong = resource.ubBuf.template GetBufferByByte<float>(PONG_G_SUB_BUF_OFFSET);
|
||||
gInputUbTensor_pong = resource.ubBuf.template GetBufferByByte<GElementInput>(PONG_G_INPUT_BUF_OFFSET);
|
||||
vNewOutputUbTensor_pong = resource.ubBuf.template GetBufferByByte<VElementOutput>(PONG_BUF_1_OFFSET);
|
||||
vNewDecayUbTensor_pong = resource.ubBuf.template GetBufferByByte<VElementOutput>(PONG_BUF_0_OFFSET);
|
||||
|
||||
shareBuffer_ = resource.ubBuf.template GetBufferByByte<uint8_t>(SHARE_BUF_OFFSET);
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
~BlockEpilogue() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void operator()(
|
||||
AscendC::GlobalTensor<VElementOutput> vnewOutput,
|
||||
AscendC::GlobalTensor<VElementOutput> vnewdecayOutput,
|
||||
AscendC::GlobalTensor<GElementInput> gInput,
|
||||
AscendC::GlobalTensor<UElementInput> uInput,
|
||||
AscendC::GlobalTensor<float> wsInput,
|
||||
uint32_t chunkSize,
|
||||
uint32_t kHeadDim,
|
||||
uint32_t vHeadDim,
|
||||
Arch::CrossCoreFlag cube1Done
|
||||
// const LayoutOutput &layoutOutput,
|
||||
// const LayoutInput &LayoutInput
|
||||
)
|
||||
{
|
||||
uint32_t mActual = chunkSize;
|
||||
uint32_t nkActual = kHeadDim;
|
||||
uint32_t nvActual = vHeadDim;
|
||||
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
uint32_t mActualPerSubBlock = CeilDiv(mActual, subBlockNum);
|
||||
uint32_t mActualThisSubBlock = (subBlockIdx == 0) ? mActualPerSubBlock : (mActual - mActualPerSubBlock);
|
||||
uint32_t mOffset = subBlockIdx * mActualPerSubBlock;
|
||||
uint32_t nOffset = 0;
|
||||
// 当前场景内部一定连续
|
||||
// k [B, H, T, D]
|
||||
// g [B, H, T]
|
||||
// 在外部offset的基础上进一步offset
|
||||
// 当前asset kdim == vHeadDim
|
||||
int64_t offsetK = mOffset * nvActual + nOffset;
|
||||
int64_t offsetD = 0; // 因为要用最后一个数减去之前所有,所以全部读入
|
||||
|
||||
uint32_t gbrcStart, gbrcRealStart, gbrcReptime, gbrcEffStart, gbrcEffEnd;
|
||||
if(subBlockIdx==0)
|
||||
{
|
||||
gbrcStart = 0;
|
||||
gbrcRealStart = 0;
|
||||
gbrcReptime = (mActualThisSubBlock + 8 - 1) / 8;
|
||||
|
||||
}
|
||||
else
|
||||
{
|
||||
gbrcStart = mActualPerSubBlock;
|
||||
gbrcRealStart = gbrcStart & ~15;
|
||||
gbrcReptime = (mActual - gbrcRealStart + 8 - 1) / 8;
|
||||
}
|
||||
gbrcEffStart = gbrcStart-gbrcRealStart;
|
||||
gbrcEffEnd = gbrcEffStart + mActualThisSubBlock;
|
||||
|
||||
AscendC::ResetMask();
|
||||
|
||||
AscendC::GlobalTensor<VElementOutput> vnewOutputThisSubBlock = vnewOutput[offsetK];
|
||||
AscendC::GlobalTensor<VElementOutput> vnewdecayOutputThisSubBlock = vnewdecayOutput[offsetK];
|
||||
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
|
||||
AscendC::GlobalTensor<UElementInput> uInputThisSubBlock = uInput[offsetK];
|
||||
AscendC::GlobalTensor<float> wsInputThisSubBlock = wsInput[offsetK];
|
||||
|
||||
pingpongFlag = isFirst ? 0 : 4;
|
||||
AscendC::LocalTensor<UElementInput> uUbTensor = isFirst ? uUbTensor_ping : uUbTensor_pong;
|
||||
AscendC::LocalTensor<float> uUbFloatTensor = isFirst ? uUbFloatTensor_ping : uUbFloatTensor_pong;
|
||||
AscendC::LocalTensor<float> wsUbTensor = isFirst ? wsUbTensor_ping : wsUbTensor_pong;
|
||||
AscendC::LocalTensor<float> gUbTensor = isFirst ? gUbTensor_ping : gUbTensor_pong;
|
||||
AscendC::LocalTensor<float> gLastUbTensor = isFirst ? gLastUbTensor_ping : gLastUbTensor_pong;
|
||||
AscendC::LocalTensor<GElementInput> gInputUbTensor = isFirst ? gInputUbTensor_ping : gInputUbTensor_pong;
|
||||
AscendC::LocalTensor<VElementOutput> vNewOutputUbTensor = isFirst ? vNewOutputUbTensor_ping : vNewOutputUbTensor_pong;
|
||||
AscendC::LocalTensor<VElementOutput> vNewDecayUbTensor = isFirst ? vNewDecayUbTensor_ping : vNewDecayUbTensor_pong;
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
|
||||
if constexpr(std::is_same<GElementInput, float>::value) {
|
||||
AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(float)), 0, 0};
|
||||
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyPad(gUbTensor, gInputThisSubBlock, gUbParams, gUbPadParams);
|
||||
} else {
|
||||
AscendC::DataCopyParams gUbParams{1, (uint16_t)(mActual * sizeof(half)), 0, 0};
|
||||
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
|
||||
AscendC::DataCopyPad(gInputUbTensor, gInputThisSubBlock, gUbParams, gUbPadParams);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
|
||||
if constexpr(!std::is_same<GElementInput, float>::value) {
|
||||
AscendC::Cast(gUbTensor, gInputUbTensor, AscendC::RoundMode::CAST_NONE, mActual);
|
||||
}
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_S>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_S>(EVENT_ID2 + pingpongFlag);
|
||||
float inputVal = gUbTensor.GetValue(mActual-1);
|
||||
AscendC::SetFlag<AscendC::HardEvent::S_V>(EVENT_ID2 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::S_V>(EVENT_ID2 + pingpongFlag);
|
||||
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Duplicate<float>(gLastUbTensor, inputVal, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::Sub<float>(gUbTensor, gLastUbTensor, gUbTensor, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::Exp(gUbTensor, gUbTensor, mActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
uint32_t dstShape_[2] = {gbrcReptime*8, nvActual};
|
||||
uint32_t srcShape_[2] = {gbrcReptime*8, 1};
|
||||
AscendC::Broadcast<float, 2, 1>(calcUbTensor, gUbTensor[gbrcRealStart], dstShape_, srcShape_, shareBuffer_);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
Arch::CrossCoreWaitFlag(cube1Done);
|
||||
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0 + pingpongFlag);
|
||||
|
||||
AscendC::DataCopy(uUbTensor, uInputThisSubBlock, mActualThisSubBlock * nvActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::Cast(uUbFloatTensor, uUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::DataCopy(wsUbTensor, wsInputThisSubBlock, mActualThisSubBlock * nvActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
|
||||
|
||||
AscendC::Sub<float>(uUbFloatTensor, uUbFloatTensor, wsUbTensor, mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Cast(vNewOutputUbTensor, uUbFloatTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(vnewOutputThisSubBlock, vNewOutputUbTensor, mActualThisSubBlock * nvActual);
|
||||
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Mul(calcUbTensor[gbrcEffStart*nvActual], uUbFloatTensor, calcUbTensor[gbrcEffStart*nvActual], mActualThisSubBlock * nvActual);
|
||||
AscendC::PipeBarrier<PIPE_V>();
|
||||
AscendC::Cast(vNewDecayUbTensor, calcUbTensor[gbrcEffStart*nvActual], AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nvActual);
|
||||
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
|
||||
AscendC::DataCopy(vnewdecayOutputThisSubBlock, vNewDecayUbTensor, mActualThisSubBlock * nvActual);
|
||||
|
||||
if (isFirst) {
|
||||
AscendC::PipeBarrier<PIPE_ALL>();
|
||||
}
|
||||
|
||||
isFirst = false;
|
||||
}
|
||||
|
||||
private:
|
||||
uint32_t pingpongFlag = 0;
|
||||
bool isFirst = true;
|
||||
|
||||
AscendC::LocalTensor<float> calcUbTensor;
|
||||
|
||||
AscendC::LocalTensor<UElementInput> uUbTensor_ping;
|
||||
AscendC::LocalTensor<float> uUbFloatTensor_ping;
|
||||
AscendC::LocalTensor<float> wsUbTensor_ping;
|
||||
AscendC::LocalTensor<float> gUbTensor_ping;
|
||||
AscendC::LocalTensor<float> gLastUbTensor_ping;
|
||||
AscendC::LocalTensor<GElementInput> gInputUbTensor_ping;
|
||||
AscendC::LocalTensor<VElementOutput> vNewOutputUbTensor_ping;
|
||||
AscendC::LocalTensor<VElementOutput> vNewDecayUbTensor_ping;
|
||||
|
||||
AscendC::LocalTensor<UElementInput> uUbTensor_pong;
|
||||
AscendC::LocalTensor<float> uUbFloatTensor_pong;
|
||||
AscendC::LocalTensor<float> wsUbTensor_pong;
|
||||
AscendC::LocalTensor<float> gUbTensor_pong;
|
||||
AscendC::LocalTensor<float> gLastUbTensor_pong;
|
||||
AscendC::LocalTensor<GElementInput> gInputUbTensor_pong;
|
||||
AscendC::LocalTensor<VElementOutput> vNewOutputUbTensor_pong;
|
||||
AscendC::LocalTensor<VElementOutput> vNewDecayUbTensor_pong;
|
||||
|
||||
AscendC::LocalTensor<uint8_t> shareBuffer_;
|
||||
|
||||
};
|
||||
}
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,35 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#ifndef CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP
|
||||
#define CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP
|
||||
|
||||
#include "catlass/catlass.hpp"
|
||||
|
||||
namespace Catlass::Epilogue {
|
||||
|
||||
struct EpilogueAtlasGDNFwdHVnew {
|
||||
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310
|
||||
using ArchTag = Arch::Ascend950;
|
||||
#else
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
#endif
|
||||
};
|
||||
|
||||
struct EpilogueAtlasGDNFwdHUpdate {
|
||||
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310
|
||||
using ArchTag = Arch::Ascend950;
|
||||
#else
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
#endif
|
||||
};
|
||||
|
||||
} // namespace Catlass::Epilogue
|
||||
|
||||
#endif // CATLASS_EPILOGUE_GDN_FWD_H_EPILOGUE_POLICIES_HPP
|
||||
@@ -0,0 +1,286 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
using namespace Catlass;
|
||||
|
||||
#ifndef CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP
|
||||
#define CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP
|
||||
|
||||
// constexpr uint32_t PING_PONG_STAGES = 1;
|
||||
constexpr uint32_t PING_PONG_STAGES = 2;
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T AlignUp(T a, T b) {
|
||||
return (b == 0) ? 0 : (a + b - 1) / b * b;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T Min(T a, T b) {
|
||||
return (a > b) ? b : a;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
CATLASS_DEVICE T Max(T a, T b) {
|
||||
return (a > b) ? a : b;
|
||||
}
|
||||
|
||||
namespace Catlass::Gemm::Block {
|
||||
|
||||
struct GDNFwdHOffsets {
|
||||
uint32_t hSrcOffset;
|
||||
uint32_t hDstOffset;
|
||||
uint32_t uvOffset;
|
||||
uint32_t wkOffset;
|
||||
uint32_t wOffset;
|
||||
uint32_t gOffset;
|
||||
uint32_t hWorkOffset;
|
||||
uint32_t vWorkOffset;
|
||||
uint32_t initialStateOffset;
|
||||
uint32_t finalStateOffset;
|
||||
bool isInitialState;
|
||||
bool isFinalState;
|
||||
uint32_t blockTokens;
|
||||
bool isDummyHead;
|
||||
// for debug
|
||||
uint32_t batchIdx;
|
||||
uint32_t headIdx;
|
||||
uint32_t chunkIdx;
|
||||
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdH {
|
||||
uint32_t batch;
|
||||
uint32_t seqlen;
|
||||
uint32_t kNumHead;
|
||||
uint32_t vNumHead;
|
||||
uint32_t kHeadDim;
|
||||
uint32_t vHeadDim;
|
||||
uint32_t chunkSize;
|
||||
uint32_t initalStateStride0;
|
||||
uint32_t vBlockSize{128};
|
||||
uint32_t isVariedLen;
|
||||
uint32_t shapeBatch;
|
||||
uint32_t tokenBatch;
|
||||
bool useInitialState;
|
||||
bool storeFinalState;
|
||||
uint32_t numSeqWorkspaceOffset;
|
||||
uint32_t numChunksWorkspaceOffset;
|
||||
|
||||
uint32_t taskIdx;
|
||||
uint32_t taskLoops;
|
||||
uint32_t cubeCoreIdx;
|
||||
uint32_t cubeCoreNum;
|
||||
uint32_t vLoops;
|
||||
uint32_t taskNum;
|
||||
uint32_t headGroups;
|
||||
uint32_t totalChunks;
|
||||
uint32_t totalTokens;
|
||||
uint32_t headInnerLoop;
|
||||
|
||||
uint32_t iterId {0};
|
||||
bool hasDummyHead;
|
||||
bool isRunning;
|
||||
bool processNewTask {true};
|
||||
bool firstLoop {true};
|
||||
bool lastLoop {false};
|
||||
GDNFwdHOffsets offsets[PING_PONG_STAGES];
|
||||
int32_t currStage{PING_PONG_STAGES - 1};
|
||||
|
||||
uint32_t vIdx;
|
||||
uint32_t batchIdx;
|
||||
uint32_t baseHeadIdx;
|
||||
uint32_t chunkIdx;
|
||||
uint32_t headInnerIdx;
|
||||
uint32_t vHeadIdx;
|
||||
uint32_t kHeadIdx;
|
||||
uint32_t shapeBatchIdx;
|
||||
uint32_t tokenBatchIdx;
|
||||
|
||||
uint32_t chunkOffset;
|
||||
uint32_t tokenOffset;
|
||||
uint32_t batchChunks;
|
||||
uint32_t batchTokens;
|
||||
|
||||
AscendC::GlobalTensor<int64_t> gmSeqlen;
|
||||
AscendC::GlobalTensor<int64_t> gmNumSeq;
|
||||
AscendC::GlobalTensor<int64_t> gmNumChunks;
|
||||
|
||||
Arch::CrossCoreFlag cube1Done{0};
|
||||
Arch::CrossCoreFlag vec1Done{1};
|
||||
Arch::CrossCoreFlag cube2Done{2};
|
||||
Arch::CrossCoreFlag vec2Done{3};
|
||||
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdH() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user, uint32_t coreIdx, uint32_t coreNum) {
|
||||
__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling);
|
||||
|
||||
batch = gdnFwdHTilingData->batch;
|
||||
seqlen = gdnFwdHTilingData->seqlen;
|
||||
kNumHead = gdnFwdHTilingData->kNumHead;
|
||||
vNumHead = gdnFwdHTilingData->vNumHead;
|
||||
kHeadDim = gdnFwdHTilingData->kHeadDim;
|
||||
vHeadDim = gdnFwdHTilingData->vHeadDim;
|
||||
chunkSize = gdnFwdHTilingData->chunkSize;
|
||||
initalStateStride0 = gdnFwdHTilingData->initalStateStride0;
|
||||
isVariedLen = gdnFwdHTilingData->isVariedLen;
|
||||
shapeBatch = gdnFwdHTilingData->shapeBatch;
|
||||
tokenBatch = gdnFwdHTilingData->tokenBatch;
|
||||
useInitialState = gdnFwdHTilingData->useInitialState;
|
||||
storeFinalState = gdnFwdHTilingData->storeFinalState;
|
||||
numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset;
|
||||
numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset;
|
||||
|
||||
gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens);
|
||||
gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset));
|
||||
gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset));
|
||||
|
||||
if (isVariedLen) {
|
||||
gmNumChunks.SetValue(0, 0);
|
||||
gmNumSeq.SetValue(0, 0);
|
||||
uint32_t actualBatch = 0;
|
||||
int64_t prevSeq = 0, currSeq;
|
||||
for (uint32_t b = 1; b <= tokenBatch; b++) {
|
||||
currSeq = gmSeqlen.GetValue(b);
|
||||
int64_t batchSeqLen = currSeq - prevSeq;
|
||||
if (batchSeqLen > 0) {
|
||||
actualBatch++;
|
||||
gmNumSeq.SetValue(actualBatch, currSeq);
|
||||
int64_t batchChunk = (batchSeqLen + chunkSize - 1) / chunkSize;
|
||||
gmNumChunks.SetValue(actualBatch, gmNumChunks.GetValue(actualBatch - 1) + batchChunk);
|
||||
}
|
||||
prevSeq = currSeq;
|
||||
}
|
||||
tokenBatch = actualBatch;
|
||||
batch = actualBatch;
|
||||
totalChunks = gmNumChunks.GetValue(tokenBatch);
|
||||
totalTokens = gmNumSeq.GetValue(tokenBatch);
|
||||
} else {
|
||||
totalChunks = (seqlen + chunkSize - 1) / chunkSize;
|
||||
totalTokens = seqlen;
|
||||
}
|
||||
|
||||
cubeCoreIdx = coreIdx;
|
||||
cubeCoreNum = coreNum;
|
||||
vLoops = vHeadDim / vBlockSize;
|
||||
taskNum = vLoops * batch * vNumHead;
|
||||
headGroups = vNumHead / kNumHead;
|
||||
hasDummyHead = (taskNum % (PING_PONG_STAGES * cubeCoreNum) <= cubeCoreNum) && (taskNum % (PING_PONG_STAGES * cubeCoreNum) > 0);
|
||||
taskLoops = (taskNum + cubeCoreNum * PING_PONG_STAGES - 1) / (cubeCoreNum * PING_PONG_STAGES);
|
||||
headInnerLoop = taskNum > cubeCoreNum ? PING_PONG_STAGES : 1;
|
||||
taskIdx = cubeCoreIdx * headInnerLoop;
|
||||
isRunning = taskIdx < taskNum;
|
||||
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void InitTask() {
|
||||
iterId++;
|
||||
currStage = (currStage + 1) % PING_PONG_STAGES;
|
||||
if (processNewTask) {
|
||||
if (taskIdx >= taskNum) {
|
||||
lastLoop = true;
|
||||
isRunning = false;
|
||||
return;
|
||||
}
|
||||
vIdx = taskIdx / (batch * vNumHead);
|
||||
batchIdx = (taskIdx - vIdx * batch * vNumHead) / vNumHead;
|
||||
baseHeadIdx = taskIdx % vNumHead;
|
||||
shapeBatchIdx = isVariedLen ? 0 : batchIdx;
|
||||
tokenBatchIdx = isVariedLen ? batchIdx : 0;
|
||||
chunkOffset = isVariedLen ? gmNumChunks.GetValue(tokenBatchIdx) : 0;
|
||||
batchChunks = isVariedLen ? (gmNumChunks.GetValue(tokenBatchIdx + 1) - chunkOffset) : totalChunks;
|
||||
tokenOffset = isVariedLen ? gmNumSeq.GetValue(tokenBatchIdx) : 0;
|
||||
batchTokens = isVariedLen ? (gmNumSeq.GetValue(tokenBatchIdx + 1) - tokenOffset) : totalTokens;
|
||||
chunkIdx = 0;
|
||||
headInnerIdx = 0;
|
||||
} else {
|
||||
chunkIdx = headInnerIdx == PING_PONG_STAGES - 1 ? chunkIdx + 1 : chunkIdx;
|
||||
headInnerIdx = (headInnerIdx + 1) % PING_PONG_STAGES;
|
||||
}
|
||||
|
||||
vHeadIdx = baseHeadIdx + headInnerIdx;
|
||||
kHeadIdx = vHeadIdx / headGroups;
|
||||
offsets[currStage].isInitialState = chunkIdx == 0;
|
||||
offsets[currStage].isFinalState = chunkIdx == (batchChunks - 1);
|
||||
offsets[currStage].initialStateOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * initalStateStride0;
|
||||
offsets[currStage].finalStateOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * vHeadDim;
|
||||
offsets[currStage].hSrcOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset + chunkIdx) * kHeadDim * vHeadDim;
|
||||
offsets[currStage].hDstOffset = offsets[currStage].hSrcOffset + kHeadDim * vHeadDim;
|
||||
offsets[currStage].uvOffset = (shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * vHeadDim;
|
||||
offsets[currStage].wkOffset = (shapeBatchIdx * kNumHead * totalTokens + kHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * kHeadDim;
|
||||
offsets[currStage].wOffset = (shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize) * kHeadDim;
|
||||
offsets[currStage].gOffset = shapeBatchIdx * vNumHead * totalTokens + vHeadIdx * totalTokens + tokenOffset + chunkIdx * chunkSize;
|
||||
offsets[currStage].hWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * kHeadDim * vHeadDim;
|
||||
offsets[currStage].vWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * vHeadDim;
|
||||
offsets[currStage].blockTokens = offsets[currStage].isFinalState ? (batchTokens - chunkIdx * chunkSize) : chunkSize;
|
||||
offsets[currStage].isDummyHead = headInnerLoop < PING_PONG_STAGES && headInnerIdx >= headInnerLoop;
|
||||
offsets[currStage].batchIdx = batchIdx;
|
||||
offsets[currStage].headIdx = vHeadIdx;
|
||||
offsets[currStage].chunkIdx = chunkIdx;
|
||||
|
||||
processNewTask = chunkIdx == batchChunks - 1 && headInnerIdx == PING_PONG_STAGES - 1;
|
||||
if (processNewTask) {
|
||||
uint32_t currLoopIdx = taskIdx / (PING_PONG_STAGES * cubeCoreNum);
|
||||
headInnerLoop = ((currLoopIdx + 2 == taskLoops) && hasDummyHead) ? 1 : PING_PONG_STAGES;
|
||||
taskIdx = (currLoopIdx + 1) * PING_PONG_STAGES * cubeCoreNum + headInnerLoop * cubeCoreIdx;
|
||||
}
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GDNFwdHOffsets& GetStage1Offsets() {
|
||||
return offsets[currStage];
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
bool NeedProcessStage1() {
|
||||
GDNFwdHOffsets& stage1Offsets = GetStage1Offsets();
|
||||
return !(lastLoop || stage1Offsets.isDummyHead);
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
GDNFwdHOffsets& GetStage2Offsets() {
|
||||
return offsets[(currStage - 1) % PING_PONG_STAGES];
|
||||
}
|
||||
|
||||
CATLASS_DEVICE
|
||||
bool NeedProcessStage2() {
|
||||
GDNFwdHOffsets& stage2Offsets = GetStage2Offsets();
|
||||
return !(iterId == 1 || (!storeFinalState && stage2Offsets.isFinalState) || stage2Offsets.isDummyHead);
|
||||
}
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdHCube : public BlockSchedulerGdnFwdH {
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdHCube() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) {
|
||||
BlockSchedulerGdnFwdH::Init(cu_seqlens, chunk_indices, tiling, user, AscendC::GetBlockIdx(), AscendC::GetBlockNum());
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
struct BlockSchedulerGdnFwdHVec : public BlockSchedulerGdnFwdH {
|
||||
CATLASS_DEVICE
|
||||
BlockSchedulerGdnFwdHVec() {}
|
||||
|
||||
CATLASS_DEVICE
|
||||
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_indices, GM_ADDR tiling, GM_ADDR user) {
|
||||
BlockSchedulerGdnFwdH::Init(cu_seqlens, chunk_indices, tiling, user, AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), AscendC::GetBlockNum());
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
} // namespace Catlass::Gemm::Block
|
||||
|
||||
#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_H_HPP
|
||||
@@ -0,0 +1,408 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310
|
||||
#define CATLASS_ARCH 3510
|
||||
|
||||
#include "catlass/arch/arch.hpp"
|
||||
#include "catlass/arch/cross_core_sync.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/debug.hpp"
|
||||
#include "catlass/epilogue/block/block_epilogue.hpp"
|
||||
#include "../../epilogue/block/block_epilogue_gdn_fwdh_update.hpp"
|
||||
#include "../../epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp"
|
||||
#include "catlass/gemm/block/block_mmad.hpp"
|
||||
#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp"
|
||||
#include "catlass/gemm/block/block_swizzle.hpp"
|
||||
#include "../block/block_scheduler_gdn_fwd_h.hpp"
|
||||
#include "catlass/gemm/dispatch_policy.hpp"
|
||||
#include "catlass/gemm/gemm_type.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "tla/tensor.hpp"
|
||||
#include "tla/layout.hpp"
|
||||
#include "tla/tensor.hpp"
|
||||
|
||||
using _0 = tla::Int<0>;
|
||||
using _1 = tla::Int<1>;
|
||||
using _2 = tla::Int<2>;
|
||||
using _4 = tla::Int<4>;
|
||||
using _8 = tla::Int<8>;
|
||||
using _16 = tla::Int<16>;
|
||||
using _32 = tla::Int<32>;
|
||||
using _64 = tla::Int<64>;
|
||||
using _128 = tla::Int<128>;
|
||||
using _256 = tla::Int<256>;
|
||||
using _512 = tla::Int<512>;
|
||||
using _1024 = tla::Int<1024>;
|
||||
using _2048 = tla::Int<2048>;
|
||||
using _4096 = tla::Int<4096>;
|
||||
using _8192 = tla::Int<8192>;
|
||||
using _16384 = tla::Int<16384>;
|
||||
using _32768 = tla::Int<32768>;
|
||||
using _65536 = tla::Int<65536>;
|
||||
|
||||
#else
|
||||
#define CATLASS_ARCH 2201
|
||||
|
||||
#include "catlass/arch/arch.hpp"
|
||||
#include "catlass/arch/cross_core_sync.hpp"
|
||||
#include "catlass/arch/resource.hpp"
|
||||
#include "catlass/catlass.hpp"
|
||||
#include "catlass/debug.hpp"
|
||||
#include "catlass/epilogue/block/block_epilogue.hpp"
|
||||
#include "../../epilogue/block/block_epilogue_gdn_fwdh_update.hpp"
|
||||
#include "../../epilogue/block/block_epilogue_gdn_fwdh_vnew.hpp"
|
||||
#include "catlass/gemm/block/block_mmad.hpp"
|
||||
#include "kernel_utils/block/block_mmad_pingpong_tla_multi.hpp"
|
||||
#include "catlass/gemm/block/block_swizzle.hpp"
|
||||
#include "../block/block_scheduler_gdn_fwd_h.hpp"
|
||||
#include "catlass/gemm/dispatch_policy.hpp"
|
||||
#include "catlass/gemm/gemm_type.hpp"
|
||||
#include "catlass/layout/layout.hpp"
|
||||
#include "catlass/gemm_coord.hpp"
|
||||
#include "tla/tensor.hpp"
|
||||
#include "tla/layout.hpp"
|
||||
#include "tla/tensor.hpp"
|
||||
#endif
|
||||
|
||||
|
||||
|
||||
#include "kernel_operator.h"
|
||||
using namespace Catlass;
|
||||
using namespace tla;
|
||||
|
||||
namespace Catlass::Gemm::Kernel {
|
||||
|
||||
template<
|
||||
typename INPUT_TYPE,
|
||||
typename G_TYPE,
|
||||
typename STATE_TYPE,
|
||||
typename WORKSPACE_TYPE
|
||||
>
|
||||
class GDNFwdHKernel {
|
||||
public:
|
||||
|
||||
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310
|
||||
using ArchTag = Arch::Ascend950;
|
||||
#else
|
||||
using ArchTag = Arch::AtlasA2;
|
||||
#endif
|
||||
using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHCube;
|
||||
using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdHVec;
|
||||
|
||||
using DispatchPolicyTla = Gemm::MmadPingpongTlaMulti<ArchTag, true, false>;
|
||||
using L1TileShapeTla = Shape<_128, _128, _128>;
|
||||
using L0TileShapeTla = L1TileShapeTla;
|
||||
|
||||
using WType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using HType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using VworkType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using KType = Gemm::GemmType<INPUT_TYPE, layout::ColumnMajor>;
|
||||
using HworkType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using VType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using GType = Gemm::GemmType<G_TYPE, layout::RowMajor>;
|
||||
using UType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
|
||||
using FinalStateType = Gemm::GemmType<STATE_TYPE, layout::RowMajor>;
|
||||
|
||||
// cube 1
|
||||
using TileCopyWH = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using BlockMmadWH = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyWH>;
|
||||
|
||||
// cube 2
|
||||
using TileCopyKV = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::ColumnMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
|
||||
using BlockMmadKV = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyKV>;
|
||||
|
||||
// vec 1
|
||||
using DispatchPolicyGDNFwdHVnew = Epilogue::EpilogueAtlasGDNFwdHVnew;
|
||||
using EpilogueGDNFwdHVnew = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdHVnew, VType, GType, UType, VworkType>;
|
||||
|
||||
// vec 2
|
||||
using DispatchPolicyGDNFwdHUpdate = Epilogue::EpilogueAtlasGDNFwdHUpdate;
|
||||
using EpilogueGDNFwdHUpdate = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdHUpdate, HType, GType, HType, HworkType, FinalStateType>;
|
||||
|
||||
using GDNFwdHOffsets = Catlass::Gemm::Block::GDNFwdHOffsets;
|
||||
|
||||
using ElementK = INPUT_TYPE;
|
||||
using ElementW = INPUT_TYPE;
|
||||
using ElementU = INPUT_TYPE;
|
||||
using ElementG = G_TYPE;
|
||||
using ElementH = INPUT_TYPE;
|
||||
using ElementV = INPUT_TYPE;
|
||||
using ElementVWork = WORKSPACE_TYPE;
|
||||
using ElementHWork = WORKSPACE_TYPE;
|
||||
using ElementInitialState = STATE_TYPE;
|
||||
using ElementFinalState = STATE_TYPE;
|
||||
|
||||
using LayoutW = Catlass::layout::RowMajor;
|
||||
using LayoutH = Catlass::layout::RowMajor;
|
||||
using LayoutV = Catlass::layout::RowMajor;
|
||||
using LayoutK = Catlass::layout::ColumnMajor;
|
||||
|
||||
|
||||
uint32_t batch;
|
||||
uint32_t seqlen;
|
||||
uint32_t kNumHead;
|
||||
uint32_t vNumHead;
|
||||
uint32_t kHeadDim;
|
||||
uint32_t vHeadDim;
|
||||
uint32_t chunkSize;
|
||||
uint32_t initalStateStride0;
|
||||
bool useInitialState;
|
||||
bool storeFinalState;
|
||||
uint32_t isVariedLen;
|
||||
uint32_t shapeBatch;
|
||||
uint32_t tokenBatch;
|
||||
uint32_t vWorkspaceOffset;
|
||||
uint32_t vUpdateWorkspaceOffset;
|
||||
uint32_t hWorkspaceOffset;
|
||||
uint32_t numSeqWorkspaceOffset;
|
||||
uint32_t numChunksWorkspaceOffset;
|
||||
|
||||
AscendC::GlobalTensor<ElementK> gmK;
|
||||
AscendC::GlobalTensor<ElementW> gmW;
|
||||
AscendC::GlobalTensor<ElementU> gmU;
|
||||
AscendC::GlobalTensor<ElementG> gmG;
|
||||
AscendC::GlobalTensor<ElementInitialState> gmInitialState;
|
||||
AscendC::GlobalTensor<ElementH> gmH;
|
||||
AscendC::GlobalTensor<ElementV> gmV;
|
||||
AscendC::GlobalTensor<ElementFinalState> gmFinalState;
|
||||
AscendC::GlobalTensor<ElementVWork> gmVWorkspace;
|
||||
AscendC::GlobalTensor<ElementV> gmVUpdateWorkspace;
|
||||
AscendC::GlobalTensor<ElementHWork> gmHWorkspace;
|
||||
|
||||
AscendC::GlobalTensor<int64_t> gmSeqlen;
|
||||
AscendC::GlobalTensor<int64_t> gmNumSeq;
|
||||
AscendC::GlobalTensor<int64_t> gmNumChunks;
|
||||
|
||||
CubeScheduler cubeBlockScheduler;
|
||||
VecScheduler vecBlockScheduler;
|
||||
|
||||
Arch::Resource<ArchTag> resource;
|
||||
|
||||
|
||||
__aicore__ inline GDNFwdHKernel() {}
|
||||
|
||||
__aicore__ inline void Init(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g, GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices,
|
||||
GM_ADDR h, GM_ADDR v_new, GM_ADDR final_state, GM_ADDR tiling, GM_ADDR user) {
|
||||
|
||||
__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling);
|
||||
|
||||
batch = gdnFwdHTilingData->batch;
|
||||
seqlen = gdnFwdHTilingData->seqlen;
|
||||
kNumHead = gdnFwdHTilingData->kNumHead;
|
||||
vNumHead = gdnFwdHTilingData->vNumHead;
|
||||
kHeadDim = gdnFwdHTilingData->kHeadDim;
|
||||
vHeadDim = gdnFwdHTilingData->vHeadDim;
|
||||
chunkSize = gdnFwdHTilingData->chunkSize;
|
||||
initalStateStride0 = gdnFwdHTilingData->initalStateStride0;
|
||||
useInitialState = gdnFwdHTilingData->useInitialState;
|
||||
storeFinalState = gdnFwdHTilingData->storeFinalState;
|
||||
isVariedLen = gdnFwdHTilingData->isVariedLen;
|
||||
shapeBatch = gdnFwdHTilingData->shapeBatch;
|
||||
tokenBatch = gdnFwdHTilingData->tokenBatch;
|
||||
vWorkspaceOffset = gdnFwdHTilingData->vWorkspaceOffset;
|
||||
vUpdateWorkspaceOffset = gdnFwdHTilingData->vUpdateWorkspaceOffset;
|
||||
hWorkspaceOffset = gdnFwdHTilingData->hWorkspaceOffset;
|
||||
numSeqWorkspaceOffset = gdnFwdHTilingData->numSeqWorkspaceOffset;
|
||||
numChunksWorkspaceOffset = gdnFwdHTilingData->numChunksWorkspaceOffset;
|
||||
|
||||
gmK.SetGlobalBuffer((__gm__ ElementK *)k);
|
||||
gmW.SetGlobalBuffer((__gm__ ElementW *)w);
|
||||
gmU.SetGlobalBuffer((__gm__ ElementU *)u);
|
||||
gmG.SetGlobalBuffer((__gm__ ElementG *)g);
|
||||
gmInitialState.SetGlobalBuffer((__gm__ ElementInitialState *)inital_state);
|
||||
gmH.SetGlobalBuffer((__gm__ ElementH *)h);
|
||||
gmV.SetGlobalBuffer((__gm__ ElementV *)v_new);
|
||||
gmFinalState.SetGlobalBuffer((__gm__ ElementFinalState *)final_state);
|
||||
gmVWorkspace.SetGlobalBuffer((__gm__ ElementVWork *)(user + vWorkspaceOffset));
|
||||
gmVUpdateWorkspace.SetGlobalBuffer((__gm__ ElementV *)(user + vUpdateWorkspaceOffset));
|
||||
gmHWorkspace.SetGlobalBuffer((__gm__ ElementHWork *)(user + hWorkspaceOffset));
|
||||
|
||||
gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens);
|
||||
gmNumSeq.SetGlobalBuffer((__gm__ int64_t *)(user + numSeqWorkspaceOffset));
|
||||
gmNumChunks.SetGlobalBuffer((__gm__ int64_t *)(user + numChunksWorkspaceOffset));
|
||||
|
||||
if ASCEND_IS_AIC {
|
||||
cubeBlockScheduler.Init(cu_seqlens, chunk_indices, tiling, user);
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
vecBlockScheduler.Init(cu_seqlens, chunk_indices, tiling, user);
|
||||
}
|
||||
}
|
||||
|
||||
__aicore__ inline void Process() {
|
||||
|
||||
if ASCEND_IS_AIC {
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
|
||||
BlockMmadWH blockMmadWH(resource);
|
||||
BlockMmadKV blockMmadKV(resource);
|
||||
|
||||
auto wLayout = tla::MakeLayout<ElementW, LayoutW>(shapeBatch * kNumHead * cubeBlockScheduler.totalTokens, kHeadDim);
|
||||
auto hLayout = tla::MakeLayout<ElementH, LayoutH>(shapeBatch * vNumHead * cubeBlockScheduler.totalChunks * kHeadDim, vHeadDim);
|
||||
auto vLayout = tla::MakeLayout<ElementVWork, LayoutV>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
|
||||
|
||||
auto kLayout = tla::MakeLayout<ElementK, LayoutK>(kHeadDim, shapeBatch * kNumHead * cubeBlockScheduler.totalTokens);
|
||||
auto vworkLayout = tla::MakeLayout<ElementV, LayoutV>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
|
||||
auto hworkLayout = tla::MakeLayout<ElementHWork, LayoutH>(coreNum * kHeadDim * PING_PONG_STAGES, vHeadDim);
|
||||
|
||||
while (cubeBlockScheduler.isRunning) {
|
||||
cubeBlockScheduler.InitTask();
|
||||
// step 1: v_work = w @ h[i]
|
||||
GDNFwdHOffsets& cube1Offsets = cubeBlockScheduler.GetStage1Offsets();
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done);
|
||||
if (cubeBlockScheduler.NeedProcessStage1()) {
|
||||
int64_t cube1OffsetW = cube1Offsets.wOffset;
|
||||
int64_t cube1OffsetH = cube1Offsets.hSrcOffset;
|
||||
int64_t cube1OffsetVwork = cube1Offsets.vWorkOffset;
|
||||
auto tensorW = tla::MakeTensor(gmW[cube1OffsetW], wLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorH = tla::MakeTensor(gmH[cube1OffsetH], hLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorV = tla::MakeTensor(gmVWorkspace[cube1OffsetVwork], vLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube1Shape {cube1Offsets.blockTokens, vHeadDim, kHeadDim};
|
||||
auto tensorBlockW = GetTile(tensorW, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k()));
|
||||
auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n()));
|
||||
auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n()));
|
||||
blockMmadWH.preSetFlags();
|
||||
blockMmadWH(tensorBlockW, tensorBlockH, tensorBlockV, cube1Shape);
|
||||
blockMmadWH.finalWaitFlags();
|
||||
}
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done);
|
||||
|
||||
if (cubeBlockScheduler.iterId > 1) {
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
|
||||
GDNFwdHOffsets& cube2Offsets = cubeBlockScheduler.GetStage2Offsets();
|
||||
if (cubeBlockScheduler.NeedProcessStage2()) {
|
||||
// step 3: h[i+1] = k.T @ v_work
|
||||
int64_t cube2OffsetK = cube2Offsets.wkOffset;
|
||||
int64_t cube2OffsetVwork = cube2Offsets.vWorkOffset;
|
||||
int64_t cube2OffsetH = cube2Offsets.hWorkOffset;
|
||||
auto tensorK = tla::MakeTensor(gmK[cube2OffsetK], kLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorVwork = tla::MakeTensor(gmVUpdateWorkspace[cube2OffsetVwork], vworkLayout, Catlass::Arch::PositionGM{});
|
||||
auto tensorHwork = tla::MakeTensor(gmHWorkspace[cube2OffsetH], hworkLayout, Catlass::Arch::PositionGM{});
|
||||
GemmCoord cube2Shape{kHeadDim, vHeadDim, cube2Offsets.blockTokens};
|
||||
auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k()));
|
||||
auto tensorBlockVwork = GetTile(tensorVwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n()));
|
||||
auto tensorBlockHwork = GetTile(tensorHwork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n()));
|
||||
blockMmadKV.preSetFlags();
|
||||
blockMmadKV(tensorBlockK, tensorBlockVwork, tensorBlockHwork, cube2Shape);
|
||||
blockMmadKV.finalWaitFlags();
|
||||
}
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done);
|
||||
}
|
||||
}
|
||||
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done);
|
||||
|
||||
}
|
||||
|
||||
if ASCEND_IS_AIV {
|
||||
uint32_t coreIdx = AscendC::GetBlockIdx();
|
||||
uint32_t coreNum = AscendC::GetBlockNum();
|
||||
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
|
||||
uint32_t subBlockNum = AscendC::GetSubBlockNum();
|
||||
|
||||
EpilogueGDNFwdHVnew epilogueGDNFwdHVnew(resource);
|
||||
|
||||
if (useInitialState) {
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensorPing = resource.ubBuf.template GetBufferByByte<ElementInitialState>(0);
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensorPong = resource.ubBuf.template GetBufferByByte<ElementInitialState>(96 * 1024);
|
||||
AscendC::LocalTensor<ElementH> hUbTensorPing = resource.ubBuf.template GetBufferByByte<ElementH>(64 * 1024);
|
||||
AscendC::LocalTensor<ElementH> hUbTensorPong = resource.ubBuf.template GetBufferByByte<ElementH>(160 * 1024);
|
||||
uint32_t totalChunks = isVariedLen ? vecBlockScheduler.totalChunks : ((seqlen + chunkSize - 1) / chunkSize);
|
||||
uint32_t stateBlockSize = kHeadDim * vHeadDim;
|
||||
uint32_t pingpongFlag = 1;
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID1);
|
||||
AscendC::DataCopyParams repeatParams = {static_cast<uint16_t>(kHeadDim), static_cast<uint16_t>(vHeadDim * sizeof(ElementInitialState) / 32),
|
||||
static_cast<uint16_t>((initalStateStride0 - vHeadDim)* sizeof(ElementInitialState) / 32), static_cast<uint16_t>(0)};
|
||||
for (uint32_t shapeBatchIdx = 0; shapeBatchIdx < shapeBatch; shapeBatchIdx++) {
|
||||
for (uint32_t vHeadIdx = 0; vHeadIdx < vNumHead; vHeadIdx++) {
|
||||
for (uint32_t tokenBatchIdx = 0; tokenBatchIdx < vecBlockScheduler.tokenBatch; tokenBatchIdx++) {
|
||||
uint32_t batchIdx = isVariedLen ? tokenBatchIdx : shapeBatchIdx;
|
||||
uint32_t chunkOffset = isVariedLen ? gmNumChunks.GetValue(tokenBatchIdx) : 0;
|
||||
uint32_t initialStateSrcOffset = (batchIdx * vNumHead + vHeadIdx) * kHeadDim * initalStateStride0;
|
||||
uint32_t hOffset = (shapeBatchIdx * vNumHead * totalChunks + vHeadIdx * totalChunks + chunkOffset) * stateBlockSize;
|
||||
AscendC::LocalTensor<ElementInitialState> stateUbTensor = pingpongFlag ? stateUbTensorPing : stateUbTensorPong;
|
||||
AscendC::LocalTensor<ElementH> hUbTensor = pingpongFlag ? hUbTensorPing : hUbTensorPong;
|
||||
auto event_id = pingpongFlag ? EVENT_ID1 : EVENT_ID0;
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(event_id);
|
||||
if constexpr(!std::is_same<ElementInitialState, ElementH>::value) {
|
||||
AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateSrcOffset], repeatParams);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(event_id);
|
||||
AscendC::Cast(hUbTensor, stateUbTensor, AscendC::RoundMode::CAST_RINT, stateBlockSize);
|
||||
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(event_id);
|
||||
AscendC::DataCopy(gmH[hOffset], hUbTensor, stateBlockSize);
|
||||
} else {
|
||||
AscendC::DataCopy(stateUbTensor, gmInitialState[initialStateSrcOffset], repeatParams);
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE2_MTE3>(event_id);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE2_MTE3>(event_id);
|
||||
AscendC::DataCopy(gmH[hOffset], stateUbTensor, stateBlockSize);
|
||||
}
|
||||
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(event_id);
|
||||
pingpongFlag = 1 - pingpongFlag;
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
|
||||
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID1);
|
||||
}
|
||||
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
|
||||
while (vecBlockScheduler.isRunning) {
|
||||
vecBlockScheduler.InitTask();
|
||||
// step 2:
|
||||
GDNFwdHOffsets& vec1Offsets = vecBlockScheduler.GetStage1Offsets();
|
||||
// gmV = gmU - gmVWorkspace
|
||||
// g_buf = gmG[-1] - gmG
|
||||
// g_buf = exp(g_buf)
|
||||
// gmVWorkspace = g_buf * gmV
|
||||
if (vecBlockScheduler.NeedProcessStage1()) {
|
||||
epilogueGDNFwdHVnew(
|
||||
gmV[vec1Offsets.uvOffset], gmVUpdateWorkspace[vec1Offsets.vWorkOffset],
|
||||
gmG[vec1Offsets.gOffset], gmU[vec1Offsets.uvOffset], gmVWorkspace[vec1Offsets.vWorkOffset],
|
||||
vec1Offsets.blockTokens, kHeadDim, vHeadDim, vecBlockScheduler.cube1Done
|
||||
);
|
||||
} else {
|
||||
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube1Done);
|
||||
}
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
|
||||
|
||||
if (vecBlockScheduler.iterId > 1) {
|
||||
GDNFwdHOffsets& vec2Offsets = vecBlockScheduler.GetStage2Offsets();
|
||||
if (vecBlockScheduler.NeedProcessStage2()) {
|
||||
// step 4: h[i+1] += h_work if i < num_chunks - 1 else None
|
||||
EpilogueGDNFwdHUpdate epilogueGDNFwdHUpdate(resource);
|
||||
epilogueGDNFwdHUpdate(
|
||||
gmH[vec2Offsets.hDstOffset], gmFinalState[vec2Offsets.finalStateOffset],
|
||||
gmG[vec2Offsets.gOffset],
|
||||
gmH[vec2Offsets.hSrcOffset],
|
||||
gmHWorkspace[vec2Offsets.hWorkOffset],
|
||||
vec2Offsets.blockTokens, kHeadDim, vHeadDim, vecBlockScheduler.cube2Done,
|
||||
(vec2Offsets.isFinalState && storeFinalState)
|
||||
);
|
||||
} else {
|
||||
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done);
|
||||
}
|
||||
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
}
|
||||
@@ -0,0 +1,95 @@
|
||||
/**
|
||||
* Copyright (c) 2026 Tianjin University, Ltd.
|
||||
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
* the BSD 3-Clause License (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.
|
||||
*/
|
||||
|
||||
/*!
|
||||
* \file chunk_gated_delta_rule_fwd_h.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
// #include "chunk_gated_delta_rule_fwd_h.h"
|
||||
#if defined(__CCE_AICORE__) && (__CCE_AICORE__ == 200)
|
||||
#include "arch20/compat_310p.h"
|
||||
#include "arch20/gemm/kernel/gdn_fwd_h_kernel.hpp"
|
||||
#else
|
||||
#include "arch22/gemm/kernel/gdn_fwd_h_kernel.hpp"
|
||||
#endif
|
||||
#include "lib/matmul_intf.h"
|
||||
|
||||
using namespace Catlass;
|
||||
|
||||
extern "C" __global__ __aicore__ void chunk_gated_delta_rule_fwd_h(GM_ADDR k, GM_ADDR w, GM_ADDR u, GM_ADDR g,
|
||||
GM_ADDR inital_state, GM_ADDR cu_seqlens, GM_ADDR chunk_indices,
|
||||
GM_ADDR h, GM_ADDR v_new, GM_ADDR final_state,
|
||||
GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2);
|
||||
|
||||
GM_ADDR user = AscendC::GetUserWorkspace(workspace);
|
||||
|
||||
__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict gdnFwdHTilingData = reinterpret_cast<__gm__ ChunkGatedDeltaRuleFwdHTilingData *__restrict>(tiling);
|
||||
|
||||
using workspaceType = float;
|
||||
// dtype: 0 - fp16, 1 - bf16, 2 - fp32
|
||||
#ifndef CATLASS_UNIFIED_CORE
|
||||
if (gdnFwdHTilingData->dataType == 1) {
|
||||
if (gdnFwdHTilingData->stateDataType == 2) {
|
||||
if (gdnFwdHTilingData->gDataType == 2) {
|
||||
using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel<bfloat16_t, float, float, workspaceType>;
|
||||
GDNFwdHKernel gdnFwdH;
|
||||
gdnFwdH.Init(k, w, u, g, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user);
|
||||
gdnFwdH.Process();
|
||||
} else {
|
||||
using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel<bfloat16_t, bfloat16_t, float, workspaceType>;
|
||||
GDNFwdHKernel gdnFwdH;
|
||||
gdnFwdH.Init(k, w, u, g, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user);
|
||||
gdnFwdH.Process();
|
||||
}
|
||||
} else {
|
||||
if (gdnFwdHTilingData->gDataType == 2) {
|
||||
using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel<bfloat16_t, float, bfloat16_t, workspaceType>;
|
||||
GDNFwdHKernel gdnFwdH;
|
||||
gdnFwdH.Init(k, w, u, g, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user);
|
||||
gdnFwdH.Process();
|
||||
} else {
|
||||
using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel<bfloat16_t, bfloat16_t, bfloat16_t, workspaceType>;
|
||||
GDNFwdHKernel gdnFwdH;
|
||||
gdnFwdH.Init(k, w, u, g, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user);
|
||||
gdnFwdH.Process();
|
||||
}
|
||||
}
|
||||
} else
|
||||
#endif
|
||||
{
|
||||
if (gdnFwdHTilingData->stateDataType == 2) {
|
||||
if (gdnFwdHTilingData->gDataType == 2) {
|
||||
using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel<half, float, float, workspaceType>;
|
||||
GDNFwdHKernel gdnFwdH;
|
||||
gdnFwdH.Init(k, w, u, g, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user);
|
||||
gdnFwdH.Process();
|
||||
} else {
|
||||
using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel<half, half, float, workspaceType>;
|
||||
GDNFwdHKernel gdnFwdH;
|
||||
gdnFwdH.Init(k, w, u, g, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user);
|
||||
gdnFwdH.Process();
|
||||
}
|
||||
} else {
|
||||
if (gdnFwdHTilingData->gDataType == 2) {
|
||||
using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel<half, float, half, workspaceType>;
|
||||
GDNFwdHKernel gdnFwdH;
|
||||
gdnFwdH.Init(k, w, u, g, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user);
|
||||
gdnFwdH.Process();
|
||||
} else {
|
||||
using GDNFwdHKernel = Catlass::Gemm::Kernel::GDNFwdHKernel<half, half, half, workspaceType>;
|
||||
GDNFwdHKernel gdnFwdH;
|
||||
gdnFwdH.Init(k, w, u, g, inital_state, cu_seqlens, chunk_indices, h, v_new, final_state, tiling, user);
|
||||
gdnFwdH.Process();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
/**
|
||||
* 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 data_copy_transpose_tiling.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <vector>
|
||||
#include <graph/tensor.h>
|
||||
#include "data_copy_transpose_tiling_def.h"
|
||||
|
||||
namespace optiling {
|
||||
|
||||
inline void GetDataCopyTransposeTiling(const ge::Shape &dstShape, const ge::Shape &srcShape, const uint32_t typeSize,
|
||||
optiling::CopyTransposeTiling &tiling)
|
||||
{
|
||||
constexpr int64_t B_INDEX = 0;
|
||||
constexpr int64_t N_INDEX = 1;
|
||||
constexpr int64_t S_INDEX = 2;
|
||||
constexpr int64_t H_INDEX = 3;
|
||||
std::vector<int64_t> dstShapeInfo = dstShape.GetDims();
|
||||
std::vector<int64_t> srcShapeInfo = srcShape.GetDims();
|
||||
|
||||
tiling.set_dstShapeB(dstShapeInfo[B_INDEX]);
|
||||
tiling.set_dstShapeN(dstShapeInfo[N_INDEX]);
|
||||
tiling.set_dstShapeS(dstShapeInfo[S_INDEX]);
|
||||
tiling.set_dstShapeH(dstShapeInfo[H_INDEX]);
|
||||
tiling.set_dstShapeHN(tiling.get_dstShapeH() / tiling.get_dstShapeN());
|
||||
|
||||
tiling.set_srcShapeB(srcShapeInfo[B_INDEX]);
|
||||
tiling.set_srcShapeN(srcShapeInfo[N_INDEX]);
|
||||
tiling.set_srcShapeS(srcShapeInfo[S_INDEX]);
|
||||
tiling.set_srcShapeHN(srcShapeInfo[H_INDEX]);
|
||||
tiling.set_originalShapeNLen(tiling.get_srcShapeHN() * typeSize);
|
||||
tiling.set_shapeSHValue(tiling.get_dstShapeS() * tiling.get_dstShapeH());
|
||||
tiling.set_shapeNsValue(tiling.get_dstShapeN() * tiling.get_dstShapeS());
|
||||
tiling.set_shapeNsnValue(tiling.get_dstShapeN() * tiling.get_srcShapeS() * tiling.get_srcShapeN());
|
||||
tiling.set_shapeBHValue(tiling.get_dstShapeB() * tiling.get_dstShapeH());
|
||||
}
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,43 @@
|
||||
/**
|
||||
* 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 data_copy_transpose_tiling_def.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
#include <register/tilingdata_base.h>
|
||||
|
||||
namespace optiling {
|
||||
|
||||
BEGIN_TILING_DATA_DEF(CopyTransposeTiling)
|
||||
TILING_DATA_FIELD_DEF(uint32_t, dstShapeB);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, dstShapeN);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, dstShapeS);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, dstShapeHN);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, dstShapeH);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, srcShapeB);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, srcShapeN);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, srcShapeS);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, srcShapeHN);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, originalShapeNLen);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, shapeSHValue);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, shapeNsValue);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, shapeNsnValue);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, invalidParamCopyTransposeTiling);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, shapeBHValue);
|
||||
TILING_DATA_FIELD_DEF(uint32_t, paramsAlign);
|
||||
END_TILING_DATA_DEF;
|
||||
REGISTER_TILING_DATA_CLASS(CopyTransposeTilingOp, CopyTransposeTiling)
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,63 @@
|
||||
#ifndef OPS_BUILT_IN_OP_TILING_ERROR_LOG_H_
|
||||
#define OPS_BUILT_IN_OP_TILING_ERROR_LOG_H_
|
||||
|
||||
#include <cstdio>
|
||||
#include <string>
|
||||
#include "toolchain/slog.h"
|
||||
|
||||
#define OP_LOGI(opname, ...)
|
||||
#define OP_LOGW(opname, ...) \
|
||||
do { \
|
||||
(void)(opname); \
|
||||
std::printf("[WARN] "); \
|
||||
std::printf(__VA_ARGS__); \
|
||||
std::printf("\n"); \
|
||||
} while (0)
|
||||
|
||||
#define OP_LOGE_WITHOUT_REPORT(opname, ...) \
|
||||
do { \
|
||||
(void)(opname); \
|
||||
std::printf("[ERRORx] "); \
|
||||
std::printf(__VA_ARGS__); \
|
||||
std::printf("\n"); \
|
||||
} while (0)
|
||||
|
||||
#define OP_LOGE(opname, ...) \
|
||||
do { \
|
||||
(void)(opname); \
|
||||
std::printf("[ERROR] "); \
|
||||
std::printf(__VA_ARGS__); \
|
||||
std::printf("\n"); \
|
||||
} while (0)
|
||||
|
||||
#define OP_LOGD(opname, ...)
|
||||
|
||||
namespace optiling {
|
||||
|
||||
#define VECTOR_INNER_ERR_REPORT_TILIING(op_name, err_msg, ...) \
|
||||
do { \
|
||||
OP_LOGE_WITHOUT_REPORT(op_name, err_msg, ##__VA_ARGS__); \
|
||||
} while (0)
|
||||
|
||||
// Modify OP_TILING_CHECK macro to ensure proper handling of expressions
|
||||
#define OP_CHECK_IF(cond, log_func, expr) \
|
||||
do { \
|
||||
if (cond) { \
|
||||
log_func; \
|
||||
expr; \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
|
||||
|
||||
#define OP_CHECK_NULL_WITH_CONTEXT(context, ptr) \
|
||||
do { \
|
||||
if ((ptr) == nullptr) { \
|
||||
OP_LOGE(context->GetNodeType(), "%s is null", #ptr); \
|
||||
return ge::GRAPH_FAILED; \
|
||||
} \
|
||||
} while (0)
|
||||
|
||||
} // namespace optiling
|
||||
|
||||
#endif // OPS_BUILT_IN_OP_TILING_ERROR_LOG_H_
|
||||
256
csrc/moe/chunk_gated_delta_rule_fwd_h/tiling_base/tiling_base.h
Normal file
256
csrc/moe/chunk_gated_delta_rule_fwd_h/tiling_base/tiling_base.h
Normal file
@@ -0,0 +1,256 @@
|
||||
/**
|
||||
* 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 tiling_base.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <sstream>
|
||||
#include <exe_graph/runtime/tiling_context.h>
|
||||
#include <graph/utils/type_utils.h>
|
||||
#include "tiling/platform/platform_ascendc.h"
|
||||
#include "error_log.h"
|
||||
|
||||
#ifdef ASCENDC_OP_TEST
|
||||
#define ASCENDC_EXTERN_C extern "C"
|
||||
#else
|
||||
#define ASCENDC_EXTERN_C
|
||||
#endif
|
||||
|
||||
namespace Ops {
|
||||
namespace Transformer {
|
||||
namespace OpTiling {
|
||||
|
||||
struct AiCoreParams {
|
||||
uint64_t ubSize = 0;
|
||||
uint64_t blockDim = 0;
|
||||
uint64_t aicNum = 0;
|
||||
uint64_t l1Size = 0;
|
||||
uint64_t l0aSize = 0;
|
||||
uint64_t l0bSize = 0;
|
||||
uint64_t l0cSize = 0;
|
||||
};
|
||||
|
||||
struct CompileInfoCommon {
|
||||
uint32_t aivNum;
|
||||
uint32_t aicNum;
|
||||
uint64_t ubSize;
|
||||
uint64_t l1Size;
|
||||
uint64_t l0aSize;
|
||||
uint64_t l0bSize;
|
||||
uint64_t l0cSize;
|
||||
uint64_t l2CacheSize;
|
||||
int64_t coreNum;
|
||||
int32_t socVersion;
|
||||
uint32_t rsvd;
|
||||
};
|
||||
|
||||
struct FlashAttentionScoreGradCompileInfo {
|
||||
uint32_t aivNum;
|
||||
uint32_t aicNum;
|
||||
uint64_t ubSize;
|
||||
uint64_t l1Size;
|
||||
uint64_t l0aSize;
|
||||
uint64_t l0bSize;
|
||||
uint64_t l0cSize;
|
||||
uint64_t l2CacheSize;
|
||||
int64_t coreNum;
|
||||
platform_ascendc::SocVersion socVersion;
|
||||
};
|
||||
|
||||
struct FACompileInfoCommon {
|
||||
uint32_t aivNum;
|
||||
uint32_t aicNum;
|
||||
uint64_t ubSize;
|
||||
uint64_t l1Size;
|
||||
uint64_t l0aSize;
|
||||
uint64_t l0bSize;
|
||||
uint64_t l0cSize;
|
||||
uint64_t l2CacheSize;
|
||||
int64_t coreNum;
|
||||
int32_t socVersion;
|
||||
uint32_t rsvd;
|
||||
};
|
||||
|
||||
class TilingBaseClass {
|
||||
public:
|
||||
explicit TilingBaseClass(gert::TilingContext* context) : context_(context)
|
||||
{}
|
||||
|
||||
virtual ~TilingBaseClass() = default;
|
||||
|
||||
// Tiling execution framework
|
||||
// 1. GRAPH_SUCCESS: Success, and no need to continue executing subsequent Tiling class implementations
|
||||
// 2. GRAPH_FAILED: Failure, abort the entire Tiling process
|
||||
// 3. GRAPH_PARAM_INVALID: This class does not support, need to continue executing other Tiling class implementations
|
||||
ge::graphStatus DoTiling()
|
||||
{
|
||||
auto ret = GetShapeAttrsInfo();
|
||||
if (ret != ge::GRAPH_SUCCESS) {
|
||||
return ret;
|
||||
}
|
||||
ret = GetPlatformInfo();
|
||||
if (ret != ge::GRAPH_SUCCESS) {
|
||||
return ret;
|
||||
}
|
||||
if (!IsCapable()) {
|
||||
return ge::GRAPH_PARAM_INVALID;
|
||||
}
|
||||
ret = DoOpTiling();
|
||||
if (ret != ge::GRAPH_SUCCESS) {
|
||||
return ret;
|
||||
}
|
||||
ret = DoLibApiTiling();
|
||||
if (ret != ge::GRAPH_SUCCESS) {
|
||||
return ret;
|
||||
}
|
||||
ret = GetWorkspaceSize();
|
||||
if (ret != ge::GRAPH_SUCCESS) {
|
||||
return ret;
|
||||
}
|
||||
ret = PostTiling();
|
||||
if (ret != ge::GRAPH_SUCCESS) {
|
||||
return ret;
|
||||
}
|
||||
context_->SetTilingKey(GetTilingKey());
|
||||
DumpTilingInfo();
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
// Update context
|
||||
virtual void Reset(gert::TilingContext* context)
|
||||
{
|
||||
context_ = context;
|
||||
}
|
||||
|
||||
protected:
|
||||
virtual bool IsCapable() = 0;
|
||||
// 1. Get platform information such as CoreNum, UB/L1/L0C resource sizes
|
||||
virtual ge::graphStatus GetPlatformInfo() = 0;
|
||||
// 2. Get INPUT/OUTPUT/ATTR information
|
||||
virtual ge::graphStatus GetShapeAttrsInfo() = 0;
|
||||
// 3. Calculate data splitting TilingData
|
||||
virtual ge::graphStatus DoOpTiling() = 0;
|
||||
// 4. Calculate high-level API TilingData
|
||||
virtual ge::graphStatus DoLibApiTiling() = 0;
|
||||
// 5. Calculate TilingKey
|
||||
[[nodiscard]] virtual uint64_t GetTilingKey() const = 0;
|
||||
// 6. Calculate Workspace size
|
||||
virtual ge::graphStatus GetWorkspaceSize() = 0;
|
||||
// 7. Save Tiling data
|
||||
virtual ge::graphStatus PostTiling() = 0;
|
||||
// 8. Dump Tiling data
|
||||
virtual void DumpTilingInfo()
|
||||
{
|
||||
int32_t enable = CheckLogLevel(static_cast<int32_t>(OP), DLOG_DEBUG);
|
||||
if (enable != 1) {
|
||||
return;
|
||||
}
|
||||
auto buf = (uint32_t*)context_->GetRawTilingData()->GetData();
|
||||
auto bufLen = context_->GetRawTilingData()->GetDataSize();
|
||||
std::ostringstream oss;
|
||||
oss << "Start to dump tiling info. tilingkey:" << context_->GetTilingKey() << ", tiling data size:" << bufLen
|
||||
<< ", content:";
|
||||
for (size_t i = 0; i < bufLen / sizeof(uint32_t); i++) {
|
||||
oss << *(buf + i) << ",";
|
||||
if (oss.str().length() > 640) { // Split according to 640 to avoid truncation
|
||||
OP_LOGD(context_, "%s", oss.str().c_str());
|
||||
oss.str("");
|
||||
}
|
||||
}
|
||||
OP_LOGD(context_, "%s", oss.str().c_str());
|
||||
}
|
||||
|
||||
static uint32_t CalcTschBlockDim(uint32_t sliceNum, uint32_t aicCoreNum, uint32_t aivCoreNum)
|
||||
{
|
||||
uint32_t ration;
|
||||
if (aicCoreNum == 0 || aivCoreNum == 0 || aicCoreNum > aivCoreNum) {
|
||||
return sliceNum;
|
||||
}
|
||||
ration = aivCoreNum / aicCoreNum;
|
||||
return (sliceNum + (ration - 1)) / ration;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
[[nodiscard]] std::string GetShapeDebugStr(const T& shape) const
|
||||
{
|
||||
std::ostringstream oss;
|
||||
oss << "[";
|
||||
if (shape.GetDimNum() > 0) {
|
||||
for (size_t i = 0; i < shape.GetDimNum() - 1; ++i) {
|
||||
oss << shape.GetDim(i) << ", ";
|
||||
}
|
||||
oss << shape.GetDim(shape.GetDimNum() - 1);
|
||||
}
|
||||
oss << "]";
|
||||
return oss.str();
|
||||
}
|
||||
|
||||
[[nodiscard]] std::string GetTensorDebugStr(
|
||||
const gert::StorageShape* shape, const gert::CompileTimeTensorDesc* tensor)
|
||||
{
|
||||
if (shape == nullptr || tensor == nullptr) {
|
||||
return "nil ";
|
||||
}
|
||||
std::ostringstream oss;
|
||||
oss << "(dtype: " << ge::TypeUtils::DataTypeToSerialString(tensor->GetDataType()) << "),";
|
||||
oss << "(shape:" << GetShapeDebugStr(shape->GetStorageShape()) << "),";
|
||||
oss << "(ori_shape:" << GetShapeDebugStr(shape->GetOriginShape()) << "),";
|
||||
oss << "(format: "
|
||||
<< ge::TypeUtils::FormatToSerialString(
|
||||
static_cast<ge::Format>(ge::GetPrimaryFormat(tensor->GetStorageFormat())))
|
||||
<< "),";
|
||||
oss << "(ori_format: " << ge::TypeUtils::FormatToSerialString(tensor->GetOriginFormat()) << ") ";
|
||||
return oss.str();
|
||||
}
|
||||
|
||||
[[nodiscard]] std::string GetTilingContextDebugStr()
|
||||
{
|
||||
std::ostringstream oss;
|
||||
for (size_t i = 0; i < context_->GetComputeNodeInfo()->GetInputsNum(); ++i) {
|
||||
oss << "input" << i << ": ";
|
||||
oss << GetTensorDebugStr(context_->GetInputShape(i), context_->GetInputDesc(i));
|
||||
}
|
||||
|
||||
for (size_t i = 0; i < context_->GetComputeNodeInfo()->GetOutputsNum(); ++i) {
|
||||
oss << "output" << i << ": ";
|
||||
oss << GetTensorDebugStr(context_->GetOutputShape(i), context_->GetOutputDesc(i));
|
||||
}
|
||||
return oss.str();
|
||||
}
|
||||
|
||||
[[nodiscard]] std::string GetTilingDataDebugStr() const
|
||||
{
|
||||
auto rawTilingData = context_->GetRawTilingData();
|
||||
auto rawTilingDataSize = rawTilingData->GetDataSize();
|
||||
auto data = reinterpret_cast<const int32_t*>(rawTilingData->GetData());
|
||||
size_t len = rawTilingDataSize / sizeof(int32_t);
|
||||
std::ostringstream oss;
|
||||
for (size_t i = 0; i < len; i++) {
|
||||
oss << data[i] << ", ";
|
||||
}
|
||||
return oss.str();
|
||||
}
|
||||
|
||||
protected:
|
||||
gert::TilingContext* context_ = nullptr;
|
||||
std::unique_ptr<platform_ascendc::PlatformAscendC> ascendcPlatform_{nullptr};
|
||||
uint32_t blockDim_{0};
|
||||
uint64_t workspaceSize_{0};
|
||||
uint64_t tilingKey_{0};
|
||||
AiCoreParams aicoreParams_;
|
||||
};
|
||||
|
||||
} // namespace OpTiling
|
||||
} // namespace Transformer
|
||||
} // namespace Ops
|
||||
@@ -0,0 +1,63 @@
|
||||
/**
|
||||
* 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 tiling_key.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
namespace Ops {
|
||||
namespace Transformer {
|
||||
namespace OpTiling {
|
||||
constexpr uint64_t RecursiveSum()
|
||||
{
|
||||
return 0;
|
||||
}
|
||||
|
||||
constexpr uint64_t kBase = 10; // Base-10 carry base
|
||||
template <typename T, typename... Args> constexpr uint64_t RecursiveSum(T templateId, Args... templateIds)
|
||||
{
|
||||
return static_cast<uint64_t>(templateId) + kBase * RecursiveSum(templateIds...);
|
||||
}
|
||||
|
||||
// TilingKey generation rules:
|
||||
// FlashAttentionScore/FlashAttentionScoreGrad assembles tiling key using decimal digits, containing the following key parameters from low to high: Ub0, Ub1,
|
||||
// Block, DataType, Format, Sparse. Specialized template Ub0, Ub1:
|
||||
// Represents the axis for UB intra-core splitting, using AxisEnum. Since we allow at most two axes to be split, UB0 and UB1 exist. If there is no UB intra-core splitting,
|
||||
// fill with AXIS_NONE. UB0 and UB1 each occupy one decimal digit;
|
||||
// Block: Represents the axis used by UB for multi-core splitting, using AxisEnum, occupies one decimal digit;
|
||||
// DataType: Represents the input/output data types supported by the current tiling key, using SupportedDtype enum, occupies one decimal digit
|
||||
// Format: Represents the Format supported by the current tiling key, using InputLayout enum, occupies one decimal digit
|
||||
// Sparse: Represents whether the current tiling key supports Sparse, using SparseCapability enum, occupies one decimal digit
|
||||
// For other specialized scenarios, define your own bit fields and values
|
||||
// usage: get tilingKey from inputted types
|
||||
// uint64_t tilingKey = GET_FLASHATTENTION_TILINGKEY(AxisEnum::AXIS_S1, AxisEnum::AXIS_S2, AxisEnum::AXIS_N2,
|
||||
// SupportedDtype::FLOAT32, InputLayout::BSH, SparseCapability::SUPPORT_ALL)
|
||||
|
||||
constexpr uint64_t TILINGKEYOFFSET = uint64_t(10000000000000000000UL); // 10^19
|
||||
template <typename... Args> constexpr uint64_t GET_TILINGKEY(Args... templateIds)
|
||||
{
|
||||
return TILINGKEYOFFSET + RecursiveSum(templateIds...);
|
||||
}
|
||||
|
||||
// usage: get tilingKey from inputted types
|
||||
// uint64_t tilingKey = TILINGKEY(S2, S1, N2, FLOAT32, BSND, ALL)
|
||||
|
||||
#define TILINGKEY(ub2, ub1, block, dtype, layout, sparse) \
|
||||
(GET_TILINGKEY(AxisEnum::ub2, AxisEnum::ub1, AxisEnum::block, DtypeEnum::dtype, LayoutEnum::layout, \
|
||||
SparseEnum::sparse))
|
||||
|
||||
} // namespace Optiling
|
||||
} // namespace Transformer
|
||||
} // namespace Ops
|
||||
@@ -0,0 +1,351 @@
|
||||
/**
|
||||
* 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 tiling_templates_registry.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <map>
|
||||
#include <string>
|
||||
#include <memory>
|
||||
#include "exe_graph/runtime/tiling_context.h"
|
||||
#include "tiling_base.h"
|
||||
#include "error_log.h"
|
||||
|
||||
namespace Ops {
|
||||
namespace Transformer {
|
||||
namespace OpTiling {
|
||||
|
||||
template <typename T>
|
||||
std::unique_ptr<TilingBaseClass> TILING_CLASS(gert::TilingContext* context)
|
||||
{
|
||||
return std::unique_ptr<T>(new (std::nothrow) T(context));
|
||||
}
|
||||
|
||||
using TilingClassCase = std::unique_ptr<TilingBaseClass> (*)(gert::TilingContext*);
|
||||
|
||||
class TilingCases {
|
||||
public:
|
||||
explicit TilingCases(std::string op_type) : op_type_(std::move(op_type))
|
||||
{}
|
||||
|
||||
template <typename T>
|
||||
void AddTiling(int32_t priority)
|
||||
{
|
||||
OP_CHECK_IF(
|
||||
cases_.find(priority) != cases_.end(), OP_LOGE(op_type_, "There are duplicate registrations."), return);
|
||||
cases_[priority] = TILING_CLASS<T>;
|
||||
OP_CHECK_IF(
|
||||
cases_[priority] == nullptr,
|
||||
OP_LOGE(op_type_, "Register op tiling func failed, please check the class name."), return);
|
||||
}
|
||||
|
||||
const std::map<int32_t, TilingClassCase>& GetTilingCases()
|
||||
{
|
||||
return cases_;
|
||||
}
|
||||
|
||||
private:
|
||||
std::map<int32_t, TilingClassCase> cases_;
|
||||
const std::string op_type_;
|
||||
};
|
||||
|
||||
// --------------------------------Interfacce with soc version --------------------------------
|
||||
class TilingRegistryNew {
|
||||
public:
|
||||
TilingRegistryNew() = default;
|
||||
|
||||
#ifdef ASCENDC_OP_TEST
|
||||
static TilingRegistryNew& GetInstance();
|
||||
#else
|
||||
static TilingRegistryNew& GetInstance()
|
||||
{
|
||||
static TilingRegistryNew registry_impl_;
|
||||
return registry_impl_;
|
||||
}
|
||||
#endif
|
||||
|
||||
std::shared_ptr<TilingCases> RegisterOp(const std::string& op_type, int32_t soc_version)
|
||||
{
|
||||
auto soc_iter = registry_map_.find(soc_version);
|
||||
if (soc_iter == registry_map_.end()) {
|
||||
std::map<std::string, std::shared_ptr<TilingCases>> op_type_map;
|
||||
op_type_map[op_type] = std::shared_ptr<TilingCases>(new (std::nothrow) TilingCases(op_type));
|
||||
registry_map_[soc_version] = op_type_map;
|
||||
} else {
|
||||
if (soc_iter->second.find(op_type) == soc_iter->second.end()) {
|
||||
soc_iter->second[op_type] = std::shared_ptr<TilingCases>(new (std::nothrow) TilingCases(op_type));
|
||||
}
|
||||
}
|
||||
|
||||
OP_CHECK_IF(
|
||||
registry_map_[soc_version][op_type] == nullptr,
|
||||
OP_LOGE(op_type, "Register tiling func failed, please check the class name."), return nullptr);
|
||||
return registry_map_[soc_version][op_type];
|
||||
}
|
||||
|
||||
ge::graphStatus DoTilingImpl(gert::TilingContext* context)
|
||||
{
|
||||
int32_t soc_version = (int32_t)platform_ascendc::SocVersion::RESERVED_VERSION;
|
||||
const char* op_type = context->GetNodeType();
|
||||
fe::PlatFormInfos* platformInfoPtr = context->GetPlatformInfo();
|
||||
if (platformInfoPtr == nullptr) {
|
||||
auto compileInfoPtr = static_cast<const CompileInfoCommon*>(context->GetCompileInfo());
|
||||
OP_CHECK_IF(
|
||||
compileInfoPtr == nullptr, OP_LOGE(op_type, "compileInfoPtr is null."), return ge::GRAPH_FAILED);
|
||||
soc_version = compileInfoPtr->socVersion;
|
||||
OP_LOGD(context, "soc version in compileInfo is %d", soc_version);
|
||||
} else {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
|
||||
soc_version = static_cast<int32_t>(ascendcPlatform.GetSocVersion());
|
||||
OP_LOGD(context, "soc version is %d", soc_version);
|
||||
if (soc_version == (int32_t)platform_ascendc::SocVersion::RESERVED_VERSION) {
|
||||
OP_LOGE(op_type, "Do op tiling failed, cannot find soc version.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
}
|
||||
auto tilingTemplateRegistryMap = GetTilingTemplates(op_type, soc_version);
|
||||
for (auto it = tilingTemplateRegistryMap.begin(); it != tilingTemplateRegistryMap.end(); ++it) {
|
||||
auto tilingTemplate = it->second(context);
|
||||
if (tilingTemplate != nullptr) {
|
||||
ge::graphStatus status = tilingTemplate->DoTiling();
|
||||
if (status != ge::GRAPH_PARAM_INVALID) {
|
||||
OP_LOGD(context, "Do general op tiling success priority=%d", it->first);
|
||||
return status;
|
||||
}
|
||||
OP_LOGD(context, "Ignore general op tiling priority=%d", it->first);
|
||||
}
|
||||
}
|
||||
OP_LOGE(op_type, "Do op tiling failed, no valid template is found.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
ge::graphStatus DoTilingImpl(gert::TilingContext* context, const std::vector<int32_t>& priorities)
|
||||
{
|
||||
int32_t soc_version;
|
||||
const char* op_type = context->GetNodeType();
|
||||
auto platformInfoPtr = context->GetPlatformInfo();
|
||||
if (platformInfoPtr == nullptr) {
|
||||
auto compileInfoPtr = reinterpret_cast<const CompileInfoCommon*>(context->GetCompileInfo());
|
||||
OP_CHECK_IF(
|
||||
compileInfoPtr == nullptr, OP_LOGE(op_type, "compileInfoPtr is null."), return ge::GRAPH_FAILED);
|
||||
soc_version = compileInfoPtr->socVersion;
|
||||
OP_LOGD(context, "soc version in compileInfo is %d", soc_version);
|
||||
} else {
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
|
||||
soc_version = static_cast<int32_t>(ascendcPlatform.GetSocVersion());
|
||||
OP_LOGD(context, "soc version is %d", soc_version);
|
||||
}
|
||||
|
||||
auto tilingTemplateRegistryMap = GetTilingTemplates(op_type, soc_version);
|
||||
for (auto priority_id : priorities) {
|
||||
auto tilingCaseIter = tilingTemplateRegistryMap.find(priority_id);
|
||||
if (tilingCaseIter != tilingTemplateRegistryMap.end()) {
|
||||
auto templateFunc = tilingCaseIter->second(context);
|
||||
if (templateFunc != nullptr) {
|
||||
ge::graphStatus status = templateFunc->DoTiling();
|
||||
if (status == ge::GRAPH_SUCCESS) {
|
||||
OP_LOGD(context, "Do general op tiling success priority=%d", priority_id);
|
||||
return status;
|
||||
}
|
||||
OP_LOGD(context, "Ignore general op tiling priority=%d", priority_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
const std::map<int32_t, TilingClassCase>& GetTilingTemplates(const std::string& op_type, int32_t soc_version)
|
||||
{
|
||||
auto soc_iter = registry_map_.find(soc_version);
|
||||
OP_CHECK_IF(
|
||||
soc_iter == registry_map_.end(),
|
||||
OP_LOGE(op_type, "Get op tiling func failed, please check the soc version %d", soc_version),
|
||||
return empty_tiling_case_);
|
||||
auto op_iter = soc_iter->second.find(op_type);
|
||||
OP_CHECK_IF(
|
||||
op_iter == soc_iter->second.end(), OP_LOGE(op_type, "Get op tiling func failed, please check the op name."),
|
||||
return empty_tiling_case_);
|
||||
return op_iter->second->GetTilingCases();
|
||||
}
|
||||
|
||||
private:
|
||||
std::map<int32_t, std::map<std::string, std::shared_ptr<TilingCases>>> registry_map_; // key is socversion
|
||||
const std::map<int32_t, TilingClassCase> empty_tiling_case_{};
|
||||
};
|
||||
|
||||
class RegisterNew {
|
||||
public:
|
||||
explicit RegisterNew(std::string op_type) : op_type_(std::move(op_type))
|
||||
{}
|
||||
|
||||
template <typename T>
|
||||
RegisterNew& tiling(int32_t priority, int32_t soc_version)
|
||||
{
|
||||
auto tilingCases = TilingRegistryNew::GetInstance().RegisterOp(op_type_, soc_version);
|
||||
OP_CHECK_IF(
|
||||
tilingCases == nullptr, OP_LOGE(op_type_, "Register op tiling failed, please the op name."), return *this);
|
||||
tilingCases->AddTiling<T>(priority);
|
||||
return *this;
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
RegisterNew& tiling(int32_t priority, const std::vector<int32_t>& soc_versions)
|
||||
{
|
||||
for (int32_t soc_version : soc_versions) {
|
||||
auto tilingCases = TilingRegistryNew::GetInstance().RegisterOp(op_type_, soc_version);
|
||||
OP_CHECK_IF(
|
||||
tilingCases == nullptr, OP_LOGE(op_type_, "Register op tiling failed, please the op name."),
|
||||
return *this);
|
||||
tilingCases->AddTiling<T>(priority);
|
||||
}
|
||||
return *this;
|
||||
}
|
||||
|
||||
private:
|
||||
const std::string op_type_;
|
||||
};
|
||||
|
||||
// --------------------------------Interfacce without soc version --------------------------------
|
||||
class TilingRegistry {
|
||||
public:
|
||||
TilingRegistry() = default;
|
||||
|
||||
#ifdef ASCENDC_OP_TEST
|
||||
static TilingRegistry& GetInstance();
|
||||
#else
|
||||
static TilingRegistry& GetInstance()
|
||||
{
|
||||
static TilingRegistry registry_impl_;
|
||||
return registry_impl_;
|
||||
}
|
||||
#endif
|
||||
|
||||
std::shared_ptr<TilingCases> RegisterOp(const std::string& op_type)
|
||||
{
|
||||
if (registry_map_.find(op_type) == registry_map_.end()) {
|
||||
registry_map_[op_type] = std::shared_ptr<TilingCases>(new (std::nothrow) TilingCases(op_type));
|
||||
}
|
||||
OP_CHECK_IF(
|
||||
registry_map_[op_type] == nullptr,
|
||||
OP_LOGE(op_type, "Register tiling func failed, please check the class name."), return nullptr);
|
||||
return registry_map_[op_type];
|
||||
}
|
||||
|
||||
ge::graphStatus DoTilingImpl(gert::TilingContext* context)
|
||||
{
|
||||
const char* op_type = context->GetNodeType();
|
||||
auto tilingTemplateRegistryMap = GetTilingTemplates(op_type);
|
||||
for (auto it = tilingTemplateRegistryMap.begin(); it != tilingTemplateRegistryMap.end(); ++it) {
|
||||
auto tilingTemplate = it->second(context);
|
||||
if (tilingTemplate != nullptr) {
|
||||
ge::graphStatus status = tilingTemplate->DoTiling();
|
||||
if (status != ge::GRAPH_PARAM_INVALID) {
|
||||
OP_LOGD(context, "Do general op tiling success priority=%d", it->first);
|
||||
return status;
|
||||
}
|
||||
OP_LOGD(context, "Ignore general op tiling priority=%d", it->first);
|
||||
}
|
||||
}
|
||||
OP_LOGE(op_type, "Do op tiling failed, no valid template is found.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
ge::graphStatus DoTilingImpl(gert::TilingContext* context, const std::vector<int32_t>& priorities)
|
||||
{
|
||||
const char* op_type = context->GetNodeType();
|
||||
auto tilingTemplateRegistryMap = GetTilingTemplates(op_type);
|
||||
for (auto priorityId : priorities) {
|
||||
auto templateFunc = tilingTemplateRegistryMap[priorityId](context);
|
||||
if (templateFunc != nullptr) {
|
||||
ge::graphStatus status = templateFunc->DoTiling();
|
||||
if (status == ge::GRAPH_SUCCESS) {
|
||||
OP_LOGD(context, "Do general op tiling success priority=%d", priorityId);
|
||||
return status;
|
||||
}
|
||||
if (status != ge::GRAPH_PARAM_INVALID) {
|
||||
OP_LOGD(context, "Do op tiling failed");
|
||||
return status;
|
||||
}
|
||||
OP_LOGD(context, "Ignore general op tiling priority=%d", priorityId);
|
||||
}
|
||||
}
|
||||
OP_LOGE(op_type, "Do op tiling failed, no valid template is found.");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
|
||||
const std::map<int32_t, TilingClassCase>& GetTilingTemplates(const std::string& op_type)
|
||||
{
|
||||
OP_CHECK_IF(
|
||||
registry_map_.find(op_type) == registry_map_.end(),
|
||||
OP_LOGE(op_type, "Get op tiling func failed, please check the op name."), return empty_tiling_case_);
|
||||
return registry_map_[op_type]->GetTilingCases();
|
||||
}
|
||||
|
||||
private:
|
||||
std::map<std::string, std::shared_ptr<TilingCases>> registry_map_;
|
||||
const std::map<int32_t, TilingClassCase> empty_tiling_case_;
|
||||
};
|
||||
|
||||
class Register {
|
||||
public:
|
||||
explicit Register(std::string op_type) : op_type_(std::move(op_type))
|
||||
{}
|
||||
|
||||
template <typename T>
|
||||
Register& tiling(int32_t priority)
|
||||
{
|
||||
auto tilingCases = TilingRegistry::GetInstance().RegisterOp(op_type_);
|
||||
OP_CHECK_IF(
|
||||
tilingCases == nullptr, OP_LOGE(op_type_, "Register op tiling failed, please the op name."), return *this);
|
||||
tilingCases->AddTiling<T>(priority);
|
||||
return *this;
|
||||
}
|
||||
|
||||
private:
|
||||
const std::string op_type_;
|
||||
};
|
||||
} // namespace OpTiling
|
||||
} // namespace Transformer
|
||||
} // namespace Ops
|
||||
|
||||
// op_type: operator name, class_name: registered tiling class, soc_version: chip version number
|
||||
// priority: priority of tiling class, smaller value means higher priority, i.e., this tiling class will be selected first
|
||||
#define REGISTER_TILING_TEMPLATE_WITH_SOCVERSION(op_type, class_name, soc_versions, priority) \
|
||||
[[maybe_unused]] uint32_t op_impl_register_template_##op_type##_##class_name##priority; \
|
||||
static Ops::Transformer::OpTiling::RegisterNew VAR_UNUSED##op_type##class_name##priority_register = \
|
||||
Ops::Transformer::OpTiling::RegisterNew(#op_type).tiling<class_name>(priority, soc_versions)
|
||||
|
||||
// op_type: operator name, class_name: registered tiling class
|
||||
// priority: priority of tiling class, smaller value means higher priority, i.e., higher probability of being selected
|
||||
#define REGISTER_TILING_TEMPLATE(op_type, class_name, priority) \
|
||||
[[maybe_unused]] uint32_t op_impl_register_template_##op_type##_##class_name##priority; \
|
||||
static Ops::Transformer::OpTiling::Register VAR_UNUSED##op_type_##class_name##priority_register = \
|
||||
Ops::Transformer::OpTiling::Register(op_type).tiling<class_name>(priority)
|
||||
|
||||
// op_type: operator name, class_name: registered tiling class
|
||||
// soc_version: SOC version, used to distinguish different SOCs
|
||||
// priority: priority of tiling class, smaller value means higher priority, i.e., this tiling class will be selected first
|
||||
#define REGISTER_TILING_TEMPLATE_NEW(op_type, class_name, soc_version, priority) \
|
||||
[[maybe_unused]] uint32_t op_impl_register_template_##op_type##_##class_name##priority; \
|
||||
static Ops::Transformer::OpTiling::RegisterNew VAR_UNUSED##op_type##class_name##priority_register = \
|
||||
Ops::Transformer::OpTiling::RegisterNew(#op_type).tiling<class_name>(priority, soc_version)
|
||||
|
||||
// op_type: operator name, class_name: registered tiling class
|
||||
// priority: priority of tiling class, smaller value means higher priority, i.e., higher probability of being selected
|
||||
// Replaces REGISTER_TILING_TEMPLATE, if op_type is a string constant, remove the quotes
|
||||
#define REGISTER_OPS_TILING_TEMPLATE(op_type, class_name, priority) \
|
||||
[[maybe_unused]] uint32_t op_impl_register_template_##op_type##_##class_name##priority; \
|
||||
static Ops::Transformer::OpTiling::Register \
|
||||
__attribute__((unused)) tiling_##op_type##_##class_name##_##priority##_register = \
|
||||
Ops::Transformer::OpTiling::Register(#op_type).tiling<class_name>(priority)
|
||||
139
csrc/moe/chunk_gated_delta_rule_fwd_h/tiling_base/tiling_type.h
Normal file
139
csrc/moe/chunk_gated_delta_rule_fwd_h/tiling_base/tiling_type.h
Normal file
@@ -0,0 +1,139 @@
|
||||
/**
|
||||
* 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 tiling_type.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
namespace optiling {
|
||||
|
||||
enum class AxisEnum {
|
||||
B = 0,
|
||||
N2 = 1,
|
||||
G = 2,
|
||||
S1 = 3,
|
||||
S2 = 4,
|
||||
D = 5,
|
||||
NONE = 9,
|
||||
};
|
||||
|
||||
enum class DtypeEnum {
|
||||
FLOAT16 = 0,
|
||||
FLOAT32 = 1,
|
||||
BFLOAT16 = 2,
|
||||
FLOAT16_PRECISION = 3,
|
||||
};
|
||||
|
||||
enum class PerformanceOrientedEnum {
|
||||
BIG_BUFFER = 1,
|
||||
BIG_DOUBLE_BUFFER = 2,
|
||||
};
|
||||
|
||||
enum class MatmulConfig {
|
||||
NULL_CONFIG = 0,
|
||||
NORMAL_CONFIG = 1,
|
||||
MDL_CONFIG = 2
|
||||
};
|
||||
|
||||
enum class PseConfig {
|
||||
NO_PSE = 0,
|
||||
EXIST_PSE = 1
|
||||
};
|
||||
|
||||
enum class AttenMaskConfig {
|
||||
NO_ATTEN_MASK = 0,
|
||||
EXIST_ATTEN_MASK = 1
|
||||
};
|
||||
|
||||
enum class DropOutConfig {
|
||||
NO_DROP_OUT = 0,
|
||||
EXIST_DROP_OUT = 1
|
||||
};
|
||||
|
||||
enum class CubeFormatEnum {
|
||||
ND = 0,
|
||||
NZ = 1
|
||||
};
|
||||
enum class LayoutEnum {
|
||||
BSND = 0,
|
||||
SBND = 1,
|
||||
BNSD = 2,
|
||||
TND = 3,
|
||||
NTD_TND = 4
|
||||
};
|
||||
|
||||
enum class CubeInputSourceEnum {
|
||||
GM = 0,
|
||||
L1 = 1
|
||||
};
|
||||
|
||||
enum class OptionEnum {
|
||||
DISABLE = 0,
|
||||
ENABLE = 1
|
||||
};
|
||||
|
||||
enum class SparseEnum {
|
||||
ALL = 0,
|
||||
NONE = 1,
|
||||
ANY = 2,
|
||||
CAUSAL = 3,
|
||||
BAND = 4,
|
||||
PREFIX = 5,
|
||||
BAND_COMPRESS = 6,
|
||||
RIGHT_DOWN_CAUSAL = 7,
|
||||
RIGHT_DOWN_CAUSAL_BAND = 8,
|
||||
BAND_LEFT_UP_CAUSAL = 9
|
||||
};
|
||||
|
||||
constexpr uint64_t RecursiveSum()
|
||||
{
|
||||
return 0;
|
||||
}
|
||||
|
||||
constexpr int64_t base10Multiplier = 10;
|
||||
|
||||
template <typename T, typename... Args> constexpr uint64_t RecursiveSum(T templateId, Args... templateIds)
|
||||
{
|
||||
return static_cast<uint64_t>(templateId) + base10Multiplier * RecursiveSum(templateIds...);
|
||||
}
|
||||
|
||||
// TilingKey generation rules:
|
||||
// FlashAttentionScore/FlashAttentionScoreGrad assembles tiling key using decimal digits, containing the following key parameters from low to high: Ub0, Ub1,
|
||||
// Block, DataType, Format, Sparse. Specialized template Ub0, Ub1:
|
||||
// Represents the axis for UB intra-core splitting, using AxisEnum. Since we allow at most two axes to be split, UB0 and UB1 exist. If there is no UB intra-core splitting,
|
||||
// fill with AXIS_NONE. UB0 and UB1 each occupy one decimal digit;
|
||||
// Block: Represents the axis used by UB for multi-core splitting, using AxisEnum, occupies one decimal digit;
|
||||
// DataType: Represents the input/output data types supported by the current tiling key, using SupportedDtype enum, occupies one decimal digit
|
||||
// Format: Represents the Format supported by the current tiling key, using InputLayout enum, occupies one decimal digit
|
||||
// Sparse: Represents whether the current tiling key supports Sparse, using SparseCapability enum, occupies one decimal digit
|
||||
// For other specialized scenarios, define your own bit fields and values
|
||||
// usage: get tilingKey from inputted types
|
||||
// uint64_t tilingKey = GET_FLASHATTENTION_TILINGKEY(AxisEnum::AXIS_S1, AxisEnum::AXIS_S2, AxisEnum::AXIS_N2,
|
||||
// SupportedDtype::FLOAT32, InputLayout::BSH, SparseCapability::SUPPORT_ALL)
|
||||
|
||||
constexpr uint64_t TILINGKEYOFFSET = uint64_t(10000000000000000000UL); // 10^19
|
||||
template <typename... Args> constexpr uint64_t GET_TILINGKEY(Args... templateIds)
|
||||
{
|
||||
return TILINGKEYOFFSET + RecursiveSum(templateIds...);
|
||||
}
|
||||
|
||||
// usage: get tilingKey from inputted types
|
||||
// uint64_t tilingKey = TILINGKEY(S2, S1, N2, FLOAT32, BSND, ALL)
|
||||
|
||||
#define TILINGKEY(ub2, ub1, block, dtype, layout, sparse) \
|
||||
(GET_TILINGKEY(AxisEnum::ub2, AxisEnum::ub1, AxisEnum::block, DtypeEnum::dtype, LayoutEnum::layout, \
|
||||
SparseEnum::sparse))
|
||||
|
||||
} // namespace optiling
|
||||
@@ -0,0 +1,30 @@
|
||||
/**
|
||||
* 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 tiling_util.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "register/op_impl_registry.h"
|
||||
|
||||
namespace Ops {
|
||||
namespace Transformer {
|
||||
namespace OpTiling {
|
||||
bool IsRegbaseSocVersion(const gert::TilingParseContext* context);
|
||||
|
||||
bool IsRegbaseSocVersion(const gert::TilingContext* context);
|
||||
|
||||
const gert::Shape& EnsureNotScalar(const gert::Shape& inShape);
|
||||
} // namespace OpTiling
|
||||
} // namespace Transformer
|
||||
} // namespace Ops
|
||||
Reference in New Issue
Block a user