已合并
fix: 文件名重复问题解决,改引用公共目录下的文件 #9572
duxinlei创建于 8月31日
fix: 文件名重复问题解决,改引用公共目录下的文件 #9572
已合并
duxinlei创建于 8月31日
共 3 个文件变更+1-381
@@ -28,7 +28,7 @@
28#include "platform/platform_infos_def.h"28#include "platform/platform_infos_def.h"
29#include "register/op_def_registry.h"29#include "register/op_def_registry.h"
30#include "tiling/tiling_api.h"30#include "tiling/tiling_api.h"
31-#include "pool_tiling_templates_registry.h"31+#include "pooling/pool_3d_common/op_host/arch35/pool_tiling_templates_registry.h"
32#include "op_host/tiling_util.h"32#include "op_host/tiling_util.h"
33 33 
34using optiling::PoolTilingRegistry;34using 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-#ifndef POOL_TILING_TEMPLATES_REGISTRY
17-#define POOL_TILING_TEMPLATES_REGISTRY
18- 
19-#include <map>
20-#include <string>
21-#include <memory>
22-#include "exe_graph/runtime/tiling_context.h"
23-#include "tiling_base.h"
24-#include "log/log.h"
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-#ifdef ASCENDC_OP_TEST
62- static PoolTilingRegistry& GetInstance();
63-#else
64- static PoolTilingRegistry& GetInstance()
65- {
66- static PoolTilingRegistry registry_impl_;
67- return registry_impl_;
68- }
69-#endif
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-#define REGISTER_POOL_TILING_TEMPLATE(op_type, class_name, priority) \
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-#define REGISTER_OPS_POOL_TILING_TEMPLATE(op_type, class_name, priority) \
165- GLOBAL_REGISTER_SYMBOL(op_type, class_name, priority, __COUNTER__, __LINE__); \
166- static PoolRegister __attribute__((unused)) \
167- tiling_##op_type##_##class_name##_##priority##_register = PoolRegister(#op_type).tiling<class_name>(priority)
168-} // namespace optiling
169-#endif // CONV_TILING_TEMPLATES_REGISTRY
@@ -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-#ifndef POOL_TILING_BASE_H__
17-#define POOL_TILING_BASE_H__
18- 
19-#include <sstream>
20-#include <exe_graph/runtime/tiling_context.h>
21-#include <graph/utils/type_utils.h>
22-#include "tiling/platform/platform_ascendc.h"
23-#include "log/log.h"
24-#include <algorithm>
25- 
26-#ifdef ASCENDC_OP_TEST
27-#define ASCENDC_EXTERN_C extern "C"
28-#else
29-#define ASCENDC_EXTERN_C
30-#endif
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-#endif // TILING_BASE_H__