init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View 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()

View 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_fwd_o_def.cpp
)
endif()
add_ops_compile_options(
OP_NAME ChunkFwdO
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_fwd_o ACLNNTYPE aclnn_exclude)
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
${CMAKE_CURRENT_SOURCE_DIR}
${CATLASS_INCLUDE_DIR_ABS}
)
endif()

View File

@@ -0,0 +1,105 @@
/**
 * 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_fwd_o_def.cpp
*\brief
*/
#include "register/op_def_registry.h"
namespace ops {
class ChunkFwdO : public OpDef {
public:
explicit ChunkFwdO(const char *name) : OpDef(name)
{
// Define inputs
this->Input("q")
.ParamType(REQUIRED)
.DataType({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})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
this->Input("k")
.ParamType(REQUIRED)
.DataType({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})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
this->Input("v")
.ParamType(REQUIRED)
.DataType({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})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
this->Input("h")
.ParamType(REQUIRED)
.DataType({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})
.UnknownShapeFormat({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_BF16, ge::DT_FLOAT16})
.Format({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})
.AutoContiguous();
this->Input("cu_seqlens")
.ParamType(OPTIONAL)
.ValueDepend(OPTIONAL)
.DataType({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})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
this->Input("chunk_offsets")
.ParamType(OPTIONAL)
.ValueDepend(OPTIONAL)
.DataType({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})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND})
.AutoContiguous();
this->Output("o")
.ParamType(REQUIRED)
.DataType({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})
.UnknownShapeFormat({ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND, ge::FORMAT_ND});
this->Attr("scale").AttrType(REQUIRED).Float(1.0);
this->Attr("chunk_size").AttrType(REQUIRED).Int(64);
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(ChunkFwdO);
} // namespace ops

View File

@@ -0,0 +1,157 @@
/**
 * 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_fwd_o_tiling.cpp
* \brief
*/
#include "chunk_fwd_o_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_Q_IDX = 0;
static constexpr size_t INPUT_K_IDX = 1;
static constexpr size_t INPUT_V_IDX = 2;
static constexpr size_t INPUT_H_IDX = 3;
static constexpr size_t INPUT_G_IDX = 4;
static constexpr size_t INPUT_SEQLENS_IDX = 5;
static constexpr size_t INPUT_CHUNK_INDICES_IDX = 6;
static constexpr size_t ATTR_SCALE_IDX = 0;
static constexpr size_t ATTR_CHUNK_SIZE_IDX = 1;
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 ChunkFwdOTilingDataPrint(gert::TilingContext *context, ChunkFwdOTilingData &tiling)
{
auto nodeName = context->GetNodeName();
OP_LOGD(nodeName, ">>>>>>>>>>>>>>> Start to print ChunkFwdO tiling data <<<<<<<<<<<<<<<<");
OP_LOGD(nodeName, "=== batch: %ld", tiling.get_shapeBatch());
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, "=== dataType: %ld", tiling.get_dataType());
OP_LOGD(nodeName, "=== isVariedLen: %ld", tiling.get_isVariedLen());
OP_LOGD(nodeName, "=== tokenBatch: %f", tiling.get_tokenBatch());
OP_LOGD(nodeName, ">>>>>>>>>>>>>>> Print ChunkFwdO tiling data end <<<<<<<<<<<<<<<<");
}
ge::graphStatus Tiling4ChunkFwdO(gert::TilingContext *context)
{
OP_LOGD(context->GetNodeName(), "Tiling4ChunkFwdO start.");
ChunkFwdOTilingData tiling;
gert::Shape qStorageShape = context->GetOptionalInputShape(INPUT_Q_IDX)->GetStorageShape();
gert::Shape vStorageShape = context->GetOptionalInputShape(INPUT_V_IDX)->GetStorageShape();
int64_t seqlen = qStorageShape.GetDim(DIM_SEQLEN);
int64_t kNumHead = qStorageShape.GetDim(DIM_HEAD_NUM);
int64_t vNumHead = vStorageShape.GetDim(DIM_HEAD_NUM);
int64_t kHeadDim = qStorageShape.GetDim(DIM_HEAD_DIM);
int64_t vHeadDim = vStorageShape.GetDim(DIM_HEAD_DIM);
int64_t isVariedLen, shapeBatch, tokenBatch;
auto cuSeqlensTensor = context->GetOptionalInputTensor(INPUT_SEQLENS_IDX);
if (cuSeqlensTensor == nullptr) {
isVariedLen = false;
shapeBatch = qStorageShape.GetDim(DIM_BATCH);
tokenBatch = 1;
} else {
isVariedLen = true;
shapeBatch = 1;
tokenBatch = cuSeqlensTensor->GetStorageShape().GetDim(DIM_BATCH) - 1;
}
auto attrPtr = context->GetAttrs();
float scale = *(attrPtr->GetAttrPointer<double>(ATTR_SCALE_IDX));
int64_t chunkSize = *(attrPtr->GetAttrPointer<int64_t>(ATTR_CHUNK_SIZE_IDX));
auto dtype = context->GetInputTensor(0)->GetDataType();
uint64_t dataType = dtype == ge::DT_BF16 ? 1 : 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;
}
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;
int64_t pingpongStages = 2;
size_t workspaceOffset = ascendcPlatform.GetLibApiWorkSpaceSize();
workspaceOffset += WORKSPACE_RSV_BYTE;
tiling.set_vWorkspaceOffset(workspaceOffset);
workspaceOffset += (aicCoreNum * chunkSize * vHeadDim * sizeof(float) * pingpongStages + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
tiling.set_hWorkspaceOffset(workspaceOffset);
workspaceOffset += (aicCoreNum * chunkSize * vHeadDim * sizeof(float) * pingpongStages + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
tiling.set_attnWorkspaceOffset(workspaceOffset);
workspaceOffset += (aicCoreNum * chunkSize * chunkSize * sizeof(float) * pingpongStages + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
tiling.set_aftermaskWorkspaceOffset(workspaceOffset);
workspaceOffset += (aicCoreNum * chunkSize * chunkSize * sizeof(float) * pingpongStages + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
tiling.set_maskWorkspaceOffset(workspaceOffset);
workspaceOffset += (chunkSize * chunkSize + GM_ALIGN) / GM_ALIGN * GM_ALIGN;
workspaceOffset += WORKSPACE_RSV_BYTE;
size_t *currentWorkspace = context->GetWorkspaceSizes(1);
currentWorkspace[0] = (workspaceOffset - 0);
tiling.set_shapeBatch(shapeBatch);
tiling.set_seqlen(seqlen);
tiling.set_kNumHead(kNumHead);
tiling.set_vNumHead(vNumHead);
tiling.set_kHeadDim(kHeadDim);
tiling.set_vHeadDim(vHeadDim);
tiling.set_scale(scale);
tiling.set_chunkSize(chunkSize);
tiling.set_isVariedLen(isVariedLen);
tiling.set_tokenBatch(tokenBatch);
tiling.set_dataType(dataType);
tiling.set_gDataType(gDataType);
tiling.SaveToBuffer(context->GetRawTilingData()->GetData(), context->GetRawTilingData()->GetCapacity());
context->GetRawTilingData()->SetDataSize(tiling.GetDataSize());
ChunkFwdOTilingDataPrint(context, tiling);
OP_LOGD(context->GetNodeName(), "Tiling4ChunkFwdO end.");
return ge::GRAPH_SUCCESS;
}
ge::graphStatus TilingPrepareForChunkFwdO(gert::TilingParseContext *context)
{
(void)context;
return ge::GRAPH_SUCCESS;
}
IMPL_OP_OPTILING(ChunkFwdO)
.Tiling(Tiling4ChunkFwdO)
.TilingParse<ChunkFwdOCompileInfo>(TilingPrepareForChunkFwdO);
} // namespace optiling

View File

@@ -0,0 +1,46 @@
/**
 * 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_fwd_o_tiling.h
* \brief
*/
#pragma once
#include <cstdint>
#include <register/tilingdata_base.h>
#include <tiling/tiling_api.h>
namespace optiling {
BEGIN_TILING_DATA_DEF(ChunkFwdOTilingData)
TILING_DATA_FIELD_DEF(int64_t, shapeBatch);
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, isVariedLen);
TILING_DATA_FIELD_DEF(int64_t, tokenBatch);
TILING_DATA_FIELD_DEF(int64_t, dataType);
TILING_DATA_FIELD_DEF(int64_t, gDataType);
TILING_DATA_FIELD_DEF(int64_t, vWorkspaceOffset);
TILING_DATA_FIELD_DEF(int64_t, hWorkspaceOffset);
TILING_DATA_FIELD_DEF(int64_t, attnWorkspaceOffset);
TILING_DATA_FIELD_DEF(int64_t, aftermaskWorkspaceOffset);
TILING_DATA_FIELD_DEF(int64_t, maskWorkspaceOffset);
TILING_DATA_FIELD_DEF(float, scale);
END_TILING_DATA_DEF;
REGISTER_TILING_DATA_CLASS(ChunkFwdO, ChunkFwdOTilingData)
struct ChunkFwdOCompileInfo {};
} // namespace optiling

View File

@@ -0,0 +1,162 @@
/**
* 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_fwd_o.h"
#include "chunk_fwd_o.h"
#include <dlfcn.h>
#include <new>
#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 ChunkFwdOParams {
const aclTensor *q = nullptr;
const aclTensor *k = nullptr;
const aclTensor *v = nullptr;
const aclTensor *h = nullptr;
const aclTensor *g = nullptr;
const aclIntArray *cuSeqlensOptional = nullptr;
const aclIntArray *chunkOffsetsOptional = nullptr;
double scale = 1.0;
int64_t chunkSize = 64;
const aclTensor *oOut = nullptr;
};
static aclnnStatus CheckNotNull(ChunkFwdOParams params)
{
CHECK_COND(params.q != nullptr, ACLNN_ERR_PARAM_NULLPTR, "q must not be nullptr.");
CHECK_COND(params.k != nullptr, ACLNN_ERR_PARAM_NULLPTR, "k must not be nullptr.");
CHECK_COND(params.v != nullptr, ACLNN_ERR_PARAM_NULLPTR, "v must not be nullptr.");
CHECK_COND(params.h != nullptr, ACLNN_ERR_PARAM_NULLPTR, "h must not be nullptr.");
CHECK_COND(params.g != nullptr, ACLNN_ERR_PARAM_NULLPTR, "g must not be nullptr.");
CHECK_COND(params.oOut != nullptr, ACLNN_ERR_PARAM_NULLPTR, "oOut must not be nullptr.");
return ACLNN_SUCCESS;
}
static aclnnStatus CheckFormat(ChunkFwdOParams params)
{
return ACLNN_SUCCESS;
}
static aclnnStatus CheckShape(ChunkFwdOParams 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(ChunkFwdOParams &params, aclOpExecutor *executorPtr)
{
CHECK_COND(DataContiguous(params.q, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
"Contiguous q failed.");
CHECK_COND(DataContiguous(params.k, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
"Contiguous k failed.");
CHECK_COND(DataContiguous(params.v, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
"Contiguous v failed.");
CHECK_COND(DataContiguous(params.h, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
"Contiguous h failed.");
CHECK_COND(DataContiguous(params.g, executorPtr) == ACLNN_SUCCESS, ACLNN_ERR_PARAM_INVALID,
"Contiguous g failed.");
return ACLNN_SUCCESS;
}
static aclnnStatus CheckDtype(ChunkFwdOParams params)
{
return ACLNN_SUCCESS;
}
static aclnnStatus CheckParams(ChunkFwdOParams params)
{
CHECK_RET(CheckNotNull(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 aclnnChunkFwdOGetWorkspaceSize(
const aclTensor *q,
const aclTensor *k,
const aclTensor *v,
const aclTensor *h,
const aclTensor *g,
const aclIntArray *cuSeqlensOptional,
const aclIntArray *chunkOffsetsOptional,
double scale,
int64_t chunkSize,
const aclTensor *oOut,
uint64_t *workspaceSize,
aclOpExecutor **executor)
{
ChunkFwdOParams params{q, k, v, h, g, cuSeqlensOptional, chunkOffsetsOptional, scale, chunkSize, oOut};
// Standard syntax, Check parameters.
L2_DFX_PHASE_1(aclnnChunkFwdO, DFX_IN(q, k, v, h, g, cuSeqlensOptional, chunkOffsetsOptional),
DFX_OUT(oOut));
// 固定写法,创建OpExecutor
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.");
auto result = l0op::ChunkFwdO(params.q, params.k, params.v, params.h, params.g, params.cuSeqlensOptional, params.chunkOffsetsOptional, params.scale, params.chunkSize, params.oOut, 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 viewCopyResult = l0op::ViewCopy(result[0], params.oOut, executorPtr);
CHECK_RET(viewCopyResult != 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 aclnnChunkFwdO(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream)
{
L2_DFX_PHASE_2(aclnnChunkFwdO);
CHECK_COND(CommonOpExecutorRun(workspace, workspaceSize, executor, stream) == ACLNN_SUCCESS, ACLNN_ERR_INNER,
"This is an error in ChunkFwdO launch aicore.");
return ACLNN_SUCCESS;
}
#ifdef __cplusplus
}
#endif

View File

@@ -0,0 +1,65 @@
/**
* 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_FWD_O_H
#define OP_API_INC_ACLNN_CHUNK_FWD_O_H
#include "aclnn/aclnn_base.h"
#ifdef __cplusplus
extern "C" {
#endif
/* function: aclnnChunkFwdOGetWorkspaceSize
* parameters :
* q : required
* k : required
* v : required
* h : required
* g : required
* cuSeqlensOptional : optional
* chunkOffsetsOptional : optional
* scale : required
* chunkSize : required
* oOut : required
* workspaceSize : size of workspace(output).
* executor : executor context(output).
*/
__attribute__((visibility("default")))
aclnnStatus aclnnChunkFwdOGetWorkspaceSize(
const aclTensor *q,
const aclTensor *k,
const aclTensor *v,
const aclTensor *h,
const aclTensor *g,
const aclIntArray *cuSeqlensOptional,
const aclIntArray *chunkOffsetsOptional,
double scale,
int64_t chunkSize,
const aclTensor *oOut,
uint64_t *workspaceSize,
aclOpExecutor **executor);
/* function: aclnnChunkFwdO
* parameters :
* workspace : workspace memory addr(input).
* workspaceSize : size of workspace(input).
* executor : executor context(input).
* stream : acl stream.
*/
__attribute__((visibility("default")))
aclnnStatus aclnnChunkFwdO(
void *workspace,
uint64_t workspaceSize,
aclOpExecutor *executor,
aclrtStream stream);
#ifdef __cplusplus
}
#endif
#endif

View File

