已合并
fast_kernel_launch <<<>>>内核调用符调用方式 适配最新卷积代码 #5671
想做鹊山祝余的舅舅创建于 6月4日
fast_kernel_launch <<<>>>内核调用符调用方式 适配最新卷积代码 #5671
已合并
想做鹊山祝余的舅舅创建于 6月4日
6 个文件变更+55-51
@@ -211,8 +211,9 @@ bool ConvBaseDeci::CheckInstrLimitsHWmode()
211 ss << "If the format of input x is NDHWC, ";211 ss << "If the format of input x is NDHWC, ";
212 ss << "the constraint of instruction %s must be met: ";212 ss << "the constraint of instruction %s must be met: ";
213 ss << "shape[%zu] * shape [%zu] ≤ %ld";213 ss << "shape[%zu] * shape [%zu] ≤ %ld";
214- vector<int64_t> outputShape = {shapeInfo_.batch,214+ vector<int64_t> outputShape = {static_cast<int64_t>(shapeInfo_.batch),
215- shapeInfo_.dout, shapeInfo_.ho, shapeInfo_.wo, shapeInfo_.co};215+ static_cast<int64_t>(shapeInfo_.dout), static_cast<int64_t>(shapeInfo_.ho),
216+ static_cast<int64_t>(shapeInfo_.wo), static_cast<int64_t>(shapeInfo_.co)};
216 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(nodeInfo_.nodeType.c_str(), "y",217 OP_LOGE_FOR_INVALID_SHAPE_WITH_REASON(nodeInfo_.nodeType.c_str(), "y",
217 VectorToString(outputShape, IntToString<int64_t>).c_str(),218 VectorToString(outputShape, IntToString<int64_t>).c_str(),
218 FormatString(ss.str().c_str(), "Fixpipe", NDHWC_W_IDX, NDHWC_C_IDX).c_str());219 FormatString(ss.str().c_str(), "Fixpipe", NDHWC_W_IDX, NDHWC_C_IDX).c_str());
@@ -48,7 +48,7 @@ __all__ = ["conv3d_custom",]
48def conv3d_custom(input: Tensor, weight: Tensor, strides: list, pads: list, dilations: list,48def conv3d_custom(input: Tensor, weight: Tensor, strides: list, pads: list, dilations: list,
49 bias: Tensor = None, enable_hf32: bool = False) -> Tensor:49 bias: Tensor = None, enable_hf32: bool = False) -> Tensor:
50 # 判断当前NPU设备类型50 # 判断当前NPU设备类型
51- if torch_npu.npu.get_device_name() == "Ascend950PR_9589":51+ if torch_npu.npu.get_device_name().startswith("Ascend950"):
52 # Ascend950设备路径:使用conv3d_v2_custom算子52 # Ascend950设备路径:使用conv3d_v2_custom算子
53 53
54 # 检查算子是否已注册54 # 检查算子是否已注册
@@ -51,6 +51,7 @@ set(CMAKE_LINKER ${BISHENG})
51# set ASCEND_INCLUDE_DIRS51# set ASCEND_INCLUDE_DIRS
52set(ASCEND_INCLUDE_DIRS52set(ASCEND_INCLUDE_DIRS
53 ${ASCEND_DIR}/include53 ${ASCEND_DIR}/include
54+ ${ASCEND_DIR}/include/op_common
54 ${ASCEND_DIR}/compiler/tikcpp/include55 ${ASCEND_DIR}/compiler/tikcpp/include
55 ${ASCEND_DIR}/compiler/ascendc/include/basic_api/impl56 ${ASCEND_DIR}/compiler/ascendc/include/basic_api/impl
56 ${ASCEND_DIR}/compiler/ascendc/include/basic_api/interface57 ${ASCEND_DIR}/compiler/ascendc/include/basic_api/interface
@@ -62,6 +63,7 @@ set(ASCEND_INCLUDE_DIRS
62 ${ASCEND_DIR}/pkg_inc63 ${ASCEND_DIR}/pkg_inc
63 ${ASCEND_DIR}/opp/built-in/op_impl/ai_core/tbe/impl/ops_nn/ascendc/conv3d_v264 ${ASCEND_DIR}/opp/built-in/op_impl/ai_core/tbe/impl/ops_nn/ascendc/conv3d_v2
64    ${ASCEND_DIR}/opp/built-in/op_impl/ai_core/tbe/impl/ops_nn/ascendc/conv3d_v2/arch3565    ${ASCEND_DIR}/opp/built-in/op_impl/ai_core/tbe/impl/ops_nn/ascendc/conv3d_v2/arch35
66+ ${ASCEND_DIR}/x86_64-linux/include/version
65    ${CMAKE_CURRENT_SOURCE_DIR}/../../conv/common/op_host/op_tiling/arch3567    ${CMAKE_CURRENT_SOURCE_DIR}/../../conv/common/op_host/op_tiling/arch35
66    ${CMAKE_CURRENT_SOURCE_DIR}/../../conv/conv3d_v2/op_host/op_tiling/arch3568    ${CMAKE_CURRENT_SOURCE_DIR}/../../conv/conv3d_v2/op_host/op_tiling/arch35
67    ${CMAKE_CURRENT_SOURCE_DIR}/../../conv69    ${CMAKE_CURRENT_SOURCE_DIR}/../../conv
@@ -13,13 +13,14 @@
13 * \brief13 * \brief
14 */14 */
15 15 
16+#include "asc_devkit_version.h"
16#include "conv3d_v2_kernel.h"17#include "conv3d_v2_kernel.h"
17#include "acl/acl.h"18#include "acl/acl.h"
18#include "conv3d_v2_template.h"19#include "conv3d_v2_template.h"
19 20 
20void Conv3dv2Template(21void Conv3dv2Template(
21 GM_ADDR x, GM_ADDR filter, GM_ADDR bias, GM_ADDR scale, GM_ADDR offset,22 GM_ADDR x, GM_ADDR filter, GM_ADDR bias, GM_ADDR scale, GM_ADDR offset,
22- GM_ADDR offset_w, GM_ADDR y, GM_ADDR workspace, Ops::NN::Conv3dV2::Conv3DV2TilingData& tiling,23+ GM_ADDR offset_w, GM_ADDR y, GM_ADDR workspace, Ops::NN::Conv3dV2::Conv3DV2TilingDataV2& tiling,
23 int8_t FmapTiling, int8_t WeightTiling, int8_t L1PingPong, int8_t L0PingPong,24 int8_t FmapTiling, int8_t WeightTiling, int8_t L1PingPong, int8_t L0PingPong,
24 int8_t OutputOrder, int8_t IterOrder,25 int8_t OutputOrder, int8_t IterOrder,
25 const std::string& dtype, int32_t numBlocks, aclrtStream stream)26 const std::string& dtype, int32_t numBlocks, aclrtStream stream)
@@ -104,7 +104,7 @@
104 */104 */
105void Conv3dv2Template(105void Conv3dv2Template(
106 GM_ADDR x, GM_ADDR filter, GM_ADDR bias, GM_ADDR scale, GM_ADDR offset,106 GM_ADDR x, GM_ADDR filter, GM_ADDR bias, GM_ADDR scale, GM_ADDR offset,
107- GM_ADDR offset_w, GM_ADDR y, GM_ADDR workspace, Ops::NN::Conv3dV2::Conv3DV2TilingData& tiling,107+ GM_ADDR offset_w, GM_ADDR y, GM_ADDR workspace, Ops::NN::Conv3dV2::Conv3DV2TilingDataV2& tiling,
108 int8_t FmapTiling, int8_t WeightTiling, int8_t L1PingPong, int8_t L0PingPong,108 int8_t FmapTiling, int8_t WeightTiling, int8_t L1PingPong, int8_t L0PingPong,
109 int8_t OutputOrder, int8_t IterOrder,109 int8_t OutputOrder, int8_t IterOrder,
110 const std::string& dtype,110 const std::string& dtype,
@@ -82,69 +82,69 @@ namespace Conv3dCustom {
82 * @param conv3dRunInfo 输出参数,3D卷积运行信息结构体82 * @param conv3dRunInfo 输出参数,3D卷积运行信息结构体
83 * @param tilingInfo 输入参数,tiling信息,包含形状、属性、块维度等信息83 * @param tilingInfo 输入参数,tiling信息,包含形状、属性、块维度等信息
84 */84 */
85-static void InitConv3dRunInfo(Ops::NN::Conv3dV2::Conv3DRunInfo& conv3dRunInfo,85+static void InitConv3dRunInfo(Ops::NN::Conv3dV2::Conv3DV2TilingDataV2& tilingData,
86 optiling::conv_ops_tiling::ConvAscendcTilingInfo& tilingInfo)86 optiling::conv_ops_tiling::ConvAscendcTilingInfo& tilingInfo)
87{87{
88 // 输入输出形状信息88 // 输入输出形状信息
89- conv3dRunInfo.batch = static_cast<uint32_t>(tilingInfo.shapeInfo.batch); // 批次大小89+ tilingData.batch = static_cast<uint32_t>(tilingInfo.shapeInfo.batch); // 批次大小
90- conv3dRunInfo.cin = static_cast<uint32_t>(tilingInfo.shapeInfo.ci); // 输入通道数90+ tilingData.cin = static_cast<uint32_t>(tilingInfo.shapeInfo.ci); // 输入通道数
91- conv3dRunInfo.din = static_cast<uint32_t>(tilingInfo.shapeInfo.di); // 输入深度91+ tilingData.din = static_cast<uint32_t>(tilingInfo.shapeInfo.di); // 输入深度
92- conv3dRunInfo.hin = static_cast<uint32_t>(tilingInfo.shapeInfo.hi); // 输入高度92+ tilingData.hin = static_cast<uint32_t>(tilingInfo.shapeInfo.hi); // 输入高度
93- conv3dRunInfo.win = static_cast<uint32_t>(tilingInfo.shapeInfo.wi); // 输入宽度93+ tilingData.win = static_cast<uint32_t>(tilingInfo.shapeInfo.wi); // 输入宽度
94- conv3dRunInfo.cout = static_cast<uint32_t>(tilingInfo.shapeInfo.co); // 输出通道数94+ tilingData.cout = static_cast<uint32_t>(tilingInfo.shapeInfo.co); // 输出通道数
95 95
96 // 卷积核大小96 // 卷积核大小
97- conv3dRunInfo.kd = static_cast<uint32_t>(tilingInfo.shapeInfo.kd); // 卷积核深度97+ tilingData.kd = static_cast<uint32_t>(tilingInfo.shapeInfo.kd); // 卷积核深度
98- conv3dRunInfo.kh = static_cast<uint32_t>(tilingInfo.shapeInfo.kh); // 卷积核高度98+ tilingData.kh = static_cast<uint32_t>(tilingInfo.shapeInfo.kh); // 卷积核高度
99- conv3dRunInfo.kw = static_cast<uint32_t>(tilingInfo.shapeInfo.kw); // 卷积核宽度99+ tilingData.kw = static_cast<uint32_t>(tilingInfo.shapeInfo.kw); // 卷积核宽度
100 100
101 // 输出形状信息101 // 输出形状信息
102- conv3dRunInfo.dout = static_cast<uint32_t>(tilingInfo.shapeInfo.dout); // 输出深度102+ tilingData.dout = static_cast<uint32_t>(tilingInfo.shapeInfo.dout); // 输出深度
103- conv3dRunInfo.hout = static_cast<uint32_t>(tilingInfo.shapeInfo.ho); // 输出高度103+ tilingData.hout = static_cast<uint32_t>(tilingInfo.shapeInfo.ho); // 输出高度
104- conv3dRunInfo.wout = static_cast<uint32_t>(tilingInfo.shapeInfo.wo); // 输出宽度104+ tilingData.wout = static_cast<uint32_t>(tilingInfo.shapeInfo.wo); // 输出宽度
105 105
106 // 块维度信息(用于并行计算)106 // 块维度信息(用于并行计算)
107- conv3dRunInfo.batchDim = tilingInfo.numBlocksRes.batchDim; // 批次维度分块数107+ tilingData.batchDim = tilingInfo.numBlocksRes.batchDim; // 批次维度分块数
108- conv3dRunInfo.doDim = tilingInfo.numBlocksRes.doDim; // 输出深度维度分块数108+ tilingData.doDim = tilingInfo.numBlocksRes.doDim; // 输出深度维度分块数
109- conv3dRunInfo.mDim = tilingInfo.numBlocksRes.mDim; // M维度(输出通道)分块数109+ tilingData.mDim = tilingInfo.numBlocksRes.mDim; // M维度(输出通道)分块数
110- conv3dRunInfo.wDim = tilingInfo.numBlocksRes.woDim; // 输出宽度维度分块数110+ tilingData.wDim = tilingInfo.numBlocksRes.woDim; // 输出宽度维度分块数
111- conv3dRunInfo.nDim = tilingInfo.numBlocksRes.nDim; // N维度(批次)分块数111+ tilingData.nDim = tilingInfo.numBlocksRes.nDim; // N维度(批次)分块数
112- conv3dRunInfo.groupDim = tilingInfo.numBlocksRes.groupDim; // 组维度分块数112+ tilingData.groupDim = tilingInfo.numBlocksRes.groupDim; // 组维度分块数
113- conv3dRunInfo.hoDim = tilingInfo.numBlocksRes.hoDim; // 输出高度维度分块数113+ tilingData.hoDim = tilingInfo.numBlocksRes.hoDim; // 输出高度维度分块数
114 114
115 // 卷积参数:步长115 // 卷积参数:步长
116- conv3dRunInfo.strideD = static_cast<uint32_t>(tilingInfo.attrInfo.strideD); // 深度方向步长116+ tilingData.strideD = static_cast<uint32_t>(tilingInfo.attrInfo.strideD); // 深度方向步长
117- conv3dRunInfo.strideH = static_cast<uint32_t>(tilingInfo.attrInfo.strideH); // 高度方向步长117+ tilingData.strideH = static_cast<uint32_t>(tilingInfo.attrInfo.strideH); // 高度方向步长
118- conv3dRunInfo.strideW = static_cast<uint32_t>(tilingInfo.attrInfo.strideW); // 宽度方向步长118+ tilingData.strideW = static_cast<uint32_t>(tilingInfo.attrInfo.strideW); // 宽度方向步长
119 119
120 // 卷积参数:膨胀(dilation)120 // 卷积参数:膨胀(dilation)
121- conv3dRunInfo.dilationD = static_cast<uint32_t>(tilingInfo.attrInfo.dilationD); // 深度方向膨胀121+ tilingData.dilationD = static_cast<uint32_t>(tilingInfo.attrInfo.dilationD); // 深度方向膨胀
122- conv3dRunInfo.dilationH = static_cast<uint32_t>(tilingInfo.attrInfo.dilationH); // 高度方向膨胀122+ tilingData.dilationH = static_cast<uint32_t>(tilingInfo.attrInfo.dilationH); // 高度方向膨胀
123- conv3dRunInfo.dilationW = static_cast<uint32_t>(tilingInfo.attrInfo.dilationW); // 宽度方向膨胀123+ tilingData.dilationW = static_cast<uint32_t>(tilingInfo.attrInfo.dilationW); // 宽度方向膨胀
124 124
125 // 卷积参数:padding125 // 卷积参数:padding
126- conv3dRunInfo.padHead = static_cast<uint32_t>(tilingInfo.attrInfo.padHead); // 头部padding126+ tilingData.padHead = static_cast<uint32_t>(tilingInfo.attrInfo.padHead); // 头部padding
127- conv3dRunInfo.padTail = static_cast<uint32_t>(tilingInfo.attrInfo.padTail); // 尾部padding127+ tilingData.padTail = static_cast<uint32_t>(tilingInfo.attrInfo.padTail); // 尾部padding
128- conv3dRunInfo.padTop = static_cast<uint32_t>(tilingInfo.attrInfo.padTop); // 顶部padding128+ tilingData.padTop = static_cast<uint32_t>(tilingInfo.attrInfo.padTop); // 顶部padding
129- conv3dRunInfo.padBottom = static_cast<uint32_t>(tilingInfo.attrInfo.padBottom); // 底部padding129+ tilingData.padBottom = static_cast<uint32_t>(tilingInfo.attrInfo.padBottom); // 底部padding
130- conv3dRunInfo.padLeft = static_cast<uint32_t>(tilingInfo.attrInfo.padLeft); // 左侧padding130+ tilingData.padLeft = static_cast<uint32_t>(tilingInfo.attrInfo.padLeft); // 左侧padding
131- conv3dRunInfo.padRight = static_cast<uint32_t>(tilingInfo.attrInfo.padRight); // 右侧padding131+ tilingData.padRight = static_cast<uint32_t>(tilingInfo.attrInfo.padRight); // 右侧padding
132 132
133 // 其他参数133 // 其他参数
134- conv3dRunInfo.groups = 1; // 分组数(当前为1,表示标准卷积)134+ tilingData.groups = 1; // 分组数(当前为1,表示标准卷积)
135- conv3dRunInfo.enlarge = 1; // 扩展标志135+ tilingData.enlarge = 1; // 扩展标志
136- conv3dRunInfo.cinOpt = static_cast<uint32_t>(tilingInfo.shapeInfo.ci); // 优化后的输入通道数136+ tilingData.cinOpt = static_cast<uint32_t>(tilingInfo.shapeInfo.ci); // 优化后的输入通道数
137- conv3dRunInfo.coutOpt = static_cast<uint32_t>(tilingInfo.shapeInfo.co); // 优化后的输出通道数137+ tilingData.coutOpt = static_cast<uint32_t>(tilingInfo.shapeInfo.co); // 优化后的输出通道数
138- conv3dRunInfo.groupOpt = 1; // 优化后的分组数138+ tilingData.groupOpt = 1; // 优化后的分组数
139- conv3dRunInfo.hasBias = static_cast<uint8_t>(tilingInfo.flagInfo.hasBias); // 是否有偏置项139+ tilingData.hasBias = static_cast<uint8_t>(tilingInfo.flagInfo.hasBias); // 是否有偏置项
140 140
141 // 根据分割模式确定hoDim141 // 根据分割模式确定hoDim
142 if (tilingInfo.flagInfo.mSplitModeFlag) {142 if (tilingInfo.flagInfo.mSplitModeFlag) {
143 // M分割模式:使用mDim作为hoDim143 // M分割模式:使用mDim作为hoDim
144- conv3dRunInfo.hoDim = static_cast<uint32_t>(tilingInfo.numBlocksRes.mDim);144+ tilingData.hoDim = static_cast<uint32_t>(tilingInfo.numBlocksRes.mDim);
145 } else {145 } else {
146 // 标准模式:使用hoDim作为hoDim146 // 标准模式:使用hoDim作为hoDim
147- conv3dRunInfo.hoDim = static_cast<uint32_t>(tilingInfo.numBlocksRes.hoDim);147+ tilingData.hoDim = static_cast<uint32_t>(tilingInfo.numBlocksRes.hoDim);
148 }148 }
149}149}
150 150 
@@ -157,11 +157,11 @@ static void InitConv3dRunInfo(Ops::NN::Conv3dV2::Conv3DRunInfo& conv3dRunInfo,
157 * @param tilingData 输出参数,tiling数据结构体157 * @param tilingData 输出参数,tiling数据结构体
158 * @param tilingInfo 输入参数,tiling信息158 * @param tilingInfo 输入参数,tiling信息
159 */159 */
160-static void InitTilingData(Ops::NN::Conv3dV2::Conv3DV2TilingData& tilingData,160+static void InitTilingData(Ops::NN::Conv3dV2::Conv3DV2TilingDataV2& tilingData,
161 optiling::conv_ops_tiling::ConvAscendcTilingInfo& tilingInfo)161 optiling::conv_ops_tiling::ConvAscendcTilingInfo& tilingInfo)
162{162{
163 // 将tiling信息转换为Conv3DRunInfo结构体163 // 将tiling信息转换为Conv3DRunInfo结构体
164- InitConv3dRunInfo(tilingData.conv3dRunInfo, tilingInfo);164+ InitConv3dRunInfo(tilingData, tilingInfo);
165}165}
166 166 
167 167 
@@ -295,7 +295,7 @@ static int32_t InitPlatformInfo(optiling::conv_ops_tiling::ConvAscendcPlatformIn
295 }295 }
296 296 
297 // 获取AI Core数量297 // 获取AI Core数量
298- platformInfo.aicariNum = ascendcPlatform->GetCoreNumAic();298+ platformInfo.aicoreNum = ascendcPlatform->GetCoreNumAic();
299 299
300 // 获取各级缓存的内存大小300 // 获取各级缓存的内存大小
301 uint64_t size {};301 uint64_t size {};
@@ -439,7 +439,7 @@ void Conv3dV2CustomApi(
439 "Failed to get block dimension information from convolution base decision");439 "Failed to get block dimension information from convolution base decision");
440 440 
441 // 初始化tiling数据结构441 // 初始化tiling数据结构
442- Ops::NN::Conv3dV2::Conv3DV2TilingData tilingData;442+ Ops::NN::Conv3dV2::Conv3DV2TilingDataV2 tilingData;
443 InitTilingData(tilingData, tilingInfo);443 InitTilingData(tilingData, tilingInfo);
444 444 
445 // 设置平台信息并获取tiling数据445 // 设置平台信息并获取tiling数据
@@ -456,8 +456,8 @@ void Conv3dV2CustomApi(
456 tilingInfo.convOpsConstParams, tilingInfo.numBlocksRes, tilingData);456 tilingInfo.convOpsConstParams, tilingInfo.numBlocksRes, tilingData);
457 457 
458 // 计算需要的AI Core数量458 // 计算需要的AI Core数量
459- uint32_t g_numBlocks = tilingData.conv3dRunInfo.batchDim * tilingData.conv3dRunInfo.doDim *459+ uint32_t g_numBlocks = tilingData.batchDim * tilingData.doDim *
460- tilingData.conv3dRunInfo.hoDim * tilingData.conv3dRunInfo.nDim;460+ tilingData.hoDim * tilingData.nDim;
461 461
462 // 获取tiling键(用于选择最优的kernel实现)462 // 获取tiling键(用于选择最优的kernel实现)
463 optiling::conv_ops_tiling::ConvTilingKeyPara tilingKeyPara {};463 optiling::conv_ops_tiling::ConvTilingKeyPara tilingKeyPara {};