已合并
卷积反向添加输入属性打印 #2268
hexinhui创建于 3月3日
卷积反向添加输入属性打印 #2268
已合并
共 5 个文件变更+194-21
| @@ -0,0 +1,83 @@ | |||
| 1 | +/** | ||
| 2 | + * Copyright (c) 2025 Huawei Technologies Co., Ltd. | ||
| 3 | + * This program is free software, you can redistribute it and/or modify it under the terms and conditions of | ||
| 4 | + * CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 5 | + * Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 6 | + * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 7 | + * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 8 | + * See LICENSE in the root of the software repository for the full text of the License. | ||
| 9 | + */ | ||
| 10 | + | ||
| 11 | +/*! | ||
| 12 | + * \file conv_tiling_debug_util.h | ||
| 13 | + * \brief Common debug utilities for conv tiling | ||
| 14 | + */ | ||
| 15 | + | ||
| 16 | + | ||
| 17 | + | ||
| 18 | + | ||
| 19 | + | ||
| 20 | + | ||
| 21 | + | ||
| 22 | + | ||
| 23 | +namespace Ops { | ||
| 24 | +namespace NN { | ||
| 25 | +namespace Conv { | ||
| 26 | + | ||
| 27 | +struct TensorInfo { | ||
| 28 | + std::vector<int64_t> shape; | ||
| 29 | + ge::Format format; | ||
| 30 | + ge::DataType dtype; | ||
| 31 | +}; | ||
| 32 | + | ||
| 33 | +template <typename T> | ||
| 34 | +std::string DebugString(const std::vector<T>& v) { | ||
| 35 | + std::ostringstream oss; | ||
| 36 | + oss << "["; | ||
| 37 | + if (v.size() > 0) { | ||
| 38 | + for (size_t i = 0; i < v.size() - 1; ++i) { | ||
| 39 | + oss << v[i] << ", "; | ||
| 40 | + } | ||
| 41 | + oss << v[v.size() - 1]; | ||
| 42 | + } | ||
| 43 | + oss << "]"; | ||
| 44 | + return oss.str(); | ||
| 45 | +} | ||
| 46 | + | ||
| 47 | +inline void DebugShape(gert::TilingContext* context, const int64_t index, std::vector<int64_t>& shape, bool isInput) { | ||
| 48 | + auto geShape = isInput ? context->GetInputShape(index)->GetStorageShape() : context->GetOutputShape(index)->GetStorageShape(); | ||
| 49 | + int32_t dimNum = geShape.GetDimNum(); | ||
| 50 | + shape.reserve(dimNum); | ||
| 51 | + for (int i = 0; i < dimNum; ++i) { | ||
| 52 | + shape.push_back(geShape.GetDim(i)); | ||
| 53 | + } | ||
| 54 | +} | ||
| 55 | + | ||
| 56 | +inline TensorInfo GetTensorInfo(gert::TilingContext* context, int64_t index, bool isInput, int64_t dimCount) { | ||
| 57 | + TensorInfo info; | ||
| 58 | + auto tensor = isInput ? context->GetInputDesc(index) : context->GetOutputDesc(index); | ||
| 59 | + OP_CHECK_IF(tensor == nullptr, OP_LOGE(context->GetNodeName(), "get tensor desc from context fail."), return info); | ||
| 60 | + | ||
| 61 | + if (dimCount == 1) { | ||
| 62 | + info.shape = {context->GetInputShape(index)->GetStorageShape().GetDim(0)}; | ||
| 63 | + } else { | ||
| 64 | + DebugShape(context, index, info.shape, isInput); | ||
| 65 | + } | ||
| 66 | + info.format = tensor->GetOriginFormat(); | ||
| 67 | + info.dtype = tensor->GetDataType(); | ||
| 68 | + return info; | ||
| 69 | +} | ||
| 70 | + | ||
| 71 | +inline std::vector<int64_t> GetAttrVector(gert::TilingContext* context, int attrIndex, int expectedSize, const char* attrName) { | ||
| 72 | + auto attrs = context->GetAttrs(); | ||
| 73 | + const auto attr = attrs->GetAttrPointer<gert::ContinuousVector>(attrIndex); | ||
| 74 | + OP_CHECK_IF(attr == nullptr, OP_LOGE(context->GetNodeName(), "get %s from context fail.", attrName), return {}); | ||
| 75 | + OP_CHECK_IF(attr->GetSize() != expectedSize, OP_LOGE(context->GetNodeName(), "%s of context dim len is invalid.", attrName), return {}); | ||
| 76 | + | ||
| 77 | + const auto data = static_cast<const int64_t *>(attr->GetData()); | ||
| 78 | + return std::vector<int64_t>(data, data + expectedSize); | ||
| 79 | +} | ||
| 80 | + | ||
| 81 | +} // namespace Conv | ||
| 82 | +} // namespace NN | ||
| 83 | +} // namespace Ops | ||
Mconv/conv3d_backprop_filter_v2/op_host/op_tiling/arch35/conv3d_backprop_filter_v2_basic_block_tiling.cpp+45-9
| @@ -25,6 +25,7 @@ | |||
| 25 | 25 | ||
| 26 | 26 | ||
| 27 | 27 | ||
| 28 | + | ||
| 28 | 29 | ||
| 29 | using Ops::NN::Optiling::RecursiveSum; | 30 | using Ops::NN::Optiling::RecursiveSum; |
| 30 | 31 | ||
| @@ -32,6 +33,14 @@ namespace { | |||
| 32 | constexpr size_t Y_INDEX = 2; | 33 | constexpr size_t Y_INDEX = 2; |
| 33 | constexpr size_t FILTER_INDEX = 0; | 34 | constexpr size_t FILTER_INDEX = 0; |
| 34 | constexpr size_t OUTPUT_BP_INDEX = 0; | 35 | constexpr size_t OUTPUT_BP_INDEX = 0; |
| 36 | +constexpr size_t FILTER_SIZE_INDEX = 1; | ||
| 37 | +const int32_t kFilterSizeDim = 1; | ||
| 38 | +const int32_t kConv3DbpDim = 5; | ||
| 39 | +const int32_t kPadsDim = 6; | ||
| 40 | +const int32_t strideIndex = 0; | ||
| 41 | +const int32_t dilationIndex = 2; | ||
| 42 | +const int32_t groupIndex = 3; | ||
| 43 | +const int32_t enabelHF32Index = 3; | ||
| 35 | } // namespace | 44 | } // namespace |
| 36 | 45 | ||
| 37 | namespace Ops { | 46 | namespace Ops { |
| @@ -1129,18 +1138,10 @@ void Conv3DDWV2BasicBlockTilingArch35::PrintTilingData() | |||
| 1129 | { | 1138 | { |
| 1130 | conv_bp_v2_kernel::TConv3DDwTiling& tiling = tilingData_.dwTiling; | 1139 | conv_bp_v2_kernel::TConv3DDwTiling& tiling = tilingData_.dwTiling; |
| 1131 | std::stringstream ss; | 1140 | std::stringstream ss; |
| 1141 | + // 删除shape stride dilation 相关打印 pads下移 | ||
Y | |||
| 1132 | ss << "batch: " << tiling.batch << " cin: " << tiling.cin << " cout: " << tiling.cout | 1142 | ss << "batch: " << tiling.batch << " cin: " << tiling.cin << " cout: " << tiling.cout |
| 1133 | << " cin1G: " << tiling.cin1G << " cout1G: " << tiling.cout1G | 1143 | << " cin1G: " << tiling.cin1G << " cout1G: " << tiling.cout1G |
| 1134 | - << " dout: " << tiling.dout << " ho: " << tiling.ho << " wo: " << tiling.wo | ||
| 1135 | - << " di: " << tiling.di << " hi: " << tiling.hi << " wi: " << tiling.wi | ||
| 1136 | - << " dk: " << tiling.dk << " hk: " << tiling.hk << " wk: " << tiling.wk | ||
| 1137 | << " group: " << tiling.group << " realGroup: " << tiling.realGroup | 1144 | << " group: " << tiling.group << " realGroup: " << tiling.realGroup |
| 1138 | - << " strideD: " << tiling.strideD << " strideH: " << tiling.strideH | ||
| 1139 | - << " strideW: " << tiling.strideW << " padFront: " << tiling.padFront | ||
| 1140 | - << " padBack: " << tiling.padBack << " padUp: " << tiling.padUp | ||
| 1141 | - << " padDown: " << tiling.padDown << " padLeft: " << tiling.padLeft | ||
| 1142 | - << " padRight: " << tiling.padRight << " dilationD: " << tiling.dilationD | ||
| 1143 | - << " dilationH: " << tiling.dilationH << " dilationW: " << tiling.dilationW | ||
| 1144 | << " channelSize: " << tiling.channelSize << " al0Pbuffer: " << tiling.al0Pbuffer | 1145 | << " channelSize: " << tiling.channelSize << " al0Pbuffer: " << tiling.al0Pbuffer |
| 1145 | << " bl0Pbuffer: " << tiling.bl0Pbuffer << " cl0Pbuffer: " << tiling.cl0Pbuffer | 1146 | << " bl0Pbuffer: " << tiling.bl0Pbuffer << " cl0Pbuffer: " << tiling.cl0Pbuffer |
| 1146 | << " al1Pbuffer: " << tiling.al1Pbuffer << " bl1Pbuffer: " << tiling.bl1Pbuffer | 1147 | << " al1Pbuffer: " << tiling.al1Pbuffer << " bl1Pbuffer: " << tiling.bl1Pbuffer |
| @@ -1154,6 +1155,41 @@ void Conv3DDWV2BasicBlockTilingArch35::PrintTilingData() | |||
| 1154 | << " splitWoSize: " << tiling.splitWo << " isSplitKernelHW: " << tiling.isSplitKernelHW | 1155 | << " splitWoSize: " << tiling.splitWo << " isSplitKernelHW: " << tiling.isSplitKernelHW |
| 1155 | << " singleCoreBatch: " << tiling.singleCoreBatch << " singleCoreCin: " << tiling.singleCoreCin; | 1156 | << " singleCoreBatch: " << tiling.singleCoreBatch << " singleCoreCin: " << tiling.singleCoreCin; |
| 1156 | OP_LOGI(opName_, "api tiling: %s", ss.str().c_str()); | 1157 | OP_LOGI(opName_, "api tiling: %s", ss.str().c_str()); |
| 1158 | + PrintInputsAttrs(tiling); | ||
| 1159 | +} | ||
| 1160 | + | ||
| 1161 | +bool Conv3DDWV2BasicBlockTilingArch35::PrintInputsAttrs(conv_bp_v2_kernel::TConv3DDwTiling& tiling){ | ||
| 1162 | + const auto op_name = context_->GetNodeName(); | ||
| 1163 | + auto inputInfo = GetTensorInfo(context_, OUTPUT_BP_INDEX, true, kConv3DbpDim); | ||
| 1164 | + auto filterSizesInfo = GetTensorInfo(context_, FILTER_SIZE_INDEX, true, kFilterSizeDim); // dw filter_size dim 1 | ||
| 1165 | + auto outBackpropInfo = GetTensorInfo(context_, Y_INDEX, true, kConv3DbpDim); | ||
| 1166 | + auto outputInfo = GetTensorInfo(context_, FILTER_INDEX, false, kConv3DbpDim); | ||
| 1167 | + | ||
| 1168 | + OP_LOGD(op_name, "input shape: %s, format: %s, dtype: %s; filter_sizes shape: %s, format: %s, dtype: %s; out_backprop shape: %s, format: %s, dtype: %s; output shape: %s, format: %s, dtype: %s;", | ||
| 1169 | + DebugString(inputInfo.shape).c_str(), ge::TypeUtils::FormatToSerialString(inputInfo.format).c_str(), | ||
| 1170 | + ge::TypeUtils::DataTypeToSerialString(inputInfo.dtype).c_str(), | ||
| 1171 | + DebugString(filterSizesInfo.shape).c_str(), ge::TypeUtils::FormatToSerialString(filterSizesInfo.format).c_str(), | ||
| 1172 | + ge::TypeUtils::DataTypeToSerialString(filterSizesInfo.dtype).c_str(), | ||
| 1173 | + DebugString(outBackpropInfo.shape).c_str(), ge::TypeUtils::FormatToSerialString(outBackpropInfo.format).c_str(), | ||
| 1174 | + ge::TypeUtils::DataTypeToSerialString(outBackpropInfo.dtype).c_str(), | ||
| 1175 | + DebugString(outputInfo.shape).c_str(), ge::TypeUtils::FormatToSerialString(outputInfo.format).c_str(), | ||
| 1176 | + ge::TypeUtils::DataTypeToSerialString(outputInfo.dtype).c_str() | ||
| 1177 | + ); | ||
| 1178 | + | ||
| 1179 | + auto stridesShape = GetAttrVector(context_, strideIndex, kConv3DbpDim, "strides"); // stride idx 0 | ||
| 1180 | + // pads打印需要修改,可能从padding获取 | ||
| 1181 | + std::vector<int64_t> padsShape{tiling.padFront, tiling.padBack, tiling.padUp, tiling.padDown, tiling.padLeft, tiling.padRight}; | ||
| 1182 | + auto dilationsShape = GetAttrVector(context_, dilationIndex, kConv3DbpDim, "dilations"); // dilation idx 2 | ||
| 1183 | + | ||
| 1184 | + auto attrs = context_->GetAttrs(); | ||
| 1185 | + const auto groups = attrs->GetAttrPointer<int64_t>(groupIndex); // groups idx 3 | ||
| 1186 | + const auto enableHf32 = attrs->GetAttrPointer<bool>(enabelHF32Index); // enable_hf32 idx 5 | ||
| 1187 | + OP_CHECK_IF(groups == nullptr, OP_LOGE(op_name, "get groups from context fail."), return false); | ||
| 1188 | + | ||
| 1189 | + OP_LOGD(op_name, "Attrs stride: %s, pads: %s, dilation: %s, groups: %ld, enable_hf32: %d.", | ||
| 1190 | + DebugString(stridesShape).c_str(), DebugString(padsShape).c_str(), DebugString(dilationsShape).c_str(), | ||
| 1191 | + *groups, *enableHf32); | ||
| 1192 | + return true; | ||
| 1157 | } | 1193 | } |
| 1158 | 1194 | ||
| 1159 | void Conv3DDWV2BasicBlockTilingArch35::PrintBasickBlockTilingData() | 1195 | void Conv3DDWV2BasicBlockTilingArch35::PrintBasickBlockTilingData() |
Mconv/conv3d_backprop_filter_v2/op_host/op_tiling/arch35/conv3d_backprop_filter_v2_basic_block_tiling.h+2-0
| @@ -197,6 +197,8 @@ protected: | |||
| 197 | 197 | ||
| 198 | void PrintBasickBlockTilingData(); | 198 | void PrintBasickBlockTilingData(); |
| 199 | 199 | ||
| 200 | + bool PrintInputsAttrs(conv_bp_v2_kernel::TConv3DDwTiling& tiling); | ||
| 201 | + | ||
| 200 | void SetBasicBlockAttrsTiling(); | 202 | void SetBasicBlockAttrsTiling(); |
| 201 | 203 | ||
| 202 | void ShrinkBaseBlock(); | 204 | void ShrinkBaseBlock(); |
Mconv/conv3d_backprop_input_v2/op_host/op_tiling/arch35/conv3d_backprop_input_v2_base_tiling.cpp+59-12
| @@ -28,6 +28,7 @@ | |||
| 28 | 28 | ||
| 29 | 29 | ||
| 30 | 30 | ||
| 31 | + | ||
| 31 | 32 | ||
| 32 | using Ops::NN::Optiling::RecursiveSum; | 33 | using Ops::NN::Optiling::RecursiveSum; |
| 33 | 34 | ||
| @@ -42,6 +43,12 @@ constexpr int32_t BUFFER_NUM_L1 = 4; | |||
| 42 | constexpr float CORE_USED_THRESHOLD = 0.6f; | 43 | constexpr float CORE_USED_THRESHOLD = 0.6f; |
| 43 | constexpr uint64_t MAX_UINT16 = 65535; | 44 | constexpr uint64_t MAX_UINT16 = 65535; |
| 44 | const int32_t FP32_FIXPIPE_BOUND_K_LIMIT = 528; // 理论值,输出fp32时,当 K >= 528 时才能不fixpipe bound | 45 | const int32_t FP32_FIXPIPE_BOUND_K_LIMIT = 528; // 理论值,输出fp32时,当 K >= 528 时才能不fixpipe bound |
| 46 | +const int32_t kInputSizeDim = 1; | ||
| 47 | +const int32_t kConv3DbpDim = 5; | ||
| 48 | +const int32_t kPadsDim = 6; | ||
| 49 | +const int32_t strideIndex = 0; | ||
| 50 | +const int32_t dilationIndex = 2; | ||
| 51 | +const int32_t groupIndex = 3; | ||
| 45 | 52 | ||
| 46 | // 0: best base block; 1: threshold base block | 53 | // 0: best base block; 1: threshold base block |
| 47 | constexpr uint32_t BASE_BLOCK_TYPE_BEST = 0; | 54 | constexpr uint32_t BASE_BLOCK_TYPE_BEST = 0; |
| @@ -214,6 +221,52 @@ ge::graphStatus Conv3DBackpropInputV2TilingArch35::DoOpTiling() | |||
| 214 | return ge::GRAPH_SUCCESS; | 221 | return ge::GRAPH_SUCCESS; |
| 215 | } | 222 | } |
| 216 | 223 | ||
| 224 | +bool Conv3DBackpropInputV2TilingArch35::PrintInputsAttrs(conv_bp_v2_kernel::TConv3DInputV2Tiling& tiling){ | ||
| 225 | + const auto op_name = context_->GetNodeName(); | ||
| 226 | + size_t weight_index = (opType_ == optiling::OpTypeV2::kConv3DTransposeV2) ? TRANSPOSE_FILTER_INDEX : FILTER_INDEX; // dx filter idx 1 | transpose filter idx 2 | ||
| 227 | + size_t dedy_x_index = (opType_ == optiling::OpTypeV2::kConv3DTransposeV2) ? TRANSPOSE_X_INDEX : OUTPUT_BP_INDEX; // dx dedy idx 2 | transpose x idx 1 | ||
| 228 | + auto inputSizeInfo = GetTensorInfo(context_, INPUT_SIZE_INDEX, true, kInputSizeDim); // input_size dim=1 | ||
| 229 | + auto weightInfo = GetTensorInfo(context_, weight_index, true, kConv3DbpDim); | ||
| 230 | + auto dedyInfo = GetTensorInfo(context_, dedy_x_index, true, kConv3DbpDim); | ||
| 231 | + auto outputInfo = GetTensorInfo(context_, Y_INDEX, false, kConv3DbpDim); | ||
| 232 | + | ||
| 233 | + OP_LOGD(op_name, "input_size shape: %s, format: %s, dtype: %s; filter shape: %s, format: %s, dtype: %s; out_backprop/x shape: %s, format: %s, dtype: %s; y shape: %s, format: %s, dtype: %s;", | ||
| 234 | + DebugString(inputSizeInfo.shape).c_str(), ge::TypeUtils::FormatToSerialString(inputSizeInfo.format).c_str(), | ||
| 235 | + ge::TypeUtils::DataTypeToSerialString(inputSizeInfo.dtype).c_str(), | ||
| 236 | + DebugString(weightInfo.shape).c_str(), ge::TypeUtils::FormatToSerialString(weightInfo.format).c_str(), | ||
| 237 | + ge::TypeUtils::DataTypeToSerialString(weightInfo.dtype).c_str(), | ||
| 238 | + DebugString(dedyInfo.shape).c_str(), ge::TypeUtils::FormatToSerialString(dedyInfo.format).c_str(), | ||
| 239 | + ge::TypeUtils::DataTypeToSerialString(dedyInfo.dtype).c_str(), | ||
| 240 | + DebugString(outputInfo.shape).c_str(), ge::TypeUtils::FormatToSerialString(outputInfo.format).c_str(), | ||
| 241 | + ge::TypeUtils::DataTypeToSerialString(outputInfo.dtype).c_str() | ||
| 242 | + ); | ||
| 243 | + | ||
| 244 | + auto stridesShape = GetAttrVector(context_, strideIndex, kConv3DbpDim, "strides"); | ||
| 245 | + // pads打印需要修改,可能从padding获取 | ||
| 246 | + std::vector<int64_t> padsShape{tiling.padFront, tiling.padBack, tiling.padUp, tiling.padDown, tiling.padLeft, tiling.padRight}; | ||
| 247 | + auto dilationsShape = GetAttrVector(context_, dilationIndex, kConv3DbpDim, "dilations"); | ||
| 248 | + | ||
| 249 | + auto attrs = context_->GetAttrs(); | ||
| 250 | + const auto groups = attrs->GetAttrPointer<int64_t>(groupIndex); | ||
| 251 | + size_t enable_hf32_index = (opType_ == optiling::OpTypeV2::kConv3DTransposeV2) ? TRANSPOSE_ENABLE_HF32_INDEX : ENABLE_HF32_INDEX; // dx hf32 idx 5 | transpose hf32 idx 7 | ||
| 252 | + const auto enableHf32 = attrs->GetAttrPointer<bool>(enable_hf32_index); | ||
| 253 | + OP_CHECK_IF(groups == nullptr, OP_LOGE(op_name, "get groups from context fail."), return false); | ||
| 254 | + if (opType_ == optiling::OpTypeV2::kConv3DTransposeV2){ | ||
| 255 | + auto output_paddingShape = GetAttrVector(context_, OUTPUT_PADDING_INDEX, kConv3DbpDim, "output_padding"); | ||
| 256 | + const auto offset = attrs->GetAttrPointer<bool>(OFFSET_X_INDEX); | ||
| 257 | + OP_LOGD(op_name, "Attrs stride: %s, pads: %s, dilation: %s, groups: %ld, enable_hf32: %d, output_padding: %s, offset_x: %ld", | ||
| 258 | + DebugString(stridesShape).c_str(), DebugString(padsShape).c_str(), DebugString(dilationsShape).c_str(), | ||
| 259 | + *groups, *enableHf32, DebugString(output_paddingShape).c_str(), *offset); | ||
| 260 | + | ||
| 261 | + } else { | ||
| 262 | + OP_LOGD(op_name, "Attrs stride: %s, pads: %s, dilation: %s, groups: %ld, enable_hf32: %d.", | ||
| 263 | + DebugString(stridesShape).c_str(), DebugString(padsShape).c_str(), DebugString(dilationsShape).c_str(), | ||
| 264 | + *groups, *enableHf32); | ||
| 265 | + } | ||
| 266 | + | ||
| 267 | + return true; | ||
| 268 | +} | ||
| 269 | + | ||
| 217 | ge::graphStatus Conv3DBackpropInputV2TilingArch35::DoLibApiTiling() | 270 | ge::graphStatus Conv3DBackpropInputV2TilingArch35::DoLibApiTiling() |
| 218 | { | 271 | { |
| 219 | SetDxTilingFromTbeTiling(); | 272 | SetDxTilingFromTbeTiling(); |
| @@ -1029,6 +1082,7 @@ void Conv3DBackpropInputV2TilingArch35::PrintTilingData() | |||
| 1029 | conv_bp_v2_kernel::Conv3DBackpropInputV2Params& params = tilingData_.params; | 1082 | conv_bp_v2_kernel::Conv3DBackpropInputV2Params& params = tilingData_.params; |
| 1030 | conv_bp_v2_kernel::TConv3DInputV2KSTiling& ksTiling = tilingData_.conv3DDxKSTiling; | 1083 | conv_bp_v2_kernel::TConv3DInputV2KSTiling& ksTiling = tilingData_.conv3DDxKSTiling; |
| 1031 | std::stringstream ss; | 1084 | std::stringstream ss; |
| 1085 | + // 删除shape stride dilation 相关打印 pads下移 | ||
| 1032 | ss << "batchDim: " << params.batchDim << " groupDim: " << params.groupDim | 1086 | ss << "batchDim: " << params.batchDim << " groupDim: " << params.groupDim |
| 1033 | << " mDim: " << params.mDim << " kDim: " << params.kDim << " nDim: " << params.nDim | 1087 | << " mDim: " << params.mDim << " kDim: " << params.kDim << " nDim: " << params.nDim |
| 1034 | << " dDim: " << params.dDim << " coreNum: " << params.coreNum | 1088 | << " dDim: " << params.dDim << " coreNum: " << params.coreNum |
| @@ -1043,22 +1097,14 @@ void Conv3DBackpropInputV2TilingArch35::PrintTilingData() | |||
| 1043 | << " enlarge: " << static_cast<uint32_t>(tiling.enlarge) | 1097 | << " enlarge: " << static_cast<uint32_t>(tiling.enlarge) |
| 1044 | << " hf32Flag: " << static_cast<uint32_t>(tiling.hf32Flag) | 1098 | << " hf32Flag: " << static_cast<uint32_t>(tiling.hf32Flag) |
| 1045 | << " initOutputFlag: " << static_cast<uint32_t>(tiling.initOutputFlag) | 1099 | << " initOutputFlag: " << static_cast<uint32_t>(tiling.initOutputFlag) |
| 1046 | - << " isBiasFullLoad: " << static_cast<uint32_t>(tiling.isBiasFullLoad) << " batch: " << tiling.batch | 1100 | + << " isBiasFullLoad: " << static_cast<uint32_t>(tiling.isBiasFullLoad) |
| 1047 | << " cin: " << tiling.cin << " cout: " << tiling.cout << " cinG: " << tiling.cinG | 1101 | << " cin: " << tiling.cin << " cout: " << tiling.cout << " cinG: " << tiling.cinG |
| 1048 | << " coutG: " << tiling.coutG << " cout1: " << tiling.cout1 << " cin1: " << tiling.cin1 | 1102 | << " coutG: " << tiling.coutG << " cout1: " << tiling.cout1 << " cin1: " << tiling.cin1 |
| 1049 | - << " cout1G: " << tiling.cout1G << " cin1G: " << tiling.cin1G << " dout: " << tiling.dout | 1103 | + << " cout1G: " << tiling.cout1G << " cin1G: " << tiling.cin1G |
| 1050 | - << " ho: " << tiling.ho << " wo: " << tiling.wo << " di: " << tiling.di | 1104 | + << " group: " << tiling.group << " oriGroup: " << tiling.oriGroup |
| 1051 | - << " hi: " << tiling.hi << " wi: " << tiling.wi << " dk: " << tiling.dk | ||
| 1052 | - << " hk: " << tiling.hk << " wk: " << tiling.wk << " group: " << tiling.group | ||
| 1053 | - << " oriGroup: " << tiling.oriGroup << " strideD: " << tiling.strideD | ||
| 1054 | - << " strideH: " << tiling.strideH << " strideW: " << tiling.strideW | ||
| 1055 | - << " padFront: " << tiling.padFront << " padBack: " << tiling.padBack | ||
| 1056 | - << " padUp: " << tiling.padUp << " padDown: " << tiling.padDown | ||
| 1057 | - << " padLeft: " << tiling.padLeft << " padRight: " << tiling.padRight | ||
| 1058 | << " backpropPadTail: " << tiling.backpropPadTail << " backpropPadUp: " << tiling.backpropPadUp | 1105 | << " backpropPadTail: " << tiling.backpropPadTail << " backpropPadUp: " << tiling.backpropPadUp |
| 1059 | << " backpropPadDown: " << tiling.backpropPadDown << " backpropPadLeft: " << tiling.backpropPadLeft | 1106 | << " backpropPadDown: " << tiling.backpropPadDown << " backpropPadLeft: " << tiling.backpropPadLeft |
| 1060 | - << " backpropPadRight: " << tiling.backpropPadRight << " dilationD: " << tiling.dilationD | 1107 | + << " backpropPadRight: " << tiling.backpropPadRight |
| 1061 | - << " dilationH: " << tiling.dilationH << " dilationW: " << tiling.dilationW | ||
| 1062 | << " singleCoreGroup: " << tiling.singleCoreGroup << " singleCoreCout: " << tiling.singleCoreCout | 1108 | << " singleCoreGroup: " << tiling.singleCoreGroup << " singleCoreCout: " << tiling.singleCoreCout |
| 1063 | << " singleCoreCin: " << tiling.singleCoreCin << " singleCoreDin: " << tiling.singleCoreDin | 1109 | << " singleCoreCin: " << tiling.singleCoreCin << " singleCoreDin: " << tiling.singleCoreDin |
| 1064 | << " baseM: " << tiling.baseM << " baseK: " << tiling.baseK << " baseN: " << tiling.baseN | 1110 | << " baseM: " << tiling.baseM << " baseK: " << tiling.baseK << " baseN: " << tiling.baseN |
| @@ -1068,6 +1114,7 @@ void Conv3DBackpropInputV2TilingArch35::PrintTilingData() | |||
| 1068 | << " enableVecTrans: " << static_cast<uint32_t>(tiling.enableVecTrans) | 1114 | << " enableVecTrans: " << static_cast<uint32_t>(tiling.enableVecTrans) |
| 1069 | << " kSCoutFullLoad: " << ksTiling.kSCoutFullLoad << " kSUseWorkSpace: " << ksTiling.kSUseWorkSpace; | 1115 | << " kSCoutFullLoad: " << ksTiling.kSCoutFullLoad << " kSUseWorkSpace: " << ksTiling.kSUseWorkSpace; |
| 1070 | OP_LOGD(opName_, "api tiling: %s", ss.str().c_str()); | 1116 | OP_LOGD(opName_, "api tiling: %s", ss.str().c_str()); |
| 1117 | + PrintInputsAttrs(tiling); | ||
| 1071 | } | 1118 | } |
| 1072 | 1119 | ||
| 1073 | REGISTER_TILING_TEMPLATE("Conv3DBackpropInputV2", Conv3DBackpropInputV2TilingArch35, 102); | 1120 | REGISTER_TILING_TEMPLATE("Conv3DBackpropInputV2", Conv3DBackpropInputV2TilingArch35, 102); |
| @@ -46,11 +46,15 @@ const size_t Y_INDEX = 0; | |||
| 46 | const size_t INPUT_SIZE_INDEX = 0; | 46 | const size_t INPUT_SIZE_INDEX = 0; |
| 47 | const size_t FILTER_INDEX = 1; | 47 | const size_t FILTER_INDEX = 1; |
| 48 | const size_t OUTPUT_BP_INDEX = 2; | 48 | const size_t OUTPUT_BP_INDEX = 2; |
| 49 | +const size_t TRANSPOSE_X_INDEX = 1; | ||
| 50 | +const size_t TRANSPOSE_FILTER_INDEX = 2; | ||
| 49 | const size_t BAIS_INDEX = 3; | 51 | const size_t BAIS_INDEX = 3; |
| 50 | const size_t SCALE_INDEX = 4; | 52 | const size_t SCALE_INDEX = 4; |
| 51 | const size_t OFFSET_W_INDEX = 4; | 53 | const size_t OFFSET_W_INDEX = 4; |
| 54 | +const size_t ENABLE_HF32_INDEX = 5; | ||
| 52 | const size_t OUTPUT_PADDING_INDEX = 5; | 55 | const size_t OUTPUT_PADDING_INDEX = 5; |
| 53 | const size_t OFFSET_X_INDEX = 6; | 56 | const size_t OFFSET_X_INDEX = 6; |
| 57 | +const size_t TRANSPOSE_ENABLE_HF32_INDEX = 5; | ||
| 54 | 58 | ||
| 55 | struct TilingValueDavid { | 59 | struct TilingValueDavid { |
| 56 | uint64_t coreNum; | 60 | uint64_t coreNum; |
| @@ -287,6 +291,7 @@ private: | |||
| 287 | bool AnalyzeFuseDtype(const bool f16flag, const ge::DataType outputBackpropDtype, | 291 | bool AnalyzeFuseDtype(const bool f16flag, const ge::DataType outputBackpropDtype, |
| 288 | const ge::DataType filterDtype, const ge::DataType yDtype) const; | 292 | const ge::DataType filterDtype, const ge::DataType yDtype) const; |
| 289 | bool CheckL0Size(uint32_t baseM, uint32_t baseN, uint32_t baseK, uint32_t l0Pbuffer = DB_ON); | 293 | bool CheckL0Size(uint32_t baseM, uint32_t baseN, uint32_t baseK, uint32_t l0Pbuffer = DB_ON); |
| 294 | + bool PrintInputsAttrs(conv_bp_v2_kernel::TConv3DInputV2Tiling& tiling); | ||
| 290 | }; | 295 | }; |
| 291 | 296 | ||
| 292 | } // namespace Conv | 297 | } // namespace Conv |


PrintInputsAttrs中重复的参数建议去掉,避免冗余打印。