@@ -0,0 +1,67 @@
/**
* 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_fwd_o.h"
using namespace op;
namespace l0op {
OP_TYPE_REGISTER(ChunkFwdO);
const std::array<const aclTensor *, 1> ChunkFwdO(
const aclTensor *q,
const aclTensor *k,
const aclTensor *v,
const aclTensor *h,
const aclTensor *g,
const aclIntArray *cuSeqlensOptional,
const aclIntArray *chunkOffsetsOptional,
double scale,
int64_t chunkSize,
const aclTensor *oOut,
aclOpExecutor *executor)
{
L0_DFX(ChunkFwdO, q, k, v, h, g, cuSeqlensOptional, chunkOffsetsOptional, scale, chunkSize, oOut);
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 *actualChunkOffsets = nullptr;
if (chunkOffsetsOptional) {
actualChunkOffsets = executor->ConvertToTensor(chunkOffsetsOptional, DataType::DT_INT64);
const_cast<aclTensor *>(actualChunkOffsets)->SetStorageFormat(Format::FORMAT_ND);
const_cast<aclTensor *>(actualChunkOffsets)->SetViewFormat(Format::FORMAT_ND);
const_cast<aclTensor *>(actualChunkOffsets)->SetOriginalFormat(Format::FORMAT_ND);
} else {
actualChunkOffsets = nullptr;
}
auto ret = ADD_TO_LAUNCHER_LIST_AICORE(ChunkFwdO,
OP_INPUT(q, k, v, h, g, actualCuSeqlens, actualChunkOffsets),
OP_OUTPUT(oOut),
OP_ATTR(scale, chunkSize));
if (ret != ACLNN_SUCCESS) {
OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ADD_TO_LAUNCHER_LIST_AICORE failed.");
return {nullptr};
}
return {oOut};
}
} // namespace l0op

View File

@@ -0,0 +1,30 @@
/**
* 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_FWD_O_H
#define OP_API_INC_LEVEL0_OP_CHUNK_FWD_O_H
#include "opdev/op_executor.h"
namespace l0op {
const std::array<const aclTensor *, 1> ChunkFwdO(
const aclTensor *q,
const aclTensor *k,
const aclTensor *v,
const aclTensor *h,
const aclTensor *g,
const aclIntArray *cuSeqlensOptional,
const aclIntArray *chunkOffsetsOptional,
double scale,
int64_t chunkSize,
const aclTensor *oOut,
aclOpExecutor *executor);
}
#endif

View File

@@ -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

View File

@@ -0,0 +1,394 @@
/**
* 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_FWDO_OUTPUT_HPP
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDO_OUTPUT_HPP
#include "catlass/catlass.hpp"
#include "catlass/arch/resource.hpp"
#include "../gdn_fwd_o_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 AInputType_,
class HInputType_
>
class BlockEpilogue <
EpilogueAtlasGDNFwdOOutput,
HOutputType_,
GInputType_,
AInputType_,
HInputType_
> {
public:
// Type aliases
using DispatchPolicy = EpilogueAtlasGDNFwdOOutput;
using ArchTag = typename DispatchPolicy::ArchTag;
using HElementOutput = typename HOutputType_::Element;
using GElementInput = typename GInputType_::Element;
using AElementInput = typename AInputType_::Element;
using HElementInput = typename HInputType_::Element;
// using CopyGmToUbInput = Tile::CopyGm2Ub<ArchTag, InputType_>;
// using CopyUbToGmOutput = Tile::CopyUb2Gm<ArchTag, OutputType_>;
static constexpr uint32_t HALF_ELENUM_PER_BLK = 16;
static constexpr uint32_t FLOAT_ELENUM_PER_BLK = 8;
static constexpr uint32_t HALF_ELENUM_PER_VECCALC = 128;
static constexpr uint32_t FLOAT_ELENUM_PER_VECCALC = 64;
static constexpr uint32_t UB_TILE_SIZE = 16384; // 64 * 128 * 2B
static constexpr uint32_t UB_LINE_SIZE = 512; // 128 * 2 * 2B
static constexpr uint32_t HALF_ELENUM_PER_LINE = 256; // 128 * 2
static constexpr uint32_t FLOAT_ELENUM_PER_LINE = 128; // 128
static constexpr uint32_t MULTIPLIER = 2;
CATLASS_DEVICE
BlockEpilogue(Arch::Resource<ArchTag> &resource)
{
constexpr uint32_t BASE = 0;
constexpr uint32_t MASK_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_SIZE = 40 * UB_LINE_SIZE;
constexpr uint32_t GBRCUP_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t FLOAT_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t HALF_UB_TENSOR_SIZE = 16 * UB_LINE_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
constexpr uint32_t G_FLOAT_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
constexpr uint32_t MASK_UB_TENSOR_OFFSET = BASE;
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_OFFSET = MASK_UB_TENSOR_OFFSET + MASK_UB_TENSOR_SIZE;
constexpr uint32_t GBRCUP_UB_TENSOR_OFFSET = GBRCLEFTCAST_UB_TENSOR_OFFSET + GBRCLEFTCAST_UB_TENSOR_SIZE;
constexpr uint32_t GCOMP_UB_TENSOR_OFFSET = GBRCUP_UB_TENSOR_OFFSET + GBRCUP_UB_TENSOR_SIZE;
constexpr uint32_t SHARE_UB_TENSOR_OFFSET = GCOMP_UB_TENSOR_OFFSET + G_FLOAT_UB_TENSOR_SIZE;
maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(MASK_UB_TENSOR_OFFSET);
gbrcLeftcastUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCLEFTCAST_UB_TENSOR_OFFSET);
gbrcUpUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCUP_UB_TENSOR_OFFSET);
gcompUbTensor = resource.ubBuf.template GetBufferByByte<float>(GCOMP_UB_TENSOR_OFFSET);
shareUbTensor = resource.ubBuf.template GetBufferByByte<uint8_t>(SHARE_UB_TENSOR_OFFSET);
constexpr uint32_t G_UB_TENSOR_OFFSET_PING = SHARE_UB_TENSOR_OFFSET + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PING = G_UB_TENSOR_OFFSET_PING + G_FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t A_UB_TENSOR_OFFSET_PING = G_HALF_UB_TENSOR_OFFSET_PING + G_HALF_UB_TENSOR_SIZE;
constexpr uint32_t H_UB_TENSOR_OFFSET_PING = A_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PING = H_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PING = OUT_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
gUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PING);
gUbFPTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
gUbBFTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
aUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PING);
hUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(H_UB_TENSOR_OFFSET_PING);
outUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PING);
outUbFPTensorPing = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
outUbBFTensorPing = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
constexpr uint32_t G_UB_TENSOR_OFFSET_PONG = OUT_HALF_UB_TENSOR_OFFSET_PING + HALF_UB_TENSOR_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PONG = G_UB_TENSOR_OFFSET_PONG + G_FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t A_UB_TENSOR_OFFSET_PONG = G_HALF_UB_TENSOR_OFFSET_PONG + G_HALF_UB_TENSOR_SIZE;
constexpr uint32_t H_UB_TENSOR_OFFSET_PONG = A_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PONG = H_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PONG = OUT_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
gUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PONG);
gUbFPTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
gUbBFTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
aUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PONG);
hUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(H_UB_TENSOR_OFFSET_PONG);
outUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PONG);
outUbFPTensorPong = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
outUbBFTensorPong = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
}
CATLASS_DEVICE
~BlockEpilogue()
{}
CATLASS_DEVICE
void operator()(
AscendC::GlobalTensor<HElementOutput> hOutput,
AscendC::GlobalTensor<GElementInput> gInput,
AscendC::GlobalTensor<AElementInput> attnInput,
AscendC::GlobalTensor<HElementInput> hInput,
float scale,
uint32_t chunkSize,
uint32_t kHeadDim,
uint32_t vHeadDim,
uint32_t &pingpongFlag
, uint32_t batchIdx, uint32_t headIdx, uint32_t chunkIdx
)
{
uint32_t mActual = chunkSize;
uint32_t nActual = vHeadDim;
uint32_t alignedM = CeilDiv(nActual, 8) * 8;
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
uint32_t subBlockNum = AscendC::GetSubBlockNum();
uint32_t blockIdx = AscendC::GetBlockIdx();
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 offsetA = mOffset * nActual + nOffset;
uint32_t gbrcStart, gbrcRealStart, gbrcRealEnd, gbrcRealProcess, gbrcEffStart, gbrcEffEnd, mulsRemain, mulsRemainIdx;
if(mActualThisSubBlock <= 32)
{
if(subBlockIdx == 0)
{
gbrcStart = 0;
gbrcRealStart = 0;
gbrcRealProcess = mActualThisSubBlock;
}
else
{
gbrcStart = mActualPerSubBlock;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActual - gbrcRealStart;
}
gbrcEffStart = gbrcStart - gbrcRealStart;
uint32_t dstShape_[2] = {gbrcRealProcess, nActual};
uint32_t srcShape_[2] = {gbrcRealProcess, 1};
AscendC::ResetMask();
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
AscendC::GlobalTensor<HElementInput> hInputThisSubBlock = hInput[gbrcStart * nActual];
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
AscendC::GlobalTensor<HElementOutput> hOutputThisSubBlock = hOutput[gbrcStart * nActual];
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
AscendC::LocalTensor<float> hUbTensor = (pingpongFlag == 0) ? hUbTensorPing : hUbTensorPong;
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
AscendC::LocalTensor<HElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
AscendC::LocalTensor<HElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
if constexpr(std::is_same<GElementInput, float>::value) {
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
} else {
AscendC::DataCopy(gUbFPTensor, gInputThisSubBlock, mActual);
}
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
if constexpr(!std::is_same<GElementInput, float>::value) {
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
AscendC::PipeBarrier<PIPE_V>();
}
AscendC::Adds(gcompUbTensor, gUbTensor, (float)0.0, mActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisSubBlock * nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisSubBlock * nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
AscendC::Exp(gcompUbTensor, gcompUbTensor, mActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstShape_, srcShape_, shareUbTensor);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::Mul(gbrcUpUbTensor, hUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*nActual], mActualThisSubBlock * nActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
AscendC::Add(gbrcUpUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisSubBlock * nActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Muls(outUbTensor, gbrcUpUbTensor, (float)scale, mActualThisSubBlock * nActual);
AscendC::PipeBarrier<PIPE_V>();
if(std::is_same<HElementOutput, half>::value)
{
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(hOutputThisSubBlock, outUbFPTensor, mActualThisSubBlock * nActual);
}
else
{
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(hOutputThisSubBlock, outUbBFTensor, mActualThisSubBlock * nActual);
}
pingpongFlag = 1 - pingpongFlag;
}
else
{
AscendC::ResetMask();
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
if constexpr(std::is_same<GElementInput, float>::value) {
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
} else {
AscendC::DataCopy(gUbFPTensor, gInputThisSubBlock, mActual);
}
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
if constexpr(!std::is_same<GElementInput, float>::value) {
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
AscendC::PipeBarrier<PIPE_V>();
}
AscendC::Adds(gcompUbTensor, gUbTensor, (float)0.0, mActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Exp(gcompUbTensor, gcompUbTensor, mActual);
AscendC::PipeBarrier<PIPE_V>();
uint32_t mActualPerStage = CeilDiv(mActualThisSubBlock, 2);
uint32_t mActualThisStage = 0;
for(uint32_t stage = 0; stage < 2; stage++)
{
if(stage == 0) mActualThisStage = mActualPerStage;
else mActualThisStage = mActualThisSubBlock - mActualPerStage;
if(subBlockIdx == 0 && stage == 0)
{
gbrcStart = 0;
gbrcRealStart = 0;
gbrcRealProcess = mActualThisStage;
}
else if(subBlockIdx == 0 && stage == 1)
{
gbrcStart = mActualPerStage;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActualThisSubBlock - gbrcRealStart;
}
else if(subBlockIdx == 1 && stage == 0)
{
gbrcStart = mActualPerSubBlock;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActualPerSubBlock + mActualThisStage - gbrcRealStart;
}
else if(subBlockIdx == 1 && stage == 1)
{
gbrcStart = mActualPerSubBlock + mActualPerStage;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActual - gbrcRealStart;
}
gbrcEffStart = gbrcStart - gbrcRealStart;
uint32_t dstShape_[2] = {gbrcRealProcess, nActual};
uint32_t srcShape_[2] = {gbrcRealProcess, 1};
AscendC::GlobalTensor<HElementOutput> hOutputThisSubBlock = hOutput[gbrcStart * nActual];
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
AscendC::GlobalTensor<HElementInput> hInputThisSubBlock = hInput[gbrcStart * nActual];
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
AscendC::LocalTensor<float> hUbTensor = (pingpongFlag == 0) ? hUbTensorPing : hUbTensorPong;
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
AscendC::LocalTensor<HElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
AscendC::LocalTensor<HElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisStage * nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisStage * nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstShape_, srcShape_, shareUbTensor);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::Mul(gbrcUpUbTensor, hUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*nActual], mActualThisStage * nActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
AscendC::Add(gbrcUpUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisStage * nActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Muls(outUbTensor, gbrcUpUbTensor, (float)scale, mActualThisStage * nActual);
AscendC::PipeBarrier<PIPE_V>();
if(std::is_same<HElementOutput, half>::value)
{
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * nActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(hOutputThisSubBlock, outUbFPTensor, mActualThisStage * nActual);
}
else
{
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * nActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(hOutputThisSubBlock, outUbBFTensor, mActualThisStage * nActual);
}
pingpongFlag = 1 - pingpongFlag;
}
}
}
private:
AscendC::LocalTensor<float> maskUbTensor;
AscendC::LocalTensor<float> gbrcLeftcastUbTensor;
AscendC::LocalTensor<float> gbrcUpUbTensor;
AscendC::LocalTensor<float> gcompUbTensor;
AscendC::LocalTensor<uint8_t> shareUbTensor;
AscendC::LocalTensor<float> gUbTensorPing;
AscendC::LocalTensor<GElementInput> gUbFPTensorPing;
AscendC::LocalTensor<GElementInput> gUbBFTensorPing;
AscendC::LocalTensor<float> aUbTensorPing;
AscendC::LocalTensor<float> hUbTensorPing;
AscendC::LocalTensor<float> outUbTensorPing;
AscendC::LocalTensor<HElementOutput> outUbFPTensorPing;
AscendC::LocalTensor<HElementOutput> outUbBFTensorPing;
AscendC::LocalTensor<float> gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensorPong;
AscendC::LocalTensor<float> aUbTensorPong;
AscendC::LocalTensor<float> hUbTensorPong;
AscendC::LocalTensor<float> outUbTensorPong;
AscendC::LocalTensor<HElementOutput> outUbFPTensorPong;
AscendC::LocalTensor<HElementOutput> outUbBFTensorPong;
};
}
#endif

View File

@@ -0,0 +1,425 @@
/**
* 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_FWDO_QKMASK_HPP
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDO_QKMASK_HPP
#include "catlass/catlass.hpp"
#include "catlass/arch/resource.hpp"
#include "../gdn_fwd_o_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 AOutputType_,
class GInputType_,
class AInputType_,
class MaskInputType_
>
class BlockEpilogue <
EpilogueAtlasGDNFwdOQkmask,
AOutputType_,
GInputType_,
AInputType_,
MaskInputType_
> {
public:
// Type aliases
using DispatchPolicy = EpilogueAtlasGDNFwdOQkmask;
using ArchTag = typename DispatchPolicy::ArchTag;
using AElementOutput = typename AOutputType_::Element;
using GElementInput = typename GInputType_::Element;
using AElementInput = typename AInputType_::Element;
using MaskElementInput = typename MaskInputType_::Element;
static constexpr uint32_t HALF_ELENUM_PER_BLK = 16;
static constexpr uint32_t FLOAT_ELENUM_PER_BLK = 8;
static constexpr uint32_t HALF_ELENUM_PER_VECCALC = 128;
static constexpr uint32_t FLOAT_ELENUM_PER_VECCALC = 64;
static constexpr uint32_t UB_TILE_SIZE = 16384; // 64 * 128 * 2B
static constexpr uint32_t UB_LINE_SIZE = 512; // 128 * 2 * 2B
static constexpr uint32_t HALF_ELENUM_PER_LINE = 256; // 128 * 2
static constexpr uint32_t FLOAT_ELENUM_PER_LINE = 128; // 128
static constexpr uint32_t MULTIPLIER = 2;
CATLASS_DEVICE
BlockEpilogue(Arch::Resource<ArchTag> &resource)
{
constexpr uint32_t BASE = 0;
constexpr uint32_t MASK_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_SIZE = 40 * UB_LINE_SIZE;
constexpr uint32_t GBRCUP_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t FLOAT_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t HALF_UB_TENSOR_SIZE = 16 * UB_LINE_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
constexpr uint32_t G_FLOAT_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
constexpr uint32_t MASK_UB_TENSOR_OFFSET = BASE;
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_OFFSET = MASK_UB_TENSOR_OFFSET + MASK_UB_TENSOR_SIZE;
constexpr uint32_t GBRCUP_UB_TENSOR_OFFSET = GBRCLEFTCAST_UB_TENSOR_OFFSET + GBRCLEFTCAST_UB_TENSOR_SIZE;
constexpr uint32_t GCOMP_UB_TENSOR_OFFSET = GBRCUP_UB_TENSOR_OFFSET + GBRCUP_UB_TENSOR_SIZE;
constexpr uint32_t SHARE_UB_TENSOR_OFFSET = GCOMP_UB_TENSOR_OFFSET + G_FLOAT_UB_TENSOR_SIZE;
maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(MASK_UB_TENSOR_OFFSET);
gbrcLeftcastUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCLEFTCAST_UB_TENSOR_OFFSET);
gbrcUpUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCUP_UB_TENSOR_OFFSET);
gcompUbTensor = resource.ubBuf.template GetBufferByByte<float>(GCOMP_UB_TENSOR_OFFSET);
shareUbTensor = resource.ubBuf.template GetBufferByByte<uint8_t>(SHARE_UB_TENSOR_OFFSET);
constexpr uint32_t G_UB_TENSOR_OFFSET_PING = SHARE_UB_TENSOR_OFFSET + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PING = G_UB_TENSOR_OFFSET_PING + G_FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t A_UB_TENSOR_OFFSET_PING = G_HALF_UB_TENSOR_OFFSET_PING + G_HALF_UB_TENSOR_SIZE;
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PING = A_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PING = OUT_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
gUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PING);
gUbFPTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
gUbBFTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
aUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PING);
outUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PING);
outUbFPTensorPing = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
outUbBFTensorPing = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
constexpr uint32_t G_UB_TENSOR_OFFSET_PONG = 32 * UB_LINE_SIZE + OUT_HALF_UB_TENSOR_OFFSET_PING + HALF_UB_TENSOR_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PONG = G_UB_TENSOR_OFFSET_PONG + G_FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t A_UB_TENSOR_OFFSET_PONG = G_HALF_UB_TENSOR_OFFSET_PONG + G_HALF_UB_TENSOR_SIZE;
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PONG = A_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PONG = OUT_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
gUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PONG);
gUbFPTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
gUbBFTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
aUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PONG);
outUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PONG);
outUbFPTensorPong = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
outUbBFTensorPong = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
}
CATLASS_DEVICE
~BlockEpilogue()
{}
CATLASS_DEVICE
void operator()(
AscendC::GlobalTensor<AElementOutput> maskOutput,
AscendC::GlobalTensor<GElementInput> gInput,
AscendC::GlobalTensor<AElementInput> attnInput,
AscendC::GlobalTensor<MaskElementInput> boolInput,
uint32_t fullChunkSize,
uint32_t chunkSize,
uint32_t kHeadDim,
uint32_t vHeadDim,
uint32_t &pingpongFlag
, uint32_t batchIdx, uint32_t headIdx, uint32_t chunkIdx
)
{
uint32_t mActual = chunkSize;
uint32_t nActual = chunkSize;
uint32_t alignedNActual = CeilDiv(nActual, 16) * 16;
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
uint32_t subBlockNum = AscendC::GetSubBlockNum();
uint32_t blockIdx = AscendC::GetBlockIdx();
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 offsetA = mOffset * nActual + nOffset;
uint16_t aInputDstStride;
if((nActual - 1) % 16 <= 7) aInputDstStride = 1;
else aInputDstStride = 0;
uint32_t gbrcStart, gbrcRealStart, gbrcRealEnd, gbrcRealProcess, gbrcEffStart, gbrcEffEnd, mulsRemain, mulsRemainIdx;
if(mActualThisSubBlock <= 32)
{ if(subBlockIdx == 0)
{
gbrcStart = 0;
gbrcRealStart = 0;
gbrcRealProcess = mActualThisSubBlock;
}
else
{
gbrcStart = mActualPerSubBlock;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActual - gbrcRealStart;
}
gbrcEffStart = gbrcStart - gbrcRealStart;
gbrcEffEnd = gbrcEffStart + mActualThisSubBlock;
uint32_t dstUpShape_[2] = {mActualThisSubBlock, alignedNActual};
uint32_t srcUpShape_[2] = {1, alignedNActual};
uint32_t dstLeftShape_[2] = {gbrcRealProcess, alignedNActual};
uint32_t srcLeftShape_[2] = {gbrcRealProcess, 1};
AscendC::ResetMask();
AscendC::GlobalTensor<AElementOutput> maskOutputThisSubBlock = maskOutput[gbrcStart * nActual];
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
AscendC::DataCopyParams aInputUbParams{(uint16_t)mActualThisSubBlock, (uint16_t)(nActual*sizeof(float)), 0, aInputDstStride};
AscendC::DataCopyPadParams aInputUbPadParams{false, 0, 0, 0};
AscendC::DataCopyExtParams aOutputUbParams{(uint16_t)mActualThisSubBlock, (uint32_t)(nActual*sizeof(half)), 0, 0, 0};
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
AscendC::LocalTensor<AElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
AscendC::LocalTensor<AElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
if constexpr(std::is_same<GElementInput, float>::value) {
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
} else {
AscendC::DataCopy(gUbFPTensor, gInputThisSubBlock, mActual);
}
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
if constexpr(!std::is_same<GElementInput, float>::value) {
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
AscendC::PipeBarrier<PIPE_V>();
}
AscendC::Adds(gcompUbTensor, gUbTensor, (float)0.0, mActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Broadcast<float, 2, 0>(gbrcUpUbTensor, gcompUbTensor, dstUpShape_, srcUpShape_, shareUbTensor);
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstLeftShape_, srcLeftShape_, shareUbTensor);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Sub(gbrcUpUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*alignedNActual], gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Mins(gbrcUpUbTensor, gbrcUpUbTensor, (float)0.0, mActualThisSubBlock * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Exp(gbrcUpUbTensor, gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
gbrcRealEnd = CeilDiv(gbrcStart + mActualThisSubBlock, 8) * 8;
AscendC::Mul(gbrcUpUbTensor[gbrcRealStart], gbrcUpUbTensor[gbrcRealStart], maskUbTensor[gbrcEffStart * 64], gbrcRealEnd - gbrcRealStart, mActualThisSubBlock,
{1, 1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(64/8)});
AscendC::PipeBarrier<PIPE_V>();
mulsRemain = alignedNActual - gbrcRealEnd;
mulsRemainIdx = gbrcRealEnd;
while(mulsRemain > 64)
{
AscendC::Muls(gbrcUpUbTensor[mulsRemainIdx], gbrcUpUbTensor[mulsRemainIdx], (float)0.0, 64, mActualThisSubBlock,
{1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8)});
mulsRemain -= 64;
mulsRemainIdx += 64;
}
AscendC::Muls(gbrcUpUbTensor[mulsRemainIdx], gbrcUpUbTensor[mulsRemainIdx], (float)0.0, mulsRemain, mActualThisSubBlock,
{1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8)});
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
if(chunkSize==fullChunkSize) AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisSubBlock*nActual);
else AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisSubBlock*nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::Mul(outUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
if(std::is_same<AElementOutput, half>::value)
{
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * alignedNActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(maskOutputThisSubBlock, outUbFPTensor, mActualThisSubBlock*nActual);
}
else
{
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * alignedNActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(maskOutputThisSubBlock, outUbBFTensor, mActualThisSubBlock*nActual);
}
pingpongFlag = 1 - pingpongFlag;
}
else // mActualThisSubBlock > 32 ; <=64
{
AscendC::ResetMask();
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
if constexpr(std::is_same<GElementInput, float>::value) {
AscendC::DataCopy(gUbTensor, gInputThisSubBlock, mActual);
} else {
AscendC::DataCopy(gUbFPTensor, gInputThisSubBlock, mActual);
}
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
if constexpr(!std::is_same<GElementInput, float>::value) {
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
AscendC::PipeBarrier<PIPE_V>();
}
AscendC::Adds(gcompUbTensor, gUbTensor, (float)0.0, mActual);
AscendC::PipeBarrier<PIPE_V>();
uint32_t mActualPerStage = CeilDiv(mActualThisSubBlock, 2);
uint32_t mActualThisStage = 0;
for(uint32_t stage = 0; stage < 2; ++stage)
{
if(stage==0) mActualThisStage = mActualPerStage;
else mActualThisStage = mActualThisSubBlock - mActualPerStage;
if(subBlockIdx == 0 && stage == 0)
{
gbrcStart = 0;
gbrcRealStart = 0;
gbrcRealProcess = mActualThisStage;
}
else if(subBlockIdx == 0 && stage == 1)
{
gbrcStart = mActualPerStage;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActualThisSubBlock - gbrcRealStart;
}
else if(subBlockIdx == 1 && stage == 0)
{
gbrcStart = mActualPerSubBlock;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActualPerSubBlock + mActualThisStage - gbrcRealStart;
}
else if(subBlockIdx == 1 && stage == 1)
{
gbrcStart = mActualPerSubBlock + mActualPerStage;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActual - gbrcRealStart;
}
gbrcEffStart = gbrcStart - gbrcRealStart;
AscendC::GlobalTensor<AElementOutput> maskOutputThisSubBlock = maskOutput[gbrcStart * nActual];
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
AscendC::DataCopyParams aInputUbParams{(uint16_t)mActualThisStage, (uint16_t)(nActual*sizeof(float)), 0, aInputDstStride};
AscendC::DataCopyPadParams aInputUbPadParams{false, 0, 0, 0};
AscendC::DataCopyExtParams aOutputUbParams{(uint16_t)mActualThisStage, (uint32_t)(nActual*sizeof(half)), 0, 0, 0};
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
AscendC::LocalTensor<AElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
AscendC::LocalTensor<AElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
if(chunkSize==fullChunkSize) AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisStage*nActual);
else AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisStage*nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
uint32_t dstUpShape_[2] = {mActualThisStage, alignedNActual};
uint32_t srcUpShape_[2] = {1, alignedNActual};
uint32_t dstLeftShape_[2] = {gbrcRealProcess, alignedNActual};
uint32_t srcLeftShape_[2] = {gbrcRealProcess, 1};
// 310P: Broadcast + gating + causal mask via row loops (strided Mul/Muls banned)
AscendC::Broadcast<float, 2, 0>(gbrcUpUbTensor, gcompUbTensor, dstUpShape_, srcUpShape_, shareUbTensor);
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstLeftShape_, srcLeftShape_, shareUbTensor);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Sub(gbrcUpUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*alignedNActual], gbrcUpUbTensor, mActualThisStage * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Mins(gbrcUpUbTensor, gbrcUpUbTensor, (float)0.0, mActualThisStage * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Exp(gbrcUpUbTensor, gbrcUpUbTensor, mActualThisStage * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
// Causal mask: zero upper triangle row by row
// Use Duplicate for count >= 8, skip for count < 8
// (near-diagonal positions have negligible impact on the causal gate)
for (uint32_t row = 0; row < mActualThisStage; ++row) {
uint32_t globalRow = gbrcStart + row;
uint32_t validCols = globalRow + 1;
if (validCols > alignedNActual) validCols = alignedNActual;
uint32_t rowOff = row * alignedNActual;
uint32_t zeroLen = alignedNActual - validCols;
if (zeroLen >= 8) {
AscendC::Duplicate<float>(gbrcUpUbTensor[rowOff + validCols], (float)0.0, zeroLen);
} else if (zeroLen > 0) {
for (uint32_t c = 0; c < zeroLen; ++c) {
gbrcUpUbTensor.SetValue(rowOff + validCols + c, (float)0.0);
}
}
}
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::Mul(outUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisStage * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
if(std::is_same<AElementOutput, half>::value)
{
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * alignedNActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(maskOutputThisSubBlock, outUbFPTensor, mActualThisStage*nActual);
}
else
{
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * alignedNActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(maskOutputThisSubBlock, outUbBFTensor, mActualThisStage*nActual);
}
pingpongFlag = 1 - pingpongFlag;
}
}
}
private:
AscendC::LocalTensor<float> maskUbTensor;
AscendC::LocalTensor<float> gbrcLeftcastUbTensor;
AscendC::LocalTensor<float> gbrcUpUbTensor;
AscendC::LocalTensor<float> gcompUbTensor;
AscendC::LocalTensor<uint8_t> shareUbTensor;
AscendC::LocalTensor<float> gUbTensorPing;
AscendC::LocalTensor<GElementInput> gUbFPTensorPing;
AscendC::LocalTensor<GElementInput> gUbBFTensorPing;
AscendC::LocalTensor<float> aUbTensorPing;
AscendC::LocalTensor<float> outUbTensorPing;
AscendC::LocalTensor<AElementOutput> outUbFPTensorPing;
AscendC::LocalTensor<AElementOutput> outUbBFTensorPing;
AscendC::LocalTensor<float> gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensorPong;
AscendC::LocalTensor<float> aUbTensorPong;
AscendC::LocalTensor<float> outUbTensorPong;
AscendC::LocalTensor<AElementOutput> outUbFPTensorPong;
AscendC::LocalTensor<AElementOutput> outUbBFTensorPong;
};
}
#endif

View File

@@ -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_O_EPILOGUE_POLICIES_HPP
#define CATLASS_EPILOGUE_GDN_FWD_O_EPILOGUE_POLICIES_HPP
#include "catlass/catlass.hpp"
namespace Catlass::Epilogue {
struct EpilogueAtlasGDNFwdOQkmask {
using ArchTag = Arch::AtlasA2;
};
struct EpilogueAtlasGDNFwdOOutput {
using ArchTag = Arch::AtlasA2;
};
} // namespace Catlass::Epilogue
#endif // CATLASS_EPILOGUE_GDN_FWD_O_EPILOGUE_POLICIES_HPP

View File

@@ -0,0 +1,274 @@
/**
 * 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_GEMM_SCHEDULER_GDN_FWD_O_HPP
#define CATLASS_GEMM_SCHEDULER_GDN_FWD_O_HPP
// constexpr uint32_t PING_PONG_STAGES = 1;
constexpr uint32_t PING_PONG_STAGES = 2;
constexpr uint32_t BYTE_SIZE_16_BIT = 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 GDNFwdOOffsets {
uint32_t qkOffset;
uint32_t ovOffset;
uint32_t hOffset;
uint32_t gOffset;
uint32_t attnWorkOffset;
uint32_t hvWorkOffset;
bool isFinalState;
uint32_t blockTokens;
uint32_t batchIdx;
uint32_t headIdx;
uint32_t chunkIdx;
};
struct BlockSchedulerGdnFwdO {
uint32_t shapeBatch;
uint32_t seqlen;
uint32_t kNumHead;
uint32_t vNumHead;
uint32_t kHeadDim;
uint32_t vHeadDim;
uint32_t chunkSize;
uint32_t isVariedLen;
uint32_t tokenBatch;
uint32_t numChunks{0};
uint32_t vBlockSize{128};
uint32_t taskIdx;
uint32_t cubeCoreIdx;
uint32_t cubeCoreNum;
uint32_t vLoops;
uint32_t taskNum;
uint32_t headGroups;
bool isRunning;
bool processNewTask {true};
bool firstLoop {true};
bool lastLoop {false};
GDNFwdOOffsets 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 batchChunkIdx;
uint32_t batchChunkStartIdx;
uint32_t tokenOffset;
uint32_t batchChunks;
uint32_t batchTokens;
AscendC::GlobalTensor<int64_t> gmSeqlen;
AscendC::GlobalTensor<int64_t> gmChunkOffsets;
Arch::CrossCoreFlag cube1Done{3};
Arch::CrossCoreFlag vec1Done{4};
Arch::CrossCoreFlag cube2Done{5};
Arch::CrossCoreFlag cube3Done{6};
Arch::CrossCoreFlag vec2Done{7};
CATLASS_DEVICE
BlockSchedulerGdnFwdO() {}
CATLASS_DEVICE
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling, uint32_t coreIdx, uint32_t coreNum) {
__gm__ ChunkFwdOTilingData *__restrict gdnFwdOTilingData = reinterpret_cast<__gm__ ChunkFwdOTilingData *__restrict>(tiling);
shapeBatch = gdnFwdOTilingData->shapeBatch;
seqlen = gdnFwdOTilingData->seqlen;
kNumHead = gdnFwdOTilingData->kNumHead;
vNumHead = gdnFwdOTilingData->vNumHead;
kHeadDim = gdnFwdOTilingData->kHeadDim;
vHeadDim = gdnFwdOTilingData->vHeadDim;
chunkSize = gdnFwdOTilingData->chunkSize;
isVariedLen = gdnFwdOTilingData->isVariedLen;
tokenBatch = gdnFwdOTilingData->tokenBatch;
gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens);
gmChunkOffsets.SetGlobalBuffer((__gm__ int64_t *)chunk_offsets);
if (isVariedLen) {
for (uint32_t b = 1; b <= tokenBatch; b++) {
numChunks += (gmSeqlen.GetValue(b) - gmSeqlen.GetValue(b - 1) + chunkSize - 1) / chunkSize;
}
} else {
numChunks = (seqlen + chunkSize - 1) / chunkSize;
}
cubeCoreIdx = coreIdx;
cubeCoreNum = coreNum;
vLoops = vHeadDim / vBlockSize;
taskNum = vLoops * shapeBatch * numChunks * vNumHead;
headGroups = vNumHead / kNumHead;
taskIdx = cubeCoreIdx * PING_PONG_STAGES;
isRunning = taskIdx < taskNum;
}
CATLASS_DEVICE
void InitTask() {
if (processNewTask) {
if (unlikely(taskIdx >= taskNum)) {
isRunning = false;
}
vIdx = taskIdx / (shapeBatch * numChunks * vNumHead);
shapeBatchIdx = (taskIdx - vIdx * shapeBatch * numChunks * vNumHead) / (numChunks * vNumHead);
chunkIdx = (taskIdx - vIdx * shapeBatch * numChunks * vNumHead - shapeBatchIdx * numChunks * vNumHead) / vNumHead;
baseHeadIdx = taskIdx % vNumHead;
tokenBatchIdx = isVariedLen ? gmChunkOffsets.GetValue(2 * chunkIdx) : 0;
batchChunkIdx = isVariedLen ? gmChunkOffsets.GetValue(2 * chunkIdx + 1) : chunkIdx;
batchChunkStartIdx = chunkIdx - batchChunkIdx;
tokenOffset = isVariedLen ? gmSeqlen.GetValue(tokenBatchIdx) : 0;
batchTokens = isVariedLen ? (gmSeqlen.GetValue(tokenBatchIdx + 1) - tokenOffset) : seqlen;
headInnerIdx = 0;
} else {
headInnerIdx = (headInnerIdx + 1) % PING_PONG_STAGES;
}
vHeadIdx = baseHeadIdx + headInnerIdx;
kHeadIdx = vHeadIdx / headGroups;
offsets[currStage].qkOffset = (shapeBatchIdx * kNumHead * seqlen + kHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize) * kHeadDim;
offsets[currStage].ovOffset = (shapeBatchIdx * vNumHead * seqlen + vHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize) * vHeadDim;
offsets[currStage].hOffset = (shapeBatchIdx * vNumHead * numChunks + vHeadIdx * numChunks + chunkIdx) * kHeadDim * vHeadDim;
offsets[currStage].gOffset = shapeBatchIdx * vNumHead * seqlen + vHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize;
offsets[currStage].attnWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * chunkSize;
offsets[currStage].hvWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * vHeadDim;
offsets[currStage].isFinalState = chunkIdx == (numChunks - 1) || (isVariedLen && gmChunkOffsets.GetValue(2 * chunkIdx + 3) == 0);
offsets[currStage].blockTokens = offsets[currStage].isFinalState ? (batchTokens - batchChunkIdx * chunkSize) : chunkSize;
offsets[currStage].batchIdx = batchIdx;
offsets[currStage].headIdx = vHeadIdx;
offsets[currStage].chunkIdx = chunkIdx;
processNewTask = headInnerIdx == PING_PONG_STAGES - 1;
if (processNewTask) {
taskIdx += PING_PONG_STAGES * cubeCoreNum;
}
currStage = (currStage + 1) % PING_PONG_STAGES;
}
};
struct BlockSchedulerGdnFwdOCube : public BlockSchedulerGdnFwdO {
CATLASS_DEVICE
BlockSchedulerGdnFwdOCube() {}
CATLASS_DEVICE
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling) {
BlockSchedulerGdnFwdO::Init(cu_seqlens, chunk_offsets, tiling, AscendC::GetBlockIdx(), AscendC::GetBlockNum());
}
CATLASS_DEVICE
bool NeedProcessCube1() {
return true;
}
CATLASS_DEVICE
GDNFwdOOffsets& GetCube1Offsets() {
return offsets[(currStage - 1) % PING_PONG_STAGES];
}
CATLASS_DEVICE
GemmCoord GetCube1Shape() {
GDNFwdOOffsets& cube1Offsets = GetCube1Offsets();
return GemmCoord{cube1Offsets.blockTokens, cube1Offsets.blockTokens, kHeadDim};
}
CATLASS_DEVICE
bool NeedProcessCube23() {
if (unlikely(firstLoop)) {
firstLoop = false;
return false;
}
return true;
}
CATLASS_DEVICE
GDNFwdOOffsets& GetCube23Offsets() {
return offsets[(currStage - 2) % PING_PONG_STAGES];
}
CATLASS_DEVICE
GemmCoord GetCube2Shape() {
GDNFwdOOffsets& cube2Offsets = GetCube23Offsets();
return GemmCoord{kHeadDim, vHeadDim, cube2Offsets.blockTokens};
}
CATLASS_DEVICE
GemmCoord GetCube3Shape() {
GDNFwdOOffsets& cube2Offsets = GetCube23Offsets();
return GemmCoord{kHeadDim, vHeadDim, cube2Offsets.blockTokens};
}
};
struct BlockSchedulerGdnFwdOVec : public BlockSchedulerGdnFwdO {
CATLASS_DEVICE
BlockSchedulerGdnFwdOVec() {}
CATLASS_DEVICE
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling) {
BlockSchedulerGdnFwdO::Init(cu_seqlens, chunk_offsets, tiling, AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), AscendC::GetBlockNum());
}
CATLASS_DEVICE
bool NeedProcessVec1() {
return isRunning;
}
CATLASS_DEVICE
bool NeedProcessVec2() {
if (unlikely(firstLoop)) {
firstLoop = false;
return false;
}
return true;
}
CATLASS_DEVICE
GDNFwdOOffsets& GetVec1Offsets() {
return offsets[(currStage - 1) % PING_PONG_STAGES];
}
CATLASS_DEVICE
GDNFwdOOffsets& GetVec2Offsets() {
return offsets[(currStage - 2) % PING_PONG_STAGES];
}
};
} // namespace Catlass::Gemm::Block
#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_O_HPP

View File

@@ -0,0 +1,551 @@
/**
* 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_fwdo_qkmask.hpp"
#include "../../epilogue/block/block_epilogue_gdn_fwdo_output.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_o.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 WORKSPACE_TYPE
>
class GDNFwdOKernel {
public:
using ArchTag = Arch::AtlasA2;
using GDNFwdOOffsets = Catlass::Gemm::Block::GDNFwdOOffsets;
using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdOCube;
using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdOVec;
using DispatchPolicyTla = Gemm::MmadPingpongTlaMulti<ArchTag, true, false>;
using L1TileShapeTla = Shape<_128, _128, _128>;
using L0TileShapeTla = L1TileShapeTla;
using QType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using KType = Gemm::GemmType<INPUT_TYPE, layout::ColumnMajor>;
using AttenType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
using AttenMaskedType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using HType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using OinterType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
using VNEWType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using GType = Gemm::GemmType<G_TYPE, layout::RowMajor>;
using OType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using MaskType = Gemm::GemmType<bool, layout::RowMajor>;
// cube 1
using TileCopyQK = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::ColumnMajor, WORKSPACE_TYPE, layout::RowMajor>;
using BlockMmadQK = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyQK>;
// cube 2
using TileCopyQH = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
using BlockMmadQH = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyQH>;
// cube 3
using TileCopyAttenVNEW = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
using BlockMmadAttenVNEW = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyAttenVNEW>;
// vec 1
using DispatchPolicyGDNFwdOQkmask = Epilogue::EpilogueAtlasGDNFwdOQkmask;
using EpilogueGDNFwdOQkmask = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdOQkmask, AttenMaskedType, GType, AttenType, MaskType>;
// vec 2
using DispatchPolicyGDNFwdOOutput = Epilogue::EpilogueAtlasGDNFwdOOutput;
using EpilogueGDNFwdOOutput = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdOOutput, OType, GType, OinterType, OinterType>;
using ElementQ = typename BlockMmadQK::ElementA;
using LayoutQ = Catlass::layout::RowMajor;
using ElementK = typename BlockMmadQK::ElementB;
using LayoutK = Catlass::layout::ColumnMajor;
using ElementAtten = typename BlockMmadQK::ElementC;
using LayoutAtten = Catlass::layout::RowMajor;
using ElementAttenMasked = typename BlockMmadQH::ElementA;
using LayoutAttenMasked = Catlass::layout::RowMajor;
using ElementH = typename BlockMmadQH::ElementB;
using LayoutH = Catlass::layout::RowMajor;
using ElementOinter = typename BlockMmadQH::ElementC;
using LayoutOinter = Catlass::layout::RowMajor;
using ElementVNEW = typename BlockMmadAttenVNEW::ElementB;
using LayoutVNEW = Catlass::layout::RowMajor;
using ElementG = G_TYPE;
using ElementMask = bool;
using L1TileShape = typename BlockMmadQK::L1TileShape;
uint32_t shapeBatch;
uint32_t seqlen;
uint32_t kNumHead;
uint32_t vNumHead;
uint32_t kHeadDim;
uint32_t vHeadDim;
uint32_t chunkSize;
float scale;
uint32_t numChunks;
uint32_t isVariedLen;
uint32_t tokenBatch;
uint32_t vWorkspaceOffset;
uint32_t hWorkspaceOffset;
uint32_t attnWorkspaceOffset;
uint32_t aftermaskWorkspaceOffset;
uint32_t maskWorkspaceOffset;
AscendC::GlobalTensor<ElementQ> gmQ;
AscendC::GlobalTensor<ElementK> gmK;
AscendC::GlobalTensor<ElementVNEW> gmV;
AscendC::GlobalTensor<ElementH> gmH;
AscendC::GlobalTensor<ElementG> gmG;
AscendC::GlobalTensor<ElementVNEW> gmO;
AscendC::GlobalTensor<ElementOinter> gmVWorkspace;
AscendC::GlobalTensor<ElementOinter> gmHWorkspace;
AscendC::GlobalTensor<ElementAtten> gmAttnWorkspace;
AscendC::GlobalTensor<ElementAttenMasked> gmAftermaskWorkspace;
AscendC::GlobalTensor<ElementMask> gmMask;
CubeScheduler cubeBlockScheduler;
VecScheduler vecBlockScheduler;
Arch::Resource<ArchTag> resource;
__aicore__ inline GDNFwdOKernel() {}
__aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR h, GM_ADDR g,
GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR o, GM_ADDR tiling, GM_ADDR user) {
__gm__ ChunkFwdOTilingData *__restrict gdnFwdOTilingData = reinterpret_cast<__gm__ ChunkFwdOTilingData *__restrict>(tiling);
shapeBatch = gdnFwdOTilingData->shapeBatch;
seqlen = gdnFwdOTilingData->seqlen;
kNumHead = gdnFwdOTilingData->kNumHead;
vNumHead = gdnFwdOTilingData->vNumHead;
kHeadDim = gdnFwdOTilingData->kHeadDim;
vHeadDim = gdnFwdOTilingData->vHeadDim;
scale = gdnFwdOTilingData->scale;
chunkSize = gdnFwdOTilingData->chunkSize;
isVariedLen = gdnFwdOTilingData->isVariedLen;
tokenBatch = gdnFwdOTilingData->tokenBatch;
vWorkspaceOffset = gdnFwdOTilingData->vWorkspaceOffset;
hWorkspaceOffset = gdnFwdOTilingData->hWorkspaceOffset;
attnWorkspaceOffset = gdnFwdOTilingData->attnWorkspaceOffset;
aftermaskWorkspaceOffset = gdnFwdOTilingData->aftermaskWorkspaceOffset;
maskWorkspaceOffset = gdnFwdOTilingData->maskWorkspaceOffset;
gmQ.SetGlobalBuffer((__gm__ ElementQ *)q);
gmK.SetGlobalBuffer((__gm__ ElementK *)k);
gmV.SetGlobalBuffer((__gm__ ElementVNEW *)v);
gmH.SetGlobalBuffer((__gm__ ElementH *)h);
gmG.SetGlobalBuffer((__gm__ ElementG *)g);
gmO.SetGlobalBuffer((__gm__ ElementVNEW *)o);
gmVWorkspace.SetGlobalBuffer((__gm__ ElementOinter *)(user + vWorkspaceOffset));
gmHWorkspace.SetGlobalBuffer((__gm__ ElementOinter *)(user + hWorkspaceOffset));
gmAttnWorkspace.SetGlobalBuffer((__gm__ ElementAtten *)(user + attnWorkspaceOffset));
gmAftermaskWorkspace.SetGlobalBuffer((__gm__ ElementAttenMasked *)(user + aftermaskWorkspaceOffset));
gmMask.SetGlobalBuffer((__gm__ ElementMask *)(user + maskWorkspaceOffset));
cubeBlockScheduler.Init(cu_seqlens, chunk_offsets, tiling);
}
__aicore__ inline void Process() {
ProcessUnifiedCore();
}
__aicore__ inline void InitCausalMask() {
AscendC::LocalTensor<float> maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(0);
// 310P: Duplicate count must be >= 8 (vector width = 8 floats).
// Build lower-triangular mask: row i has 1.0 in cols [0..i], 0.0 elsewhere.
// Fill all 1.0 first, then zero the upper triangle with count >= 8.
AscendC::Duplicate<float>(maskUbTensor, (float)1.0, 64 * 64);
AscendC::PipeBarrier<PIPE_V>();
for (uint32_t i = 0; i < 64; ++i) {
uint32_t zeroStart = i + 1;
uint32_t zeroLen = 64 - zeroStart;
if (zeroLen >= 8) {
AscendC::Duplicate<float>(maskUbTensor[i * 64 + zeroStart], (float)0.0, zeroLen);
} else {
for (uint32_t j = 0; j < zeroLen; ++j) {
maskUbTensor.SetValue(i * 64 + zeroStart + j, (float)0.0);
}
}
}
AscendC::PipeBarrier<PIPE_V>();
}
__aicore__ inline void ProcessUnifiedCore() {
uint32_t coreNum = AscendC::GetBlockNum();
BlockMmadQK blockMmadQK(resource);
BlockMmadQH blockMmadQH(resource);
BlockMmadAttenVNEW blockMmadAttenVNEW(resource);
auto qLayout = tla::MakeLayout<ElementQ, LayoutQ>(shapeBatch * kNumHead * seqlen, kHeadDim);
auto kLayout = tla::MakeLayout<ElementK, LayoutK>(kHeadDim, shapeBatch * kNumHead * seqlen);
auto hLayout = tla::MakeLayout<ElementH, LayoutH>(shapeBatch * vNumHead * seqlen * kHeadDim, vHeadDim);
auto ointerLayout = tla::MakeLayout<ElementOinter, LayoutOinter>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
auto vnewLayout = tla::MakeLayout<ElementVNEW, LayoutVNEW>(shapeBatch * vNumHead * seqlen, vHeadDim);
bool needRun = false;
uint32_t pingpongFlag = 0;
while (cubeBlockScheduler.isRunning) {
cubeBlockScheduler.InitTask();
if (cubeBlockScheduler.isRunning) {
// CUBE1: attn = q @ k.T
GDNFwdOOffsets& cube1Offsets = cubeBlockScheduler.GetCube1Offsets();
auto attenLayout = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, cube1Offsets.blockTokens);
auto tensorQ = tla::MakeTensor(gmQ[cube1Offsets.qkOffset], qLayout, Catlass::Arch::PositionGM{});
auto tensorK = tla::MakeTensor(gmK[cube1Offsets.qkOffset], kLayout, Catlass::Arch::PositionGM{});
auto tensorAttn = tla::MakeTensor(gmAttnWorkspace[cube1Offsets.attnWorkOffset], attenLayout, Catlass::Arch::PositionGM{});
GemmCoord cube1Shape{cube1Offsets.blockTokens, cube1Offsets.blockTokens, kHeadDim};
auto tensorBlockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k()));
auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n()));
auto tensorBlockAttn = GetTile(tensorAttn, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n()));
blockMmadQK.preSetFlags();
blockMmadQK(tensorBlockQ, tensorBlockK, tensorBlockAttn, cube1Shape);
blockMmadQK.finalWaitFlags();
// Re-init causal mask after cube (cube overwrites UB[0])
InitCausalMask();
// VEC1: qkmask epilogue
EpilogueGDNFwdOQkmask epilogueGDNFwdOQkmask(resource);
epilogueGDNFwdOQkmask(
gmAftermaskWorkspace[cube1Offsets.attnWorkOffset],
gmG[cube1Offsets.gOffset], gmAttnWorkspace[cube1Offsets.attnWorkOffset], gmMask,
chunkSize, cube1Offsets.blockTokens, kHeadDim, vHeadDim, pingpongFlag,
cube1Offsets.batchIdx, cube1Offsets.headIdx, cube1Offsets.chunkIdx
);
}
// GM fence: ensure Vec1 MTE3 writes are committed before Cube3 MTE2 reads
AscendC::PipeBarrier<PIPE_ALL>();
if (needRun) {
GDNFwdOOffsets& prevOffsets = cubeBlockScheduler.GetCube23Offsets();
// CUBE2: h_work = q @ h
auto tensorQ2 = tla::MakeTensor(gmQ[prevOffsets.qkOffset], qLayout, Catlass::Arch::PositionGM{});
auto tensorH = tla::MakeTensor(gmH[prevOffsets.hOffset], hLayout, Catlass::Arch::PositionGM{});
auto tensorHWork = tla::MakeTensor(gmHWorkspace[prevOffsets.hvWorkOffset], ointerLayout, Catlass::Arch::PositionGM{});
GemmCoord cube2Shape{prevOffsets.blockTokens, vHeadDim, kHeadDim};
auto tensorBlockQ2 = GetTile(tensorQ2, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k()));
auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n()));
auto tensorBlockHWork = GetTile(tensorHWork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n()));
blockMmadQH.preSetFlags();
blockMmadQH(tensorBlockQ2, tensorBlockH, tensorBlockHWork, cube2Shape);
blockMmadQH.finalWaitFlags();
// CUBE3: v_work = attn_masked @ v
auto attenLayout3 = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, prevOffsets.blockTokens);
auto tensorAttnMask = tla::MakeTensor(gmAftermaskWorkspace[prevOffsets.attnWorkOffset], attenLayout3, Catlass::Arch::PositionGM{});
auto tensorV = tla::MakeTensor(gmV[prevOffsets.ovOffset], vnewLayout, Catlass::Arch::PositionGM{});
auto tensorVWork = tla::MakeTensor(gmVWorkspace[prevOffsets.hvWorkOffset], ointerLayout, Catlass::Arch::PositionGM{});
GemmCoord cube3Shape{prevOffsets.blockTokens, vHeadDim, prevOffsets.blockTokens};
auto tensorBlockAttnMask = GetTile(tensorAttnMask, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.k()));
auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.k(), cube3Shape.n()));
auto tensorBlockVWork = GetTile(tensorVWork, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.n()));
blockMmadAttenVNEW.preSetFlags();
blockMmadAttenVNEW(tensorBlockAttnMask, tensorBlockV, tensorBlockVWork, cube3Shape);
blockMmadAttenVNEW.finalWaitFlags();
// GM fence: ensure Cube2/3 L0C→UB→MTE3→GM writes are committed
AscendC::PipeBarrier<PIPE_ALL>();
// VEC2 inline for 310P: o = scale * (v_work + exp(g) * h_work)
// The epilogue class uses event-based MTE2 sync that breaks after cube matmul on 310P.
{
constexpr uint32_t STAGE_ROWS = 32;
uint32_t bt = prevOffsets.blockTokens;
uint32_t stageCnt = STAGE_ROWS * vHeadDim;
// UB layout: vwUb[0..stageCnt), hwUb[stageCnt..2*stageCnt), gUb[2*stageCnt..+64)
AscendC::LocalTensor<float> vwUb = resource.ubBuf.template GetBufferByByte<float>(0);
AscendC::LocalTensor<float> hwUb = resource.ubBuf.template GetBufferByByte<float>(stageCnt * sizeof(float));
AscendC::LocalTensor<float> gUb = resource.ubBuf.template GetBufferByByte<float>(stageCnt * sizeof(float) * 2);
// outUb (half) after gUb, aligned to 512B
constexpr uint32_t G_RESERVE = 512;
AscendC::LocalTensor<ElementVNEW> outUb = resource.ubBuf.template GetBufferByByte<ElementVNEW>(
stageCnt * sizeof(float) * 2 + G_RESERVE);
for (uint32_t row = 0; row < bt; row += STAGE_ROWS) {
uint32_t rows = (row + STAGE_ROWS <= bt) ? STAGE_ROWS : (bt - row);
uint32_t elems = rows * vHeadDim;
uint32_t gmOff = row * vHeadDim;
// Load v_work, h_work, g from GM
AscendC::DataCopy(vwUb, gmVWorkspace[prevOffsets.hvWorkOffset + gmOff], elems);
AscendC::DataCopy(hwUb, gmHWorkspace[prevOffsets.hvWorkOffset + gmOff], elems);
// Load g (may be float or half)
if constexpr (std::is_same<ElementG, float>::value) {
AscendC::DataCopy(gUb, gmG[prevOffsets.gOffset + row], rows);
} else {
AscendC::LocalTensor<ElementG> gTyped = resource.ubBuf.template GetBufferByByte<ElementG>(
stageCnt * sizeof(float) * 2 + 256);
AscendC::DataCopy(gTyped, gmG[prevOffsets.gOffset + row], rows);
AscendC::PipeBarrier<PIPE_ALL>();
AscendC::Cast(gUb, gTyped, AscendC::RoundMode::CAST_NONE, rows);
}
AscendC::PipeBarrier<PIPE_ALL>();
// exp(g)
AscendC::Exp(gUb, gUb, rows);
AscendC::PipeBarrier<PIPE_V>();
// Broadcast exp(g) into gBrc: each row r gets exp(g[r]) repeated Dv times
// gBrc lives after outUb in UB
AscendC::LocalTensor<float> gBrc = resource.ubBuf.template GetBufferByByte<float>(
stageCnt * sizeof(float) * 2 + G_RESERVE + stageCnt * sizeof(ElementVNEW));
{
uint32_t dstShape[2] = {rows, vHeadDim};
uint32_t srcShape[2] = {rows, 1};
// Broadcast needs a shared temp buffer — use space after gBrc
AscendC::LocalTensor<uint8_t> brcTmp = resource.ubBuf.template GetBufferByByte<uint8_t>(
stageCnt * sizeof(float) * 2 + G_RESERVE + stageCnt * sizeof(ElementVNEW) + elems * sizeof(float));
AscendC::Broadcast<float, 2, 1>(gBrc, gUb, dstShape, srcShape, brcTmp);
}
AscendC::PipeBarrier<PIPE_V>();
AscendC::Mul(hwUb, hwUb, gBrc, elems);
AscendC::PipeBarrier<PIPE_V>();
// v_work + exp(g)*h_work
AscendC::Add(vwUb, vwUb, hwUb, elems);
AscendC::PipeBarrier<PIPE_V>();
// * scale
AscendC::Muls(vwUb, vwUb, (float)scale, elems);
AscendC::PipeBarrier<PIPE_V>();
// Cast to output dtype
AscendC::Cast(outUb, vwUb, AscendC::RoundMode::CAST_NONE, elems);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0);
AscendC::DataCopyParams cp{1, static_cast<uint16_t>(elems * sizeof(ElementVNEW) / 32), 0, 0};
AscendC::DataCopy(gmO[prevOffsets.ovOffset + gmOff], outUb, cp);
AscendC::PipeBarrier<PIPE_ALL>();
}
}
}
needRun = true;
}
}
__aicore__ inline void ProcessSplitCore() {
if ASCEND_IS_AIC {
uint32_t coreIdx = AscendC::GetBlockIdx();
uint32_t coreNum = AscendC::GetBlockNum();
BlockMmadQK blockMmadQK(resource);
BlockMmadQH blockMmadQH(resource);
BlockMmadAttenVNEW blockMmadAttenVNEW(resource);
auto qLayout = tla::MakeLayout<ElementQ, LayoutQ>(shapeBatch * kNumHead * seqlen, kHeadDim);
auto kLayout = tla::MakeLayout<ElementK, LayoutK>(kHeadDim, shapeBatch * kNumHead * seqlen);
auto hLayout = tla::MakeLayout<ElementH, LayoutH>(shapeBatch * vNumHead * seqlen * kHeadDim, vHeadDim);
auto ointerLayout = tla::MakeLayout<ElementOinter, LayoutOinter>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
auto vnewLayout = tla::MakeLayout<ElementVNEW, LayoutVNEW>(shapeBatch * vNumHead * seqlen, vHeadDim);
bool needRun = false;
bool isFirstC3 = true;
while (cubeBlockScheduler.isRunning) {
cubeBlockScheduler.InitTask();
if (cubeBlockScheduler.isRunning && coreIdx < coreNum) {
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
GDNFwdOOffsets& cube1Offsets = cubeBlockScheduler.GetCube1Offsets();
int64_t cube1OffsetQ = cube1Offsets.qkOffset;
int64_t cube1OffsetK = cube1Offsets.qkOffset;
int64_t cube1OffsetAttn = cube1Offsets.attnWorkOffset;
auto attenLayout = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, cube1Offsets.blockTokens);
auto tensorQ = tla::MakeTensor(gmQ[cube1OffsetQ], qLayout, Catlass::Arch::PositionGM{});
auto tensorK = tla::MakeTensor(gmK[cube1OffsetK], kLayout, Catlass::Arch::PositionGM{});
auto tensorAttn = tla::MakeTensor(gmAttnWorkspace[cube1OffsetAttn], attenLayout, Catlass::Arch::PositionGM{});
GemmCoord cube1Shape{cube1Offsets.blockTokens, cube1Offsets.blockTokens, kHeadDim};
auto tensorBlockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k()));
auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n()));
auto tensorBlockAttn = GetTile(tensorAttn, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n()));
blockMmadQK.preSetFlags();
blockMmadQK(tensorBlockQ, tensorBlockK, tensorBlockAttn, cube1Shape);
blockMmadQK.finalWaitFlags();
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done);
}
// AscendC::PipeBarrier<PIPE_ALL>();
if (needRun && coreIdx < coreNum) {
if(!cubeBlockScheduler.isRunning) Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done);
GDNFwdOOffsets& cube2Offsets = cubeBlockScheduler.GetCube23Offsets();
int64_t cube2OffsetQ = cube2Offsets.qkOffset;
int64_t cube2OffsetH = cube2Offsets.hOffset;
int64_t cube2OffsetHWork = cube2Offsets.hvWorkOffset;
auto tensorQ = tla::MakeTensor(gmQ[cube2OffsetQ], qLayout, Catlass::Arch::PositionGM{});
auto tensorH = tla::MakeTensor(gmH[cube2OffsetH], hLayout, Catlass::Arch::PositionGM{});
auto tensorHWork = tla::MakeTensor(gmHWorkspace[cube2OffsetHWork], ointerLayout, Catlass::Arch::PositionGM{});
GemmCoord cube2Shape{cube2Offsets.blockTokens, vHeadDim, kHeadDim};
auto tensorBlockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k()));
auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n()));
auto tensorBlockHWork = GetTile(tensorHWork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n()));
blockMmadQH.preSetFlags();
blockMmadQH(tensorBlockQ, tensorBlockH, tensorBlockHWork, cube2Shape);
blockMmadQH.finalWaitFlags();
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done);
}
if (needRun && coreIdx < coreNum) {
GDNFwdOOffsets& cube3Offsets = cubeBlockScheduler.GetCube23Offsets();
if(isFirstC3) Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
int64_t cube3OffsetAttnMask = cube3Offsets.attnWorkOffset;
int64_t cube3OffsetV = cube3Offsets.ovOffset;
int64_t cube3OffsetVWork = cube3Offsets.hvWorkOffset;
auto attenLayout = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, cube3Offsets.blockTokens);
auto tensorAttnMask = tla::MakeTensor(gmAftermaskWorkspace[cube3OffsetAttnMask], attenLayout, Catlass::Arch::PositionGM{});
auto tensorV = tla::MakeTensor(gmV[cube3OffsetV], vnewLayout, Catlass::Arch::PositionGM{});
auto tensorVWork = tla::MakeTensor(gmVWorkspace[cube3OffsetVWork], ointerLayout, Catlass::Arch::PositionGM{});
GemmCoord cube3Shape{cube3Offsets.blockTokens, vHeadDim, cube3Offsets.blockTokens};
auto tensorBlockAttnMask = GetTile(tensorAttnMask, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.k()));
auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.k(), cube3Shape.n()));
auto tensorBlockVWork = GetTile(tensorVWork, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.n()));
blockMmadAttenVNEW.preSetFlags();
blockMmadAttenVNEW(tensorBlockAttnMask, tensorBlockV, tensorBlockVWork, cube3Shape);
blockMmadAttenVNEW.finalWaitFlags();
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube3Done);
isFirstC3 = false;
}
needRun = true;
// AscendC::PipeBarrier<PIPE_ALL>();
}
if (coreIdx < coreNum) {
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();
AscendC::LocalTensor<float> maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(0);
AscendC::Duplicate<float>(maskUbTensor, (float)0.0, 64*64);
AscendC::PipeBarrier<PIPE_V>();
for(uint32_t i = 0; i < 64; ++ i) AscendC::Duplicate<float>(maskUbTensor[i * 64], (float)1.0, i + 1);
AscendC::PipeBarrier<PIPE_V>();
bool needRun = false;
uint32_t pingpongFlag = 0;
if (coreIdx < coreNum * subBlockNum) {
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
}
while (vecBlockScheduler.isRunning) {
vecBlockScheduler.InitTask();
if (vecBlockScheduler.isRunning && coreIdx < coreNum * subBlockNum) {
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube1Done);
GDNFwdOOffsets& vec1Offsets = vecBlockScheduler.GetVec1Offsets();
int64_t vec1OffsetAttnMask = vec1Offsets.attnWorkOffset;
int64_t vec1OffsetG = vec1Offsets.gOffset;
int64_t vec1OffsetAttn = vec1Offsets.attnWorkOffset;
EpilogueGDNFwdOQkmask epilogueGDNFwdOQkmask(resource);
epilogueGDNFwdOQkmask(
gmAftermaskWorkspace[vec1OffsetAttnMask],
gmG[vec1OffsetG], gmAttnWorkspace[vec1OffsetAttn], gmMask,
chunkSize, vec1Offsets.blockTokens, kHeadDim, vHeadDim, pingpongFlag, vec1Offsets.batchIdx, vec1Offsets.headIdx, vec1Offsets.chunkIdx
);
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
}
// AscendC::PipeBarrier<PIPE_ALL>();
if (needRun && coreIdx < coreNum * subBlockNum) {
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done);
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube3Done);
GDNFwdOOffsets& vec2Offsets = vecBlockScheduler.GetVec2Offsets();
int64_t vec2OffsetO = vec2Offsets.ovOffset;
int64_t vec2OffsetG = vec2Offsets.gOffset;
int64_t vec2OffsetVWork = vec2Offsets.hvWorkOffset;
int64_t vec2OffsetHWork = vec2Offsets.hvWorkOffset;
EpilogueGDNFwdOOutput epilogueGDNFwdOOutput(resource);
epilogueGDNFwdOOutput(
gmO[vec2OffsetO],
gmG[vec2OffsetG], gmVWorkspace[vec2OffsetVWork], gmHWorkspace[vec2OffsetHWork],
scale, vec2Offsets.blockTokens, kHeadDim, vHeadDim, pingpongFlag, vec2Offsets.batchIdx, vec2Offsets.headIdx, vec2Offsets.chunkIdx
);
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
}
// AscendC::PipeBarrier<PIPE_ALL>();
needRun = true;
}
}
}
};
}

View File

@@ -0,0 +1,394 @@
/**
* 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_FWDO_OUTPUT_HPP
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDO_OUTPUT_HPP
#include "catlass/catlass.hpp"
#include "catlass/arch/resource.hpp"
#include "../gdn_fwd_o_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 AInputType_,
class HInputType_
>
class BlockEpilogue <
EpilogueAtlasGDNFwdOOutput,
HOutputType_,
GInputType_,
AInputType_,
HInputType_
> {
public:
// Type aliases
using DispatchPolicy = EpilogueAtlasGDNFwdOOutput;
using ArchTag = typename DispatchPolicy::ArchTag;
using HElementOutput = typename HOutputType_::Element;
using GElementInput = typename GInputType_::Element;
using AElementInput = typename AInputType_::Element;
using HElementInput = typename HInputType_::Element;
// using CopyGmToUbInput = Tile::CopyGm2Ub<ArchTag, InputType_>;
// using CopyUbToGmOutput = Tile::CopyUb2Gm<ArchTag, OutputType_>;
static constexpr uint32_t HALF_ELENUM_PER_BLK = 16;
static constexpr uint32_t FLOAT_ELENUM_PER_BLK = 8;
static constexpr uint32_t HALF_ELENUM_PER_VECCALC = 128;
static constexpr uint32_t FLOAT_ELENUM_PER_VECCALC = 64;
static constexpr uint32_t UB_TILE_SIZE = 16384; // 64 * 128 * 2B
static constexpr uint32_t UB_LINE_SIZE = 512; // 128 * 2 * 2B
static constexpr uint32_t HALF_ELENUM_PER_LINE = 256; // 128 * 2
static constexpr uint32_t FLOAT_ELENUM_PER_LINE = 128; // 128
static constexpr uint32_t MULTIPLIER = 2;
CATLASS_DEVICE
BlockEpilogue(Arch::Resource<ArchTag> &resource)
{
constexpr uint32_t BASE = 0;
constexpr uint32_t MASK_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_SIZE = 40 * UB_LINE_SIZE;
constexpr uint32_t GBRCUP_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t FLOAT_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t HALF_UB_TENSOR_SIZE = 16 * UB_LINE_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
constexpr uint32_t G_FLOAT_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
constexpr uint32_t MASK_UB_TENSOR_OFFSET = BASE;
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_OFFSET = MASK_UB_TENSOR_OFFSET + MASK_UB_TENSOR_SIZE;
constexpr uint32_t GBRCUP_UB_TENSOR_OFFSET = GBRCLEFTCAST_UB_TENSOR_OFFSET + GBRCLEFTCAST_UB_TENSOR_SIZE;
constexpr uint32_t GCOMP_UB_TENSOR_OFFSET = GBRCUP_UB_TENSOR_OFFSET + GBRCUP_UB_TENSOR_SIZE;
constexpr uint32_t SHARE_UB_TENSOR_OFFSET = GCOMP_UB_TENSOR_OFFSET + G_FLOAT_UB_TENSOR_SIZE;
maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(MASK_UB_TENSOR_OFFSET);
gbrcLeftcastUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCLEFTCAST_UB_TENSOR_OFFSET);
gbrcUpUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCUP_UB_TENSOR_OFFSET);
gcompUbTensor = resource.ubBuf.template GetBufferByByte<float>(GCOMP_UB_TENSOR_OFFSET);
shareUbTensor = resource.ubBuf.template GetBufferByByte<uint8_t>(SHARE_UB_TENSOR_OFFSET);
constexpr uint32_t G_UB_TENSOR_OFFSET_PING = SHARE_UB_TENSOR_OFFSET + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PING = G_UB_TENSOR_OFFSET_PING + G_FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t A_UB_TENSOR_OFFSET_PING = G_HALF_UB_TENSOR_OFFSET_PING + G_HALF_UB_TENSOR_SIZE;
constexpr uint32_t H_UB_TENSOR_OFFSET_PING = A_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PING = H_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PING = OUT_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
gUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PING);
gUbFPTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
gUbBFTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
aUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PING);
hUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(H_UB_TENSOR_OFFSET_PING);
outUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PING);
outUbFPTensorPing = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
outUbBFTensorPing = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
constexpr uint32_t G_UB_TENSOR_OFFSET_PONG = OUT_HALF_UB_TENSOR_OFFSET_PING + HALF_UB_TENSOR_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PONG = G_UB_TENSOR_OFFSET_PONG + G_FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t A_UB_TENSOR_OFFSET_PONG = G_HALF_UB_TENSOR_OFFSET_PONG + G_HALF_UB_TENSOR_SIZE;
constexpr uint32_t H_UB_TENSOR_OFFSET_PONG = A_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PONG = H_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PONG = OUT_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
gUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PONG);
gUbFPTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
gUbBFTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
aUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PONG);
hUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(H_UB_TENSOR_OFFSET_PONG);
outUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PONG);
outUbFPTensorPong = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
outUbBFTensorPong = resource.ubBuf.template GetBufferByByte<HElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
}
CATLASS_DEVICE
~BlockEpilogue()
{}
CATLASS_DEVICE
void operator()(
AscendC::GlobalTensor<HElementOutput> hOutput,
AscendC::GlobalTensor<GElementInput> gInput,
AscendC::GlobalTensor<AElementInput> attnInput,
AscendC::GlobalTensor<HElementInput> hInput,
float scale,
uint32_t chunkSize,
uint32_t kHeadDim,
uint32_t vHeadDim,
uint32_t &pingpongFlag
, uint32_t batchIdx, uint32_t headIdx, uint32_t chunkIdx
)
{
uint32_t mActual = chunkSize;
uint32_t nActual = vHeadDim;
uint32_t alignedM = CeilDiv(nActual, 8) * 8;
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
uint32_t subBlockNum = AscendC::GetSubBlockNum();
uint32_t blockIdx = AscendC::GetBlockIdx();
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 offsetA = mOffset * nActual + nOffset;
uint32_t gbrcStart, gbrcRealStart, gbrcRealEnd, gbrcRealProcess, gbrcEffStart, gbrcEffEnd, mulsRemain, mulsRemainIdx;
if(mActualThisSubBlock <= 32)
{
if(subBlockIdx == 0)
{
gbrcStart = 0;
gbrcRealStart = 0;
gbrcRealProcess = mActualThisSubBlock;
}
else
{
gbrcStart = mActualPerSubBlock;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActual - gbrcRealStart;
}
gbrcEffStart = gbrcStart - gbrcRealStart;
uint32_t dstShape_[2] = {gbrcRealProcess, nActual};
uint32_t srcShape_[2] = {gbrcRealProcess, 1};
AscendC::ResetMask();
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
AscendC::GlobalTensor<HElementInput> hInputThisSubBlock = hInput[gbrcStart * nActual];
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
AscendC::GlobalTensor<HElementOutput> hOutputThisSubBlock = hOutput[gbrcStart * nActual];
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
AscendC::LocalTensor<float> hUbTensor = (pingpongFlag == 0) ? hUbTensorPing : hUbTensorPong;
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
AscendC::LocalTensor<HElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
AscendC::LocalTensor<HElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
if constexpr(std::is_same<GElementInput, float>::value) {
AscendC::DataCopyPad(gUbTensor, gInputThisSubBlock, gfloatUbParams, gUbPadParams);
} else {
AscendC::DataCopyPad(gUbFPTensor, gInputThisSubBlock, ghalfUbParams, gUbPadParams);
}
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
if constexpr(!std::is_same<GElementInput, float>::value) {
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
AscendC::PipeBarrier<PIPE_V>();
}
AscendC::Copy(gcompUbTensor, gUbTensor, 64, 2, {1, 1, 8, 8});
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisSubBlock * nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisSubBlock * nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
AscendC::Exp(gcompUbTensor, gcompUbTensor, mActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstShape_, srcShape_, shareUbTensor);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::Mul(gbrcUpUbTensor, hUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*nActual], mActualThisSubBlock * nActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
AscendC::Add(gbrcUpUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisSubBlock * nActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Muls(outUbTensor, gbrcUpUbTensor, (float)scale, mActualThisSubBlock * nActual);
AscendC::PipeBarrier<PIPE_V>();
if(std::is_same<HElementOutput, half>::value)
{
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * nActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(hOutputThisSubBlock, outUbFPTensor, mActualThisSubBlock * nActual);
}
else
{
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * nActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(hOutputThisSubBlock, outUbBFTensor, mActualThisSubBlock * nActual);
}
pingpongFlag = 1 - pingpongFlag;
}
else
{
AscendC::ResetMask();
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
if constexpr(std::is_same<GElementInput, float>::value) {
AscendC::DataCopyPad(gUbTensor, gInputThisSubBlock, gfloatUbParams, gUbPadParams);
} else {
AscendC::DataCopyPad(gUbFPTensor, gInputThisSubBlock, ghalfUbParams, gUbPadParams);
}
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
if constexpr(!std::is_same<GElementInput, float>::value) {
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
AscendC::PipeBarrier<PIPE_V>();
}
AscendC::Copy(gcompUbTensor, gUbTensor, 64, 2, {1, 1, 8, 8});
AscendC::PipeBarrier<PIPE_V>();
AscendC::Exp(gcompUbTensor, gcompUbTensor, mActual);
AscendC::PipeBarrier<PIPE_V>();
uint32_t mActualPerStage = CeilDiv(mActualThisSubBlock, 2);
uint32_t mActualThisStage = 0;
for(uint32_t stage = 0; stage < 2; stage++)
{
if(stage == 0) mActualThisStage = mActualPerStage;
else mActualThisStage = mActualThisSubBlock - mActualPerStage;
if(subBlockIdx == 0 && stage == 0)
{
gbrcStart = 0;
gbrcRealStart = 0;
gbrcRealProcess = mActualThisStage;
}
else if(subBlockIdx == 0 && stage == 1)
{
gbrcStart = mActualPerStage;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActualThisSubBlock - gbrcRealStart;
}
else if(subBlockIdx == 1 && stage == 0)
{
gbrcStart = mActualPerSubBlock;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActualPerSubBlock + mActualThisStage - gbrcRealStart;
}
else if(subBlockIdx == 1 && stage == 1)
{
gbrcStart = mActualPerSubBlock + mActualPerStage;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActual - gbrcRealStart;
}
gbrcEffStart = gbrcStart - gbrcRealStart;
uint32_t dstShape_[2] = {gbrcRealProcess, nActual};
uint32_t srcShape_[2] = {gbrcRealProcess, 1};
AscendC::GlobalTensor<HElementOutput> hOutputThisSubBlock = hOutput[gbrcStart * nActual];
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
AscendC::GlobalTensor<HElementInput> hInputThisSubBlock = hInput[gbrcStart * nActual];
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
AscendC::LocalTensor<float> hUbTensor = (pingpongFlag == 0) ? hUbTensorPing : hUbTensorPong;
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
AscendC::LocalTensor<HElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
AscendC::LocalTensor<HElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::DataCopy(hUbTensor, hInputThisSubBlock, mActualThisStage * nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID2 + pingpongFlag);
AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisStage * nActual);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstShape_, srcShape_, shareUbTensor);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::Mul(gbrcUpUbTensor, hUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*nActual], mActualThisStage * nActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID2 + pingpongFlag);
AscendC::Add(gbrcUpUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisStage * nActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Muls(outUbTensor, gbrcUpUbTensor, (float)scale, mActualThisStage * nActual);
AscendC::PipeBarrier<PIPE_V>();
if(std::is_same<HElementOutput, half>::value)
{
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * nActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(hOutputThisSubBlock, outUbFPTensor, mActualThisStage * nActual);
}
else
{
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisStage * nActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::DataCopy(hOutputThisSubBlock, outUbBFTensor, mActualThisStage * nActual);
}
pingpongFlag = 1 - pingpongFlag;
}
}
}
private:
AscendC::LocalTensor<float> maskUbTensor;
AscendC::LocalTensor<float> gbrcLeftcastUbTensor;
AscendC::LocalTensor<float> gbrcUpUbTensor;
AscendC::LocalTensor<float> gcompUbTensor;
AscendC::LocalTensor<uint8_t> shareUbTensor;
AscendC::LocalTensor<float> gUbTensorPing;
AscendC::LocalTensor<GElementInput> gUbFPTensorPing;
AscendC::LocalTensor<GElementInput> gUbBFTensorPing;
AscendC::LocalTensor<float> aUbTensorPing;
AscendC::LocalTensor<float> hUbTensorPing;
AscendC::LocalTensor<float> outUbTensorPing;
AscendC::LocalTensor<HElementOutput> outUbFPTensorPing;
AscendC::LocalTensor<HElementOutput> outUbBFTensorPing;
AscendC::LocalTensor<float> gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensorPong;
AscendC::LocalTensor<float> aUbTensorPong;
AscendC::LocalTensor<float> hUbTensorPong;
AscendC::LocalTensor<float> outUbTensorPong;
AscendC::LocalTensor<HElementOutput> outUbFPTensorPong;
AscendC::LocalTensor<HElementOutput> outUbBFTensorPong;
};
}
#endif

View File

@@ -0,0 +1,428 @@
/**
* 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_FWDO_QKMASK_HPP
#define CATLASS_EPILOGUE_BLOCK_BLOCK_EPILOGUE_GDN_FWDO_QKMASK_HPP
#include "catlass/catlass.hpp"
#include "catlass/arch/resource.hpp"
#include "../gdn_fwd_o_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 AOutputType_,
class GInputType_,
class AInputType_,
class MaskInputType_
>
class BlockEpilogue <
EpilogueAtlasGDNFwdOQkmask,
AOutputType_,
GInputType_,
AInputType_,
MaskInputType_
> {
public:
// Type aliases
using DispatchPolicy = EpilogueAtlasGDNFwdOQkmask;
using ArchTag = typename DispatchPolicy::ArchTag;
using AElementOutput = typename AOutputType_::Element;
using GElementInput = typename GInputType_::Element;
using AElementInput = typename AInputType_::Element;
using MaskElementInput = typename MaskInputType_::Element;
static constexpr uint32_t HALF_ELENUM_PER_BLK = 16;
static constexpr uint32_t FLOAT_ELENUM_PER_BLK = 8;
static constexpr uint32_t HALF_ELENUM_PER_VECCALC = 128;
static constexpr uint32_t FLOAT_ELENUM_PER_VECCALC = 64;
static constexpr uint32_t UB_TILE_SIZE = 16384; // 64 * 128 * 2B
static constexpr uint32_t UB_LINE_SIZE = 512; // 128 * 2 * 2B
static constexpr uint32_t HALF_ELENUM_PER_LINE = 256; // 128 * 2
static constexpr uint32_t FLOAT_ELENUM_PER_LINE = 128; // 128
static constexpr uint32_t MULTIPLIER = 2;
CATLASS_DEVICE
BlockEpilogue(Arch::Resource<ArchTag> &resource)
{
constexpr uint32_t BASE = 0;
constexpr uint32_t MASK_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_SIZE = 40 * UB_LINE_SIZE;
constexpr uint32_t GBRCUP_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t FLOAT_UB_TENSOR_SIZE = 32 * UB_LINE_SIZE;
constexpr uint32_t HALF_UB_TENSOR_SIZE = 16 * UB_LINE_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
constexpr uint32_t G_FLOAT_UB_TENSOR_SIZE = 2 * UB_LINE_SIZE;
constexpr uint32_t MASK_UB_TENSOR_OFFSET = BASE;
constexpr uint32_t GBRCLEFTCAST_UB_TENSOR_OFFSET = MASK_UB_TENSOR_OFFSET + MASK_UB_TENSOR_SIZE;
constexpr uint32_t GBRCUP_UB_TENSOR_OFFSET = GBRCLEFTCAST_UB_TENSOR_OFFSET + GBRCLEFTCAST_UB_TENSOR_SIZE;
constexpr uint32_t GCOMP_UB_TENSOR_OFFSET = GBRCUP_UB_TENSOR_OFFSET + GBRCUP_UB_TENSOR_SIZE;
constexpr uint32_t SHARE_UB_TENSOR_OFFSET = GCOMP_UB_TENSOR_OFFSET + G_FLOAT_UB_TENSOR_SIZE;
maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(MASK_UB_TENSOR_OFFSET);
gbrcLeftcastUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCLEFTCAST_UB_TENSOR_OFFSET);
gbrcUpUbTensor = resource.ubBuf.template GetBufferByByte<float>(GBRCUP_UB_TENSOR_OFFSET);
gcompUbTensor = resource.ubBuf.template GetBufferByByte<float>(GCOMP_UB_TENSOR_OFFSET);
shareUbTensor = resource.ubBuf.template GetBufferByByte<uint8_t>(SHARE_UB_TENSOR_OFFSET);
constexpr uint32_t G_UB_TENSOR_OFFSET_PING = SHARE_UB_TENSOR_OFFSET + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PING = G_UB_TENSOR_OFFSET_PING + G_FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t A_UB_TENSOR_OFFSET_PING = G_HALF_UB_TENSOR_OFFSET_PING + G_HALF_UB_TENSOR_SIZE;
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PING = A_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PING = OUT_UB_TENSOR_OFFSET_PING + FLOAT_UB_TENSOR_SIZE;
gUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PING);
gUbFPTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
gUbBFTensorPing = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PING);
aUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PING);
outUbTensorPing = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PING);
outUbFPTensorPing = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
outUbBFTensorPing = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PING);
constexpr uint32_t G_UB_TENSOR_OFFSET_PONG = 32 * UB_LINE_SIZE + OUT_HALF_UB_TENSOR_OFFSET_PING + HALF_UB_TENSOR_SIZE;
constexpr uint32_t G_HALF_UB_TENSOR_OFFSET_PONG = G_UB_TENSOR_OFFSET_PONG + G_FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t A_UB_TENSOR_OFFSET_PONG = G_HALF_UB_TENSOR_OFFSET_PONG + G_HALF_UB_TENSOR_SIZE;
constexpr uint32_t OUT_UB_TENSOR_OFFSET_PONG = A_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
constexpr uint32_t OUT_HALF_UB_TENSOR_OFFSET_PONG = OUT_UB_TENSOR_OFFSET_PONG + FLOAT_UB_TENSOR_SIZE;
gUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(G_UB_TENSOR_OFFSET_PONG);
gUbFPTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
gUbBFTensorPong = resource.ubBuf.template GetBufferByByte<GElementInput>(G_HALF_UB_TENSOR_OFFSET_PONG);
aUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(A_UB_TENSOR_OFFSET_PONG);
outUbTensorPong = resource.ubBuf.template GetBufferByByte<float>(OUT_UB_TENSOR_OFFSET_PONG);
outUbFPTensorPong = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
outUbBFTensorPong = resource.ubBuf.template GetBufferByByte<AElementOutput>(OUT_HALF_UB_TENSOR_OFFSET_PONG);
}
CATLASS_DEVICE
~BlockEpilogue()
{}
CATLASS_DEVICE
void operator()(
AscendC::GlobalTensor<AElementOutput> maskOutput,
AscendC::GlobalTensor<GElementInput> gInput,
AscendC::GlobalTensor<AElementInput> attnInput,
AscendC::GlobalTensor<MaskElementInput> boolInput,
uint32_t fullChunkSize,
uint32_t chunkSize,
uint32_t kHeadDim,
uint32_t vHeadDim,
uint32_t &pingpongFlag
, uint32_t batchIdx, uint32_t headIdx, uint32_t chunkIdx
)
{
uint32_t mActual = chunkSize;
uint32_t nActual = chunkSize;
uint32_t alignedNActual = CeilDiv(nActual, 16) * 16;
uint32_t subBlockIdx = AscendC::GetSubBlockIdx();
uint32_t subBlockNum = AscendC::GetSubBlockNum();
uint32_t blockIdx = AscendC::GetBlockIdx();
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 offsetA = mOffset * nActual + nOffset;
uint16_t aInputDstStride;
if((nActual - 1) % 16 <= 7) aInputDstStride = 1;
else aInputDstStride = 0;
uint32_t gbrcStart, gbrcRealStart, gbrcRealEnd, gbrcRealProcess, gbrcEffStart, gbrcEffEnd, mulsRemain, mulsRemainIdx;
if(mActualThisSubBlock <= 32)
{ if(subBlockIdx == 0)
{
gbrcStart = 0;
gbrcRealStart = 0;
gbrcRealProcess = mActualThisSubBlock;
}
else
{
gbrcStart = mActualPerSubBlock;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActual - gbrcRealStart;
}
gbrcEffStart = gbrcStart - gbrcRealStart;
gbrcEffEnd = gbrcEffStart + mActualThisSubBlock;
uint32_t dstUpShape_[2] = {mActualThisSubBlock, alignedNActual};
uint32_t srcUpShape_[2] = {1, alignedNActual};
uint32_t dstLeftShape_[2] = {gbrcRealProcess, alignedNActual};
uint32_t srcLeftShape_[2] = {gbrcRealProcess, 1};
AscendC::ResetMask();
AscendC::GlobalTensor<AElementOutput> maskOutputThisSubBlock = maskOutput[gbrcStart * nActual];
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
AscendC::DataCopyParams aInputUbParams{(uint16_t)mActualThisSubBlock, (uint16_t)(nActual*sizeof(float)), 0, aInputDstStride};
AscendC::DataCopyPadParams aInputUbPadParams{false, 0, 0, 0};
AscendC::DataCopyExtParams aOutputUbParams{(uint16_t)mActualThisSubBlock, (uint32_t)(nActual*sizeof(half)), 0, 0, 0};
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
AscendC::LocalTensor<AElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
AscendC::LocalTensor<AElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
if constexpr(std::is_same<GElementInput, float>::value) {
AscendC::DataCopyPad(gUbTensor, gInputThisSubBlock, gfloatUbParams, gUbPadParams);
} else {
AscendC::DataCopyPad(gUbFPTensor, gInputThisSubBlock, ghalfUbParams, gUbPadParams);
}
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
if constexpr(!std::is_same<GElementInput, float>::value) {
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
AscendC::PipeBarrier<PIPE_V>();
}
AscendC::Copy(gcompUbTensor, gUbTensor, 64, 2, {1, 1, 8, 8});
AscendC::PipeBarrier<PIPE_V>();
AscendC::Broadcast<float, 2, 0>(gbrcUpUbTensor, gcompUbTensor, dstUpShape_, srcUpShape_, shareUbTensor);
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstLeftShape_, srcLeftShape_, shareUbTensor);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Sub(gbrcUpUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*alignedNActual], gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Mins(gbrcUpUbTensor, gbrcUpUbTensor, (float)0.0, mActualThisSubBlock * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Exp(gbrcUpUbTensor, gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
gbrcRealEnd = CeilDiv(gbrcStart + mActualThisSubBlock, 8) * 8;
AscendC::Mul(gbrcUpUbTensor[gbrcRealStart], gbrcUpUbTensor[gbrcRealStart], maskUbTensor[gbrcEffStart * 64], gbrcRealEnd - gbrcRealStart, mActualThisSubBlock,
{1, 1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(64/8)});
AscendC::PipeBarrier<PIPE_V>();
mulsRemain = alignedNActual - gbrcRealEnd;
mulsRemainIdx = gbrcRealEnd;
while(mulsRemain > 64)
{
AscendC::Muls(gbrcUpUbTensor[mulsRemainIdx], gbrcUpUbTensor[mulsRemainIdx], (float)0.0, 64, mActualThisSubBlock,
{1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8)});
mulsRemain -= 64;
mulsRemainIdx += 64;
}
AscendC::Muls(gbrcUpUbTensor[mulsRemainIdx], gbrcUpUbTensor[mulsRemainIdx], (float)0.0, mulsRemain, mActualThisSubBlock,
{1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8)});
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
if(chunkSize==fullChunkSize) AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisSubBlock*nActual);
else AscendC::DataCopyPad(aUbTensor, attnInputThisSubBlock, aInputUbParams, aInputUbPadParams);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::Mul(outUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisSubBlock * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
if(std::is_same<AElementOutput, half>::value)
{
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisSubBlock * alignedNActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
if(chunkSize==fullChunkSize) AscendC::DataCopy(maskOutputThisSubBlock, outUbFPTensor, mActualThisSubBlock*nActual);
else AscendC::DataCopyPad(maskOutputThisSubBlock, outUbFPTensor, aOutputUbParams);
}
else
{
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisSubBlock * alignedNActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
if(chunkSize==fullChunkSize) AscendC::DataCopy(maskOutputThisSubBlock, outUbBFTensor, mActualThisSubBlock*nActual);
else AscendC::DataCopyPad(maskOutputThisSubBlock, outUbBFTensor, aOutputUbParams);
}
pingpongFlag = 1 - pingpongFlag;
}
else // mActualThisSubBlock > 32 ; <=64
{
AscendC::ResetMask();
AscendC::GlobalTensor<GElementInput> gInputThisSubBlock = gInput;
AscendC::DataCopyParams gfloatUbParams{1, (uint16_t)(mActual*sizeof(float)), 0, 0};
AscendC::DataCopyParams ghalfUbParams{1, (uint16_t)(mActual*sizeof(half)), 0, 0};
AscendC::DataCopyPadParams gUbPadParams{false, 0, 0, 0};
AscendC::LocalTensor<float> gUbTensor = (pingpongFlag == 0) ? gUbTensorPing : gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensor = (pingpongFlag == 0) ? gUbFPTensorPing : gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensor = (pingpongFlag == 0) ? gUbBFTensorPing : gUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
AscendC::SetFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::MTE3_MTE2>(EVENT_ID0);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID0 + pingpongFlag);
if constexpr(std::is_same<GElementInput, float>::value) {
AscendC::DataCopyPad(gUbTensor, gInputThisSubBlock, gfloatUbParams, gUbPadParams);
} else {
AscendC::DataCopyPad(gUbFPTensor, gInputThisSubBlock, ghalfUbParams, gUbPadParams);
}
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID0 + pingpongFlag);
if constexpr(!std::is_same<GElementInput, float>::value) {
AscendC::Cast(gUbTensor, gUbFPTensor, AscendC::RoundMode::CAST_NONE, mActual);
AscendC::PipeBarrier<PIPE_V>();
}
AscendC::Copy(gcompUbTensor, gUbTensor, 64, 2, {1, 1, 8, 8});
AscendC::PipeBarrier<PIPE_V>();
uint32_t mActualPerStage = CeilDiv(mActualThisSubBlock, 2);
uint32_t mActualThisStage = 0;
for(uint32_t stage = 0; stage < 2; ++stage)
{
if(stage==0) mActualThisStage = mActualPerStage;
else mActualThisStage = mActualThisSubBlock - mActualPerStage;
if(subBlockIdx == 0 && stage == 0)
{
gbrcStart = 0;
gbrcRealStart = 0;
gbrcRealProcess = mActualThisStage;
}
else if(subBlockIdx == 0 && stage == 1)
{
gbrcStart = mActualPerStage;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActualThisSubBlock - gbrcRealStart;
}
else if(subBlockIdx == 1 && stage == 0)
{
gbrcStart = mActualPerSubBlock;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActualPerSubBlock + mActualThisStage - gbrcRealStart;
}
else if(subBlockIdx == 1 && stage == 1)
{
gbrcStart = mActualPerSubBlock + mActualPerStage;
gbrcRealStart = gbrcStart & ~7;
gbrcRealProcess = mActual - gbrcRealStart;
}
gbrcEffStart = gbrcStart - gbrcRealStart;
AscendC::GlobalTensor<AElementOutput> maskOutputThisSubBlock = maskOutput[gbrcStart * nActual];
AscendC::GlobalTensor<AElementInput> attnInputThisSubBlock = attnInput[gbrcStart * nActual];
AscendC::DataCopyParams aInputUbParams{(uint16_t)mActualThisStage, (uint16_t)(nActual*sizeof(float)), 0, aInputDstStride};
AscendC::DataCopyPadParams aInputUbPadParams{false, 0, 0, 0};
AscendC::DataCopyExtParams aOutputUbParams{(uint16_t)mActualThisStage, (uint32_t)(nActual*sizeof(half)), 0, 0, 0};
AscendC::LocalTensor<float> aUbTensor = (pingpongFlag == 0) ? aUbTensorPing : aUbTensorPong;
AscendC::LocalTensor<float> outUbTensor = (pingpongFlag == 0) ? outUbTensorPing : outUbTensorPong;
AscendC::LocalTensor<AElementOutput> outUbFPTensor = (pingpongFlag == 0) ? outUbFPTensorPing : outUbFPTensorPong;
AscendC::LocalTensor<AElementOutput> outUbBFTensor = (pingpongFlag == 0) ? outUbBFTensorPing : outUbBFTensorPong;
AscendC::SetFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE2>(EVENT_ID1 + pingpongFlag);
if(chunkSize==fullChunkSize) AscendC::DataCopy(aUbTensor, attnInputThisSubBlock, mActualThisStage*nActual);
else AscendC::DataCopyPad(aUbTensor, attnInputThisSubBlock, aInputUbParams, aInputUbPadParams);
AscendC::SetFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
uint32_t dstUpShape_[2] = {mActualThisStage, alignedNActual};
uint32_t srcUpShape_[2] = {1, alignedNActual};
uint32_t dstLeftShape_[2] = {gbrcRealProcess, alignedNActual};
uint32_t srcLeftShape_[2] = {gbrcRealProcess, 1};
AscendC::Broadcast<float, 2, 0>(gbrcUpUbTensor, gcompUbTensor, dstUpShape_, srcUpShape_, shareUbTensor);
AscendC::Broadcast<float, 2, 1>(gbrcLeftcastUbTensor, gcompUbTensor[gbrcRealStart], dstLeftShape_, srcLeftShape_, shareUbTensor);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Sub(gbrcUpUbTensor, gbrcLeftcastUbTensor[gbrcEffStart*alignedNActual], gbrcUpUbTensor, mActualThisStage * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Mins(gbrcUpUbTensor, gbrcUpUbTensor, (float)0.0, mActualThisStage * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
AscendC::Exp(gbrcUpUbTensor, gbrcUpUbTensor, mActualThisStage * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
gbrcRealEnd = CeilDiv(gbrcStart + mActualThisStage, 8) * 8;
AscendC::Mul(gbrcUpUbTensor[gbrcRealStart], gbrcUpUbTensor[gbrcRealStart], maskUbTensor[gbrcEffStart * 64], gbrcRealEnd - gbrcRealStart, mActualThisStage,
{1, 1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(64/8)});
AscendC::PipeBarrier<PIPE_V>();
mulsRemain = alignedNActual - gbrcRealEnd;
mulsRemainIdx = gbrcRealEnd;
while(mulsRemain > 64)
{
AscendC::Muls(gbrcUpUbTensor[mulsRemainIdx], gbrcUpUbTensor[mulsRemainIdx], (float)0.0, 64, mActualThisStage,
{1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8)});
mulsRemain -= 64;
mulsRemainIdx += 64;
}
AscendC::Muls(gbrcUpUbTensor[mulsRemainIdx], gbrcUpUbTensor[mulsRemainIdx], (float)0.0, mulsRemain, mActualThisStage,
{1, 1, static_cast<uint8_t>(alignedNActual/8), static_cast<uint8_t>(alignedNActual/8)});
AscendC::PipeBarrier<PIPE_V>();
AscendC::WaitFlag<AscendC::HardEvent::MTE2_V>(EVENT_ID1 + pingpongFlag);
AscendC::Mul(outUbTensor, aUbTensor, gbrcUpUbTensor, mActualThisStage * alignedNActual);
AscendC::PipeBarrier<PIPE_V>();
if(std::is_same<AElementOutput, half>::value)
{
AscendC::Cast(outUbFPTensor, outUbTensor, AscendC::RoundMode::CAST_NONE, mActualThisStage * alignedNActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
if(chunkSize==fullChunkSize) AscendC::DataCopy(maskOutputThisSubBlock, outUbFPTensor, mActualThisStage*nActual);
else AscendC::DataCopyPad(maskOutputThisSubBlock, outUbFPTensor, aOutputUbParams);
}
else
{
AscendC::Cast(outUbBFTensor, outUbTensor, AscendC::RoundMode::CAST_RINT, mActualThisStage * alignedNActual);
AscendC::SetFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
AscendC::WaitFlag<AscendC::HardEvent::V_MTE3>(EVENT_ID0 + pingpongFlag);
if(chunkSize==fullChunkSize) AscendC::DataCopy(maskOutputThisSubBlock, outUbBFTensor, mActualThisStage*nActual);
else AscendC::DataCopyPad(maskOutputThisSubBlock, outUbBFTensor, aOutputUbParams);
}
pingpongFlag = 1 - pingpongFlag;
}
}
}
private:
AscendC::LocalTensor<float> maskUbTensor;
AscendC::LocalTensor<float> gbrcLeftcastUbTensor;
AscendC::LocalTensor<float> gbrcUpUbTensor;
AscendC::LocalTensor<float> gcompUbTensor;
AscendC::LocalTensor<uint8_t> shareUbTensor;
AscendC::LocalTensor<float> gUbTensorPing;
AscendC::LocalTensor<GElementInput> gUbFPTensorPing;
AscendC::LocalTensor<GElementInput> gUbBFTensorPing;
AscendC::LocalTensor<float> aUbTensorPing;
AscendC::LocalTensor<float> outUbTensorPing;
AscendC::LocalTensor<AElementOutput> outUbFPTensorPing;
AscendC::LocalTensor<AElementOutput> outUbBFTensorPing;
AscendC::LocalTensor<float> gUbTensorPong;
AscendC::LocalTensor<GElementInput> gUbFPTensorPong;
AscendC::LocalTensor<GElementInput> gUbBFTensorPong;
AscendC::LocalTensor<float> aUbTensorPong;
AscendC::LocalTensor<float> outUbTensorPong;
AscendC::LocalTensor<AElementOutput> outUbFPTensorPong;
AscendC::LocalTensor<AElementOutput> outUbBFTensorPong;
};
}
#endif

View File

@@ -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_O_EPILOGUE_POLICIES_HPP
#define CATLASS_EPILOGUE_GDN_FWD_O_EPILOGUE_POLICIES_HPP
#include "catlass/catlass.hpp"
namespace Catlass::Epilogue {
struct EpilogueAtlasGDNFwdOQkmask {
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310
using ArchTag = Arch::Ascend950;
#else
using ArchTag = Arch::AtlasA2;
#endif
};
struct EpilogueAtlasGDNFwdOOutput {
#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_O_EPILOGUE_POLICIES_HPP

View File

@@ -0,0 +1,275 @@
/**
 * 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_GEMM_SCHEDULER_GDN_FWD_O_HPP
#define CATLASS_GEMM_SCHEDULER_GDN_FWD_O_HPP
// constexpr uint32_t PING_PONG_STAGES = 1;
constexpr uint32_t PING_PONG_STAGES = 2;
constexpr uint32_t BYTE_SIZE_16_BIT = 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 GDNFwdOOffsets {
uint32_t qkOffset;
uint32_t ovOffset;
uint32_t hOffset;
uint32_t gOffset;
uint32_t attnWorkOffset;
uint32_t hvWorkOffset;
bool isFinalState;
uint32_t blockTokens;
// for debug
uint32_t batchIdx;
uint32_t headIdx;
uint32_t chunkIdx;
};
struct BlockSchedulerGdnFwdO {
uint32_t shapeBatch;
uint32_t seqlen;
uint32_t kNumHead;
uint32_t vNumHead;
uint32_t kHeadDim;
uint32_t vHeadDim;
uint32_t chunkSize;
uint32_t isVariedLen;
uint32_t tokenBatch;
uint32_t numChunks{0};
uint32_t vBlockSize{128};
uint32_t taskIdx;
uint32_t cubeCoreIdx;
uint32_t cubeCoreNum;
uint32_t vLoops;
uint32_t taskNum;
uint32_t headGroups;
bool isRunning;
bool processNewTask {true};
bool firstLoop {true};
bool lastLoop {false};
GDNFwdOOffsets 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 batchChunkIdx;
uint32_t batchChunkStartIdx;
uint32_t tokenOffset;
uint32_t batchChunks;
uint32_t batchTokens;
AscendC::GlobalTensor<int64_t> gmSeqlen;
AscendC::GlobalTensor<int64_t> gmChunkOffsets;
Arch::CrossCoreFlag cube1Done{3};
Arch::CrossCoreFlag vec1Done{4};
Arch::CrossCoreFlag cube2Done{5};
Arch::CrossCoreFlag cube3Done{6};
Arch::CrossCoreFlag vec2Done{7};
CATLASS_DEVICE
BlockSchedulerGdnFwdO() {}
CATLASS_DEVICE
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling, uint32_t coreIdx, uint32_t coreNum) {
__gm__ ChunkFwdOTilingData *__restrict gdnFwdOTilingData = reinterpret_cast<__gm__ ChunkFwdOTilingData *__restrict>(tiling);
shapeBatch = gdnFwdOTilingData->shapeBatch;
seqlen = gdnFwdOTilingData->seqlen;
kNumHead = gdnFwdOTilingData->kNumHead;
vNumHead = gdnFwdOTilingData->vNumHead;
kHeadDim = gdnFwdOTilingData->kHeadDim;
vHeadDim = gdnFwdOTilingData->vHeadDim;
chunkSize = gdnFwdOTilingData->chunkSize;
isVariedLen = gdnFwdOTilingData->isVariedLen;
tokenBatch = gdnFwdOTilingData->tokenBatch;
gmSeqlen.SetGlobalBuffer((__gm__ int64_t *)cu_seqlens);
gmChunkOffsets.SetGlobalBuffer((__gm__ int64_t *)chunk_offsets);
if (isVariedLen) {
for (uint32_t b = 1; b <= tokenBatch; b++) {
numChunks += (gmSeqlen.GetValue(b) - gmSeqlen.GetValue(b - 1) + chunkSize - 1) / chunkSize;
}
} else {
numChunks = (seqlen + chunkSize - 1) / chunkSize;
}
cubeCoreIdx = coreIdx;
cubeCoreNum = coreNum;
vLoops = vHeadDim / vBlockSize;
taskNum = vLoops * shapeBatch * numChunks * vNumHead;
headGroups = vNumHead / kNumHead;
taskIdx = cubeCoreIdx * PING_PONG_STAGES;
isRunning = taskIdx < taskNum;
}
CATLASS_DEVICE
void InitTask() {
if (processNewTask) {
if (unlikely(taskIdx >= taskNum)) {
isRunning = false;
}
vIdx = taskIdx / (shapeBatch * numChunks * vNumHead);
shapeBatchIdx = (taskIdx - vIdx * shapeBatch * numChunks * vNumHead) / (numChunks * vNumHead);
chunkIdx = (taskIdx - vIdx * shapeBatch * numChunks * vNumHead - shapeBatchIdx * numChunks * vNumHead) / vNumHead;
baseHeadIdx = taskIdx % vNumHead;
tokenBatchIdx = isVariedLen ? gmChunkOffsets.GetValue(2 * chunkIdx) : 0;
batchChunkIdx = isVariedLen ? gmChunkOffsets.GetValue(2 * chunkIdx + 1) : chunkIdx;
batchChunkStartIdx = chunkIdx - batchChunkIdx;
tokenOffset = isVariedLen ? gmSeqlen.GetValue(tokenBatchIdx) : 0;
batchTokens = isVariedLen ? (gmSeqlen.GetValue(tokenBatchIdx + 1) - tokenOffset) : seqlen;
headInnerIdx = 0;
} else {
headInnerIdx = (headInnerIdx + 1) % PING_PONG_STAGES;
}
vHeadIdx = baseHeadIdx + headInnerIdx;
kHeadIdx = vHeadIdx / headGroups;
offsets[currStage].qkOffset = (shapeBatchIdx * kNumHead * seqlen + kHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize) * kHeadDim;
offsets[currStage].ovOffset = (shapeBatchIdx * vNumHead * seqlen + vHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize) * vHeadDim;
offsets[currStage].hOffset = (shapeBatchIdx * vNumHead * numChunks + vHeadIdx * numChunks + chunkIdx) * kHeadDim * vHeadDim;
offsets[currStage].gOffset = shapeBatchIdx * vNumHead * seqlen + vHeadIdx * seqlen + tokenOffset + batchChunkIdx * chunkSize;
offsets[currStage].attnWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * chunkSize;
offsets[currStage].hvWorkOffset = (cubeCoreIdx * PING_PONG_STAGES + currStage) * chunkSize * vHeadDim;
offsets[currStage].isFinalState = chunkIdx == (numChunks - 1) || (isVariedLen && gmChunkOffsets.GetValue(2 * chunkIdx + 3) == 0);
offsets[currStage].blockTokens = offsets[currStage].isFinalState ? (batchTokens - batchChunkIdx * chunkSize) : chunkSize;
offsets[currStage].batchIdx = batchIdx;
offsets[currStage].headIdx = vHeadIdx;
offsets[currStage].chunkIdx = chunkIdx;
processNewTask = headInnerIdx == PING_PONG_STAGES - 1;
if (processNewTask) {
taskIdx += PING_PONG_STAGES * cubeCoreNum;
}
currStage = (currStage + 1) % PING_PONG_STAGES;
}
};
struct BlockSchedulerGdnFwdOCube : public BlockSchedulerGdnFwdO {
CATLASS_DEVICE
BlockSchedulerGdnFwdOCube() {}
CATLASS_DEVICE
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling) {
BlockSchedulerGdnFwdO::Init(cu_seqlens, chunk_offsets, tiling, AscendC::GetBlockIdx(), AscendC::GetBlockNum());
}
CATLASS_DEVICE
bool NeedProcessCube1() {
return true;
}
CATLASS_DEVICE
GDNFwdOOffsets& GetCube1Offsets() {
return offsets[(currStage - 1) % PING_PONG_STAGES];
}
CATLASS_DEVICE
GemmCoord GetCube1Shape() {
GDNFwdOOffsets& cube1Offsets = GetCube1Offsets();
return GemmCoord{cube1Offsets.blockTokens, cube1Offsets.blockTokens, kHeadDim};
}
CATLASS_DEVICE
bool NeedProcessCube23() {
if (unlikely(firstLoop)) {
firstLoop = false;
return false;
}
return true;
}
CATLASS_DEVICE
GDNFwdOOffsets& GetCube23Offsets() {
return offsets[(currStage - 2) % PING_PONG_STAGES];
}
CATLASS_DEVICE
GemmCoord GetCube2Shape() {
GDNFwdOOffsets& cube2Offsets = GetCube23Offsets();
return GemmCoord{kHeadDim, vHeadDim, cube2Offsets.blockTokens};
}
CATLASS_DEVICE
GemmCoord GetCube3Shape() {
GDNFwdOOffsets& cube2Offsets = GetCube23Offsets();
return GemmCoord{kHeadDim, vHeadDim, cube2Offsets.blockTokens};
}
};
struct BlockSchedulerGdnFwdOVec : public BlockSchedulerGdnFwdO {
CATLASS_DEVICE
BlockSchedulerGdnFwdOVec() {}
CATLASS_DEVICE
void Init(GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR tiling) {
BlockSchedulerGdnFwdO::Init(cu_seqlens, chunk_offsets, tiling, AscendC::GetBlockIdx() / AscendC::GetSubBlockNum(), AscendC::GetBlockNum());
}
CATLASS_DEVICE
bool NeedProcessVec1() {
return isRunning;
}
CATLASS_DEVICE
bool NeedProcessVec2() {
if (unlikely(firstLoop)) {
firstLoop = false;
return false;
}
return true;
}
CATLASS_DEVICE
GDNFwdOOffsets& GetVec1Offsets() {
return offsets[(currStage - 1) % PING_PONG_STAGES];
}
CATLASS_DEVICE
GDNFwdOOffsets& GetVec2Offsets() {
return offsets[(currStage - 2) % PING_PONG_STAGES];
}
};
} // namespace Catlass::Gemm::Block
#endif // CATLASS_GEMM_SCHEDULER_GDN_FWD_O_HPP

View File

@@ -0,0 +1,404 @@
/**
* 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_fwdo_qkmask.hpp"
#include "../../epilogue/block/block_epilogue_gdn_fwdo_output.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_o.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_fwdo_qkmask.hpp"
#include "../../epilogue/block/block_epilogue_gdn_fwdo_output.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_o.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;
// template <>
namespace Catlass::Gemm::Kernel {
template<
typename INPUT_TYPE,
typename G_TYPE,
typename WORKSPACE_TYPE
>
class GDNFwdOKernel {
public:
#if defined(__CCE_AICORE__) && __CCE_AICORE__ == 310
using ArchTag = Arch::Ascend950;
#else
using ArchTag = Arch::AtlasA2;
#endif
using GDNFwdOOffsets = Catlass::Gemm::Block::GDNFwdOOffsets;
using CubeScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdOCube;
using VecScheduler = typename Catlass::Gemm::Block::BlockSchedulerGdnFwdOVec;
using DispatchPolicyTla = Gemm::MmadPingpongTlaMulti<ArchTag, true, false>;
using L1TileShapeTla = Shape<_128, _128, _128>;
using L0TileShapeTla = L1TileShapeTla;
using QType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using KType = Gemm::GemmType<INPUT_TYPE, layout::ColumnMajor>;
using AttenType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
using AttenMaskedType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using HType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using OinterType = Gemm::GemmType<WORKSPACE_TYPE, layout::RowMajor>;
using VNEWType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using GType = Gemm::GemmType<G_TYPE, layout::RowMajor>;
using OType = Gemm::GemmType<INPUT_TYPE, layout::RowMajor>;
using MaskType = Gemm::GemmType<bool, layout::RowMajor>;
// cube 1
using TileCopyQK = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::ColumnMajor, WORKSPACE_TYPE, layout::RowMajor>;
using BlockMmadQK = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyQK>;
// cube 2
using TileCopyQH = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
using BlockMmadQH = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyQH>;
// cube 3
using TileCopyAttenVNEW = Catlass::Gemm::Tile::PackedTileCopyTla<ArchTag, INPUT_TYPE, layout::RowMajor, INPUT_TYPE, layout::RowMajor, WORKSPACE_TYPE, layout::RowMajor>;
using BlockMmadAttenVNEW = Gemm::Block::BlockMmadTla<DispatchPolicyTla, L1TileShapeTla, L0TileShapeTla, INPUT_TYPE, INPUT_TYPE, WORKSPACE_TYPE, void, TileCopyAttenVNEW>;
// vec 1
using DispatchPolicyGDNFwdOQkmask = Epilogue::EpilogueAtlasGDNFwdOQkmask;
using EpilogueGDNFwdOQkmask = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdOQkmask, AttenMaskedType, GType, AttenType, MaskType>;
// vec 2
using DispatchPolicyGDNFwdOOutput = Epilogue::EpilogueAtlasGDNFwdOOutput;
using EpilogueGDNFwdOOutput = Epilogue::Block::BlockEpilogue<DispatchPolicyGDNFwdOOutput, OType, GType, OinterType, OinterType>;
using ElementQ = typename BlockMmadQK::ElementA;
using LayoutQ = Catlass::layout::RowMajor;
using ElementK = typename BlockMmadQK::ElementB;
using LayoutK = Catlass::layout::ColumnMajor;
using ElementAtten = typename BlockMmadQK::ElementC;
using LayoutAtten = Catlass::layout::RowMajor;
using ElementAttenMasked = typename BlockMmadQH::ElementA;
using LayoutAttenMasked = Catlass::layout::RowMajor;
using ElementH = typename BlockMmadQH::ElementB;
using LayoutH = Catlass::layout::RowMajor;
using ElementOinter = typename BlockMmadQH::ElementC;
using LayoutOinter = Catlass::layout::RowMajor;
using ElementVNEW = typename BlockMmadAttenVNEW::ElementB;
using LayoutVNEW = Catlass::layout::RowMajor;
using ElementG = G_TYPE;
using ElementMask = bool;
using L1TileShape = typename BlockMmadQK::L1TileShape;
uint32_t shapeBatch;
uint32_t seqlen;
uint32_t kNumHead;
uint32_t vNumHead;
uint32_t kHeadDim;
uint32_t vHeadDim;
uint32_t chunkSize;
float scale;
uint32_t numChunks;
uint32_t isVariedLen;
uint32_t tokenBatch;
uint32_t vWorkspaceOffset;
uint32_t hWorkspaceOffset;
uint32_t attnWorkspaceOffset;
uint32_t aftermaskWorkspaceOffset;
uint32_t maskWorkspaceOffset;
AscendC::GlobalTensor<ElementQ> gmQ;
AscendC::GlobalTensor<ElementK> gmK;
AscendC::GlobalTensor<ElementVNEW> gmV;
AscendC::GlobalTensor<ElementH> gmH;
AscendC::GlobalTensor<ElementG> gmG;
AscendC::GlobalTensor<ElementVNEW> gmO;
AscendC::GlobalTensor<ElementOinter> gmVWorkspace;
AscendC::GlobalTensor<ElementOinter> gmHWorkspace;
AscendC::GlobalTensor<ElementAtten> gmAttnWorkspace;
AscendC::GlobalTensor<ElementAttenMasked> gmAftermaskWorkspace;
AscendC::GlobalTensor<ElementMask> gmMask;
CubeScheduler cubeBlockScheduler;
VecScheduler vecBlockScheduler;
Arch::Resource<ArchTag> resource;
__aicore__ inline GDNFwdOKernel() {}
__aicore__ inline void Init(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR h, GM_ADDR g,
GM_ADDR cu_seqlens, GM_ADDR chunk_offsets, GM_ADDR o, GM_ADDR tiling, GM_ADDR user) {
__gm__ ChunkFwdOTilingData *__restrict gdnFwdOTilingData = reinterpret_cast<__gm__ ChunkFwdOTilingData *__restrict>(tiling);
shapeBatch = gdnFwdOTilingData->shapeBatch;
seqlen = gdnFwdOTilingData->seqlen;
kNumHead = gdnFwdOTilingData->kNumHead;
vNumHead = gdnFwdOTilingData->vNumHead;
kHeadDim = gdnFwdOTilingData->kHeadDim;
vHeadDim = gdnFwdOTilingData->vHeadDim;
scale = gdnFwdOTilingData->scale;
chunkSize = gdnFwdOTilingData->chunkSize;
isVariedLen = gdnFwdOTilingData->isVariedLen;
tokenBatch = gdnFwdOTilingData->tokenBatch;
vWorkspaceOffset = gdnFwdOTilingData->vWorkspaceOffset;
hWorkspaceOffset = gdnFwdOTilingData->hWorkspaceOffset;
attnWorkspaceOffset = gdnFwdOTilingData->attnWorkspaceOffset;
aftermaskWorkspaceOffset = gdnFwdOTilingData->aftermaskWorkspaceOffset;
maskWorkspaceOffset = gdnFwdOTilingData->maskWorkspaceOffset;
gmQ.SetGlobalBuffer((__gm__ ElementQ *)q);
gmK.SetGlobalBuffer((__gm__ ElementK *)k);
gmV.SetGlobalBuffer((__gm__ ElementVNEW *)v);
gmH.SetGlobalBuffer((__gm__ ElementH *)h);
gmG.SetGlobalBuffer((__gm__ ElementG *)g);
gmO.SetGlobalBuffer((__gm__ ElementVNEW *)o);
gmVWorkspace.SetGlobalBuffer((__gm__ ElementOinter *)(user + vWorkspaceOffset));
gmHWorkspace.SetGlobalBuffer((__gm__ ElementOinter *)(user + hWorkspaceOffset));
gmAttnWorkspace.SetGlobalBuffer((__gm__ ElementAtten *)(user + attnWorkspaceOffset));
gmAftermaskWorkspace.SetGlobalBuffer((__gm__ ElementAttenMasked *)(user + aftermaskWorkspaceOffset));
gmMask.SetGlobalBuffer((__gm__ ElementMask *)(user + maskWorkspaceOffset));
if ASCEND_IS_AIC {
cubeBlockScheduler.Init(cu_seqlens, chunk_offsets, tiling);
}
if ASCEND_IS_AIV {
vecBlockScheduler.Init(cu_seqlens, chunk_offsets, tiling);
}
}
__aicore__ inline void Process() {
if ASCEND_IS_AIC {
uint32_t coreIdx = AscendC::GetBlockIdx();
uint32_t coreNum = AscendC::GetBlockNum();
BlockMmadQK blockMmadQK(resource);
BlockMmadQH blockMmadQH(resource);
BlockMmadAttenVNEW blockMmadAttenVNEW(resource);
auto qLayout = tla::MakeLayout<ElementQ, LayoutQ>(shapeBatch * kNumHead * seqlen, kHeadDim);
auto kLayout = tla::MakeLayout<ElementK, LayoutK>(kHeadDim, shapeBatch * kNumHead * seqlen);
auto hLayout = tla::MakeLayout<ElementH, LayoutH>(shapeBatch * vNumHead * seqlen * kHeadDim, vHeadDim);
auto ointerLayout = tla::MakeLayout<ElementOinter, LayoutOinter>(coreNum * chunkSize * PING_PONG_STAGES, vHeadDim);
auto vnewLayout = tla::MakeLayout<ElementVNEW, LayoutVNEW>(shapeBatch * vNumHead * seqlen, vHeadDim);
bool needRun = false;
bool isFirstC3 = true;
while (cubeBlockScheduler.isRunning) {
cubeBlockScheduler.InitTask();
if (cubeBlockScheduler.isRunning && coreIdx < coreNum) {
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
GDNFwdOOffsets& cube1Offsets = cubeBlockScheduler.GetCube1Offsets();
int64_t cube1OffsetQ = cube1Offsets.qkOffset;
int64_t cube1OffsetK = cube1Offsets.qkOffset;
int64_t cube1OffsetAttn = cube1Offsets.attnWorkOffset;
auto attenLayout = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, cube1Offsets.blockTokens);
auto tensorQ = tla::MakeTensor(gmQ[cube1OffsetQ], qLayout, Catlass::Arch::PositionGM{});
auto tensorK = tla::MakeTensor(gmK[cube1OffsetK], kLayout, Catlass::Arch::PositionGM{});
auto tensorAttn = tla::MakeTensor(gmAttnWorkspace[cube1OffsetAttn], attenLayout, Catlass::Arch::PositionGM{});
GemmCoord cube1Shape{cube1Offsets.blockTokens, cube1Offsets.blockTokens, kHeadDim};
auto tensorBlockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.k()));
auto tensorBlockK = GetTile(tensorK, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.k(), cube1Shape.n()));
auto tensorBlockAttn = GetTile(tensorAttn, tla::MakeCoord(0, 0), tla::MakeShape(cube1Shape.m(), cube1Shape.n()));
blockMmadQK.preSetFlags();
blockMmadQK(tensorBlockQ, tensorBlockK, tensorBlockAttn, cube1Shape);
blockMmadQK.finalWaitFlags();
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube1Done);
}
// AscendC::PipeBarrier<PIPE_ALL>();
if (needRun && coreIdx < coreNum) {
if(!cubeBlockScheduler.isRunning) Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec2Done);
GDNFwdOOffsets& cube2Offsets = cubeBlockScheduler.GetCube23Offsets();
int64_t cube2OffsetQ = cube2Offsets.qkOffset;
int64_t cube2OffsetH = cube2Offsets.hOffset;
int64_t cube2OffsetHWork = cube2Offsets.hvWorkOffset;
auto tensorQ = tla::MakeTensor(gmQ[cube2OffsetQ], qLayout, Catlass::Arch::PositionGM{});
auto tensorH = tla::MakeTensor(gmH[cube2OffsetH], hLayout, Catlass::Arch::PositionGM{});
auto tensorHWork = tla::MakeTensor(gmHWorkspace[cube2OffsetHWork], ointerLayout, Catlass::Arch::PositionGM{});
GemmCoord cube2Shape{cube2Offsets.blockTokens, vHeadDim, kHeadDim};
auto tensorBlockQ = GetTile(tensorQ, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.k()));
auto tensorBlockH = GetTile(tensorH, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.k(), cube2Shape.n()));
auto tensorBlockHWork = GetTile(tensorHWork, tla::MakeCoord(0, 0), tla::MakeShape(cube2Shape.m(), cube2Shape.n()));
blockMmadQH.preSetFlags();
blockMmadQH(tensorBlockQ, tensorBlockH, tensorBlockHWork, cube2Shape);
blockMmadQH.finalWaitFlags();
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube2Done);
}
if (needRun && coreIdx < coreNum) {
GDNFwdOOffsets& cube3Offsets = cubeBlockScheduler.GetCube23Offsets();
if(isFirstC3) Arch::CrossCoreWaitFlag(cubeBlockScheduler.vec1Done);
int64_t cube3OffsetAttnMask = cube3Offsets.attnWorkOffset;
int64_t cube3OffsetV = cube3Offsets.ovOffset;
int64_t cube3OffsetVWork = cube3Offsets.hvWorkOffset;
auto attenLayout = tla::MakeLayout<ElementAtten, LayoutAtten>(coreNum * chunkSize * PING_PONG_STAGES, cube3Offsets.blockTokens);
auto tensorAttnMask = tla::MakeTensor(gmAftermaskWorkspace[cube3OffsetAttnMask], attenLayout, Catlass::Arch::PositionGM{});
auto tensorV = tla::MakeTensor(gmV[cube3OffsetV], vnewLayout, Catlass::Arch::PositionGM{});
auto tensorVWork = tla::MakeTensor(gmVWorkspace[cube3OffsetVWork], ointerLayout, Catlass::Arch::PositionGM{});
GemmCoord cube3Shape{cube3Offsets.blockTokens, vHeadDim, cube3Offsets.blockTokens};
auto tensorBlockAttnMask = GetTile(tensorAttnMask, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.k()));
auto tensorBlockV = GetTile(tensorV, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.k(), cube3Shape.n()));
auto tensorBlockVWork = GetTile(tensorVWork, tla::MakeCoord(0, 0), tla::MakeShape(cube3Shape.m(), cube3Shape.n()));
blockMmadAttenVNEW.preSetFlags();
blockMmadAttenVNEW(tensorBlockAttnMask, tensorBlockV, tensorBlockVWork, cube3Shape);
blockMmadAttenVNEW.finalWaitFlags();
Arch::CrossCoreSetFlag<0x2, PIPE_FIX>(cubeBlockScheduler.cube3Done);
isFirstC3 = false;
}
needRun = true;
// AscendC::PipeBarrier<PIPE_ALL>();
}
if (coreIdx < coreNum) {
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();
AscendC::LocalTensor<float> maskUbTensor = resource.ubBuf.template GetBufferByByte<float>(0);
AscendC::Duplicate<float>(maskUbTensor, (float)0.0, 64*64);
AscendC::PipeBarrier<PIPE_V>();
for(uint32_t i = 0; i < 64; ++ i) AscendC::Duplicate<float>(maskUbTensor[i * 64], (float)1.0, i + 1);
AscendC::PipeBarrier<PIPE_V>();
bool needRun = false;
uint32_t pingpongFlag = 0;
if (coreIdx < coreNum * subBlockNum) {
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
}
while (vecBlockScheduler.isRunning) {
vecBlockScheduler.InitTask();
if (vecBlockScheduler.isRunning && coreIdx < coreNum * subBlockNum) {
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube1Done);
GDNFwdOOffsets& vec1Offsets = vecBlockScheduler.GetVec1Offsets();
int64_t vec1OffsetAttnMask = vec1Offsets.attnWorkOffset;
int64_t vec1OffsetG = vec1Offsets.gOffset;
int64_t vec1OffsetAttn = vec1Offsets.attnWorkOffset;
EpilogueGDNFwdOQkmask epilogueGDNFwdOQkmask(resource);
epilogueGDNFwdOQkmask(
gmAftermaskWorkspace[vec1OffsetAttnMask],
gmG[vec1OffsetG], gmAttnWorkspace[vec1OffsetAttn], gmMask,
chunkSize, vec1Offsets.blockTokens, kHeadDim, vHeadDim, pingpongFlag, vec1Offsets.batchIdx, vec1Offsets.headIdx, vec1Offsets.chunkIdx
);
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec1Done);
}
// AscendC::PipeBarrier<PIPE_ALL>();
if (needRun && coreIdx < coreNum * subBlockNum) {
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube2Done);
Arch::CrossCoreWaitFlag(vecBlockScheduler.cube3Done);
GDNFwdOOffsets& vec2Offsets = vecBlockScheduler.GetVec2Offsets();
int64_t vec2OffsetO = vec2Offsets.ovOffset;
int64_t vec2OffsetG = vec2Offsets.gOffset;
int64_t vec2OffsetVWork = vec2Offsets.hvWorkOffset;
int64_t vec2OffsetHWork = vec2Offsets.hvWorkOffset;
EpilogueGDNFwdOOutput epilogueGDNFwdOOutput(resource);
epilogueGDNFwdOOutput(
gmO[vec2OffsetO],
gmG[vec2OffsetG], gmVWorkspace[vec2OffsetVWork], gmHWorkspace[vec2OffsetHWork],
scale, vec2Offsets.blockTokens, kHeadDim, vHeadDim, pingpongFlag, vec2Offsets.batchIdx, vec2Offsets.headIdx, vec2Offsets.chunkIdx
);
Arch::CrossCoreSetFlag<0x2, PIPE_MTE3>(vecBlockScheduler.vec2Done);
}
// AscendC::PipeBarrier<PIPE_ALL>();
needRun = true;
}
}
}
};
}

View File

@@ -0,0 +1,69 @@
/**
 * 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_fwd_o.cpp
* \brief
*/
// #include "chunk_fwd_o.h"
#if defined(__CCE_AICORE__) && (__CCE_AICORE__ == 200)
#include "arch20/compat_310p.h"
#include "arch20/gemm/kernel/gdn_fwd_o_kernel.hpp"
#else
#include "arch22/gemm/kernel/gdn_fwd_o_kernel.hpp"
#endif
#include "lib/matmul_intf.h"
using namespace Catlass;
extern "C" __global__ __aicore__ void chunk_fwd_o(GM_ADDR q, GM_ADDR k, GM_ADDR v, GM_ADDR h,
GM_ADDR g, GM_ADDR cu_seqlens, GM_ADDR chunk_offsets,
GM_ADDR o, GM_ADDR workspace, GM_ADDR tiling)
{
#ifdef CATLASS_UNIFIED_CORE
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC);
#else
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIC_1_2);
#endif
GM_ADDR user = AscendC::GetUserWorkspace(workspace);
__gm__ ChunkFwdOTilingData *__restrict gdnFwdOTilingData = reinterpret_cast<__gm__ ChunkFwdOTilingData *__restrict>(tiling);
using workspaceType = float;
// dtype: 0 - fp16, 1 - bf16, 2 - fp32
#ifndef CATLASS_UNIFIED_CORE
if (gdnFwdOTilingData->dataType == 1) {
if (gdnFwdOTilingData->gDataType == 2) {
using GDNFwdOKernel = Catlass::Gemm::Kernel::GDNFwdOKernel<bfloat16_t, float, workspaceType>;
GDNFwdOKernel gdnFwdO;
gdnFwdO.Init(q, k, v, h, g, cu_seqlens, chunk_offsets, o, tiling, user);
gdnFwdO.Process();
} else {
using GDNFwdOKernel = Catlass::Gemm::Kernel::GDNFwdOKernel<bfloat16_t, bfloat16_t, workspaceType>;
GDNFwdOKernel gdnFwdO;
gdnFwdO.Init(q, k, v, h, g, cu_seqlens, chunk_offsets, o, tiling, user);
gdnFwdO.Process();
}
} else
#endif
{
if (gdnFwdOTilingData->gDataType == 2) {
using GDNFwdOKernel = Catlass::Gemm::Kernel::GDNFwdOKernel<half, float, workspaceType>;
GDNFwdOKernel gdnFwdO;
gdnFwdO.Init(q, k, v, h, g, cu_seqlens, chunk_offsets, o, tiling, user);
gdnFwdO.Process();
} else {
using GDNFwdOKernel = Catlass::Gemm::Kernel::GDNFwdOKernel<half, half, workspaceType>;
GDNFwdOKernel gdnFwdO;
gdnFwdO.Init(q, k, v, h, g, cu_seqlens, chunk_offsets, o, tiling, user);
gdnFwdO.Process();
}
}
}

View File

@@ -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

View File

@@ -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

View File

@@ -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_

View 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

View File

@@ -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

View File

@@ -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)

View 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

View File

@@ -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