已合并
fix: 文件名重复问题解决,改引用公共目录下的文件 #9572
duxinlei创建于 8月31日
fix: 文件名重复问题解决,改引用公共目录下的文件 #9572
已合并
共 3 个文件变更+1-381
| @@ -28,7 +28,7 @@ | |||
| 28 | 28 | ||
| 29 | 29 | ||
| 30 | 30 | ||
| 31 | -#include "pool_tiling_templates_registry.h" | 31 | +#include "pooling/pool_3d_common/op_host/arch35/pool_tiling_templates_registry.h" |
| 32 | 32 | ||
| 33 | 33 | ||
| 34 | using optiling::PoolTilingRegistry; | 34 | using optiling::PoolTilingRegistry; |
| @@ -1,169 +0,0 @@ | |||
| 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 pool_tiling_templates_registry.h | ||
| 13 | - * \brief | ||
| 14 | - */ | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | -namespace optiling { | ||
| 27 | -template <typename T> | ||
| 28 | -std::unique_ptr<TilingBaseClass> TILING_CLASS(gert::TilingContext* context) | ||
| 29 | -{ | ||
| 30 | - return std::make_unique<T>(T(context)); | ||
| 31 | -} | ||
| 32 | - | ||
| 33 | -using TilingClassCase = std::unique_ptr<TilingBaseClass> (*)(gert::TilingContext*); | ||
| 34 | - | ||
| 35 | -class PoolTilingCases { | ||
| 36 | -public: | ||
| 37 | - explicit PoolTilingCases(std::string op_type) : op_type_(std::move(op_type)) {} | ||
| 38 | - | ||
| 39 | - template <typename T> | ||
| 40 | - void AddTiling(int32_t priority) | ||
| 41 | - { | ||
| 42 | - OP_CHECK_IF(cases_.find(priority) != cases_.end(), OP_LOGE(op_type_, "There are duplicate registrations."), | ||
| 43 | - return); | ||
| 44 | - cases_[priority] = TILING_CLASS<T>; | ||
| 45 | - OP_CHECK_IF(cases_[priority] == nullptr, | ||
| 46 | - OP_LOGE(op_type_, "PoolRegister op tiling func failed, please check the class name."), return); | ||
| 47 | - } | ||
| 48 | - | ||
| 49 | - const std::map<int32_t, TilingClassCase>& GetPoolTilingCases() { return cases_; } | ||
| 50 | - | ||
| 51 | -private: | ||
| 52 | - std::map<int32_t, TilingClassCase> cases_; | ||
| 53 | - const std::string op_type_; | ||
| 54 | -}; | ||
| 55 | - | ||
| 56 | -// --------------------------------Interfacce without soc version -------------------------------- | ||
| 57 | -class PoolTilingRegistry { | ||
| 58 | -public: | ||
| 59 | - PoolTilingRegistry() = default; | ||
| 60 | - | ||
| 61 | - | ||
| 62 | - static PoolTilingRegistry& GetInstance(); | ||
| 63 | - | ||
| 64 | - static PoolTilingRegistry& GetInstance() | ||
| 65 | - { | ||
| 66 | - static PoolTilingRegistry registry_impl_; | ||
| 67 | - return registry_impl_; | ||
| 68 | - } | ||
| 69 | - | ||
| 70 | - | ||
| 71 | - std::shared_ptr<PoolTilingCases> RegisterOp(const std::string& op_type) | ||
| 72 | - { | ||
| 73 | - if (registry_map_.find(op_type) == registry_map_.end()) { | ||
| 74 | - registry_map_[op_type] = std::make_shared<PoolTilingCases>(PoolTilingCases(op_type)); | ||
| 75 | - } | ||
| 76 | - OP_CHECK_IF(registry_map_[op_type] == nullptr, | ||
| 77 | - OP_LOGE(op_type, "PoolRegister tiling func failed, please check the class name."), return nullptr); | ||
| 78 | - return registry_map_[op_type]; | ||
| 79 | - } | ||
| 80 | - | ||
| 81 | - ge::graphStatus DoTilingImpl(gert::TilingContext* context) | ||
| 82 | - { | ||
| 83 | - const char* op_type = context->GetNodeType(); | ||
| 84 | - auto tilingTemplateRegistryMap = GetTilingTemplates(op_type); | ||
| 85 | - for (auto it = tilingTemplateRegistryMap.begin(); it != tilingTemplateRegistryMap.end(); ++it) { | ||
| 86 | - auto tilingTemplate = it->second(context); | ||
| 87 | - if (tilingTemplate != nullptr) { | ||
| 88 | - ge::graphStatus status = tilingTemplate->DoTiling(); | ||
| 89 | - if (status != ge::GRAPH_PARAM_INVALID) { | ||
| 90 | - OP_LOGD(context, "Do general op tiling success priority=%d", it->first); | ||
| 91 | - return status; | ||
| 92 | - } | ||
| 93 | - OP_LOGD(context, "Ignore general op tiling priority=%d", it->first); | ||
| 94 | - } | ||
| 95 | - } | ||
| 96 | - OP_LOGE(op_type, "Do op tiling failed, no valid template is found."); | ||
| 97 | - return ge::GRAPH_FAILED; | ||
| 98 | - } | ||
| 99 | - | ||
| 100 | - ge::graphStatus DoTilingImpl(gert::TilingContext* context, const std::vector<int32_t>& priorities) | ||
| 101 | - { | ||
| 102 | - const char* op_type = context->GetNodeType(); | ||
| 103 | - auto tilingTemplateRegistryMap = GetTilingTemplates(op_type); | ||
| 104 | - for (auto priorityId : priorities) { | ||
| 105 | - auto templateFunc = tilingTemplateRegistryMap[priorityId](context); | ||
| 106 | - if (templateFunc != nullptr) { | ||
| 107 | - ge::graphStatus status = templateFunc->DoTiling(); | ||
| 108 | - if (status == ge::GRAPH_SUCCESS) { | ||
| 109 | - OP_LOGD(context, "Do general op tiling success priority=%d", priorityId); | ||
| 110 | - return status; | ||
| 111 | - } | ||
| 112 | - if (status != ge::GRAPH_PARAM_INVALID) { | ||
| 113 | - OP_LOGD(context, "Do op tiling failed"); | ||
| 114 | - return status; | ||
| 115 | - } | ||
| 116 | - OP_LOGD(context, "Ignore general op tiling priority=%d", priorityId); | ||
| 117 | - } | ||
| 118 | - } | ||
| 119 | - OP_LOGE(op_type, "Do op tiling failed, no valid template is found."); | ||
| 120 | - return ge::GRAPH_FAILED; | ||
| 121 | - } | ||
| 122 | - | ||
| 123 | - const std::map<int32_t, TilingClassCase>& GetTilingTemplates(const std::string& op_type) | ||
| 124 | - { | ||
| 125 | - OP_CHECK_IF(registry_map_.find(op_type) == registry_map_.end(), | ||
| 126 | - OP_LOGE(op_type, "Get op tiling func failed, please check the op name."), | ||
| 127 | - return empty_tiling_case_); | ||
| 128 | - return registry_map_[op_type]->GetPoolTilingCases(); | ||
| 129 | - } | ||
| 130 | - | ||
| 131 | -private: | ||
| 132 | - std::map<std::string, std::shared_ptr<PoolTilingCases>> registry_map_; | ||
| 133 | - const std::map<int32_t, TilingClassCase> empty_tiling_case_; | ||
| 134 | -}; | ||
| 135 | - | ||
| 136 | -class PoolRegister { | ||
| 137 | -public: | ||
| 138 | - explicit PoolRegister(std::string op_type) : op_type_(std::move(op_type)) {} | ||
| 139 | - | ||
| 140 | - template <typename T> | ||
| 141 | - PoolRegister& tiling(int32_t priority) | ||
| 142 | - { | ||
| 143 | - auto PoolTilingCases = PoolTilingRegistry::GetInstance().RegisterOp(op_type_); | ||
| 144 | - OP_CHECK_IF(PoolTilingCases == nullptr, OP_LOGE(op_type_, "PoolRegister op tiling failed, please the op name."), | ||
| 145 | - return *this); | ||
| 146 | - PoolTilingCases->AddTiling<T>(priority); | ||
| 147 | - return *this; | ||
| 148 | - } | ||
| 149 | - | ||
| 150 | -private: | ||
| 151 | - const std::string op_type_; | ||
| 152 | -}; | ||
| 153 | - | ||
| 154 | -// op_type: 算子名称, class_name: 注册的 tiling 类, | ||
| 155 | -// priority: tiling 类的优先级, 越小表示优先级越高, 即被选中的概率越大 | ||
| 156 | - | ||
| 157 | - GLOBAL_REGISTER_STR_SYMBOL(op_type, class_name, priority, __COUNTER__, __LINE__); \ | ||
| 158 | - static PoolRegister VAR_UNUSED##op_type_##class_name##priority_register = PoolRegister(op_type) \ | ||
| 159 | - .tiling<class_name>(priority) | ||
| 160 | - | ||
| 161 | -// op_type: 算子名称, class_name: 注册的 tiling 类, | ||
| 162 | -// priority: tiling 类的优先级, 越小表示优先级越高, 即被选中的概率越大 | ||
| 163 | -// 取代 REGISTER_TILING_TEMPLATE , 传入的op_type如果是字符串常量,需要去掉引号 | ||
| 164 | - | ||
| 165 | - GLOBAL_REGISTER_SYMBOL(op_type, class_name, priority, __COUNTER__, __LINE__); \ | ||
| 166 | - static PoolRegister __attribute__((unused)) \ | ||
| 167 | - tiling_##op_type##_##class_name##_# | ||
| 168 | -} // namespace optiling | ||
| 169 | - | ||
| @@ -1,211 +0,0 @@ | |||
| 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 tiling_base.h | ||
| 13 | - * \brief | ||
| 14 | - */ | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - | ||
| 27 | - | ||
| 28 | - | ||
| 29 | - | ||
| 30 | - | ||
| 31 | - | ||
| 32 | -namespace optiling { | ||
| 33 | - | ||
| 34 | -struct AiCoreParams { | ||
| 35 | - uint64_t ubSize = 0UL; | ||
| 36 | - uint64_t blockDim = 0UL; | ||
| 37 | - uint64_t numBlocks = 0UL; | ||
| 38 | - uint64_t aicNum = 0UL; | ||
| 39 | - uint64_t l1Size = 0UL; | ||
| 40 | - uint64_t l0aSize = 0UL; | ||
| 41 | - uint64_t l0bSize = 0UL; | ||
| 42 | - uint64_t l0cSize = 0UL; | ||
| 43 | -}; | ||
| 44 | - | ||
| 45 | -class TilingBaseClass { | ||
| 46 | -public: | ||
| 47 | - explicit TilingBaseClass(gert::TilingContext* context) : context_(context) {} | ||
| 48 | - | ||
| 49 | - virtual ~TilingBaseClass() = default; | ||
| 50 | - | ||
| 51 | - // Tiling执行框架 | ||
| 52 | - // 1、GRAPH_SUCCESS: 成功,并且不需要继续执行后续Tiling类的实现 | ||
| 53 | - // 2、GRAPH_FAILED: 失败,中止整个Tiling流程 | ||
| 54 | - // 3、GRAPH_PARAM_INVALID: 本类不支持,需要继续往下执行其他Tiling类的实现 | ||
| 55 | - ge::graphStatus DoTiling() | ||
| 56 | - { | ||
| 57 | - auto ret = GetShapeAttrsInfo(); | ||
| 58 | - if (ret != ge::GRAPH_SUCCESS) { | ||
| 59 | - return ret; | ||
| 60 | - } | ||
| 61 | - ret = GetPlatformInfo(); | ||
| 62 | - if (ret != ge::GRAPH_SUCCESS) { | ||
| 63 | - return ret; | ||
| 64 | - } | ||
| 65 | - if (!IsCapable()) { | ||
| 66 | - return ge::GRAPH_PARAM_INVALID; | ||
| 67 | - } | ||
| 68 | - ret = DoOpTiling(); | ||
| 69 | - if (ret != ge::GRAPH_SUCCESS) { | ||
| 70 | - return ret; | ||
| 71 | - } | ||
| 72 | - ret = DoLibApiTiling(); | ||
| 73 | - if (ret != ge::GRAPH_SUCCESS) { | ||
| 74 | - return ret; | ||
| 75 | - } | ||
| 76 | - ret = GetWorkspaceSize(); | ||
| 77 | - if (ret != ge::GRAPH_SUCCESS) { | ||
| 78 | - return ret; | ||
| 79 | - } | ||
| 80 | - ret = PostTiling(); | ||
| 81 | - if (ret != ge::GRAPH_SUCCESS) { | ||
| 82 | - return ret; | ||
| 83 | - } | ||
| 84 | - context_->SetTilingKey(GetTilingKey()); | ||
| 85 | - DumpTilingInfo(); | ||
| 86 | - return ge::GRAPH_SUCCESS; | ||
| 87 | - } | ||
| 88 | - | ||
| 89 | - // 更新 context | ||
| 90 | - virtual void Reset(gert::TilingContext* context) { context_ = context; } | ||
| 91 | - | ||
| 92 | -protected: | ||
| 93 | - virtual bool IsCapable() = 0; | ||
| 94 | - // 1、获取平台信息比如CoreNum、UB/L1/L0C资源大小 | ||
| 95 | - virtual ge::graphStatus GetPlatformInfo() = 0; | ||
| 96 | - // 2、获取INPUT/OUTPUT/ATTR信息 | ||
| 97 | - virtual ge::graphStatus GetShapeAttrsInfo() = 0; | ||
| 98 | - // 3、计算数据切分TilingData | ||
| 99 | - virtual ge::graphStatus DoOpTiling() = 0; | ||
| 100 | - // 4、计算高阶API的TilingData | ||
| 101 | - virtual ge::graphStatus DoLibApiTiling() = 0; | ||
| 102 | - // 5、计算TilingKey | ||
| 103 | - [[nodiscard]] virtual uint64_t GetTilingKey() const = 0; | ||
| 104 | - // 6、计算Workspace 大小 | ||
| 105 | - virtual ge::graphStatus GetWorkspaceSize() = 0; | ||
| 106 | - // 7、保存Tiling数据 | ||
| 107 | - virtual ge::graphStatus PostTiling() = 0; | ||
| 108 | - // 8、Dump Tiling数据 | ||
| 109 | - virtual void DumpTilingInfo() { OP_LOGD(context_, "%ld", DefaultTilingInfoDump()); } | ||
| 110 | - | ||
| 111 | - int64_t DefaultTilingInfoDump() | ||
| 112 | - { | ||
| 113 | - auto buf = (uint32_t*)context_->GetRawTilingData()->GetData(); | ||
| 114 | - auto bufLen = context_->GetRawTilingData()->GetDataSize(); | ||
| 115 | - std::ostringstream oss; | ||
| 116 | - oss << "Start to dump tiling info. tilingkey:" << context_->GetTilingKey() << ", tiling data size:" << bufLen | ||
| 117 | - << ", content:"; | ||
| 118 | - for (size_t i = 0; i < bufLen / sizeof(uint32_t); i++) { | ||
| 119 | - oss << *(buf + i) << ","; | ||
| 120 | - if (oss.str().length() > 640) { // Split according to 640 to avoid truncation | ||
| 121 | - OP_LOGD(context_, "%s", oss.str().c_str()); | ||
| 122 | - oss.str(""); | ||
| 123 | - } | ||
| 124 | - } | ||
| 125 | - OP_LOGD(context_, "%s", oss.str().c_str()); | ||
| 126 | - return 0; | ||
| 127 | - } | ||
| 128 | - | ||
| 129 | - static uint32_t CalcTschBlockDim(uint32_t sliceNum, uint32_t aicCoreNum, uint32_t aivCoreNum) | ||
| 130 | - { | ||
| 131 | - uint32_t ration; | ||
| 132 | - if (aicCoreNum == 0 || aivCoreNum == 0 || aicCoreNum > aivCoreNum) { | ||
| 133 | - return sliceNum; | ||
| 134 | - } | ||
| 135 | - ration = aivCoreNum / aicCoreNum; | ||
| 136 | - return (sliceNum + (ration - 1)) / ration; | ||
| 137 | - } | ||
| 138 | - | ||
| 139 | - template <typename T> | ||
| 140 | - [[nodiscard]] std::string GetShapeDebugStr(const T& shape) const | ||
| 141 | - { | ||
| 142 | - std::ostringstream oss; | ||
| 143 | - oss << "["; | ||
| 144 | - if (shape.GetDimNum() > 0) { | ||
| 145 | - for (size_t i = 0; i < shape.GetDimNum() - 1; ++i) { | ||
| 146 | - oss << shape.GetDim(i) << ", "; | ||
| 147 | - } | ||
| 148 | - oss << shape.GetDim(shape.GetDimNum() - 1); | ||
| 149 | - } | ||
| 150 | - oss << "]"; | ||
| 151 | - return oss.str(); | ||
| 152 | - } | ||
| 153 | - | ||
| 154 | - [[nodiscard]] std::string GetTensorDebugStr(const gert::StorageShape* shape, | ||
| 155 | - const gert::CompileTimeTensorDesc* tensor) | ||
| 156 | - { | ||
| 157 | - if (shape == nullptr || tensor == nullptr) { | ||
| 158 | - return "nil "; | ||
| 159 | - } | ||
| 160 | - std::ostringstream oss; | ||
| 161 | - oss << "(dtype: " << ge::TypeUtils::DataTypeToSerialString(tensor->GetDataType()) << "),"; | ||
| 162 | - oss << "(shape:" << GetShapeDebugStr(shape->GetStorageShape()) << "),"; | ||
| 163 | - oss << "(ori_shape:" << GetShapeDebugStr(shape->GetOriginShape()) << "),"; | ||
| 164 | - oss << "(format: " | ||
| 165 | - << ge::TypeUtils::FormatToSerialString( | ||
| 166 | - static_cast<ge::Format>(ge::GetPrimaryFormat(tensor->GetStorageFormat()))) | ||
| 167 | - << "),"; | ||
| 168 | - oss << "(ori_format: " << ge::TypeUtils::FormatToSerialString(tensor->GetOriginFormat()) << ") "; | ||
| 169 | - return oss.str(); | ||
| 170 | - } | ||
| 171 | - | ||
| 172 | - [[nodiscard]] std::string GetTilingContextDebugStr() | ||
| 173 | - { | ||
| 174 | - std::ostringstream oss; | ||
| 175 | - for (size_t i = 0; i < context_->GetComputeNodeInfo()->GetInputsNum(); ++i) { | ||
| 176 | - oss << "input" << i << ": "; | ||
| 177 | - oss << GetTensorDebugStr(context_->GetInputShape(i), context_->GetInputDesc(i)); | ||
| 178 | - } | ||
| 179 | - | ||
| 180 | - for (size_t i = 0; i < context_->GetComputeNodeInfo()->GetOutputsNum(); ++i) { | ||
| 181 | - oss << "output" << i << ": "; | ||
| 182 | - oss << GetTensorDebugStr(context_->GetOutputShape(i), context_->GetOutputDesc(i)); | ||
| 183 | - } | ||
| 184 | - return oss.str(); | ||
| 185 | - } | ||
| 186 | - | ||
| 187 | - [[nodiscard]] std::string GetTilingDataDebugStr() const | ||
| 188 | - { | ||
| 189 | - auto rawTilingData = context_->GetRawTilingData(); | ||
| 190 | - auto rawTilingDataSize = rawTilingData->GetDataSize(); | ||
| 191 | - auto data = reinterpret_cast<const int32_t*>(rawTilingData->GetData()); | ||
| 192 | - size_t len = rawTilingDataSize / sizeof(int32_t); | ||
| 193 | - std::ostringstream oss; | ||
| 194 | - for (size_t i = 0; i < len; i++) { | ||
| 195 | - oss << data[i] << ", "; | ||
| 196 | - } | ||
| 197 | - return oss.str(); | ||
| 198 | - } | ||
| 199 | - | ||
| 200 | -protected: | ||
| 201 | - gert::TilingContext* context_ = nullptr; | ||
| 202 | - std::unique_ptr<platform_ascendc::PlatformAscendC> ascendcPlatform_{nullptr}; | ||
| 203 | - uint32_t blockDim_{0}; | ||
| 204 | - uint64_t workspaceSize_{0}; | ||
| 205 | - uint64_t tilingKey_{0}; | ||
| 206 | - AiCoreParams aicoreParams_; | ||
| 207 | -}; | ||
| 208 | - | ||
| 209 | -} // namespace optiling | ||
| 210 | - | ||
| 211 | - | ||