19
csrc/moe/causal_conv1d/CMakeLists.txt
Normal file
19
csrc/moe/causal_conv1d/CMakeLists.txt
Normal file
@@ -0,0 +1,19 @@
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
# Copyright (c) 2025 Huawei Technologies Co., Ltd.
|
||||
# This program is free software, you can redistribute it and/or modify it under the terms and conditions of
|
||||
# CANN Open Software License Agreement Version 2.0 (the "License").
|
||||
# Please refer to the License for details. You may not use this file except in compliance with the License.
|
||||
# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
|
||||
# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
|
||||
# See LICENSE in the root of the software repository for the full text of the License.
|
||||
# -----------------------------------------------------------------------------------------------------------
|
||||
|
||||
file(GLOB CURRENT_DIRS RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} ${CMAKE_CURRENT_SOURCE_DIR}/*)
|
||||
if(NOT ENABLE_TEST AND NOT BENCHMARK)
|
||||
list(REMOVE_ITEM CURRENT_DIRS tests)
|
||||
endif()
|
||||
foreach(SUB_DIR ${CURRENT_DIRS})
|
||||
if(EXISTS "${CMAKE_CURRENT_SOURCE_DIR}/${SUB_DIR}/CMakeLists.txt")
|
||||
add_subdirectory(${SUB_DIR})
|
||||
endif()
|
||||
endforeach()
|
||||
22
csrc/moe/causal_conv1d/op_host/CMakeLists.txt
Normal file
22
csrc/moe/causal_conv1d/op_host/CMakeLists.txt
Normal file
@@ -0,0 +1,22 @@
|
||||
add_op_to_compiled_list()
|
||||
|
||||
if (BUILD_OPEN_PROJECT)
|
||||
target_sources(op_host_aclnn PRIVATE
|
||||
causal_conv1d_def.cpp
|
||||
)
|
||||
endif()
|
||||
|
||||
add_ops_compile_options(
|
||||
OP_NAME CausalConv1d
|
||||
OPTIONS
|
||||
--cce-auto-sync=on
|
||||
-Wno-deprecated-declarations
|
||||
)
|
||||
|
||||
if (NOT BUILD_OPS_RTY_KERNEL)
|
||||
add_modules_sources(OPTYPE causal_conv1d ACLNNTYPE aclnn)
|
||||
target_include_directories(${OPHOST_NAME}_tiling_obj PRIVATE
|
||||
${CMAKE_CURRENT_SOURCE_DIR}
|
||||
)
|
||||
endif()
|
||||
|
||||
90
csrc/moe/causal_conv1d/op_host/causal_conv1d_def.cpp
Normal file
90
csrc/moe/causal_conv1d/op_host/causal_conv1d_def.cpp
Normal file
@@ -0,0 +1,90 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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 causal_conv1d_def.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_def_registry.h"
|
||||
|
||||
namespace ops {
|
||||
|
||||
class CausalConv1d : public OpDef {
|
||||
public:
|
||||
explicit CausalConv1d(const char* name) : OpDef(name)
|
||||
{
|
||||
this->Input("x")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("weight")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("bias")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("convStates")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("queryStartLoc")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_INT32, ge::DT_INT64})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("cacheIndices")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_INT32, ge::DT_INT64})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("initialStateMode")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_BOOL, ge::DT_INT32, ge::DT_INT64})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
this->Input("numAcceptedTokens")
|
||||
.ParamType(OPTIONAL)
|
||||
.DataTypeList({ge::DT_INT32, ge::DT_INT64})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
|
||||
this->Output("y")
|
||||
.ParamType(REQUIRED)
|
||||
.DataType({ge::DT_FLOAT16, ge::DT_BF16})
|
||||
.FormatList({ge::FORMAT_ND})
|
||||
.AutoContiguous();
|
||||
|
||||
this->Attr("activationMode").AttrType(OPTIONAL).Int(0);
|
||||
this->Attr("padSlotId").AttrType(OPTIONAL).Int(-1);
|
||||
this->Attr("runMode").AttrType(OPTIONAL).Int(0);
|
||||
|
||||
OpAICoreConfig aicoreConfig;
|
||||
aicoreConfig.DynamicCompileStaticFlag(true)
|
||||
.DynamicFormatFlag(false)
|
||||
.DynamicRankSupportFlag(true)
|
||||
.DynamicShapeSupportFlag(true)
|
||||
.NeedCheckSupportFlag(false)
|
||||
.PrecisionReduceFlag(true)
|
||||
.ExtendCfgInfo("coreType.value", "AiCore");
|
||||
this->AICore().AddConfig("ascend910b", aicoreConfig);
|
||||
this->AICore().AddConfig("ascend910_93", aicoreConfig);
|
||||
this->AICore().AddConfig("ascend950", aicoreConfig);
|
||||
}
|
||||
};
|
||||
OP_ADD(CausalConv1d);
|
||||
|
||||
} // namespace ops
|
||||
42
csrc/moe/causal_conv1d/op_host/causal_conv1d_infershape.cpp
Normal file
42
csrc/moe/causal_conv1d/op_host/causal_conv1d_infershape.cpp
Normal file
@@ -0,0 +1,42 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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 causal_conv1d_infershape.cpp
|
||||
* \brief
|
||||
*/
|
||||
#include "register/op_impl_registry.h"
|
||||
#include "tiling_base/error_log.h"
|
||||
|
||||
using namespace ge;
|
||||
|
||||
namespace ops {
|
||||
static constexpr int64_t IDX_0 = 0;
|
||||
|
||||
static ge::graphStatus InferShapeCausalConv1d(gert::InferShapeContext* context)
|
||||
{
|
||||
OP_LOGD(context->GetNodeName(), "Begin to do InferShapeCausalConv1d");
|
||||
|
||||
// get input shapes
|
||||
const gert::Shape* xShape = context->GetInputShape(IDX_0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, xShape);
|
||||
|
||||
// get output shapes
|
||||
gert::Shape* yShape = context->GetOutputShape(IDX_0);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, yShape);
|
||||
*yShape = *xShape;
|
||||
|
||||
OP_LOGD(context->GetNodeName(), "End to do InferShapeCausalConv1d");
|
||||
return GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_INFERSHAPE(CausalConv1d).InferShape(InferShapeCausalConv1d);
|
||||
} // namespace ops
|
||||
167
csrc/moe/causal_conv1d/op_host/causal_conv1d_tiling.cpp
Normal file
167
csrc/moe/causal_conv1d/op_host/causal_conv1d_tiling.cpp
Normal file
@@ -0,0 +1,167 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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 causal_conv1d_tiling.cpp
|
||||
*/
|
||||
|
||||
#include "tiling_base/tiling_templates_registry.h"
|
||||
#include "causal_conv1d_tiling_utils.h"
|
||||
#include "causal_conv1d_tiling_planner.h"
|
||||
#include "causal_conv1d_tiling_validation.h"
|
||||
|
||||
namespace optiling {
|
||||
|
||||
using namespace Ops::Transformer::OpTiling;
|
||||
using namespace causal_conv1d_host;
|
||||
|
||||
static ge::graphStatus CausalConv1dTilingFunc(gert::TilingContext *context)
|
||||
{
|
||||
uint64_t ubSize = 0;
|
||||
uint32_t coreNum = 0;
|
||||
OP_CHECK_IF(GetPlatformInfo(context, ubSize, coreNum) != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context, "GetPlatformInfo error"), return ge::GRAPH_FAILED);
|
||||
|
||||
CausalConv1dTilingData *tiling = context->GetTilingData<CausalConv1dTilingData>();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, tiling);
|
||||
OP_CHECK_IF(memset_s(tiling, sizeof(CausalConv1dTilingData), 0, sizeof(CausalConv1dTilingData)) != EOK,
|
||||
OP_LOGE(context, "set tiling data error"), return ge::GRAPH_FAILED);
|
||||
|
||||
CausalConv1dAttrInfo attrInfo;
|
||||
OP_CHECK_IF(GetAttrsInfo(context, attrInfo) != ge::GRAPH_SUCCESS, OP_LOGE(context, "GetAttrsInfo error"),
|
||||
return ge::GRAPH_FAILED);
|
||||
bool hasBias = false;
|
||||
OP_CHECK_IF(GetShapeDtypeInfo(context, attrInfo, *tiling, hasBias) != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context, "GetShapeDtypeInfo error"), return ge::GRAPH_FAILED);
|
||||
|
||||
const int64_t &dim = tiling->dim;
|
||||
const int64_t &batch = tiling->batch;
|
||||
OP_CHECK_IF(dim <= 0 || batch <= 0, OP_LOGE(context, "dim/batch must be positive"), return ge::GRAPH_FAILED);
|
||||
|
||||
const uint32_t runModeKey = static_cast<uint32_t>(attrInfo.runMode);
|
||||
const bool &isFn = (runModeKey == CAUSAL_CONV1D_TPL_RUN_MODE_FN);
|
||||
const bool &hasActivation = (attrInfo.activationMode != 0);
|
||||
const char *plannerModeTag = "update";
|
||||
DimTileChoice baseDimChoice;
|
||||
FnExecutionPlan fnExecutionPlan = FN_EXECUTION_PLAN_INVALID;
|
||||
FnHostPlan fnHostPlan;
|
||||
const int64_t *qslData = nullptr;
|
||||
|
||||
if (isFn) {
|
||||
fnHostPlan = ChooseFnHostPlan(context, *tiling, ubSize, coreNum);
|
||||
plannerModeTag = GetFnTilingCaseName(fnHostPlan.caseKind);
|
||||
baseDimChoice = fnHostPlan.baseDimChoice;
|
||||
fnExecutionPlan = fnHostPlan.executionPlan;
|
||||
} else {
|
||||
baseDimChoice = ChooseCanonicalUpdateBaseDimChoice(context, tiling->batch, tiling->dim, coreNum);
|
||||
}
|
||||
|
||||
OP_CHECK_IF(baseDimChoice.baseDim <= 0 || baseDimChoice.baseDimCnt <= 0 || baseDimChoice.gridSize <= 0,
|
||||
OP_LOGE(context, "invalid dim tile size selection"), return ge::GRAPH_FAILED);
|
||||
|
||||
int64_t effectiveGridSize = baseDimChoice.gridSize;
|
||||
|
||||
if (isFn) {
|
||||
OP_CHECK_IF(fnHostPlan.caseKind == FN_TILING_CASE_INVALID || fnExecutionPlan == FN_EXECUTION_PLAN_INVALID ||
|
||||
!fnHostPlan.tokenBlockChoice.enabled || fnHostPlan.tokenBlockChoice.tokenBlockSize <= 0 ||
|
||||
fnHostPlan.tokenBlockChoice.tokenBlockCnt <= 0 || fnHostPlan.tokenBlockChoice.gridSize <= 0 ||
|
||||
fnHostPlan.tokenCoreMapping.tokenCoreBudget <= 0 || fnHostPlan.tokenCoreMapping.blockDim <= 0,
|
||||
OP_LOGE(context, "runMode=0 must resolve a valid unified token tiling plan"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
tiling->tokenBlockSize = fnHostPlan.tokenBlockChoice.tokenBlockSize;
|
||||
tiling->tokenBlockCnt = fnHostPlan.tokenBlockChoice.tokenBlockCnt;
|
||||
effectiveGridSize = fnHostPlan.tokenBlockChoice.gridSize;
|
||||
if (tiling->inputMode == 0) {
|
||||
fnHostPlan.tokenSeqRangePlan =
|
||||
BuildFnTokenSeqRangePlan(qslData, tiling->batch, tiling->tokenBlockSize, tiling->tokenBlockCnt);
|
||||
if (fnHostPlan.tokenSeqRangePlan.enabled) {
|
||||
tiling->hasExplicitTokenSeqRanges = 1;
|
||||
tiling->explicitTokenSeqRangeCount = fnHostPlan.tokenSeqRangePlan.rangeCount;
|
||||
for (int64_t i = 0; i < fnHostPlan.tokenSeqRangePlan.rangeCount; ++i) {
|
||||
tiling->tokenTileStartSeq[i] = fnHostPlan.tokenSeqRangePlan.tokenTileStartSeq[i];
|
||||
tiling->tokenTileEndSeq[i] = fnHostPlan.tokenSeqRangePlan.tokenTileEndSeq[i];
|
||||
}
|
||||
} else if (qslData != nullptr && tiling->tokenBlockCnt > MAX_FN_TOKEN_SEQ_RANGE_COUNT) {
|
||||
OP_LOGD(context,
|
||||
"FnTokenSeqRanges disabled: tokenBlockCnt[%ld] exceeds fixed tiling capacity[%ld].",
|
||||
tiling->tokenBlockCnt, MAX_FN_TOKEN_SEQ_RANGE_COUNT);
|
||||
}
|
||||
}
|
||||
OP_LOGD(context,
|
||||
"FnHostPlan(case=%s): inputMode[%ld], dim[%ld], cuSeqlen[%ld], baseDim[%ld], baseDimCnt[%ld], "
|
||||
"tokenCoreBudget[%ld], tokenBlockSize[%ld], tokenBlockCnt[%ld], tokenBlocksPerCore[%ld], "
|
||||
"tokenCoreTailCnt[%ld], explicitSeqRanges[%ld], baseGrid[%ld], phase1Grid[%ld], mappedBlockDim[%ld].",
|
||||
plannerModeTag, tiling->inputMode, tiling->dim, tiling->cuSeqlen, baseDimChoice.baseDim,
|
||||
baseDimChoice.baseDimCnt, fnHostPlan.tokenCoreMapping.tokenCoreBudget,
|
||||
fnHostPlan.tokenBlockChoice.tokenBlockSize, fnHostPlan.tokenBlockChoice.tokenBlockCnt,
|
||||
fnHostPlan.tokenCoreMapping.tokenBlocksPerCore, fnHostPlan.tokenCoreMapping.tokenCoreTailCnt,
|
||||
tiling->hasExplicitTokenSeqRanges,
|
||||
baseDimChoice.gridSize, fnHostPlan.tokenBlockChoice.gridSize, fnHostPlan.tokenCoreMapping.blockDim);
|
||||
}
|
||||
|
||||
uint32_t blockDim =
|
||||
(effectiveGridSize < static_cast<int64_t>(coreNum)) ? static_cast<uint32_t>(effectiveGridSize) : coreNum;
|
||||
if (isFn) {
|
||||
const int64_t mappedBlockDim = (effectiveGridSize < fnHostPlan.tokenCoreMapping.blockDim) ? effectiveGridSize : fnHostPlan.tokenCoreMapping.blockDim;
|
||||
OP_CHECK_IF(mappedBlockDim <= 0, OP_LOGE(context, "invalid mapped blockDim for runMode=0"),
|
||||
return ge::GRAPH_FAILED);
|
||||
blockDim = static_cast<uint32_t>(mappedBlockDim);
|
||||
}
|
||||
|
||||
OP_LOGD(context,
|
||||
"Tiling result: mode[%s], batch[%ld], dim[%ld], baseDim[%ld], baseDimCnt[%ld], gridSize[%ld], "
|
||||
"effectiveGrid[%ld], blockDim[%u], coreNum[%u], tokenTiling[%ld,%ld], hasActivation[%d], hasBias[%d], "
|
||||
"fnPlan[%ld].",
|
||||
plannerModeTag, batch, dim, baseDimChoice.baseDim, baseDimChoice.baseDimCnt, baseDimChoice.gridSize,
|
||||
effectiveGridSize, blockDim, coreNum, tiling->tokenBlockSize, tiling->tokenBlockCnt,
|
||||
static_cast<int32_t>(hasActivation), static_cast<int32_t>(hasBias), static_cast<int64_t>(fnExecutionPlan));
|
||||
|
||||
context->SetBlockDim(blockDim);
|
||||
tiling->baseDim = baseDimChoice.baseDim;
|
||||
tiling->baseDimCnt = baseDimChoice.baseDimCnt;
|
||||
const uint32_t fnPlanKey = NormalizeFnPlanTilingKey(runModeKey, fnExecutionPlan);
|
||||
const uint32_t widthKey = NormalizeWidthTilingKey(runModeKey, static_cast<int32_t>(tiling->width));
|
||||
if (isFn && tiling->hasInitialStateMode != 0) {
|
||||
constexpr int64_t kDtypeSize = 2;
|
||||
constexpr int64_t kSyncBytesPerBlock = 32;
|
||||
const int64_t historyCount = (tiling->width - 1 > 0) ? tiling->width - 1 : 0;
|
||||
const int64_t syncWorkspaceSize = static_cast<int64_t>(blockDim) * kSyncBytesPerBlock;
|
||||
const int64_t snapshotWorkspaceSize = tiling->batch * historyCount * tiling->dim * kDtypeSize;
|
||||
const int64_t workspaceSize =
|
||||
ASCENDC_RESERVED_WORKSPACE_SIZE + syncWorkspaceSize + snapshotWorkspaceSize;
|
||||
OP_CHECK_IF(SetWorkspaceSize(context, static_cast<size_t>(workspaceSize)) != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context, "SetWorkspaceSize error"), return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(context->SetScheduleMode(1) != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context, "SetScheduleMode(1) error"), return ge::GRAPH_FAILED);
|
||||
tiling->hasInitStateWorkspace = 1;
|
||||
} else {
|
||||
OP_CHECK_IF(SetWorkspaceSize(context, 0) != ge::GRAPH_SUCCESS, OP_LOGE(context, "SetWorkspaceSize error"),
|
||||
return ge::GRAPH_FAILED);
|
||||
tiling->hasInitStateWorkspace = 0;
|
||||
}
|
||||
|
||||
const uint64_t tilingKey = GET_TPL_TILING_KEY(runModeKey, widthKey, fnPlanKey);
|
||||
context->SetTilingKey(tilingKey);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
static ge::graphStatus TilingParseForCausalConv1d(gert::TilingParseContext *context)
|
||||
{
|
||||
OP_LOGD(context, "Enter TilingParseForCausalConv1d.");
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
IMPL_OP_OPTILING(CausalConv1d)
|
||||
.Tiling(CausalConv1dTilingFunc)
|
||||
.TilingParse<CausalConv1dCompileInfo>(TilingParseForCausalConv1d);
|
||||
|
||||
}
|
||||
312
csrc/moe/causal_conv1d/op_host/causal_conv1d_tiling_planner.h
Normal file
312
csrc/moe/causal_conv1d/op_host/causal_conv1d_tiling_planner.h
Normal file
@@ -0,0 +1,312 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_TILING_PLANNER_H
|
||||
#define CAUSAL_CONV1D_TILING_PLANNER_H
|
||||
|
||||
#include "causal_conv1d_tiling_utils.h"
|
||||
#include "../op_kernel/causal_conv1d_tiling_data.h"
|
||||
|
||||
namespace optiling::causal_conv1d_host {
|
||||
|
||||
using namespace Ops::Transformer::OpTiling;
|
||||
|
||||
inline DimTileChoice ChooseCanonicalUpdateBaseDimChoice(gert::TilingContext *context, int64_t batch, int64_t dim,
|
||||
uint32_t coreNum)
|
||||
{
|
||||
const int64_t candidates[] = {4096, 2048, 1024, 512, 384, 192};
|
||||
|
||||
auto chooseOnce = [&](bool requireExactDiv) -> DimTileChoice {
|
||||
DimTileChoice bestOver;
|
||||
int64_t bestOverGap = std::numeric_limits<int64_t>::max();
|
||||
DimTileChoice bestUnder;
|
||||
|
||||
for (int64_t baseDim : candidates) {
|
||||
if (baseDim <= 0) {
|
||||
continue;
|
||||
}
|
||||
if (requireExactDiv && (dim % baseDim != 0)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int64_t baseDimCnt = requireExactDiv ? (dim / baseDim) : CeilDivInt64(dim, baseDim);
|
||||
const int64_t gridSize = batch * baseDimCnt;
|
||||
if (gridSize <= 0) {
|
||||
continue;
|
||||
}
|
||||
|
||||
OP_LOGD(context,
|
||||
"DimTile(update) candidate[%s]: baseDim[%ld], baseDimCnt[%ld], gridSize[%ld], coreNum[%u].",
|
||||
requireExactDiv ? "exact" : "tail", baseDim, baseDimCnt, gridSize, coreNum);
|
||||
if (gridSize >= static_cast<int64_t>(coreNum)) {
|
||||
const int64_t gap = gridSize - static_cast<int64_t>(coreNum);
|
||||
if (gap < bestOverGap) {
|
||||
// bestOver = {baseDim, baseDimCnt, gridSize};
|
||||
bestOver.baseDim = baseDim;
|
||||
bestOver.baseDimCnt = baseDimCnt;
|
||||
bestOver.gridSize = gridSize;
|
||||
bestOverGap = gap;
|
||||
}
|
||||
} else if (gridSize > bestUnder.gridSize ||
|
||||
(gridSize == bestUnder.gridSize && baseDim < bestUnder.baseDim)) {
|
||||
// bestUnder = {baseDim, baseDimCnt, gridSize};
|
||||
bestUnder.baseDim = baseDim;
|
||||
bestUnder.baseDimCnt = baseDimCnt;
|
||||
bestUnder.gridSize = gridSize;
|
||||
}
|
||||
}
|
||||
|
||||
return (bestOver.baseDim != 0) ? bestOver : bestUnder;
|
||||
};
|
||||
|
||||
DimTileChoice result = chooseOnce(true);
|
||||
if (result.baseDim == 0) {
|
||||
result = chooseOnce(false);
|
||||
}
|
||||
OP_LOGD(context, "DimTile(update) chosen: baseDim[%ld], baseDimCnt[%ld], gridSize[%ld].", result.baseDim,
|
||||
result.baseDimCnt, result.gridSize);
|
||||
return result;
|
||||
}
|
||||
|
||||
inline int64_t ResolveFnTokenCoreBudget(int64_t baseDimCnt, FnExecutionPlan fnExecutionPlan, uint32_t coreNum)
|
||||
{
|
||||
if (baseDimCnt <= 0 || coreNum == 0 || fnExecutionPlan == FN_EXECUTION_PLAN_INVALID) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
int64_t tokenCoreBudget = static_cast<int64_t>(coreNum);
|
||||
if (fnExecutionPlan == FN_EXECUTION_PLAN_CUTBSD) {
|
||||
tokenCoreBudget = std::max<int64_t>(1, tokenCoreBudget / baseDimCnt);
|
||||
}
|
||||
return tokenCoreBudget;
|
||||
}
|
||||
|
||||
inline VarlenTokenTileChoice ChooseFnTokenBlockChoice(int64_t cuSeqlen, int64_t baseDimCnt,
|
||||
FnExecutionPlan fnExecutionPlan, uint32_t coreNum);
|
||||
|
||||
inline int64_t ComputeFnUbLimitedBaseDim(uint64_t ubSize)
|
||||
{
|
||||
if (ubSize <= static_cast<uint64_t>(FN_UB_RESERVED_BYTES)) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
const int64_t bytesPerElem = (RING_SLOT_CNT * BF16_FP16_ELEM_BYTES) + (FN_OUT_SLOT_CNT * BF16_FP16_ELEM_BYTES) +
|
||||
(FN_CALC_FP32_SLOT_CNT * static_cast<int64_t>(sizeof(float)));
|
||||
const int64_t budgetBytes = static_cast<int64_t>(ubSize) - FN_UB_RESERVED_BYTES;
|
||||
const int64_t ubLimitedBaseDim = AlignDownInt64(budgetBytes / bytesPerElem, DIM_ALIGN_ELEMS);
|
||||
return std::min<int64_t>(MAX_DIM_TILE_SIZE, ubLimitedBaseDim);
|
||||
}
|
||||
|
||||
inline DimTileChoice ChooseFnTokenFirstBaseDimChoice(int64_t dim)
|
||||
{
|
||||
if (dim <= 0 || dim > MAX_DIM_TILE_SIZE) {
|
||||
return {};
|
||||
}
|
||||
DimTileChoice choice;
|
||||
choice.baseDim = dim;
|
||||
choice.baseDimCnt = 1;
|
||||
choice.gridSize = 1;
|
||||
return choice;
|
||||
}
|
||||
|
||||
inline DimTileChoice ChooseFnTokenDimCoSplitBaseDimChoice(gert::TilingContext *context, int64_t dim, uint64_t ubSize,
|
||||
uint32_t coreNum)
|
||||
{
|
||||
if (dim <= 0) {
|
||||
return {};
|
||||
}
|
||||
|
||||
const int64_t ubLimitedBaseDim = ComputeFnUbLimitedBaseDim(ubSize);
|
||||
if (ubLimitedBaseDim <= 0) {
|
||||
OP_LOGD(context, "FnDimCoSplit: UB budget is too small to form a valid baseDim.");
|
||||
return {};
|
||||
}
|
||||
|
||||
DimTileChoice result;
|
||||
result.baseDim = ubLimitedBaseDim;
|
||||
result.baseDimCnt = CeilDivInt64(dim, result.baseDim);
|
||||
result.gridSize = result.baseDimCnt;
|
||||
|
||||
if (coreNum == 0 || result.baseDimCnt <= 1 || result.baseDimCnt >= static_cast<int64_t>(coreNum) ||
|
||||
(coreNum % result.baseDimCnt == 0)) {
|
||||
OP_LOGD(context,
|
||||
"FnDimCoSplit: dim[%ld], ubLimitedBaseDim[%ld], baseDimCnt[%ld], coreNum[%u], adjusted[%d].", dim,
|
||||
result.baseDim, result.baseDimCnt, coreNum, 0);
|
||||
return result;
|
||||
}
|
||||
|
||||
int64_t adjustedBaseDimCnt = result.baseDimCnt;
|
||||
while (adjustedBaseDimCnt < static_cast<int64_t>(coreNum) && (coreNum % adjustedBaseDimCnt != 0)) {
|
||||
++adjustedBaseDimCnt;
|
||||
}
|
||||
|
||||
if (adjustedBaseDimCnt >= static_cast<int64_t>(coreNum)) {
|
||||
OP_LOGD(context,
|
||||
"FnDimCoSplit: keep baseDimCnt[%ld] because no divisible adjustment exists under coreNum[%u].",
|
||||
result.baseDimCnt, coreNum);
|
||||
return result;
|
||||
}
|
||||
|
||||
const int64_t adjustedBaseDim = AlignUpInt64(CeilDivInt64(dim, adjustedBaseDimCnt), DIM_ALIGN_ELEMS);
|
||||
if (adjustedBaseDim <= 0 || adjustedBaseDim > ubLimitedBaseDim || adjustedBaseDim > MAX_DIM_TILE_SIZE) {
|
||||
OP_LOGD(context,
|
||||
"FnDimCoSplit: rejected adjusted baseDim[%ld] with baseDimCnt[%ld], ubLimitedBaseDim[%ld].",
|
||||
adjustedBaseDim, adjustedBaseDimCnt, ubLimitedBaseDim);
|
||||
return result;
|
||||
}
|
||||
|
||||
result.baseDim = adjustedBaseDim;
|
||||
result.baseDimCnt = CeilDivInt64(dim, result.baseDim);
|
||||
result.gridSize = result.baseDimCnt;
|
||||
OP_LOGD(context,
|
||||
"FnDimCoSplit: dim[%ld], ubLimitedBaseDim[%ld], adjustedBaseDim[%ld], baseDimCnt[%ld], coreNum[%u].",
|
||||
dim, ubLimitedBaseDim, result.baseDim, result.baseDimCnt, coreNum);
|
||||
return result;
|
||||
}
|
||||
|
||||
inline TokenCoreMappingChoice BuildFnTokenCoreMappingChoice(int64_t tokenBlockCnt, int64_t baseDimCnt,
|
||||
FnExecutionPlan fnExecutionPlan, uint32_t coreNum)
|
||||
{
|
||||
TokenCoreMappingChoice mapping;
|
||||
mapping.tokenCoreBudget = ResolveFnTokenCoreBudget(baseDimCnt, fnExecutionPlan, coreNum);
|
||||
if (tokenBlockCnt <= 0 || mapping.tokenCoreBudget <= 0 || baseDimCnt <= 0) {
|
||||
return mapping;
|
||||
}
|
||||
|
||||
mapping.tokenBlocksPerCore = CeilDivInt64(tokenBlockCnt, mapping.tokenCoreBudget);
|
||||
mapping.tokenCoreTailCnt =
|
||||
tokenBlockCnt - (std::max<int64_t>(0, mapping.tokenBlocksPerCore - 1) * mapping.tokenCoreBudget);
|
||||
if (mapping.tokenCoreTailCnt <= 0) {
|
||||
mapping.tokenCoreTailCnt = mapping.tokenCoreBudget;
|
||||
}
|
||||
mapping.blockDim = mapping.tokenCoreBudget * baseDimCnt;
|
||||
return mapping;
|
||||
}
|
||||
|
||||
inline FnTokenSeqRangePlan BuildFnTokenSeqRangePlan(const int64_t *qslData, int64_t batch, int64_t tokenBlockSize,
|
||||
int64_t tokenBlockCnt)
|
||||
{
|
||||
FnTokenSeqRangePlan plan;
|
||||
if (qslData == nullptr || batch <= 0 || tokenBlockSize <= 0 || tokenBlockCnt <= 0 ||
|
||||
tokenBlockCnt > MAX_FN_TOKEN_SEQ_RANGE_COUNT) {
|
||||
return plan;
|
||||
}
|
||||
|
||||
plan.enabled = true;
|
||||
plan.rangeCount = tokenBlockCnt;
|
||||
int64_t seq = 0;
|
||||
for (int64_t tokenTileId = 0; tokenTileId < tokenBlockCnt; ++tokenTileId) {
|
||||
const int64_t tokenStart = tokenTileId * tokenBlockSize;
|
||||
const int64_t tokenEnd = tokenStart + tokenBlockSize;
|
||||
|
||||
while (seq < batch && qslData[seq + 1] <= tokenStart) {
|
||||
++seq;
|
||||
}
|
||||
|
||||
int64_t endSeq = seq;
|
||||
while (endSeq < batch && qslData[endSeq] < tokenEnd) {
|
||||
++endSeq;
|
||||
}
|
||||
|
||||
plan.tokenTileStartSeq[tokenTileId] = seq;
|
||||
plan.tokenTileEndSeq[tokenTileId] = endSeq;
|
||||
}
|
||||
return plan;
|
||||
}
|
||||
|
||||
inline VarlenTokenTileChoice ChooseUnifiedFnTokenBlockPlan(gert::TilingContext *context,
|
||||
const CausalConv1dTilingData &tiling,
|
||||
const DimTileChoice &baseDimChoice,
|
||||
FnExecutionPlan fnExecutionPlan,
|
||||
uint32_t coreNum)
|
||||
{
|
||||
VarlenTokenTileChoice tokenBlockChoice;
|
||||
if ((tiling.inputMode != 0 && tiling.inputMode != 1) || tiling.batch <= 0 || tiling.cuSeqlen <= 0 ||
|
||||
baseDimChoice.baseDimCnt <= 0 || coreNum == 0 || fnExecutionPlan == FN_EXECUTION_PLAN_INVALID) {
|
||||
return tokenBlockChoice;
|
||||
}
|
||||
if (tiling.hasNumAcceptedTokens != 0) {
|
||||
OP_LOGD(context, "Varlen token tiling disabled: speculative decode still uses the existing seq mapping.");
|
||||
return tokenBlockChoice;
|
||||
}
|
||||
|
||||
tokenBlockChoice = ChooseFnTokenBlockChoice(tiling.cuSeqlen, baseDimChoice.baseDimCnt, fnExecutionPlan, coreNum);
|
||||
|
||||
OP_LOGD(context,
|
||||
"FnTokenTile(plan=%ld): cuSeqlen[%ld], baseDimCnt[%ld], tokenBlockSize[%ld], "
|
||||
"tokenBlockCnt[%ld], gridSize[%ld].",
|
||||
static_cast<int64_t>(fnExecutionPlan), tiling.cuSeqlen, baseDimChoice.baseDimCnt,
|
||||
tokenBlockChoice.tokenBlockSize, tokenBlockChoice.tokenBlockCnt, tokenBlockChoice.gridSize);
|
||||
return tokenBlockChoice;
|
||||
}
|
||||
|
||||
inline VarlenTokenTileChoice ChooseFnTokenBlockChoice(int64_t cuSeqlen, int64_t baseDimCnt,
|
||||
FnExecutionPlan fnExecutionPlan, uint32_t coreNum)
|
||||
{
|
||||
VarlenTokenTileChoice tokenBlockChoice;
|
||||
const int64_t tokenCoreBudget = ResolveFnTokenCoreBudget(baseDimCnt, fnExecutionPlan, coreNum);
|
||||
if (cuSeqlen <= 0 || tokenCoreBudget <= 0) {
|
||||
return tokenBlockChoice;
|
||||
}
|
||||
|
||||
tokenBlockChoice.enabled = true;
|
||||
const int64_t idealBlockSize = CeilDivInt64(cuSeqlen, tokenCoreBudget);
|
||||
tokenBlockChoice.tokenBlockSize = (idealBlockSize > 0) ? idealBlockSize : 1;
|
||||
tokenBlockChoice.tokenBlockCnt = CeilDivInt64(cuSeqlen, tokenBlockChoice.tokenBlockSize);
|
||||
tokenBlockChoice.gridSize = tokenBlockChoice.tokenBlockCnt * baseDimCnt;
|
||||
return tokenBlockChoice;
|
||||
}
|
||||
|
||||
inline FnHostPlan ChooseFnHostPlan(gert::TilingContext *context, const CausalConv1dTilingData &tiling, uint64_t ubSize,
|
||||
uint32_t coreNum)
|
||||
{
|
||||
FnHostPlan plan;
|
||||
if ((tiling.inputMode != 0 && tiling.inputMode != 1) || tiling.batch <= 0 || tiling.cuSeqlen <= 0 ||
|
||||
tiling.dim <= 0 || coreNum == 0) {
|
||||
return plan;
|
||||
}
|
||||
|
||||
if (tiling.dim <= MAX_DIM_TILE_SIZE) {
|
||||
plan.caseKind = FN_TILING_CASE_TOKEN_FIRST;
|
||||
plan.executionPlan = FN_EXECUTION_PLAN_CUTBS;
|
||||
plan.baseDimChoice = ChooseFnTokenFirstBaseDimChoice(tiling.dim);
|
||||
} else {
|
||||
plan.caseKind = FN_TILING_CASE_TOKEN_DIM_CO_SPLIT;
|
||||
plan.executionPlan = FN_EXECUTION_PLAN_CUTBSD;
|
||||
plan.baseDimChoice = ChooseFnTokenDimCoSplitBaseDimChoice(context, tiling.dim, ubSize, coreNum);
|
||||
}
|
||||
|
||||
if (plan.baseDimChoice.baseDim <= 0 || plan.baseDimChoice.baseDimCnt <= 0) {
|
||||
return {};
|
||||
}
|
||||
|
||||
plan.baseDimChoice.gridSize = tiling.batch * plan.baseDimChoice.baseDimCnt;
|
||||
plan.tokenBlockChoice =
|
||||
ChooseUnifiedFnTokenBlockPlan(context, tiling, plan.baseDimChoice, plan.executionPlan, coreNum);
|
||||
if (!plan.tokenBlockChoice.enabled || plan.tokenBlockChoice.tokenBlockSize <= 0 ||
|
||||
plan.tokenBlockChoice.tokenBlockCnt <= 0 || plan.tokenBlockChoice.gridSize <= 0) {
|
||||
return {};
|
||||
}
|
||||
|
||||
plan.tokenCoreMapping = BuildFnTokenCoreMappingChoice(plan.tokenBlockChoice.tokenBlockCnt,
|
||||
plan.baseDimChoice.baseDimCnt, plan.executionPlan, coreNum);
|
||||
if (plan.tokenCoreMapping.tokenCoreBudget <= 0 || plan.tokenCoreMapping.blockDim <= 0) {
|
||||
return {};
|
||||
}
|
||||
if (plan.tokenCoreMapping.blockDim > static_cast<int64_t>(coreNum)) {
|
||||
plan.tokenCoreMapping.blockDim = static_cast<int64_t>(coreNum);
|
||||
}
|
||||
return plan;
|
||||
}
|
||||
|
||||
} // namespace optiling::causal_conv1d_host
|
||||
|
||||
#endif // CAUSAL_CONV1D_TILING_PLANNER_H
|
||||
165
csrc/moe/causal_conv1d/op_host/causal_conv1d_tiling_utils.h
Normal file
165
csrc/moe/causal_conv1d/op_host/causal_conv1d_tiling_utils.h
Normal file
@@ -0,0 +1,165 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_TILING_UTILS_H
|
||||
#define CAUSAL_CONV1D_TILING_UTILS_H
|
||||
|
||||
#include "tiling_base/tiling_util.h"
|
||||
#include "../op_kernel/causal_conv1d_tiling_key.h"
|
||||
|
||||
namespace optiling::causal_conv1d_host {
|
||||
|
||||
constexpr uint32_t X_INDEX = 0;
|
||||
constexpr uint32_t WEIGHT_INDEX = 1;
|
||||
constexpr uint32_t BIAS_INDEX = 2;
|
||||
constexpr uint32_t CONV_STATES_INDEX = 3;
|
||||
constexpr uint32_t QUERY_START_LOC_INDEX = 4;
|
||||
constexpr uint32_t CACHE_INDICES_INDEX = 5;
|
||||
constexpr uint32_t INITIAL_STATE_MODE_INDEX = 6;
|
||||
constexpr uint32_t NUM_ACCEPTED_TOKENS_INDEX = 7;
|
||||
|
||||
constexpr int32_t ATTR_ACTIVATION_MODE_INDEX = 0;
|
||||
constexpr int32_t ATTR_PAD_SLOT_ID_INDEX = 1;
|
||||
constexpr int32_t ATTR_RUN_MODE_INDEX = 2;
|
||||
constexpr int64_t ASCENDC_RESERVED_WORKSPACE_SIZE = 16 * 1024 * 1024;
|
||||
|
||||
struct CausalConv1dCompileInfo {
|
||||
uint64_t ubSize = 0;
|
||||
uint32_t coreNum = 0;
|
||||
};
|
||||
|
||||
struct CausalConv1dAttrInfo {
|
||||
int64_t activationMode = 0;
|
||||
int64_t padSlotId = -1;
|
||||
int64_t runMode = 0;
|
||||
};
|
||||
|
||||
struct DimTileChoice {
|
||||
int64_t baseDim = 0;
|
||||
int64_t baseDimCnt = 0;
|
||||
int64_t gridSize = 0;
|
||||
};
|
||||
|
||||
struct VarlenTokenTileChoice {
|
||||
bool enabled = false;
|
||||
int64_t tokenBlockSize = 0;
|
||||
int64_t tokenBlockCnt = 0;
|
||||
int64_t gridSize = 0;
|
||||
};
|
||||
|
||||
enum FnTilingCaseKind : int64_t {
|
||||
FN_TILING_CASE_INVALID = 0,
|
||||
FN_TILING_CASE_TOKEN_FIRST = 1,
|
||||
FN_TILING_CASE_TOKEN_DIM_CO_SPLIT = 2,
|
||||
};
|
||||
|
||||
struct TokenCoreMappingChoice {
|
||||
int64_t tokenCoreBudget = 0;
|
||||
int64_t tokenBlocksPerCore = 0;
|
||||
int64_t tokenCoreTailCnt = 0;
|
||||
int64_t blockDim = 0;
|
||||
};
|
||||
|
||||
constexpr int64_t MAX_FN_TOKEN_SEQ_RANGE_COUNT = 128;
|
||||
|
||||
struct FnTokenSeqRangePlan {
|
||||
bool enabled = false;
|
||||
int64_t rangeCount = 0;
|
||||
int64_t tokenTileStartSeq[MAX_FN_TOKEN_SEQ_RANGE_COUNT] = {};
|
||||
int64_t tokenTileEndSeq[MAX_FN_TOKEN_SEQ_RANGE_COUNT] = {};
|
||||
};
|
||||
|
||||
struct FnHostPlan {
|
||||
FnTilingCaseKind caseKind = FN_TILING_CASE_INVALID;
|
||||
FnExecutionPlan executionPlan = FN_EXECUTION_PLAN_INVALID;
|
||||
DimTileChoice baseDimChoice;
|
||||
VarlenTokenTileChoice tokenBlockChoice;
|
||||
TokenCoreMappingChoice tokenCoreMapping;
|
||||
FnTokenSeqRangePlan tokenSeqRangePlan;
|
||||
};
|
||||
|
||||
constexpr int64_t DIM_ALIGN_BYTES = 32;
|
||||
constexpr int64_t BF16_FP16_ELEM_BYTES = 2;
|
||||
constexpr int64_t DIM_ALIGN_ELEMS = DIM_ALIGN_BYTES / BF16_FP16_ELEM_BYTES;
|
||||
constexpr int64_t MAX_DIM_TILE_SIZE = 4096;
|
||||
constexpr int64_t FN_UB_RESERVED_BYTES = 512;
|
||||
constexpr int64_t RING_SLOT_CNT = 5;
|
||||
constexpr int64_t FN_OUT_SLOT_CNT = 2;
|
||||
constexpr int64_t FN_CALC_FP32_SLOT_CNT = 8;
|
||||
|
||||
inline uint32_t NormalizeFnPlanTilingKey(uint32_t runModeKey, FnExecutionPlan fnExecutionPlan)
|
||||
{
|
||||
if (runModeKey != CAUSAL_CONV1D_TPL_RUN_MODE_FN) {
|
||||
return CAUSAL_CONV1D_TPL_FN_PLAN_INVALID;
|
||||
}
|
||||
switch (fnExecutionPlan) {
|
||||
case FN_EXECUTION_PLAN_CUTBS:
|
||||
return CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS;
|
||||
case FN_EXECUTION_PLAN_CUTBSD:
|
||||
return CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD;
|
||||
default:
|
||||
return CAUSAL_CONV1D_TPL_FN_PLAN_INVALID;
|
||||
}
|
||||
}
|
||||
|
||||
inline uint32_t NormalizeWidthTilingKey(uint32_t runModeKey, int32_t width)
|
||||
{
|
||||
if (runModeKey != CAUSAL_CONV1D_TPL_RUN_MODE_FN) {
|
||||
return CAUSAL_CONV1D_TPL_WIDTH_RUNTIME;
|
||||
}
|
||||
switch (width) {
|
||||
case 2:
|
||||
return CAUSAL_CONV1D_TPL_WIDTH_2;
|
||||
case 3:
|
||||
return CAUSAL_CONV1D_TPL_WIDTH_3;
|
||||
case 4:
|
||||
return CAUSAL_CONV1D_TPL_WIDTH_4;
|
||||
default:
|
||||
return CAUSAL_CONV1D_TPL_WIDTH_RUNTIME;
|
||||
}
|
||||
}
|
||||
|
||||
inline int64_t CeilDivInt64(int64_t x, int64_t y)
|
||||
{
|
||||
return (x + y - 1) / y;
|
||||
}
|
||||
|
||||
inline int64_t AlignDownInt64(int64_t value, int64_t align)
|
||||
{
|
||||
if (align <= 0 || value <= 0) {
|
||||
return 0;
|
||||
}
|
||||
return (value / align) * align;
|
||||
}
|
||||
|
||||
inline int64_t AlignUpInt64(int64_t value, int64_t align)
|
||||
{
|
||||
if (align <= 0 || value <= 0) {
|
||||
return 0;
|
||||
}
|
||||
return CeilDivInt64(value, align) * align;
|
||||
}
|
||||
|
||||
inline const char *GetFnTilingCaseName(FnTilingCaseKind caseKind)
|
||||
{
|
||||
switch (caseKind) {
|
||||
case FN_TILING_CASE_TOKEN_FIRST:
|
||||
return "token_first";
|
||||
case FN_TILING_CASE_TOKEN_DIM_CO_SPLIT:
|
||||
return "token_dim_co_split";
|
||||
default:
|
||||
return "invalid";
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace optiling::causal_conv1d_host
|
||||
|
||||
#endif // CAUSAL_CONV1D_TILING_UTILS_H
|
||||
385
csrc/moe/causal_conv1d/op_host/causal_conv1d_tiling_validation.h
Normal file
385
csrc/moe/causal_conv1d/op_host/causal_conv1d_tiling_validation.h
Normal file
@@ -0,0 +1,385 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_TILING_VALIDATION_H
|
||||
#define CAUSAL_CONV1D_TILING_VALIDATION_H
|
||||
|
||||
#include "tiling_base/tiling_util.h"
|
||||
#include "causal_conv1d_tiling_utils.h"
|
||||
#include "../op_kernel/causal_conv1d_tiling_data.h"
|
||||
|
||||
namespace optiling::causal_conv1d_host {
|
||||
|
||||
using namespace Ops::Transformer::OpTiling;
|
||||
|
||||
inline ge::graphStatus GetPlatformInfo(gert::TilingContext *context, uint64_t &ubSize, uint32_t &coreNum)
|
||||
{
|
||||
fe::PlatFormInfos *platformInfoPtr = context->GetPlatformInfo();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, platformInfoPtr);
|
||||
auto ascendcPlatform = platform_ascendc::PlatformAscendC(platformInfoPtr);
|
||||
coreNum = ascendcPlatform.GetCoreNumAiv();
|
||||
OP_CHECK_IF(coreNum == 0, OP_LOGE(context, "coreNum is 0"), return ge::GRAPH_FAILED);
|
||||
ascendcPlatform.GetCoreMemSize(platform_ascendc::CoreMemType::UB, ubSize);
|
||||
OP_CHECK_IF(ubSize == 0, OP_LOGE(context, "ubSize is 0"), return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
inline ge::graphStatus SetWorkspaceSize(gert::TilingContext *context, size_t workspaceSize)
|
||||
{
|
||||
size_t *currentWorkspace = context->GetWorkspaceSizes(1);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, currentWorkspace);
|
||||
currentWorkspace[0] = workspaceSize;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
inline ge::graphStatus GetAttrsInfo(gert::TilingContext *context, CausalConv1dAttrInfo &attrInfo)
|
||||
{
|
||||
auto attrs = context->GetAttrs();
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, attrs);
|
||||
|
||||
const int64_t *activationModePtr = attrs->GetAttrPointer<int64_t>(ATTR_ACTIVATION_MODE_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, activationModePtr);
|
||||
attrInfo.activationMode = *activationModePtr;
|
||||
OP_CHECK_IF(attrInfo.activationMode != 0 && attrInfo.activationMode != 1,
|
||||
OP_LOGE(context, "activationMode only supports 0/1"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
const int64_t *padSlotIdPtr = attrs->GetAttrPointer<int64_t>(ATTR_PAD_SLOT_ID_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, padSlotIdPtr);
|
||||
attrInfo.padSlotId = *padSlotIdPtr;
|
||||
|
||||
const int64_t *runModePtr = attrs->GetAttrPointer<int64_t>(ATTR_RUN_MODE_INDEX);
|
||||
attrInfo.runMode = (runModePtr == nullptr) ? 0 : *runModePtr;
|
||||
OP_CHECK_IF(attrInfo.runMode != 0 && attrInfo.runMode != 1, OP_LOGE(context, "runMode only supports 0/1"),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
inline ge::graphStatus ValidateAlignedDim(gert::TilingContext *context, int64_t dim)
|
||||
{
|
||||
OP_CHECK_IF(dim % DIM_ALIGN_ELEMS != 0,
|
||||
OP_LOGE(context,
|
||||
"dim must satisfy dim %% %ld == 0 for causal_conv1d; "
|
||||
"x/weight/convStates last dimension and bias length must all use the same aligned dim, "
|
||||
"got dim=%ld.",
|
||||
DIM_ALIGN_ELEMS, dim),
|
||||
return ge::GRAPH_FAILED);
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
inline ge::graphStatus GetShapeDtypeInfo(gert::TilingContext *context, const CausalConv1dAttrInfo &attrInfo,
|
||||
CausalConv1dTilingData &tiling, bool &hasBias)
|
||||
{
|
||||
const bool isDecodeMode = (attrInfo.runMode == 1);
|
||||
tiling.activationMode = attrInfo.activationMode;
|
||||
tiling.padSlotId = attrInfo.padSlotId;
|
||||
|
||||
auto xShapePtr = context->GetInputShape(X_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, xShapePtr);
|
||||
auto xShape = EnsureNotScalar(xShapePtr->GetStorageShape());
|
||||
|
||||
int64_t dim = 0;
|
||||
int64_t cuSeqlen = 0;
|
||||
int64_t seqLen = 0;
|
||||
int64_t batch = 0;
|
||||
int64_t inputMode = 0;
|
||||
|
||||
if (xShape.GetDimNum() == 2) {
|
||||
if (isDecodeMode) {
|
||||
inputMode = 2;
|
||||
batch = xShape.GetDim(0);
|
||||
dim = xShape.GetDim(1);
|
||||
seqLen = 1;
|
||||
cuSeqlen = batch;
|
||||
OP_CHECK_IF(batch <= 0 || dim <= 0, OP_LOGE(context, "invalid x shape for 2D decode mode"),
|
||||
return ge::GRAPH_FAILED);
|
||||
} else {
|
||||
inputMode = 0;
|
||||
cuSeqlen = xShape.GetDim(0);
|
||||
dim = xShape.GetDim(1);
|
||||
seqLen = 0;
|
||||
OP_CHECK_IF(dim <= 0 || cuSeqlen < 0, OP_LOGE(context, "invalid x shape for 2D varlen mode"),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
} else if (xShape.GetDimNum() == 3) {
|
||||
inputMode = 1;
|
||||
batch = xShape.GetDim(0);
|
||||
seqLen = xShape.GetDim(1);
|
||||
dim = xShape.GetDim(2);
|
||||
cuSeqlen = batch * seqLen;
|
||||
OP_CHECK_IF(batch <= 0 || dim <= 0 || seqLen <= 0, OP_LOGE(context, "invalid x shape for 3D batch mode"),
|
||||
return ge::GRAPH_FAILED);
|
||||
} else {
|
||||
OP_LOGE(context, "x must be 2D (cu_seqlen, dim) or 3D (batch, seqlen, dim)");
|
||||
return ge::GRAPH_FAILED;
|
||||
}
|
||||
OP_CHECK_IF(ValidateAlignedDim(context, dim) != ge::GRAPH_SUCCESS,
|
||||
OP_LOGE(context, "dim alignment validation failed"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto wShapePtr = context->GetInputShape(WEIGHT_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, wShapePtr);
|
||||
auto wShape = EnsureNotScalar(wShapePtr->GetStorageShape());
|
||||
OP_CHECK_IF(wShape.GetDimNum() != 2, OP_LOGE(context, "weight must be 2D: (width, dim)"), return ge::GRAPH_FAILED);
|
||||
const int64_t width = wShape.GetDim(0);
|
||||
const int64_t wDim = wShape.GetDim(1);
|
||||
OP_CHECK_IF(wDim != dim, OP_LOGE(context, "weight.shape[1] must equal dim"), return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(width < 2 || width > 4, OP_LOGE(context, "Only support width in [2,4] now, actually is %ld.", width),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto sShapePtr = context->GetInputShape(CONV_STATES_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, sShapePtr);
|
||||
auto sShape = EnsureNotScalar(sShapePtr->GetStorageShape());
|
||||
OP_CHECK_IF(sShape.GetDimNum() != 3, OP_LOGE(context, "convStates must be 3D: (num_cache_lines, state_len, dim)"),
|
||||
return ge::GRAPH_FAILED);
|
||||
const int64_t numCacheLines = sShape.GetDim(0);
|
||||
const int64_t stateLen = sShape.GetDim(1);
|
||||
const int64_t sDim = sShape.GetDim(2);
|
||||
OP_CHECK_IF(numCacheLines <= 0, OP_LOGE(context, "convStates.shape[0] (num_cache_lines) must be > 0"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(sDim != dim, OP_LOGE(context, "convStates.shape[2] must equal dim"), return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(stateLen < (width - 1), OP_LOGE(context, "convStates.shape[1] must be >= width-1"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto qslShapePtr = context->GetOptionalInputShape(QUERY_START_LOC_INDEX);
|
||||
const gert::CompileTimeTensorDesc *qslDesc = context->GetOptionalInputDesc(QUERY_START_LOC_INDEX);
|
||||
bool qslAbsent = true;
|
||||
int64_t qslSize = 0;
|
||||
if (qslShapePtr != nullptr) {
|
||||
const auto qslStorageShape = qslShapePtr->GetStorageShape();
|
||||
const int64_t qslDimNum = qslStorageShape.GetDimNum();
|
||||
qslAbsent = (qslDimNum == 0) || (qslDimNum == 1 && qslStorageShape.GetDim(0) <= 0);
|
||||
if (!qslAbsent) {
|
||||
auto qslShape = EnsureNotScalar(qslStorageShape);
|
||||
OP_CHECK_IF(qslShape.GetDimNum() != 1, OP_LOGE(context, "queryStartLoc must be 1D"),
|
||||
return ge::GRAPH_FAILED);
|
||||
qslSize = qslShape.GetDim(0);
|
||||
OP_CHECK_IF(qslSize < 1, OP_LOGE(context, "queryStartLoc.size must be >= 1"), return ge::GRAPH_FAILED);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, qslDesc);
|
||||
const ge::DataType qslDtype = qslDesc->GetDataType();
|
||||
OP_CHECK_IF(qslDtype != ge::DT_INT32 && qslDtype != ge::DT_INT64,
|
||||
OP_LOGE(context, "queryStartLoc dtype must be int32 or int64"), return ge::GRAPH_FAILED);
|
||||
}
|
||||
}
|
||||
|
||||
if (qslAbsent) {
|
||||
OP_CHECK_IF(inputMode == 0, OP_LOGE(context, "queryStartLoc is required in 2D varlen mode (inputMode=0)"),
|
||||
return ge::GRAPH_FAILED);
|
||||
qslSize = batch + 1;
|
||||
}
|
||||
tiling.hasQueryStartLoc = qslAbsent ? 0 : 1;
|
||||
|
||||
OP_CHECK_IF(cuSeqlen > static_cast<int64_t>(std::numeric_limits<int32_t>::max()),
|
||||
OP_LOGE(context, "cuSeqlen is too large for int32 indexing, got %ld", cuSeqlen),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (!qslAbsent && isDecodeMode && inputMode == 2) {
|
||||
const int64_t batchFromQsl = qslSize - 1;
|
||||
if (batchFromQsl != batch) {
|
||||
inputMode = 0;
|
||||
cuSeqlen = xShape.GetDim(0);
|
||||
batch = batchFromQsl;
|
||||
seqLen = 0;
|
||||
OP_CHECK_IF(dim <= 0 || cuSeqlen < 0 || batch < 0,
|
||||
OP_LOGE(context, "invalid x/queryStartLoc shapes for 2D varlen decode mode"),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
}
|
||||
|
||||
if (inputMode == 0) {
|
||||
batch = qslSize - 1;
|
||||
}
|
||||
if (!qslAbsent && (inputMode == 1 || inputMode == 2)) {
|
||||
OP_CHECK_IF(qslSize != batch + 1, OP_LOGE(context, "queryStartLoc.size must equal batch + 1"),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
if (isDecodeMode) {
|
||||
const int64_t decodeSeqLen = (inputMode == 1) ? seqLen : 1;
|
||||
OP_CHECK_IF(decodeSeqLen < 1, OP_LOGE(context, "decode mode requires seqlen >= 1, actual is %ld", decodeSeqLen),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
tiling.hasCacheIndices = 0;
|
||||
tiling.cacheIndicesStride = 1;
|
||||
bool ciAbsent = true;
|
||||
auto ciShapePtr = context->GetOptionalInputShape(CACHE_INDICES_INDEX);
|
||||
if (ciShapePtr != nullptr) {
|
||||
const auto ciStorageShape = ciShapePtr->GetStorageShape();
|
||||
const int64_t ciDimNum = ciStorageShape.GetDimNum();
|
||||
ciAbsent = (ciDimNum == 0) || (ciDimNum == 1 && ciStorageShape.GetDim(0) <= 0);
|
||||
if (!ciAbsent) {
|
||||
auto ciShape = EnsureNotScalar(ciStorageShape);
|
||||
// Spec decode passes cache indices as [batch, num_spec + 1];
|
||||
// kernels read the first column with cacheIndicesStride.
|
||||
OP_CHECK_IF(ciShape.GetDimNum() != 1 && ciShape.GetDimNum() != 2,
|
||||
OP_LOGE(context, "cacheIndices must be 1D or 2D"), return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(ciShape.GetDim(0) != batch, OP_LOGE(context, "cacheIndices first dim must equal batch"),
|
||||
return ge::GRAPH_FAILED);
|
||||
if (ciShape.GetDimNum() == 2) {
|
||||
OP_CHECK_IF(ciShape.GetDim(1) <= 0, OP_LOGE(context, "cacheIndices second dim must be positive"),
|
||||
return ge::GRAPH_FAILED);
|
||||
tiling.cacheIndicesStride = ciShape.GetDim(1);
|
||||
}
|
||||
tiling.hasCacheIndices = 1;
|
||||
}
|
||||
}
|
||||
if (ciAbsent) {
|
||||
OP_CHECK_IF(numCacheLines < batch,
|
||||
OP_LOGE(context,
|
||||
"cacheIndices is absent, requires convStates.shape[0] (num_cache_lines) >= batch for "
|
||||
"identity mapping, got num_cache_lines=%ld batch=%ld",
|
||||
numCacheLines, batch),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
tiling.hasInitialStateMode = 0;
|
||||
auto ismShapePtr = context->GetOptionalInputShape(INITIAL_STATE_MODE_INDEX);
|
||||
if (ismShapePtr != nullptr) {
|
||||
const auto ismStorageShape = ismShapePtr->GetStorageShape();
|
||||
const int64_t ismDimNum = ismStorageShape.GetDimNum();
|
||||
const bool ismAbsent = (ismDimNum == 0) || (ismDimNum == 1 && ismStorageShape.GetDim(0) <= 0);
|
||||
if (!ismAbsent) {
|
||||
OP_CHECK_IF(isDecodeMode,
|
||||
OP_LOGE(context, "initialStateMode is only supported in runMode=0 (fn/prefill)"),
|
||||
return ge::GRAPH_FAILED);
|
||||
auto ismShape = EnsureNotScalar(ismStorageShape);
|
||||
OP_CHECK_IF(ismShape.GetDimNum() != 1, OP_LOGE(context, "initialStateMode must be 1D"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(ismShape.GetDim(0) != batch, OP_LOGE(context, "initialStateMode.size must equal batch"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
tiling.hasInitialStateMode = 1;
|
||||
}
|
||||
}
|
||||
|
||||
tiling.hasNumAcceptedTokens = 0;
|
||||
auto natShapePtr = context->GetOptionalInputShape(NUM_ACCEPTED_TOKENS_INDEX);
|
||||
if (natShapePtr != nullptr) {
|
||||
const auto natStorageShape = natShapePtr->GetStorageShape();
|
||||
const int64_t natDimNum = natStorageShape.GetDimNum();
|
||||
const bool natAbsent = (natDimNum == 0) || (natDimNum == 1 && natStorageShape.GetDim(0) <= 0);
|
||||
if (!natAbsent) {
|
||||
OP_CHECK_IF(!isDecodeMode,
|
||||
OP_LOGE(context, "numAcceptedTokens is only supported in runMode=1 (decode/update)"),
|
||||
return ge::GRAPH_FAILED);
|
||||
auto natShape = EnsureNotScalar(natStorageShape);
|
||||
OP_CHECK_IF(natShape.GetDimNum() != 1, OP_LOGE(context, "numAcceptedTokens must be 1D"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(natShape.GetDim(0) != batch, OP_LOGE(context, "numAcceptedTokens.size must equal batch"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (inputMode == 1) {
|
||||
const int64_t reqStateLen = (width - 1) + (seqLen - 1);
|
||||
OP_CHECK_IF(stateLen < reqStateLen,
|
||||
OP_LOGE(context,
|
||||
"spec decode requires stateLen >= (width-1) + (seqlen-1), got stateLen=%ld req=%ld",
|
||||
stateLen, reqStateLen),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
tiling.hasNumAcceptedTokens = 1;
|
||||
}
|
||||
}
|
||||
|
||||
tiling.hasBias = 0;
|
||||
hasBias = false;
|
||||
auto biasShapePtr = context->GetOptionalInputShape(BIAS_INDEX);
|
||||
if (biasShapePtr != nullptr) {
|
||||
const auto biasStorageShape = biasShapePtr->GetStorageShape();
|
||||
const int64_t biasDimNum = biasStorageShape.GetDimNum();
|
||||
const bool biasAbsent = (biasDimNum == 0) || (biasDimNum == 1 && biasStorageShape.GetDim(0) <= 0);
|
||||
if (!biasAbsent) {
|
||||
auto biasShape = EnsureNotScalar(biasStorageShape);
|
||||
OP_CHECK_IF(biasShape.GetDimNum() != 1, OP_LOGE(context, "bias must be 1D: (dim,)"),
|
||||
return ge::GRAPH_FAILED);
|
||||
OP_CHECK_IF(biasShape.GetDim(0) != dim, OP_LOGE(context, "bias.size must equal dim"),
|
||||
return ge::GRAPH_FAILED);
|
||||
tiling.hasBias = 1;
|
||||
hasBias = true;
|
||||
}
|
||||
}
|
||||
|
||||
auto xDesc = context->GetInputDesc(X_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, xDesc);
|
||||
const ge::DataType xDtype = xDesc->GetDataType();
|
||||
OP_CHECK_IF(xDtype != ge::DT_BF16 && xDtype != ge::DT_FLOAT16,
|
||||
OP_LOGE(context, "x dtype only supports bf16/fp16"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
auto wDesc = context->GetInputDesc(WEIGHT_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, wDesc);
|
||||
OP_CHECK_IF(wDesc->GetDataType() != xDtype, OP_LOGE(context, "weight dtype must equal x dtype"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (hasBias) {
|
||||
auto biasDesc = context->GetOptionalInputDesc(BIAS_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, biasDesc);
|
||||
OP_CHECK_IF(biasDesc->GetDataType() != xDtype, OP_LOGE(context, "bias dtype must equal x dtype"),
|
||||
return ge::GRAPH_FAILED);
|
||||
}
|
||||
|
||||
auto sDesc = context->GetInputDesc(CONV_STATES_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, sDesc);
|
||||
OP_CHECK_IF(sDesc->GetDataType() != xDtype, OP_LOGE(context, "convStates dtype must equal x dtype"),
|
||||
return ge::GRAPH_FAILED);
|
||||
|
||||
if (!qslAbsent) {
|
||||
auto qslDesc2 = context->GetOptionalInputDesc(QUERY_START_LOC_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, qslDesc2);
|
||||
const ge::DataType qslDtype = qslDesc2->GetDataType();
|
||||
OP_CHECK_IF(qslDtype != ge::DT_INT32 && qslDtype != ge::DT_INT64,
|
||||
OP_LOGE(context, "queryStartLoc dtype must be int32 or int64"), return ge::GRAPH_FAILED);
|
||||
tiling.queryStartLocUseInt64 = (qslDtype == ge::DT_INT64) ? 1 : 0;
|
||||
}
|
||||
tiling.cacheIndicesUseInt64 = 0;
|
||||
if (tiling.hasCacheIndices == 1) {
|
||||
auto ciDesc = context->GetOptionalInputDesc(CACHE_INDICES_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, ciDesc);
|
||||
const ge::DataType ciDtype = ciDesc->GetDataType();
|
||||
OP_CHECK_IF(ciDtype != ge::DT_INT32 && ciDtype != ge::DT_INT64,
|
||||
OP_LOGE(context, "cacheIndices dtype must be int32 or int64"), return ge::GRAPH_FAILED);
|
||||
tiling.cacheIndicesUseInt64 = (ciDtype == ge::DT_INT64) ? 1 : 0;
|
||||
}
|
||||
if (tiling.hasInitialStateMode == 1) {
|
||||
auto ismDesc = context->GetOptionalInputDesc(INITIAL_STATE_MODE_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, ismDesc);
|
||||
const ge::DataType ismDtype = ismDesc->GetDataType();
|
||||
OP_CHECK_IF(ismDtype != ge::DT_BOOL && ismDtype != ge::DT_INT32 && ismDtype != ge::DT_INT64,
|
||||
OP_LOGE(context, "initialStateMode dtype must be bool, int32 or int64"), return ge::GRAPH_FAILED);
|
||||
tiling.initialStateModeDtype =
|
||||
(ismDtype == ge::DT_INT64) ? 2 : ((ismDtype == ge::DT_INT32) ? 1 : 0);
|
||||
}
|
||||
tiling.numAcceptedTokensUseInt64 = 0;
|
||||
if (tiling.hasNumAcceptedTokens == 1) {
|
||||
OP_CHECK_IF(width != 4, OP_LOGE(context, "numAcceptedTokens is only supported for width=4 currently"),
|
||||
return ge::GRAPH_FAILED);
|
||||
auto natDesc = context->GetOptionalInputDesc(NUM_ACCEPTED_TOKENS_INDEX);
|
||||
OP_CHECK_NULL_WITH_CONTEXT(context, natDesc);
|
||||
const ge::DataType natDtype = natDesc->GetDataType();
|
||||
OP_CHECK_IF(natDtype != ge::DT_INT32 && natDtype != ge::DT_INT64,
|
||||
OP_LOGE(context, "numAcceptedTokens dtype must be int32 or int64"), return ge::GRAPH_FAILED);
|
||||
tiling.numAcceptedTokensUseInt64 = (natDtype == ge::DT_INT64) ? 1 : 0;
|
||||
}
|
||||
|
||||
tiling.dim = dim;
|
||||
tiling.cuSeqlen = cuSeqlen;
|
||||
tiling.seqLen = seqLen;
|
||||
tiling.inputMode = inputMode;
|
||||
tiling.width = width;
|
||||
tiling.stateLen = stateLen;
|
||||
tiling.numCacheLines = numCacheLines;
|
||||
tiling.batch = batch;
|
||||
return ge::GRAPH_SUCCESS;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
#endif
|
||||
61
csrc/moe/causal_conv1d/op_host/math_util.h
Normal file
61
csrc/moe/causal_conv1d/op_host/math_util.h
Normal file
@@ -0,0 +1,61 @@
|
||||
/**
|
||||
* 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 math_util.h
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#ifndef TILING_MATMUL_MATH_UTIL_H
|
||||
#define TILING_MATMUL_MATH_UTIL_H
|
||||
|
||||
#include <array>
|
||||
#include <cstdint>
|
||||
#include <vector>
|
||||
#include <utility>
|
||||
namespace matmul_tiling {
|
||||
class MathUtil {
|
||||
public:
|
||||
static bool IsEqual(float leftValue, float rightValue);
|
||||
template<typename T>
|
||||
static auto CeilDivision(T num1, T num2) -> T
|
||||
{
|
||||
if (num2 == 0) {
|
||||
return 0;
|
||||
}
|
||||
return static_cast<T>((static_cast<int64_t>(num1) + static_cast<int64_t>(num2) - 1) /
|
||||
static_cast<int64_t>(num2));
|
||||
}
|
||||
template<typename T>
|
||||
static auto Align(T num1, T num2) -> T
|
||||
{
|
||||
return CeilDivision(num1, num2) * num2;
|
||||
}
|
||||
static int32_t AlignDown(int32_t num1, int32_t num2);
|
||||
static bool CheckMulOverflow(int32_t a, int32_t b, int32_t &c);
|
||||
static int32_t MapShape(int32_t shape, bool roundUpFlag = true);
|
||||
static void AddFactor(std::vector<int32_t> &dimsFactors, int32_t dim);
|
||||
static void GetFactorCnt(const int32_t shape, int32_t &factorCnt, const int32_t factorStart,
|
||||
const int32_t factorEnd);
|
||||
static void GetFactorLayerCnt(const int32_t shape, int32_t &factorCnt, const int32_t factorStart,
|
||||
const int32_t factorEnd);
|
||||
static bool CheckFactorNumSatisfy(const int32_t dim);
|
||||
static int32_t FindBestSingleCore(const int32_t oriShape, const int32_t mappedShape, const int32_t coreNum,
|
||||
bool isKDim);
|
||||
static void GetFactors(std::vector<int32_t> &factorList, int32_t srcNum, int32_t minFactor, int32_t maxFactor);
|
||||
static void GetFactors(std::vector<int32_t> &factorList, int32_t srcNum, int32_t maxFactor);
|
||||
static void GetBlockFactors(std::vector<int32_t> &factorList, const int32_t oriShape, const int32_t mpShape,
|
||||
const int32_t coreNum, const int32_t maxNum);
|
||||
static int32_t GetNonFactorMap(std::vector<int32_t> &factorList, int32_t srcNum, int32_t maxFactor);
|
||||
static std::vector<std::pair<int, int>> GetFactorPairs(int32_t num);
|
||||
static std::pair<int32_t, int32_t> DivideIntoMainAndTail(int32_t num, int32_t divisor);
|
||||
};
|
||||
} // namespace matmul_tiling
|
||||
#endif // _MATH_UTIL_H_
|
||||
176
csrc/moe/causal_conv1d/op_kernel/arch35/causal_conv1d_regbase.h
Normal file
176
csrc/moe/causal_conv1d/op_kernel/arch35/causal_conv1d_regbase.h
Normal file
@@ -0,0 +1,176 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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 causal_conv1d_regbase.h
|
||||
* \brief
|
||||
*/
|
||||
#ifndef CAUSAL_CONV1D_REGBASE_H
|
||||
#define CAUSAL_CONV1D_REGBASE_H
|
||||
|
||||
namespace NsCausalConv1d {
|
||||
using namespace AscendC;
|
||||
using namespace AscendC::MicroAPI;
|
||||
|
||||
constexpr uint16_t V_LENGTH = VECTOR_REG_WIDTH / sizeof(float);
|
||||
|
||||
constexpr CastTrait castTraitB16ToB32 = {
|
||||
RegLayout::ZERO, SatMode::UNKNOWN, MaskMergeMode::ZEROING, RoundMode::UNKNOWN};
|
||||
|
||||
template <typename T, bool hasActivation>
|
||||
__aicore__ inline void ComputeFnRollingOutputRegbase(LocalTensor<T> ring, LocalTensor<float> currF,
|
||||
LocalTensor<float> state0F, LocalTensor<float> weightF, uint32_t dataCount)
|
||||
{
|
||||
__ubuf__ T* ringAddr = (__ubuf__ T*)ring.GetPhyAddr();
|
||||
__ubuf__ float* currFAddr = (__ubuf__ float*)currF.GetPhyAddr();
|
||||
__ubuf__ float* state0FAddr = (__ubuf__ float*)state0F.GetPhyAddr();
|
||||
__ubuf__ float* weightFAddr = (__ubuf__ float*)weightF.GetPhyAddr();
|
||||
|
||||
uint16_t colLoopTimes = static_cast<uint16_t>(Ceil(dataCount, V_LENGTH));
|
||||
__VEC_SCOPE__
|
||||
{
|
||||
RegTensor<T> ring;
|
||||
RegTensor<float> currF;
|
||||
RegTensor<float> state0F;
|
||||
RegTensor<float> weightF;
|
||||
RegTensor<float> tmp;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(dataCount);
|
||||
DataCopy<T, LoadDist::DIST_UNPACK_B16>(ring, ringAddr + j * V_LENGTH);
|
||||
DataCopy(state0F, state0FAddr + j * V_LENGTH);
|
||||
DataCopy(weightF, weightFAddr + j * V_LENGTH);
|
||||
Cast<float, T, castTraitB16ToB32>(currF, ring, pregLoop);
|
||||
Mul(currF, currF, weightF, pregLoop);
|
||||
Add(state0F, state0F, currF, pregLoop);
|
||||
if constexpr (hasActivation) {
|
||||
Muls(tmp, state0F, -1.0f, pregLoop);
|
||||
Exp(tmp, tmp, pregLoop);
|
||||
Adds(tmp, tmp, 1.0f, pregLoop);
|
||||
Div(currF, state0F, tmp, pregLoop);
|
||||
DataCopy(currFAddr + j * V_LENGTH, currF, pregLoop);
|
||||
} else {
|
||||
DataCopy(state0FAddr + j * V_LENGTH, state0F, pregLoop);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static __simd_vf__ inline void AdvanceFnLocalPartialsWidthTwo(__ubuf__ T* ringAddr, __ubuf__ float* weight0FAddr,
|
||||
__ubuf__ float* state0FAddr, uint32_t dataCount, uint16_t colLoopTimes)
|
||||
{
|
||||
RegTensor<T> ring;
|
||||
RegTensor<float> currF;
|
||||
RegTensor<float> weight0F;
|
||||
RegTensor<float> state0F;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(dataCount);
|
||||
DataCopy<T, LoadDist::DIST_UNPACK_B16>(ring, ringAddr + j * V_LENGTH);
|
||||
DataCopy(weight0F, weight0FAddr + j * V_LENGTH);
|
||||
Cast<float, T, castTraitB16ToB32>(currF, ring, pregLoop);
|
||||
Mul(state0F, currF, weight0F, pregLoop);
|
||||
DataCopy(state0FAddr + j * V_LENGTH, state0F, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static __simd_vf__ inline void AdvanceFnLocalPartialsWidthThree(__ubuf__ T* ringAddr, __ubuf__ float* weight0FAddr,
|
||||
__ubuf__ float* weight1FAddr, __ubuf__ float* state0FAddr, __ubuf__ float* state1FAddr, uint32_t dataCount,
|
||||
uint16_t colLoopTimes)
|
||||
{
|
||||
RegTensor<T> ring;
|
||||
RegTensor<float> currF;
|
||||
RegTensor<float> weight0F;
|
||||
RegTensor<float> weight1F;
|
||||
RegTensor<float> state0F;
|
||||
RegTensor<float> state1F;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(dataCount);
|
||||
DataCopy<T, LoadDist::DIST_UNPACK_B16>(ring, ringAddr + j * V_LENGTH);
|
||||
DataCopy(state1F, state1FAddr + j * V_LENGTH);
|
||||
Cast<float, T, castTraitB16ToB32>(currF, ring, pregLoop);
|
||||
DataCopy(weight1F, weight1FAddr + j * V_LENGTH);
|
||||
Mul(state0F, currF, weight1F, pregLoop);
|
||||
DataCopy(weight0F, weight0FAddr + j * V_LENGTH);
|
||||
Add(state0F, state0F, state1F, pregLoop);
|
||||
Mul(state1F, currF, weight0F, pregLoop);
|
||||
DataCopy(state0FAddr + j * V_LENGTH, state0F, pregLoop);
|
||||
DataCopy(state1FAddr + j * V_LENGTH, state1F, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
static __simd_vf__ inline void AdvanceFnLocalPartialsWidthFour(__ubuf__ T* ringAddr, __ubuf__ float* weight0FAddr,
|
||||
__ubuf__ float* weight1FAddr, __ubuf__ float* weight2FAddr, __ubuf__ float* state0FAddr, __ubuf__ float* state1FAddr,
|
||||
__ubuf__ float* state2FAddr, uint32_t dataCount, uint16_t colLoopTimes)
|
||||
{
|
||||
RegTensor<T> ring;
|
||||
RegTensor<float> currF;
|
||||
RegTensor<float> weight0F;
|
||||
RegTensor<float> weight1F;
|
||||
RegTensor<float> weight2F;
|
||||
RegTensor<float> state0F;
|
||||
RegTensor<float> state1F;
|
||||
RegTensor<float> state2F;
|
||||
MaskReg pregLoop;
|
||||
for (uint16_t j = 0; j < colLoopTimes; j++) {
|
||||
pregLoop = UpdateMask<float>(dataCount);
|
||||
DataCopy<T, LoadDist::DIST_UNPACK_B16>(ring, ringAddr + j * V_LENGTH);
|
||||
DataCopy(state1F, state1FAddr + j * V_LENGTH);
|
||||
DataCopy(state2F, state2FAddr + j * V_LENGTH);
|
||||
Cast<float, T, castTraitB16ToB32>(currF, ring, pregLoop);
|
||||
DataCopy(weight2F, weight2FAddr + j * V_LENGTH);
|
||||
Mul(state0F, currF, weight2F, pregLoop);
|
||||
DataCopy(weight1F, weight1FAddr + j * V_LENGTH);
|
||||
Add(state0F, state0F, state1F, pregLoop);
|
||||
Mul(state1F, currF, weight1F, pregLoop);
|
||||
DataCopy(weight0F, weight0FAddr + j * V_LENGTH);
|
||||
Add(state1F, state1F, state2F, pregLoop);
|
||||
Mul(state2F, currF, weight0F, pregLoop);
|
||||
DataCopy(state0FAddr + j * V_LENGTH, state0F, pregLoop);
|
||||
DataCopy(state1FAddr + j * V_LENGTH, state1F, pregLoop);
|
||||
DataCopy(state2FAddr + j * V_LENGTH, state2F, pregLoop);
|
||||
}
|
||||
}
|
||||
|
||||
template <typename T, int32_t kTemplateWidth>
|
||||
__aicore__ inline void AdvanceFnLocalPartialsRegbase(LocalTensor<T> ring, LocalTensor<float> weightF,
|
||||
LocalTensor<float> state0F, LocalTensor<float> state1F, LocalTensor<float> state2F, uint32_t dataCount,
|
||||
uint32_t weightStep)
|
||||
{
|
||||
uint16_t colLoopTimes = static_cast<uint16_t>(Ceil(dataCount, V_LENGTH));
|
||||
|
||||
__ubuf__ T* ringAddr = (__ubuf__ T*)ring.GetPhyAddr();
|
||||
__ubuf__ float* weight0FAddr = (__ubuf__ float*)weightF.GetPhyAddr();
|
||||
__ubuf__ float* state0FAddr = (__ubuf__ float*)state0F.GetPhyAddr();
|
||||
if constexpr (kTemplateWidth == 2) {
|
||||
AscendC::VF_CALL<AdvanceFnLocalPartialsWidthTwo<T>>(ringAddr, weight0FAddr, state0FAddr, dataCount, colLoopTimes);
|
||||
} else if constexpr (kTemplateWidth == 3) {
|
||||
__ubuf__ float* weight1FAddr = weight0FAddr + weightStep;
|
||||
__ubuf__ float* state1FAddr = (__ubuf__ float*)state1F.GetPhyAddr();
|
||||
AscendC::VF_CALL<AdvanceFnLocalPartialsWidthThree<T>>(ringAddr, weight0FAddr, weight1FAddr, state0FAddr,
|
||||
state1FAddr, dataCount, colLoopTimes);
|
||||
} else if constexpr (kTemplateWidth == 4) {
|
||||
__ubuf__ float* weight1FAddr = weight0FAddr + weightStep;
|
||||
__ubuf__ float* weight2FAddr = weight1FAddr + weightStep;
|
||||
__ubuf__ float* state1FAddr = (__ubuf__ float*)state1F.GetPhyAddr();
|
||||
__ubuf__ float* state2FAddr = (__ubuf__ float*)state2F.GetPhyAddr();
|
||||
AscendC::VF_CALL<AdvanceFnLocalPartialsWidthFour<T>>(ringAddr, weight0FAddr, weight1FAddr, weight2FAddr,
|
||||
state0FAddr, state1FAddr, state2FAddr, dataCount, colLoopTimes);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace NsCausalConv1d
|
||||
|
||||
#endif // CAUSAL_CONV1D_REGBASE_H
|
||||
57
csrc/moe/causal_conv1d/op_kernel/causal_conv1d.cpp
Normal file
57
csrc/moe/causal_conv1d/op_kernel/causal_conv1d.cpp
Normal file
@@ -0,0 +1,57 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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 causal_conv1d.cpp
|
||||
* \brief
|
||||
*/
|
||||
|
||||
#include "causal_conv1d_fn.h"
|
||||
#include "causal_conv1d_update.h"
|
||||
|
||||
namespace {
|
||||
|
||||
template <typename T, uint32_t runModeKey, uint32_t widthKey, uint32_t fnPlanKey>
|
||||
__aicore__ inline void RunCausalConv1d(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates,
|
||||
GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode,
|
||||
GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace,
|
||||
const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
if constexpr (runModeKey == CAUSAL_CONV1D_TPL_RUN_MODE_FN) {
|
||||
NsCausalConv1d::RunCausalConv1dFn<T, widthKey, fnPlanKey>(
|
||||
x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, workspace,
|
||||
tilingData);
|
||||
} else {
|
||||
NsCausalConv1d::RunCausalConv1dUpdate<T>(
|
||||
x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, workspace,
|
||||
tilingData);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
template <uint32_t runModeKey, uint32_t widthKey, uint32_t fnPlanKey>
|
||||
__global__ __aicore__ void causal_conv1d(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates,
|
||||
GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode,
|
||||
GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace, GM_ADDR tiling)
|
||||
{
|
||||
REGISTER_TILING_DEFAULT(CausalConv1dTilingData);
|
||||
GET_TILING_DATA(tilingData, tiling);
|
||||
KERNEL_TASK_TYPE_DEFAULT(KERNEL_TYPE_MIX_AIV_1_0);
|
||||
GM_ADDR userWorkspace = workspace;
|
||||
if (workspace != nullptr) {
|
||||
userWorkspace = AscendC::GetUserWorkspace(workspace);
|
||||
}
|
||||
|
||||
RunCausalConv1d<DTYPE_X, runModeKey, widthKey, fnPlanKey>(
|
||||
x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y,
|
||||
userWorkspace, &tilingData);
|
||||
}
|
||||
1088
csrc/moe/causal_conv1d/op_kernel/causal_conv1d.h
Normal file
1088
csrc/moe/causal_conv1d/op_kernel/causal_conv1d.h
Normal file
File diff suppressed because it is too large
Load Diff
66
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_common.h
Normal file
66
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_common.h
Normal file
@@ -0,0 +1,66 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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 causal_conv1d_common.h
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_COMMON_H
|
||||
#define CAUSAL_CONV1D_COMMON_H
|
||||
|
||||
#include "kernel_operator.h"
|
||||
|
||||
namespace NsCausalConv1dCommon {
|
||||
|
||||
constexpr int32_t MAX_WIDTH = 4;
|
||||
constexpr int32_t MAX_BLOCK_DIM = 4096;
|
||||
constexpr int32_t RING_SLOTS = 5;
|
||||
|
||||
__aicore__ inline int32_t SlotCurr(int32_t t)
|
||||
{
|
||||
return (t + 3) % RING_SLOTS;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t SlotHist(int32_t t, int32_t i)
|
||||
{
|
||||
return (t + 3 - i) % RING_SLOTS;
|
||||
}
|
||||
|
||||
__aicore__ inline int32_t SlotPrefetch(int32_t t)
|
||||
{
|
||||
return (t + 4) % RING_SLOTS;
|
||||
}
|
||||
|
||||
struct CalcBufLayout {
|
||||
AscendC::LocalTensor<float> weightF;
|
||||
AscendC::LocalTensor<float> biasF;
|
||||
AscendC::LocalTensor<float> accF;
|
||||
AscendC::LocalTensor<float> tmpF;
|
||||
AscendC::LocalTensor<float> currF;
|
||||
|
||||
__aicore__ inline CalcBufLayout() = default;
|
||||
|
||||
__aicore__ static inline CalcBufLayout FromCalcBuf(AscendC::TBuf<AscendC::QuePosition::VECCALC> &calcBuf)
|
||||
{
|
||||
CalcBufLayout layout;
|
||||
AscendC::LocalTensor<float> calc = calcBuf.template Get<float>();
|
||||
layout.weightF = calc;
|
||||
layout.biasF = calc[MAX_WIDTH * MAX_BLOCK_DIM];
|
||||
layout.accF = layout.biasF[MAX_BLOCK_DIM];
|
||||
layout.tmpF = layout.accF[MAX_BLOCK_DIM];
|
||||
layout.currF = layout.tmpF[MAX_BLOCK_DIM];
|
||||
return layout;
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace NsCausalConv1dCommon
|
||||
|
||||
#endif // CAUSAL_CONV1D_COMMON_H
|
||||
93
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_fn.h
Normal file
93
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_fn.h
Normal file
@@ -0,0 +1,93 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_FN_H
|
||||
#define CAUSAL_CONV1D_FN_H
|
||||
|
||||
#include "causal_conv1d.h"
|
||||
|
||||
namespace NsCausalConv1d {
|
||||
|
||||
template <typename T, uint32_t widthKey, uint32_t fnPlanKey>
|
||||
class CausalConv1dFn : public CausalConv1d<T, CAUSAL_CONV1D_TPL_RUN_MODE_FN, widthKey, fnPlanKey> {
|
||||
public:
|
||||
__aicore__ inline void Init(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates, GM_ADDR queryStartLoc,
|
||||
GM_ADDR cacheIndices, GM_ADDR initialStateMode, GM_ADDR numAcceptedTokens, GM_ADDR y,
|
||||
GM_ADDR workspace, const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
(void)numAcceptedTokens;
|
||||
this->ResetRuntimeState(tilingData);
|
||||
this->xGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(x));
|
||||
this->weightGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(weight));
|
||||
this->biasGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(bias));
|
||||
this->convStatesGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(convStates));
|
||||
if (tilingData->hasQueryStartLoc != 0) {
|
||||
if (tilingData->queryStartLocUseInt64 != 0) {
|
||||
this->queryStartLocGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(queryStartLoc));
|
||||
} else {
|
||||
this->queryStartLocGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(queryStartLoc));
|
||||
}
|
||||
}
|
||||
if (tilingData->hasCacheIndices != 0) {
|
||||
if (tilingData->cacheIndicesUseInt64 != 0) {
|
||||
this->cacheIndicesGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(cacheIndices));
|
||||
} else {
|
||||
this->cacheIndicesGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(cacheIndices));
|
||||
}
|
||||
}
|
||||
if (tilingData->hasInitialStateMode != 0) {
|
||||
if (tilingData->initialStateModeDtype == 2) {
|
||||
this->initialStateModeGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(initialStateMode));
|
||||
} else if (tilingData->initialStateModeDtype == 1) {
|
||||
this->initialStateModeGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(initialStateMode));
|
||||
} else {
|
||||
this->initialStateModeGmBool.SetGlobalBuffer(reinterpret_cast<__gm__ bool *>(initialStateMode));
|
||||
}
|
||||
}
|
||||
this->yGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(y));
|
||||
if (tilingData->hasInitStateWorkspace != 0) {
|
||||
const uint64_t syncElems =
|
||||
static_cast<uint64_t>(GetBlockNum()) * INIT_STATE_SYNCALL_NEED_SIZE;
|
||||
const uint64_t syncBytes = syncElems * sizeof(int32_t);
|
||||
const uint64_t workspaceElems =
|
||||
static_cast<uint64_t>(tilingData->batch) *
|
||||
static_cast<uint64_t>(tilingData->width - 1) *
|
||||
static_cast<uint64_t>(tilingData->dim);
|
||||
this->initStateSyncGm_.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(workspace), syncElems);
|
||||
auto *workspaceBytes = reinterpret_cast<__gm__ uint8_t *>(workspace);
|
||||
this->initStateWorkspaceGm_.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(workspaceBytes + syncBytes),
|
||||
workspaceElems);
|
||||
}
|
||||
this->InitSharedBuffersAndEvents();
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
this->ProcessVarlenTokenTiled();
|
||||
this->ReleaseEvents();
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T, uint32_t widthKey, uint32_t fnPlanKey>
|
||||
__aicore__ inline void RunCausalConv1dFn(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates,
|
||||
GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode,
|
||||
GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace,
|
||||
const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
CausalConv1dFn<T, widthKey, fnPlanKey> op;
|
||||
op.Init(x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, workspace,
|
||||
tilingData);
|
||||
op.Process();
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
#endif
|
||||
306
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_fn_tasks.h
Normal file
306
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_fn_tasks.h
Normal file
@@ -0,0 +1,306 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_FN_TASKS_H
|
||||
#define CAUSAL_CONV1D_FN_TASKS_H
|
||||
|
||||
struct FnDirectBlockTask {
|
||||
bool valid = false;
|
||||
int32_t tokenTileId = 0;
|
||||
int32_t baseDimIdx = 0;
|
||||
int32_t tokenStart = 0;
|
||||
int32_t tokenEnd = 0;
|
||||
int32_t channelStart = 0;
|
||||
int32_t baseDimSize = 0;
|
||||
};
|
||||
|
||||
__aicore__ inline FnDirectBlockTask ResolveFnDirectBlockTask(int32_t blockIdx, int32_t tokenBlockCnt, int32_t tokenBlockSize,
|
||||
int32_t cuSeqlen, int32_t baseDimCnt, int32_t baseDim,
|
||||
int32_t dim)
|
||||
{
|
||||
FnDirectBlockTask task;
|
||||
if (blockIdx < 0 || tokenBlockCnt <= 0 || tokenBlockSize <= 0 || cuSeqlen <= 0 || baseDimCnt <= 0 || baseDim <= 0 ||
|
||||
dim <= 0) {
|
||||
return task;
|
||||
}
|
||||
|
||||
const int64_t phase1Grid = static_cast<int64_t>(tokenBlockCnt) * baseDimCnt;
|
||||
if (phase1Grid <= 0 || static_cast<int64_t>(blockIdx) >= phase1Grid) {
|
||||
return task;
|
||||
}
|
||||
|
||||
task.tokenTileId = blockIdx / baseDimCnt;
|
||||
task.baseDimIdx = blockIdx % baseDimCnt;
|
||||
task.channelStart = task.baseDimIdx * baseDim;
|
||||
if (task.channelStart >= dim) {
|
||||
return task;
|
||||
}
|
||||
|
||||
task.baseDimSize = (task.channelStart + baseDim <= dim) ? baseDim : (dim - task.channelStart);
|
||||
task.tokenStart = task.tokenTileId * tokenBlockSize;
|
||||
if (task.tokenStart >= cuSeqlen) {
|
||||
return task;
|
||||
}
|
||||
|
||||
const int32_t tokenEndRaw = task.tokenStart + tokenBlockSize;
|
||||
task.tokenEnd = (tokenEndRaw <= cuSeqlen) ? tokenEndRaw : cuSeqlen;
|
||||
if (task.baseDimSize <= 0 || task.tokenEnd <= task.tokenStart) {
|
||||
return {};
|
||||
}
|
||||
|
||||
task.valid = true;
|
||||
return task;
|
||||
}
|
||||
|
||||
__aicore__ inline bool IsFnInitStateSnapshotOwnerBlock(const FnDirectBlockTask &task)
|
||||
{
|
||||
return task.valid && task.tokenTileId == 0;
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline int32_t CAUSAL_CONV1D_CLASS::FindVarlenSeqByToken(int32_t tokenIdx) const
|
||||
{
|
||||
int32_t left = 0;
|
||||
int32_t right = static_cast<int32_t>(tilingData_->batch);
|
||||
while (left < right) {
|
||||
const int32_t mid = left + ((right - left) >> 1);
|
||||
const int32_t endVal = ReadQueryStartLocValue(mid + 1);
|
||||
if (tokenIdx < endVal) {
|
||||
right = mid;
|
||||
} else {
|
||||
left = mid + 1;
|
||||
}
|
||||
}
|
||||
return left;
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline bool CAUSAL_CONV1D_CLASS::ResolveExplicitTokenTileSeqRange(int32_t tokenTileId, int32_t &startSeq,
|
||||
int32_t &endSeq) const
|
||||
{
|
||||
if (!HasExplicitFnTokenSeqRanges() || tokenTileId < 0 || tokenTileId >= tilingData_->explicitTokenSeqRangeCount) {
|
||||
return false;
|
||||
}
|
||||
startSeq = static_cast<int32_t>(tilingData_->tokenTileStartSeq[tokenTileId]);
|
||||
endSeq = static_cast<int32_t>(tilingData_->tokenTileEndSeq[tokenTileId]);
|
||||
return (startSeq >= 0) && (endSeq >= startSeq);
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void CAUSAL_CONV1D_CLASS::InitRingSeqSplit(int32_t seq, int32_t cacheIdx, bool hasInit,
|
||||
int32_t seqStart, int32_t tileStart, int32_t tileLen,
|
||||
int32_t channelStart, int32_t baseDim, int32_t dim)
|
||||
{
|
||||
const int32_t stateLen = tilingData_->stateLen;
|
||||
const int32_t width = static_cast<int32_t>(tilingData_->width);
|
||||
const int32_t historyCount = width - 1;
|
||||
const int32_t ringStart = MAX_WIDTH - width;
|
||||
const int32_t historyStartTok = tileStart - historyCount;
|
||||
LocalTensor<T> ring = inBuf.Get<T>();
|
||||
bool hasGmHistoryCopy = false;
|
||||
bool hasVectorInit = false;
|
||||
const int64_t stateBaseOffset = static_cast<int64_t>(cacheIdx) * stateLen * dim + channelStart;
|
||||
int64_t xHistoryOffset = static_cast<int64_t>(historyStartTok) * dim + channelStart;
|
||||
|
||||
for (int32_t i = 0; i < ringStart; ++i) {
|
||||
Duplicate(ring[i * MAX_BLOCK_DIM], static_cast<T>(0), baseDim);
|
||||
hasVectorInit = true;
|
||||
}
|
||||
|
||||
for (int32_t i = 0, srcTok = historyStartTok; i < historyCount; ++i, ++srcTok, xHistoryOffset += dim) {
|
||||
LocalTensor<T> histSlot = ring[(ringStart + i) * MAX_BLOCK_DIM];
|
||||
if (srcTok >= seqStart) {
|
||||
DataCopy(histSlot, xGm[xHistoryOffset], baseDim);
|
||||
hasGmHistoryCopy = true;
|
||||
} else if (hasInit) {
|
||||
const int32_t statePos = srcTok - seqStart + historyCount;
|
||||
const int64_t stateOffset = stateBaseOffset + static_cast<int64_t>(statePos) * dim;
|
||||
if (tilingData_->hasInitStateWorkspace != 0) {
|
||||
const int64_t snapshotOffset =
|
||||
(static_cast<int64_t>(seq) * historyCount + statePos) * dim + channelStart;
|
||||
DataCopy(histSlot, initStateWorkspaceGm_[snapshotOffset], baseDim);
|
||||
} else {
|
||||
DataCopy(histSlot, convStatesGm[stateOffset], baseDim);
|
||||
}
|
||||
hasGmHistoryCopy = true;
|
||||
} else {
|
||||
Duplicate(histSlot, static_cast<T>(0), baseDim);
|
||||
hasVectorInit = true;
|
||||
}
|
||||
}
|
||||
|
||||
if (hasGmHistoryCopy) {
|
||||
SetFlag<HardEvent::MTE2_V>(stateMte2ToVEvent_);
|
||||
WaitFlag<HardEvent::MTE2_V>(stateMte2ToVEvent_);
|
||||
}
|
||||
if (hasVectorInit) {
|
||||
PipeBarrier<PIPE_V>();
|
||||
}
|
||||
|
||||
if (tileLen > 0) {
|
||||
const int32_t slot0 = SlotCurr(0);
|
||||
const int64_t xOffset = static_cast<int64_t>(tileStart) * dim + channelStart;
|
||||
DataCopy(ring[slot0 * MAX_BLOCK_DIM], xGm[xOffset], baseDim);
|
||||
SetFlag<HardEvent::MTE2_V>(inputMte2ToVEvent_[slot0]);
|
||||
}
|
||||
|
||||
if (tileLen > 1) {
|
||||
SetFlag<HardEvent::V_MTE2>(inputVToMte2Event_);
|
||||
}
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void CAUSAL_CONV1D_CLASS::ProcessFnChunk(int32_t seq, int32_t cacheIdx, bool hasInit,
|
||||
int32_t seqStart, int32_t seqLen, int32_t chunkStart,
|
||||
int32_t chunkLen, int32_t channelStart, int32_t baseDim,
|
||||
int32_t dim)
|
||||
{
|
||||
LoadWeightAndBias(channelStart, baseDim);
|
||||
InitRingSeqSplit(seq, cacheIdx, hasInit, seqStart, chunkStart, chunkLen, channelStart, baseDim, dim);
|
||||
|
||||
RunSeq(chunkStart, chunkLen, channelStart, baseDim, dim);
|
||||
|
||||
MaybeWriteBackSeqSplitTailChunk(chunkStart, chunkLen, seqStart, seqLen, cacheIdx, channelStart, baseDim, dim);
|
||||
DrainTaskMte3();
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void
|
||||
CAUSAL_CONV1D_CLASS::MaybeWriteBackSeqSplitTailChunk(int32_t chunkStart, int32_t chunkLen, int32_t seqStart,
|
||||
int32_t seqLen, int32_t cacheIdx, int32_t channelStart,
|
||||
int32_t baseDim, int32_t dim)
|
||||
{
|
||||
if (chunkStart + chunkLen != seqStart + seqLen) {
|
||||
return;
|
||||
}
|
||||
|
||||
DrainTaskMte3();
|
||||
WriteBackState(cacheIdx, chunkLen, channelStart, baseDim, dim);
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void CAUSAL_CONV1D_CLASS::PrefetchInitStatesToWorkspace(int32_t channelStart, int32_t baseDimSize)
|
||||
{
|
||||
if (tilingData_->hasInitStateWorkspace == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
const int32_t dim = tilingData_->dim;
|
||||
const int32_t historyCount = static_cast<int32_t>(tilingData_->width - 1);
|
||||
const int32_t batch = tilingData_->batch;
|
||||
const bool hasCacheIndices = (tilingData_->hasCacheIndices != 0);
|
||||
const bool hasInitialStateMode = (tilingData_->hasInitialStateMode != 0);
|
||||
LocalTensor<T> tmpBuf = inBuf.Get<T>()[0 * MAX_BLOCK_DIM];
|
||||
|
||||
for (int32_t seq = 0; seq < batch; ++seq) {
|
||||
if (!ResolveSeqHasInit(seq, hasInitialStateMode)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
int32_t cacheIdx = 0;
|
||||
if (!ResolveSeqCacheIndex(seq, hasCacheIndices, cacheIdx)) {
|
||||
continue;
|
||||
}
|
||||
|
||||
const int64_t stateBaseOffset = static_cast<int64_t>(cacheIdx) * tilingData_->stateLen * dim + channelStart;
|
||||
const int64_t snapshotBaseOffset = static_cast<int64_t>(seq) * historyCount * dim + channelStart;
|
||||
for (int32_t statePos = 0; statePos < historyCount; ++statePos) {
|
||||
const int64_t stateOffset = stateBaseOffset + static_cast<int64_t>(statePos) * dim;
|
||||
const int64_t snapshotOffset = snapshotBaseOffset + static_cast<int64_t>(statePos) * dim;
|
||||
DataCopy(tmpBuf, convStatesGm[stateOffset], baseDimSize);
|
||||
SetFlag<HardEvent::MTE2_MTE3>(initSnapshotMte2ToMte3Event_);
|
||||
WaitFlag<HardEvent::MTE2_MTE3>(initSnapshotMte2ToMte3Event_);
|
||||
DataCopy(initStateWorkspaceGm_[snapshotOffset], tmpBuf, baseDimSize);
|
||||
SetFlag<HardEvent::MTE3_MTE2>(initSnapshotMte3ToMte2Event_);
|
||||
WaitFlag<HardEvent::MTE3_MTE2>(initSnapshotMte3ToMte2Event_);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <CAUSAL_CONV1D_TEMPLATE_ARGS>
|
||||
__aicore__ inline void CAUSAL_CONV1D_CLASS::ProcessVarlenTokenTiled()
|
||||
{
|
||||
const int32_t dim = tilingData_->dim;
|
||||
const int32_t batch = tilingData_->batch;
|
||||
const int32_t seqLen = tilingData_->seqLen;
|
||||
const int32_t cuSeqlen = tilingData_->cuSeqlen;
|
||||
const int32_t baseDim = static_cast<int32_t>(tilingData_->baseDim);
|
||||
const int32_t baseDimCnt = static_cast<int32_t>(tilingData_->baseDimCnt);
|
||||
const int32_t tokenBlockSize = static_cast<int32_t>(tilingData_->tokenBlockSize);
|
||||
const int32_t tokenBlockCnt = static_cast<int32_t>(tilingData_->tokenBlockCnt);
|
||||
const bool hasCacheIndices = (tilingData_->hasCacheIndices != 0);
|
||||
const bool hasInitialStateMode = (tilingData_->hasInitialStateMode != 0);
|
||||
const bool isVarlenMode = (tilingData_->inputMode == 0);
|
||||
|
||||
const int32_t blockIdx = static_cast<int32_t>(GetBlockIdx());
|
||||
const auto blockTask = ResolveFnDirectBlockTask(blockIdx, tokenBlockCnt, tokenBlockSize, cuSeqlen, baseDimCnt,
|
||||
baseDim, dim);
|
||||
if (tilingData_->hasInitStateWorkspace != 0) {
|
||||
if (IsFnInitStateSnapshotOwnerBlock(blockTask)) {
|
||||
PrefetchInitStatesToWorkspace(blockTask.channelStart, blockTask.baseDimSize);
|
||||
}
|
||||
SyncAll();
|
||||
}
|
||||
if (!blockTask.valid) {
|
||||
return;
|
||||
}
|
||||
|
||||
int32_t seq = 0;
|
||||
int32_t seqUpperBound = batch;
|
||||
if (isVarlenMode) {
|
||||
if (!ResolveExplicitTokenTileSeqRange(blockTask.tokenTileId, seq, seqUpperBound)) {
|
||||
seq = FindVarlenSeqByToken(blockTask.tokenStart);
|
||||
}
|
||||
} else {
|
||||
seq = (seqLen > 0) ? (blockTask.tokenStart / seqLen) : 0;
|
||||
}
|
||||
|
||||
int32_t cursor = blockTask.tokenStart;
|
||||
while (cursor < blockTask.tokenEnd && seq < seqUpperBound) {
|
||||
int32_t seqStart = 0;
|
||||
int32_t curSeqLen = 0;
|
||||
if (!ResolveSeqTaskWindow(seq, tilingData_->inputMode, seqLen, seqStart, curSeqLen)) {
|
||||
++seq;
|
||||
continue;
|
||||
}
|
||||
const int32_t curSeqEnd = seqStart + curSeqLen;
|
||||
if (cursor < seqStart) {
|
||||
cursor = seqStart;
|
||||
}
|
||||
if (cursor >= curSeqEnd) {
|
||||
++seq;
|
||||
continue;
|
||||
}
|
||||
|
||||
const int32_t tileEnd = (blockTask.tokenEnd <= curSeqEnd) ? blockTask.tokenEnd : curSeqEnd;
|
||||
const int32_t tileLen = tileEnd - cursor;
|
||||
if (tileLen <= 0) {
|
||||
++seq;
|
||||
continue;
|
||||
}
|
||||
|
||||
int32_t cacheIdx = 0;
|
||||
if (!ResolveSeqCacheIndex(seq, hasCacheIndices, cacheIdx)) {
|
||||
cursor = tileEnd;
|
||||
++seq;
|
||||
continue;
|
||||
}
|
||||
|
||||
const bool hasInit = ResolveSeqHasInit(seq, hasInitialStateMode);
|
||||
ProcessFnChunk(seq, cacheIdx, hasInit, seqStart, curSeqLen, cursor, tileLen, blockTask.channelStart,
|
||||
blockTask.baseDimSize, dim);
|
||||
|
||||
cursor = tileEnd;
|
||||
++seq;
|
||||
}
|
||||
}
|
||||
|
||||
#endif
|
||||
68
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_tiling_data.h
Normal file
68
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_tiling_data.h
Normal file
@@ -0,0 +1,68 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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 causal_conv1d_tiling_data.h
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_TILING_DATA_H_
|
||||
#define CAUSAL_CONV1D_TILING_DATA_H_
|
||||
|
||||
#include <cstdint>
|
||||
|
||||
enum FnExecutionPlan : int64_t {
|
||||
FN_EXECUTION_PLAN_INVALID = 0,
|
||||
FN_EXECUTION_PLAN_CUTBS = 1,
|
||||
FN_EXECUTION_PLAN_CUTBSD = 2,
|
||||
};
|
||||
|
||||
inline constexpr int64_t ResolveFnExecutionPlan(int64_t baseDimCnt)
|
||||
{
|
||||
return (baseDimCnt <= 0) ? FN_EXECUTION_PLAN_INVALID
|
||||
: (baseDimCnt <= 1) ? FN_EXECUTION_PLAN_CUTBS
|
||||
: FN_EXECUTION_PLAN_CUTBSD;
|
||||
}
|
||||
|
||||
|
||||
struct CausalConv1dTilingData {
|
||||
int64_t dim;
|
||||
int64_t cuSeqlen;
|
||||
int64_t seqLen;
|
||||
int64_t inputMode;
|
||||
|
||||
int64_t width;
|
||||
|
||||
int64_t stateLen;
|
||||
int64_t numCacheLines;
|
||||
int64_t batch;
|
||||
int64_t activationMode;
|
||||
int64_t padSlotId;
|
||||
int64_t hasBias;
|
||||
int64_t baseDim;
|
||||
int64_t baseDimCnt;
|
||||
int64_t hasNumAcceptedTokens;
|
||||
int64_t hasQueryStartLoc;
|
||||
int64_t hasCacheIndices;
|
||||
int64_t hasInitialStateMode;
|
||||
int64_t queryStartLocUseInt64;
|
||||
int64_t cacheIndicesStride;
|
||||
int64_t cacheIndicesUseInt64;
|
||||
int64_t initialStateModeDtype;
|
||||
int64_t numAcceptedTokensUseInt64;
|
||||
int64_t tokenBlockSize;
|
||||
int64_t tokenBlockCnt;
|
||||
int64_t hasExplicitTokenSeqRanges;
|
||||
int64_t explicitTokenSeqRangeCount;
|
||||
int64_t tokenTileStartSeq[128];
|
||||
int64_t tokenTileEndSeq[128];
|
||||
int64_t hasInitStateWorkspace;
|
||||
};
|
||||
#endif // CAUSAL_CONV1D_TILING_DATA_H_
|
||||
66
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_tiling_key.h
Normal file
66
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_tiling_key.h
Normal file
@@ -0,0 +1,66 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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 causal_conv1d_tiling_key.h
|
||||
* \brief causal_conv1d tiling key declare
|
||||
*/
|
||||
|
||||
#ifndef __CAUSAL_CONV1D_TILING_KEY_H__
|
||||
#define __CAUSAL_CONV1D_TILING_KEY_H__
|
||||
|
||||
#include "causal_conv1d_tiling_data.h"
|
||||
#include "ascendc/host_api/tiling/template_argument.h"
|
||||
|
||||
#define CAUSAL_CONV1D_TPL_RUN_MODE_FN 0
|
||||
#define CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE 1
|
||||
#define CAUSAL_CONV1D_TPL_WIDTH_RUNTIME 0
|
||||
#define CAUSAL_CONV1D_TPL_WIDTH_2 1
|
||||
#define CAUSAL_CONV1D_TPL_WIDTH_3 2
|
||||
#define CAUSAL_CONV1D_TPL_WIDTH_4 3
|
||||
#define CAUSAL_CONV1D_TPL_FN_PLAN_INVALID 0
|
||||
#define CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS 1
|
||||
#define CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD 2
|
||||
ASCENDC_TPL_ARGS_DECL(CausalConv1d,
|
||||
ASCENDC_TPL_UINT_DECL(runModeKey, 1, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_RUN_MODE_FN,
|
||||
CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE),
|
||||
ASCENDC_TPL_UINT_DECL(widthKey, 2, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_WIDTH_RUNTIME,
|
||||
CAUSAL_CONV1D_TPL_WIDTH_2, CAUSAL_CONV1D_TPL_WIDTH_3,
|
||||
CAUSAL_CONV1D_TPL_WIDTH_4),
|
||||
ASCENDC_TPL_UINT_DECL(fnPlanKey, 2, ASCENDC_TPL_UI_LIST, CAUSAL_CONV1D_TPL_FN_PLAN_INVALID,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS, CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD));
|
||||
|
||||
#define CAUSAL_CONV1D_TPL_SEL_ENTRY(RUN_MODE, WIDTH, FN_PLAN) \
|
||||
ASCENDC_TPL_ARGS_SEL(ASCENDC_TPL_UINT_SEL(runModeKey, ASCENDC_TPL_UI_LIST, RUN_MODE), \
|
||||
ASCENDC_TPL_UINT_SEL(widthKey, ASCENDC_TPL_UI_LIST, WIDTH), \
|
||||
ASCENDC_TPL_UINT_SEL(fnPlanKey, ASCENDC_TPL_UI_LIST, FN_PLAN), \
|
||||
ASCENDC_TPL_TILING_STRUCT_SEL(CausalConv1dTilingData))
|
||||
|
||||
// Keep entries in encoded tiling-key order: real-device sub-kernel dispatch is sensitive to declaration order.
|
||||
ASCENDC_TPL_SEL(
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE, CAUSAL_CONV1D_TPL_WIDTH_RUNTIME,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_INVALID),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_2,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_3,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_4,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBS),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_2,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_3,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD),
|
||||
CAUSAL_CONV1D_TPL_SEL_ENTRY(CAUSAL_CONV1D_TPL_RUN_MODE_FN, CAUSAL_CONV1D_TPL_WIDTH_4,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_CUTBSD));
|
||||
|
||||
#undef CAUSAL_CONV1D_TPL_SEL_ENTRY
|
||||
|
||||
#endif // __CAUSAL_CONV1D_TILING_KEY_H__
|
||||
91
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_update.h
Normal file
91
csrc/moe/causal_conv1d/op_kernel/causal_conv1d_update.h
Normal file
@@ -0,0 +1,91 @@
|
||||
/**
|
||||
* This program is free software, you can redistribute it and/or modify it.
|
||||
* 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.
|
||||
*/
|
||||
|
||||
#ifndef CAUSAL_CONV1D_UPDATE_H
|
||||
#define CAUSAL_CONV1D_UPDATE_H
|
||||
|
||||
#include "causal_conv1d.h"
|
||||
|
||||
namespace NsCausalConv1d {
|
||||
|
||||
template <typename T>
|
||||
class CausalConv1dUpdate
|
||||
: public CausalConv1d<T, CAUSAL_CONV1D_TPL_RUN_MODE_UPDATE, CAUSAL_CONV1D_TPL_WIDTH_RUNTIME,
|
||||
CAUSAL_CONV1D_TPL_FN_PLAN_INVALID> {
|
||||
public:
|
||||
__aicore__ inline void Init(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates, GM_ADDR queryStartLoc,
|
||||
GM_ADDR cacheIndices, GM_ADDR, GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace,
|
||||
const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
(void)workspace;
|
||||
this->ResetRuntimeState(tilingData);
|
||||
this->xGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(x));
|
||||
this->weightGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(weight));
|
||||
this->biasGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(bias));
|
||||
this->convStatesGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(convStates));
|
||||
if (tilingData->hasQueryStartLoc != 0) {
|
||||
if (tilingData->queryStartLocUseInt64 != 0) {
|
||||
this->queryStartLocGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(queryStartLoc));
|
||||
} else {
|
||||
this->queryStartLocGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(queryStartLoc));
|
||||
}
|
||||
}
|
||||
if (tilingData->hasCacheIndices != 0) {
|
||||
if (tilingData->cacheIndicesUseInt64 != 0) {
|
||||
this->cacheIndicesGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(cacheIndices));
|
||||
} else {
|
||||
this->cacheIndicesGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(cacheIndices));
|
||||
}
|
||||
}
|
||||
if (tilingData->hasNumAcceptedTokens != 0) {
|
||||
if (tilingData->numAcceptedTokensUseInt64 != 0) {
|
||||
this->numAcceptedTokensGmInt64.SetGlobalBuffer(reinterpret_cast<__gm__ int64_t *>(numAcceptedTokens));
|
||||
} else {
|
||||
this->numAcceptedTokensGmInt32.SetGlobalBuffer(reinterpret_cast<__gm__ int32_t *>(numAcceptedTokens));
|
||||
}
|
||||
}
|
||||
this->yGm.SetGlobalBuffer(reinterpret_cast<__gm__ T *>(y));
|
||||
this->InitSharedBuffersAndEvents();
|
||||
}
|
||||
|
||||
__aicore__ inline void Process()
|
||||
{
|
||||
const CausalConv1dTilingData *tilingData = this->GetTilingData();
|
||||
const int32_t dim = tilingData->dim;
|
||||
const int32_t baseDimCnt = static_cast<int32_t>(tilingData->baseDimCnt);
|
||||
const int32_t width = static_cast<int32_t>(tilingData->width);
|
||||
const int32_t baseDim = static_cast<int32_t>(tilingData->baseDim);
|
||||
if (baseDim <= 0 || baseDimCnt <= 0 || baseDim > MAX_BLOCK_DIM || width < 2 || width > MAX_WIDTH || dim <= 0 ||
|
||||
tilingData->batch <= 0) {
|
||||
this->ReleaseEvents();
|
||||
return;
|
||||
}
|
||||
|
||||
this->ProcessDefault();
|
||||
this->ReleaseEvents();
|
||||
}
|
||||
};
|
||||
|
||||
template <typename T>
|
||||
__aicore__ inline void RunCausalConv1dUpdate(GM_ADDR x, GM_ADDR weight, GM_ADDR bias, GM_ADDR convStates,
|
||||
GM_ADDR queryStartLoc, GM_ADDR cacheIndices, GM_ADDR initialStateMode,
|
||||
GM_ADDR numAcceptedTokens, GM_ADDR y, GM_ADDR workspace,
|
||||
const CausalConv1dTilingData *tilingData)
|
||||
{
|
||||
CausalConv1dUpdate<T> op;
|
||||
op.Init(x, weight, bias, convStates, queryStartLoc, cacheIndices, initialStateMode, numAcceptedTokens, y, workspace,
|
||||
tilingData);
|
||||
op.Process();
|
||||
}
|
||||
|
||||
} // namespace NsCausalConv1d
|
||||
|
||||
#endif // CAUSAL_CONV1D_UPDATE_H
|
||||
Reference in New Issue
Block a user