已合并
refactor: 面向VV融合的求解器,删除L0/L2求解器 #1684
refactor: 面向VV融合的求解器,删除L0/L2求解器 #1684
已合并
gcw_NLOuEjCz创建于 22 天前
42 个文件变更+75-7108
@@ -38,7 +38,7 @@ using ge::SymbolicUtils;
38 38 
39namespace att {39namespace att {
40using Expr = af::Expression;40using Expr = af::Expression;
41-enum SolverType : uint32_t { L0_TILE = 0, L2_TILE, SEARCH_TILE, ERROR };41+enum SolverType : uint32_t { SEARCH_TILE, ERROR };
42 42 
43enum class HardwareDef {43enum class HardwareDef {
44 GM = 0,44 GM = 0,
@@ -1,892 +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-#ifndef ATT_L0_SOLVER_H_
11-#define ATT_L0_SOLVER_H_
12-#include <sstream>
13-#include <string>
14-#include "util/base_types_printer.h"
15- 
16-namespace att {
17-inline std::string GenVarDef() {
18- std::string strs = "";
19- strs += AddAnotationLine(" L0Var的备选值的个数\n", "");
20- strs += "static const uint32_t candidate_size = 7u;\n";
21- strs += AddAnotationLine(" L0Var的备选值\n", "");
22- strs += "static const uint32_t candidate_value[] = {16u, 32u, 64u, 128u,\n";
23- strs += " 256u, 512u, 1024u};\n";
24- strs += AddAnotationLine(" 表达L0的求解值至少要满足核数的比例,可手动修改\n", "");
25- strs += "static const double CORE_NUM_RATIO = 0.6f;\n";
26- strs += AddAnotationLine(" 表达L0的求解值pad之后的值不允许超过原始值大小的倍数,可手动修改\n", "");
27- strs += "static const uint32_t UPPER_BOUND_RATIO = 2u;\n";
28- strs += AddAnotationLine(" 表达最大L0Var的个数\n", "");
29- strs += "static const uint32_t MAX_L0_VAR_NUM = 3u;\n";
30- strs += "\n";
31- return strs;
32-}
33- 
34-inline std::string GenL0VarDef() {
35- std::string strs = "";
36- strs += AddAnotationLine(" L0相关变量的数据结构\n", "");
37- strs += "struct L0Var {\n";
38- strs += AddAnotationLine(" 最大值,初始化为输入原始轴的大小\n", " ");
39- strs += " uint32_t max_value{0u};\n";
40- strs += AddAnotationLine(" 是否绑多核\n", " ");
41- strs += " bool bind_multicore{false};\n";
42- strs += " bool is_innermost{false};\n";
43- strs += AddAnotationLine(" 对齐值\n", " ");
44- strs += " uint32_t align{0u};\n";
45- strs += AddAnotationLine(" 提示当前L0Var的最佳对齐值,通常源于父轴的对齐值,\n", " ");
46- strs += AddAnotationLine(" 举例,stepm是basem的父轴,stepm的对齐要求是256和basem,那么basem的prompt_align就是256,\n",
47- " ");
48- strs += AddAnotationLine(
49- " 同时约束L0Var的取值必须是256对齐或者是256的因子,因为这样父轴stepm才能既满足256也满足basem对齐\n", " ");
50- strs += " uint32_t prompt_align{0u};\n";
51- strs += AddAnotationLine(" L0变量的索引\n", " ");
52- strs += " uint32_t idx;\n";
53- strs += AddAnotationLine(" L0变量的值\n", " ");
54- strs += " uint32_t value{0u};\n";
55- strs += "};\n";
56- strs += "\n";
57- return strs;
58-}
59- 
60-inline std::string GenL0Input() {
61- std::string strs = "";
62- strs += AddAnotationLine(" 求解器接收的输入\n", "");
63- strs += "struct L0TileInput {\n";
64- strs += AddAnotationLine(" 待求解L0变量的集合\n", " ");
65- strs += " L0Var *l0_vars{nullptr};\n";
66- strs += AddAnotationLine(" 待求解L0变量的数量\n", " ");
67- strs += " uint32_t size;\n";
68- strs += AddAnotationLine(" 核数\n", " ");
69- strs += " uint32_t core_num;\n";
70- strs += "};\n";
71- strs += "\n";
72- return strs;
73-}
74- 
75-inline std::string GenL0VarCmpAnnotation() {
76- std::string strs = "";
77- strs += " * 比较两个 L0Var 类型变量的大小\n";
78- strs += " *\n";
79- strs +=
80- " * 这个函数用来比较两个 L0Var "
81- "类型变量的大小。它遵循特定的比较逻辑:\n";
82- strs +=
83- " * 1. 如果 a 变量绑定到多核且 b 变量没有绑定多核,则 a 被认为大于 "
84- "b,函数返回\n";
85- strs += " * true。\n";
86- strs +=
87- " * 2. 如果 a 变量没有绑定多核且 b 变量绑定多核,则 a 被认为小于 "
88- "b,函数返回\n";
89- strs += " * false。\n";
90- strs +=
91- " * 3. 如果 a 和 b 变量都绑定多核或都未绑定多核,则比较它们的 "
92- "prompt_align\n";
93- strs += " * 属性,prompt_align 属性值大的变量被认为更大。\n";
94- strs += " *\n";
95- strs += " * @param a 第一个要比较的 L0Var 变量。\n";
96- strs += " * @param b 第二个要比较的 L0Var 变量。\n";
97- strs += " * @return 如果 a 大于 b,返回 true;如果 a 小于 b,返回 false。\n";
98- return AddAnotationBlock(strs, "");
99-}
100- 
101-inline std::string GenL0VarCmp() {
102- std::string strs = "";
103- strs += GenL0VarCmpAnnotation();
104- strs += "static bool L0VarCmp(L0Var a, L0Var b) {\n";
105- strs += " if (a.bind_multicore && !b.bind_multicore) {\n";
106- strs += " return true;\n";
107- strs += " }\n";
108- strs += " if (!a.bind_multicore && b.bind_multicore) {\n";
109- strs += " return false;\n";
110- strs += " }\n";
111- strs += " if (a.is_innermost && !b.is_innermost) {\n";
112- strs += " return true;\n";
113- strs += " }\n";
114- strs += " if (!a.is_innermost && b.is_innermost) {\n";
115- strs += " return false;\n";
116- strs += " }\n";
117- strs += " return a.prompt_align > b.prompt_align;\n";
118- strs += "}\n";
119- return strs;
120-}
121- 
122-inline std::string GenL0SolverAnnotaion() {
123- std::string annotations = "";
124- annotations += " * L0 求解器类\n";
125- return AddAnotationBlock(annotations, "");
126-}
127- 
128-inline std::string GenL0SolverGenAnnotaion() {
129- std::string annotations = "";
130- annotations += " * 构造函数\n";
131- annotations += " *\n";
132- annotations += " * @param input 一个 L0TileInput 结构体,包含了 L0 变量相关的信息\n";
133- annotations += " *\n";
134- annotations += " * 这个构造函数初始化了 L0TileSolver 对象\n";
135- return AddAnotationBlock(annotations, " ");
136-}
137- 
138-inline std::string GenL0SolverDegenDef() {
139- std::string strs = "";
140- std::string annotations = "";
141- annotations += " * 析构函数\n";
142- annotations += " *\n";
143- annotations += " * 当 L0TileSolver 对象被销毁时,析构函数被调用\n";
144- annotations += " * 用来释放使用 new 运算符动态分配的内存,确保没有内存泄漏\n";
145- strs += AddAnotationBlock(annotations, " ");
146- strs += " ~L0TileSolver() {\n";
147- strs += " if (sortedvars_ != nullptr) {\n";
148- strs += " delete[] sortedvars_;\n";
149- strs += " }\n";
150- strs += " if (output_ != nullptr) {\n";
151- strs += " delete[] output_;\n";
152- strs += " }\n";
153- strs += " if (input_.l0_vars != nullptr) {\n";
154- strs += " delete[] input_.l0_vars;\n";
155- strs += " }\n";
156- strs += " }\n";
157- return strs;
158-}
159- 
160-inline std::string GenRunAnnotaion() {
161- std::string annotations = "";
162- annotations += " * 运行求解器\n";
163- annotations += " *\n";
164- annotations += " * @return 如果求解成功,返回 true;否则返回 false\n";
165- annotations += " *\n";
166- annotations += " * 这个方法是算法的入口点,调用它会启动求解过程\n";
167- annotations += " * 成功与否取决于 CheckBufferUseValid() 方法的返回值\n";
168- return AddAnotationBlock(annotations, " ");
169-}
170- 
171-inline std::string GenGetOutputAnnotaion() {
172- std::string annotations = "";
173- annotations += " * 获取优化结果\n";
174- annotations += " *\n";
175- annotations += " * @return 指向求解结果数据的指针\n";
176- annotations += " *\n";
177- annotations += " * 如果有求解结果,这个方法将返回一个指向结果数据的指针\n";
178- annotations += " * 结果数据的内存使用完后,应该使用 delete[] 释放内存\n";
179- return AddAnotationBlock(annotations, " ");
180-}
181- 
182-inline std::string GenCheckBufferUseValidAnnotaion() {
183- std::string annotations = "";
184- annotations += " * 检查是否满足buffer约束\n";
185- annotations += " *\n";
186- annotations += " * @return 如果满足,返回 true;否则返回 false\n";
187- annotations += " *\n";
188- annotations += " * 这个纯虚函数要求派生类提供实现,以确保缓冲区使用是有效的\n";
189- annotations += " * 在 L0TileSolver 类中,它是一个抽象方法,需要在子类中实现\n";
190- return AddAnotationBlock(annotations, " ");
191-}
192- 
193-inline std::string GenCheckInputDef() {
194- std::string strs = "";
195- std::string annotations = "";
196- annotations += " * 检查输入数据的完整性和正确性\n";
197- annotations += " *\n";
198- annotations += " * @return 如果输入数据有效,返回 true;否则返回 false\n";
199- annotations += " *\n";
200- annotations += " * 这个私有方法检查输入数据的格式和逻辑,确保它们适用求解算法\n";
201- strs += AddAnotationBlock(annotations, " ");
202- strs += " bool CheckInput();\n";
203- strs += "\n";
204- return strs;
205-}
206- 
207-inline std::string GenInitInputDef() {
208- std::string strs = "";
209- std::string annotations = "";
210- annotations += " * 使用输入数据初始化算法所需的内部数据结构\n";
211- annotations += " *\n";
212- annotations +=
213- " * 这个方法根据输入的 L0TileInput "
214- "结构体中的数据,初始化算法所需的内部数据结构\n";
215- annotations += " * 确保 sortedvars_ 和 output_ 成员变量被正确初始化\n";
216- strs += AddAnotationBlock(annotations, " ");
217- strs += " void InitInput();\n";
218- strs += "\n";
219- return strs;
220-}
221- 
222-inline std::string GenCheckOutputDef() {
223- std::string strs = "";
224- std::string annotations = "";
225- annotations += " * 检查算法运行的结果,确保它们符合预期\n";
226- annotations += " *\n";
227- annotations += " * @return 如果输出数据有效,返回 true;否则返回 false\n";
228- annotations += " *\n";
229- annotations += " * 这个方法检查运行算法后得到的结果,确保它们在逻辑上是合理的\n";
230- strs += AddAnotationBlock(annotations, " ");
231- strs += " bool CheckOutput();\n";
232- strs += "\n";
233- return strs;
234-}
235- 
236-inline std::string GenUpdateAlignAnnotaion() {
237- std::string annotations = "";
238- annotations += " * 更新算法执行过程中的对齐设置\n";
239- annotations += " *\n";
240- annotations +=
241- " * "
242- "这个方法根据算法执行过程中的数据更新对齐提示值,确保结果按照预期的方式"
243- "对齐\n";
244- return AddAnotationBlock(annotations, " ");
245-}
246- 
247-inline std::string GenBestAlignAnnotaion() {
248- std::string strs = "";
249- strs += " * 为指定索引的 L0 变量获取最佳对齐值\n";
250- strs += " *\n";
251- strs += " * @param i 想要的变量索引值\n";
252- strs += " * @return 最佳对齐值\n";
253- strs += " *\n";
254- strs += " * 这个方法计算并返回给定索引值的 L0 变量的最佳对齐值,\n";
255- strs += " * 确保变量以最恰当的方式对齐,从而提高效率或者减少资源浪费\n";
256- return AddAnotationBlock(strs, " ");
257-}
258- 
259-inline std::string GenIterativeRunAnnotaion() {
260- std::string strs = "";
261- strs += " * 为 L0 变量找到最优值进行迭代运行\n";
262- strs += " *\n";
263- strs += " * @param loop_id 当前循环索引,表示正在处理的 L0 变量的位置\n";
264- strs += " * @param best_var_value 一个指针,指向用于存储每个 L0\n";
265- strs += " * 变量迄今为止找到的最佳值的数组\n";
266- strs += " *\n";
267- strs += " * 这个函数使用递归方法来遍历 L0\n";
268- strs +=
269- " * "
270- "变量的所有可能值。对于每个值,它检查是否满足约束条件,如小于上界且满足"
271- "对齐要求。如果满足这些条件,它将继续下一个循环或者递归调用自身来处理下"
272- "一个\n";
273- strs += " * L0 变量。如果是最后一个 L0\n";
274- strs +=
275- " * "
276- "变量,它将检查当前组合是否满足优化条件,如核心数量和数据处理量。如果满"
277- "足条件,它将当前组合存储为最优解。\n";
278- strs += " *\n";
279- strs +=
280- " * 请注意,这个函数没有返回值,而是将最优解存储在传入的 "
281- "best_var_value\n";
282- strs += " * 数组中。\n";
283- return AddAnotationBlock(strs, " ");
284-}
285- 
286-inline std::string GenMaxCoreNumAnnotaion() {
287- std::string strs = "";
288- strs += " * 根据 L0 变量信息和总核心数计算可以分配的最大核心数\n";
289- strs += " *\n";
290- strs += " * @param l0_vars 一个指向 L0Var 结构体数组的指针\n";
291- strs += " * @param core_num 可用于分配的总核心数\n";
292- strs += " * @return 可以分配的最大核心数\n";
293- strs += " *\n";
294- strs +=
295- " * 这个函数计算在给定 L0 "
296- "变量信息和总核心数的情况下,可以分配的最大核心数。\n";
297- strs += " * 它遍历输入的 L0Var\n";
298- strs +=
299- " * "
300- "结构体数组,对于每个变量,根据其是否绑定多核心以及最大、当前和提示对齐"
301- "值计算所需的块数。\n";
302- strs += " * 通过将所有变量的块数相乘,得到总块数。\n";
303- strs +=
304- " * "
305- "最大核心数是总块数和总核心数中的最小值,以确保核心数不会超过可用资源。"
306- "\n";
307- strs += " *\n";
308- strs += " * 返回值表示可以分配给 L0\n";
309- strs += " * 变量的最大核心数,这对于在多核心系统中进行资源分配是有用的。\n";
310- return AddAnotationBlock(strs, " ");
311-}
312- 
313-inline std::string GenGetMacUseAnnotaion() {
314- std::string strs = "";
315- strs += " * 计算所有 L0 变量值的乘积,作为mac计算量的度量\n";
316- strs += " *\n";
317- strs += " * @return mac计算量\n";
318- strs += " *\n";
319- strs += " * 这个函数计算所有 L0 变量值的乘积,结果是一个数字。\n";
320- strs += " * 这个数字可以作为数据处理量的度量,例如在评估算法性能时。\n";
321- strs +=
322- " * 通过不断更新 usage 变量,乘法操作确保了所有 L0 "
323- "变量的影响都被计入。\n";
324- strs +=
325- " * 最终 usage 变量中的值就是所有 L0 "
326- "变量值的乘积,代表了整体的数据处理量。\n";
327- strs +=
328- " * "
329- "返回值可以帮助了解算法处理的数据量,从而对算法的效率和扩展性有更直观的"
330- "认识。\n";
331- return AddAnotationBlock(strs, " ");
332-}
333- 
334-inline std::string GenSortedvarsAnnotaion() {
335- std::string strs = "";
336- strs += " * 用于排序的 L0Var 对象数组\n";
337- return AddAnotationBlock(strs, " ");
338-}
339- 
340-inline std::string GenMaxcoreAnnotaion() {
341- std::string strs = "";
342- strs += " * 最大核心数\n";
343- return AddAnotationBlock(strs, " ");
344-}
345- 
346-inline std::string GenMacuseAnnotaion() {
347- std::string strs = "";
348- strs += " * 最大 MAC 使用量\n";
349- return AddAnotationBlock(strs, " ");
350-}
351- 
352-inline std::string GenL0TileSolver() {
353- std::string strs = "";
354- strs += GenL0SolverAnnotaion();
355- strs += "class L0TileSolver {\n";
356- strs += "public:\n";
357- strs += GenL0SolverGenAnnotaion();
358- strs += " explicit L0TileSolver(L0TileInput input) : input_(input) {}\n";
359- strs += " L0TileSolver() {};\n";
360- strs += GenL0SolverDegenDef();
361- strs += GenRunAnnotaion();
362- strs += " bool Run();\n";
363- strs += GenGetOutputAnnotaion();
364- strs += " uint32_t *GetOutput() { return output_; }\n";
365- strs += "\n";
366- strs += "protected:\n";
367- strs += GenCheckBufferUseValidAnnotaion();
368- strs += " virtual bool CheckBufferUseValid() = 0;\n";
369- strs += " L0TileInput input_;\n";
370- strs += " uint32_t *output_{nullptr};\n";
371- strs += "\n";
372- strs += "private:\n";
373- strs += GenCheckInputDef();
374- strs += GenInitInputDef();
375- strs += GenCheckOutputDef();
376- strs += GenUpdateAlignAnnotaion();
377- strs += " void UpdateAlign();\n";
378- strs += "\n";
379- strs += GenBestAlignAnnotaion();
380- strs += " uint32_t GetBestAlign(uint32_t i) const;\n";
381- strs += "\n";
382- strs += GenIterativeRunAnnotaion();
383- strs += " void IterativeRun(uint32_t loop_id, uint32_t *best_var_value);\n";
384- strs += "\n";
385- strs += GenMaxCoreNumAnnotaion();
386- strs +=
387- " int32_t MaxCoreNum(const L0Var *l0_vars, const uint32_t "
388- "&core_num);\n";
389- strs += "\n";
390- strs += GenGetMacUseAnnotaion();
391- strs += " uint32_t GetMacUse() const;\n";
392- strs += GenSortedvarsAnnotaion();
393- strs += " L0Var *sortedvars_{nullptr};\n";
394- strs += GenMaxcoreAnnotaion();
395- strs += " int64_t max_corenum_{-1};\n";
396- strs += GenMacuseAnnotaion();
397- strs += " int32_t max_macuse_{-1};\n";
398- strs += "};\n";
399- strs += "\n";
400- return strs;
401-}
402- 
403-inline std::string GenGetBestAlignFuncAnnotation() {
404- std::string strs = "";
405- strs += " * 获取给定索引的 L0 变量的最佳对齐值\n";
406- strs += " *\n";
407- strs += " * @param i 想要的变量索引值\n";
408- strs += " * @return 最佳对齐值\n";
409- strs += " *\n";
410- strs += " * 这个方法为给定索引的 L0\n";
411- strs +=
412- " * "
413- "变量计算最佳对齐值。它考虑到变量的最大、当前和提示对齐值,以确保数据存"
414- "储和访问的效率。\n";
415- strs +=
416- " * "
417- "根据变量的原始值(ori_"
418- "value),它首先确定最小和最大对齐值的范围。然后,通过在这个范围内以二"
419- "的幂次方递增,它找到最大的满足条件的值。\n";
420- strs +=
421- " * "
422- "如果没有找到这样的值,它将返回最小对齐值。如果在范围内找到了一个值,它"
423- "将返回这个值的二分之一,作为最佳对齐值。\n";
424- strs += " * 这个最佳对齐值可以用于确保数据以最有效的方式存储\n";
425- return AddAnotationBlock(strs, "");
426-}
427- 
428-inline std::string GenGetBestAlignFunc() {
429- std::string strs = "";
430- strs += GenGetBestAlignFuncAnnotation();
431- strs += "uint32_t L0TileSolver::GetBestAlign(uint32_t i) const {\n";
432- strs += " uint32_t ori_value = input_.l0_vars[i].max_value;\n";
433- strs += " uint32_t min_align = input_.l0_vars[i].align;\n";
434- strs += " uint32_t max_align = input_.l0_vars[i].prompt_align;\n";
435- strs += " uint32_t ori_align = min_align;\n";
436- strs += " uint32_t max_value = std::min(ori_value, max_align);\n";
437- strs += " while (ori_align <= max_value) {\n";
438- strs += " ori_align = ori_align << 1;\n";
439- strs += " }\n";
440- strs += " if (ori_align == min_align) {\n";
441- strs += " return min_align;\n";
442- strs += " }\n";
443- strs += " return std::max(1u, ori_align >> 1);\n";
444- strs += "}\n";
445- strs += "\n";
446- return strs;
447-}
448- 
449-inline std::string GenMaxCoreNumFuncAnnotation() {
450- std::string strs = "";
451- strs += " * 根据给定的 L0 变量信息计算可以分配的最大核心数\n";
452- strs += " *\n";
453- strs += " * @param l0_vars 指向 L0Var 结构数组的指针\n";
454- strs += " * @param core_num 总核心数\n";
455- strs += " * @return 可以分配的最大核心数\n";
456- strs += " *\n";
457- strs +=
458- " * 这个函数遍历 L0Var 结构数组,根据每个变量的 bind_multicore "
459- "属性以及\n";
460- strs += " * max_value、value 和 prompt_align\n";
461- strs +=
462- " * "
463- "的值来计算每个变量所需的块数。对于绑定多核心的变量,块数计算方式为:("
464- "max_value\n";
465- strs += " * + max(value,prompt_align)-1)/\n";
466- strs += " * max(value,prompt_align)。对于未绑定多核心的变量,块数为 1。\n";
467- strs += " * 总块数通过将所有变量的块数相乘得到。之后,通过比较总块数和\n";
468- strs +=
469- " * "
470- "core_"
471- "num,返回两者中的最小值,作为可以分配的最大核心数。如果总块数超过了\n";
472- strs += " * core_num,那么系统的核心数将成为瓶颈,因此需要将 core_num\n";
473- strs += " * 设置为最大核心数。如果总块数小于等于\n";
474- strs += " * core_num,那么总块数就是可以分配的最大核心数。\n";
475- return AddAnotationBlock(strs, "");
476-}
477- 
478-inline std::string GenMaxCoreNumFunc() {
479- std::string strs = "";
480- strs += GenMaxCoreNumFuncAnnotation();
481- strs += "int32_t L0TileSolver::MaxCoreNum(const L0Var *l0_vars,\n";
482- strs += " const uint32_t &core_num) {\n";
483- strs += " uint32_t total_block_size = 1u;\n";
484- strs += " for (uint32_t i = 0u; i < input_.size; i++) {\n";
485- strs += " auto var = l0_vars[i];\n";
486- strs += " uint32_t block_num =\n";
487- strs += " var.bind_multicore\n";
488- strs +=
489- " ? ((var.max_value + std::max(var.value, var.prompt_align) "
490- "- 1)) /\n";
491- strs += " std::max(var.value, var.prompt_align)\n";
492- strs += " : 1;\n";
493- strs += " total_block_size *= block_num;\n";
494- strs += " }\n";
495- strs += " int64_t max_core_num =\n";
496- strs += " total_block_size > core_num ? core_num : total_block_size;\n";
497- strs += " return max_core_num;\n";
498- strs += "}\n";
499- strs += "\n";
500- return strs;
501-}
502- 
503-inline std::string GenGetMacUseFuncAnnotation() {
504- std::string strs = "";
505- strs += " * 计算所有 L0 变量值的乘积,作为数据处理量的度量\n";
506- strs += " *\n";
507- strs += " * @return 数据处理量\n";
508- strs += " *\n";
509- strs +=
510- " * 这个函数遍历 L0TileInput 结构体中的所有 L0Var 对象,计算它们的 "
511- "value\n";
512- strs += " * 属性的乘积。这个乘积代表了所有 L0\n";
513- strs +=
514- " * 变量值的联合效应,或者说数据处理量的一个度量。 通过不断更新 "
515- "usage\n";
516- strs += " * 变量,乘法操作确保了所有 L0 变量的贡献都被包含在内。最终 usage\n";
517- strs += " * 变量中的值就是所有 L0 变量值的乘积。\n";
518- strs +=
519- " * "
520- "返回值可以帮助评估算法在处理给定输入数据时的效率,以及比较不同算法或优"
521- "化策略的数据处理量。\n";
522- return AddAnotationBlock(strs, "");
523-}
524- 
525-inline std::string GenGetMacUseFunc() {
526- std::string strs = "";
527- strs += GenGetMacUseFuncAnnotation();
528- strs += "uint32_t L0TileSolver::GetMacUse() const {\n";
529- strs += " uint32_t usage = 1u;\n";
530- strs += " for (uint32_t j = 0; j < input_.size; j++) {\n";
531- strs += " usage *= input_.l0_vars[j].value;\n";
532- strs += " }\n";
533- strs += " return usage;\n";
534- strs += "}\n";
535- strs += "\n";
536- return strs;
537-}
538- 
539-inline std::string GenIterativeRunFuncAnnotation() {
540- std::string strs = "";
541- strs += " * 为 L0 变量找到最优值进行迭代运行\n";
542- strs += " *\n";
543- strs += " * @param loop_id 当前循环索引,表示正在处理的 L0 变量的位置\n";
544- strs +=
545- " * @param best_var_value 一个指针,指向用于存储每个 L0 "
546- "变量迄今为止找到的最佳值的数组\n";
547- strs += " *\n";
548- strs +=
549- " * 这个函数使用递归方法来遍历 L0 "
550- "变量的所有可能值。对于每个值,它检查是否满足约束条件,如小于上界且满足"
551- "对齐要求。\n";
552- strs +=
553- " * 如果满足这些条件,它将继续下一个循环或者递归调用自身来处理下一个L0 "
554- "变量。\n";
555- strs +=
556- " * 如果是最后一个 L0 "
557- "变量,它将检查当前组合是否满足优化条件,如核心数量和数据处理量。如果满"
558- "足条件,它将当前组合存储为最优解。\n";
559- strs += " *\n";
560- strs +=
561- " * 请注意,这个函数没有返回值,而是将最优解存储在传入的 "
562- "best_var_value 数组中。\n";
563- return AddAnotationBlock(strs, "");
564-}
565- 
566-inline std::string GenIterativeRunFunc() {
567- std::string strs = "";
568- strs += GenIterativeRunFuncAnnotation();
569- strs +=
570- "void L0TileSolver::IterativeRun(uint32_t loop_id, uint32_t "
571- "*best_var_value) {\n";
572- strs += " for (uint32_t i = 0u; i < candidate_size; i++) {\n";
573- strs += " uint32_t candi_value = candidate_value[i];\n";
574- strs += " const auto &l0_tile = sortedvars_[loop_id];\n";
575- strs += " // L0Var的上限\n";
576- strs += " uint32_t upper_bound = l0_tile.max_value * UPPER_BOUND_RATIO;\n";
577- strs += " if (candi_value >= upper_bound) {\n";
578- strs += " continue;\n";
579- strs += " }\n";
580- strs += " // 必须满足prompt_align对齐或者是prompt_align的因子\n";
581- strs += " if ((candi_value % l0_tile.prompt_align != 0) &&\n";
582- strs += " (l0_tile.prompt_align % candi_value != 0)) {\n";
583- strs += " continue;\n";
584- strs += " }\n";
585- strs += " auto idx = l0_tile.idx;\n";
586- strs += " input_.l0_vars[idx].value = candi_value;\n";
587- strs += " // 终止条件为遍历到最后一个变量\n";
588- strs += " if (loop_id == input_.size - 1) {\n";
589- strs += " if (!CheckBufferUseValid()) {\n";
590- strs += " break;\n";
591- strs += " }\n";
592- strs += " int32_t usage = GetMacUse();\n";
593- strs +=
594- " int32_t core_num = MaxCoreNum(input_.l0_vars, "
595- "input_.core_num);\n";
596- strs +=
597- " // "
598- "最大核数如果满足核数*系数(默认0."
599- "6),则比较mac利用率即可,否则需要比较核数的使用和mac利用率\n";
600- strs += " if (((core_num >= max_corenum_) ||\n";
601- strs += " (core_num >=\n";
602- strs +=
603- " static_cast<int32_t>(input_.core_num * CORE_NUM_RATIO))) "
604- "&&\n";
605- strs += " (usage >= max_macuse_)) {\n";
606- strs += " max_corenum_ = core_num;\n";
607- strs += " max_macuse_ = usage;\n";
608- strs += " for (uint32_t k = 0u; k < input_.size; k++) {\n";
609- strs += " best_var_value[k] = input_.l0_vars[k].value;\n";
610- strs += " }\n";
611- strs += " }\n";
612- strs += " } else {\n";
613- strs += " IterativeRun(loop_id + 1, best_var_value);\n";
614- strs += " }\n";
615- strs += " }\n";
616- strs += "}\n";
617- strs += "\n";
618- return strs;
619-}
620- 
621-inline std::string GenUpdateAlignFuncAnnotation() {
622- std::string strs = "";
623- strs += " * 更新 L0Var 对象的对齐值\n";
624- strs += " *\n";
625- strs += " * 这个函数用于更新 L0Var 对象的 prompt_align\n";
626- strs +=
627- " * 值,以确保它们在内存中按照最优方式对齐。它遍历 input_ 对象中的 "
628- "l0_vars\n";
629- strs += " * 数组,为每个 L0Var 对象计算并设置最佳的对齐值。\n";
630- strs += " *\n";
631- strs += " * @param 无\n";
632- strs += " * @return 无\n";
633- return AddAnotationBlock(strs, "");
634-}
635- 
636-inline std::string GenUpdateAlignFunc() {
637- std::string strs = "";
638- strs += GenUpdateAlignFuncAnnotation();
639- strs += "void L0TileSolver::UpdateAlign() {\n";
640- strs += " for (uint32_t i = 0u; i < input_.size; i++) {\n";
641- strs += " uint32_t best_align = GetBestAlign(i);\n";
642- strs += " input_.l0_vars[i].prompt_align = best_align;\n";
643- strs += " }\n";
644- strs += "}\n";
645- strs += "\n";
646- return strs;
647-}
648- 
649-inline std::string GenCheckInputFuncAnnotation() {
650- std::string strs = "";
651- strs += " * 检查输入数据的有效性\n";
652- strs += " *\n";
653- strs +=
654- " * 这个函数用来检查 L0TileSolver "
655- "类的输入数据是否有效。它验证以下几个方面:\n";
656- strs += " * - 基础变量指针(l0_vars)是否为空。\n";
657- strs += " * - 输入数据的大小(size)是否为0,表示没有 L0 参数需要求解。\n";
658- strs +=
659- " * - "
660- "输入数据的大小(size)是否超过最大支持的参数数量(MAX_L0_VAR_"
661- "NUM)。\n";
662- strs += " * - 核心数量(core_num)是否为0。\n";
663- strs +=
664- " * - 对于输入数据中的每个 L0Var 对象(通过索引 i "
665- "访问),它检查几个属性:\n";
666- strs += " * - max_value、align 和 prompt_align 是否都不等于0。\n";
667- strs += " * - align 是否不大于 prompt_align。\n";
668- strs += " *\n";
669- strs +=
670- " * 如果以上任何一个条件不满足,函数将通过 OP_LOG "
671- "宏记录一条错误消息,并返回\n";
672- strs +=
673- " * false,表示输入无效。如果所有条件都满足,函数返回 "
674- "true,表示输入有效。\n";
675- strs += " *\n";
676- strs += " * @return 如果输入数据有效,则返回 true;否则返回 false。\n";
677- return AddAnotationBlock(strs, "");
678-}
679- 
680-inline std::string GenCheckInputFunc() {
681- std::string strs = "";
682- strs += GenCheckInputFuncAnnotation();
683- strs += "bool L0TileSolver::CheckInput() {\n";
684- strs += " if (input_.l0_vars == nullptr) {\n";
685- strs += " OP_LOGW(OP_NAME, \"Input basevar is null\");\n";
686- strs += " return false;\n";
687- strs += " }\n";
688- strs += " if (input_.size == 0u) {\n";
689- strs += " OP_LOGW(OP_NAME, \"Size is 0, no l0 arg to be solved\");\n";
690- strs += " return false;\n";
691- strs += " }\n";
692- strs += " if (input_.size > MAX_L0_VAR_NUM) {\n";
693- strs += " OP_LOGW(OP_NAME, \"L0 solver does not support more than 3 input args\");\n";
694- strs += " return false;\n";
695- strs += " }\n";
696- strs += " if (input_.core_num == 0) {\n";
697- strs += " OP_LOGW(OP_NAME, \"Corenum is 0\");\n";
698- strs += " return false;\n";
699- strs += " }\n";
700- strs += " for (uint32_t i = 0u; i < input_.size; i++) {\n";
701- strs += " auto var = input_.l0_vars[i];\n";
702- strs +=
703- " if ((var.max_value == 0) || (var.align == 0) || (var.prompt_align "
704- "== 0)) {\n";
705- strs += " OP_LOGW(OP_NAME, \"Input [%u] exists 0\", i);\n";
706- strs += " return false;\n";
707- strs += " }\n";
708- strs += " if (var.align > var.prompt_align) {\n";
709- strs += " OP_LOGW(OP_NAME, \"Input [%u] align is larger than prompt align\", i);\n";
710- strs += " return false;\n";
711- strs += " }\n";
712- strs += " }\n";
713- strs += " return true;\n";
714- strs += "}\n";
715- strs += "\n";
716- return strs;
717-}
718- 
719-inline std::string GenInitInputFuncAnnotation() {
720- std::string strs = "";
721- strs += " * 初始化 L0Var 对象数组\n";
722- strs += " *\n";
723- strs +=
724- " * 这个函数用于初始化 L0TileSolver 类的 input_ 对象中的 l0_vars "
725- "数组。它遍历\n";
726- strs += " * l0_vars 数组中的每一个元素,对于每个元素,执行以下操作:\n";
727- strs += " * 1. 通过访问索引 i 对应的 L0Var 对象的引用 var,重置其 max_value\n";
728- strs +=
729- " * 属性。具体重置方式是,先将 max_value 增加 align 属性值减 1,再除以 "
730- "align\n";
731- strs +=
732- " * 属性值,最后乘以 align 属性值。这样做的目的可能是为了确保 "
733- "max_value 是 align\n";
734- strs += " * 的整数倍。\n";
735- strs +=
736- " * 2. 将当前循环的索引值 i 设置为 var 的 idx "
737- "属性。这可能是为了标记每个 L0Var\n";
738- strs += " * 对象在数组中的位置,以便后续处理。\n";
739- strs += " *\n";
740- strs += " * @param 无\n";
741- strs += " * @return 无\n";
742- return AddAnotationBlock(strs, "");
743-}
744- 
745-inline std::string GenInitInputFunc() {
746- std::string strs = "";
747- strs += GenInitInputFuncAnnotation();
748- strs += "void L0TileSolver::InitInput() {\n";
749- strs += " for (uint32_t i = 0u; i < input_.size; i++) {\n";
750- strs += " auto &var = input_.l0_vars[i];\n";
751- strs +=
752- " var.max_value = (var.max_value + var.align - 1) / var.align * "
753- "var.align;\n";
754- strs += " var.idx = i;\n";
755- strs += " }\n";
756- strs += "}\n";
757- strs += "\n";
758- return strs;
759-}
760- 
761-inline std::string GenCheckOutputFuncAnnotation() {
762- std::string strs = "";
763- strs += " * 检查输出数据的有效性\n";
764- strs += " *\n";
765- strs +=
766- " * 这个函数用于检查 L0TileSolver 类的输出数据是否有效。它首先检查 "
767- "output_\n";
768- strs +=
769- " * 指针是否为空。如果 output_ 指针为空,通过 OP_LOG "
770- "宏记录一条错误消息,并返回\n";
771- strs += " * false,表示输出无效。\n";
772- strs += " *\n";
773- strs +=
774- " * 接着,函数遍历 output_ "
775- "数组中的每个元素。对于每个元素,它检查其值是否为\n";
776- strs += " * 0。如果发现任何一个元素的值为 0,函数会通过 OP_LOG\n";
777- strs +=
778- " * 宏记录相应的错误消息,并返回 "
779- "false,表示输出数据中存在无效的元素。\n";
780- strs += " *\n";
781- strs += " * 如果输出数据有效,即 output_ 指针不为空且 output_ 数组中没有 0\n";
782- strs += " * 值元素,函数返回 true。\n";
783- strs += " *\n";
784- strs += " * @return 如果输出数据有效,则返回 true;否则返回 false。\n";
785- return AddAnotationBlock(strs, "");
786-}
787- 
788-inline std::string GenCheckOutputFunc() {
789- std::string strs = "";
790- strs += GenCheckOutputFuncAnnotation();
791- strs += "bool L0TileSolver::CheckOutput() {\n";
792- strs += " if (output_ == nullptr) {\n";
793- strs += " OP_LOGW(OP_NAME, \"Output is null\");\n";
794- strs += " return false;\n";
795- strs += " }\n";
796- strs += " for (uint32_t i = 0u; i < input_.size; i++) {\n";
797- strs += " if (output_[i] == 0u) {\n";
798- strs += " OP_LOGW(OP_NAME, \"Output [%u] is 0\", i);\n";
799- strs += " return false;\n";
800- strs += " }\n";
801- strs += " }\n";
802- strs += " return true;\n";
803- strs += "}\n";
804- strs += "\n";
805- return strs;
806-}
807- 
808-inline std::string GenRunFuncAnnotation() {
809- std::string strs = "";
810- strs += " * 执行 L0TileSolver 类的主要流程\n";
811- strs += " * @return 如果所有操作成功并且输出有效,则返回 true;否则返回false\n";
812- return AddAnotationBlock(strs, "");
813-}
814- 
815-inline std::string GenRunFunc() {
816- std::string strs = "";
817- strs += GenRunFuncAnnotation();
818- strs += "bool L0TileSolver::Run() {\n";
819- strs += " // 检查输入数据的有效性\n";
820- strs += " if (!CheckInput()) {\n";
821- strs += " // 如果输入检查失败,则记录一条错误日志,并返回 false\n";
822- strs += " OP_LOGW(OP_NAME, \"Check input failed\");\n";
823- strs += " return false;\n";
824- strs += " }\n";
825- strs += " // 初始化输入数据\n";
826- strs += " InitInput();\n";
827- strs += " // 更新 L0Var 对象的对齐值\n";
828- strs += " UpdateAlign();\n";
829- strs += " output_ = new (std::nothrow) uint32_t[input_.size]();\n";
830- strs += " bool is_fast_mode = true;\n";
831- strs += " for (uint32_t i=0u; i < input_.size; i++) {\n";
832- strs += " auto &var = input_.l0_vars[i];\n";
833- strs += " uint32_t upper_bound = var.max_value * UPPER_BOUND_RATIO;\n";
834- strs += " if ((var.value == 0u) || (var.value > upper_bound)) {\n";
835- strs += " is_fast_mode = false;\n";
836- strs += " break;\n";
837- strs += " }\n";
838- strs += " }\n";
839- strs += " if (is_fast_mode && CheckBufferUseValid()) {\n";
840- strs += " for (uint32_t k=0u; k < input_.size; k++) {\n";
841- strs += " output_[k] = input_.l0_vars[k].value;\n";
842- strs += " }\n";
843- strs += " } else {\n";
844- strs += " // 为排序后的变量申请内存,并初始化为 0\n";
845- strs += " sortedvars_ = new (std::nothrow) L0Var[input_.size];\n";
846- strs += " // 将输入数据复制到新的内存中\n";
847- strs += " std::copy(input_.l0_vars, input_.l0_vars + input_.size, sortedvars_);\n";
848- strs += " // 根据比较函数对变量进行排序\n";
849- strs += " std::sort(sortedvars_, sortedvars_ + input_.size, L0VarCmp);\n";
850- strs += " // 调用 IterativeRun 函数,传递参数 0 和 output_ 数组的指针\n";
851- strs += " IterativeRun(0u, output_);\n";
852- strs += " }\n";
853- strs += " // 检查输出数据的有效性\n";
854- strs += " if (!CheckOutput()) {\n";
855- strs += " // 如果输出检查失败,则记录一条错误日志,并返回 false\n";
856- strs += " OP_LOGW(OP_NAME, \"Check output failed\");\n";
857- strs += " return false;\n";
858- strs += " }\n";
859- strs += " // 如果所有操作都成功,返回 true\n";
860- strs += " return true;\n";
861- strs += "}\n";
862- return strs;
863-}
864- 
865-inline std::string GetL0SolverHead() {
866- std::string strs = "";
867- strs += GenVarDef();
868- strs += GenL0VarDef();
869- strs += GenL0Input();
870- strs += GenL0VarCmp();
871- strs += GenL0TileSolver();
872- return strs;
873-}
874- 
875-inline std::string GetL0SolverFunc() {
876- std::string strs = "";
877- strs += GenGetBestAlignFunc();
878- strs += GenMaxCoreNumFunc();
879- strs += GenGetMacUseFunc();
880- strs += GenIterativeRunFunc();
881- strs += GenUpdateAlignFunc();
882- strs += GenCheckInputFunc();
883- strs += GenInitInputFunc();
884- strs += GenCheckOutputFunc();
885- strs += GenRunFunc();
886- return strs;
887-}
888- 
889-inline const std::string L0_SOLVER_CODE_HEAD = GetL0SolverHead();
890-inline const std::string L0_SOLVER_CODE_FUNC = GetL0SolverFunc();
891-} // namespace att
892-#endif
@@ -1,390 +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-#ifndef ATT_L2_SOLVER_H_
11-#define ATT_L2_SOLVER_H_
12-#include <sstream>
13-#include <string>
14- 
15-namespace att {
16-inline std::string GenCeilDivision() {
17- std::string strs;
18- strs += "// L2的占用经验值大小\n";
19- strs += "const uint32_t EMPIRIC_L2_SIZE = 128 * 1024 * 1024u;\n";
20- strs += "inline uint32_t CeilDivision(uint32_t a, uint32_t b) {\n";
21- strs += " if (b == 0) {\n";
22- strs += " return 0;\n";
23- strs += " }\n";
24- strs += " return uint32_t((a + b - 1) / b);\n";
25- strs += "}\n";
26- strs += "\n";
27- return strs;
28-}
29- 
30-inline std::string GenL2Var() {
31- std::string strs;
32- strs += "// 每个L2变量的信息\n";
33- strs += "struct L2Var {\n";
34- strs += " // 最大值,初始化为原始输入的大小\n";
35- strs += " uint32_t max_value{0};\n";
36- strs += " // 对气值\n";
37- strs += " uint32_t align{0};\n";
38- strs += " // 对应的L0基本块的大小,举例,TileL2M对应的基本块basem\n";
39- strs += " uint32_t base_val{0};\n";
40- strs += " // 当前变量的值\n";
41- strs += " uint32_t value{0};\n";
42- strs += "};\n";
43- strs += "\n";
44- return strs;
45-}
46- 
47-inline std::string L2TileInput() {
48- std::string strs;
49- strs += "// 求解器的输入\n";
50- strs += "struct L2TileInput {\n";
51- strs += " // L2变量的集合\n";
52- strs += " L2Var *l2_vars{nullptr};\n";
53- strs += " // L2变量的个数\n";
54- strs += " uint32_t size{0};\n";
55- strs += " // 核数\n";
56- strs += " uint32_t core_num{0};\n";
57- strs += " // l2的大小,默认为经验值的大小\n";
58- strs += " uint32_t l2_size{0};\n";
59- strs += "};\n";
60- strs += "\n";
61- return strs;
62-}
63- 
64-inline std::string GenL2TileSolverAnnotation() {
65- std::string strs;
66- strs += "// L2求解器的适用范围如下\n";
67- strs += "// 举例如下,一个TileL2M * TileL2N的结果矩阵,每个小方格表示一个basem * basen的基本块\n";
68- strs +=
69- "// TileL2M 和 "
70- "TileL2N大小的结果矩阵存储在L2中,以基本块为粒度分多核,如下图所示,假设有4个核,每个核计算两个基本块\n";
71- strs += "// \n";
72- strs += "// tileL2N basen\n";
73- strs += "// .-------.-------.-------.-------.\n";
74- strs += "// | core0 | core0 | core1 | core1 | ->basem\n";
75- strs += "// tileL2M <- '-------'-------'-------'-------'\n";
76- strs += "// | core2 | core2 | core3 | core3 |\n";
77- strs += "// '-------'-------'-------'-------'\n";
78- return strs;
79-}
80- 
81-inline std::string GenDeconstructL2TileSolver() {
82- std::string strs;
83- strs += " // 析构函数,用于清理堆上分配的内存\n";
84- strs += " ~L2TileSolver() {\n";
85- strs += " // 如果 blocknum_per_tile_ 指针不为空,则释放其指向的内存\n";
86- strs += " if (blocknum_per_tile_!= nullptr) {\n";
87- strs += " delete[] blocknum_per_tile_;\n";
88- strs += " }\n";
89- strs += " // 如果 size_per_tile_ 指针不为空,则释放其指向的内存\n";
90- strs += " if (size_per_tile_!= nullptr) {\n";
91- strs += " delete[] size_per_tile_;\n";
92- strs += " }\n";
93- strs += " // 如果 tilenum_ 指针不为空,则释放其指向的内存\n";
94- strs += " if (tilenum_!= nullptr) {\n";
95- strs += " delete[] tilenum_;\n";
96- strs += " }\n";
97- strs += " // 如果 total_blocknum_ 指针不为空,则释放其指向的内存\n";
98- strs += " if (total_blocknum_!= nullptr) {\n";
99- strs += " delete[] total_blocknum_;\n";
100- strs += " }\n";
101- strs += " if (input_.l2_vars != nullptr) {\n";
102- strs += " delete[] input_.l2_vars;\n";
103- strs += " }\n";
104- strs += " }\n";
105- return strs;
106-}
107- 
108-inline std::string GenL2TileSolver() {
109- std::string strs;
110- strs += GenL2TileSolverAnnotation();
111- strs += "class L2TileSolver {\n";
112- strs += "public:\n";
113- strs += " // 构造函数,接受 L2TileInput 类型的参数 input\n";
114- strs += " explicit L2TileSolver(L2TileInput input) : input_(input) {};\n";
115- strs += " // 无参构造函数\n";
116- strs += " L2TileSolver() {}\n";
117- strs += GenDeconstructL2TileSolver();
118- strs += " // Run() 成员函数,返回布尔值,可能用于指示某个操作的成功或失败\n";
119- strs += " bool Run();\n";
120- strs += " // GetL2Tile() 成员函数,返回 uint32_t 类型的指针\n";
121- strs += " uint32_t *GetL2Tile() { return size_per_tile_; }\n";
122- strs += "\n";
123- strs += "protected:\n";
124- strs += " // 纯虚函数,需要在子类中实现,用于获取 L2 的使用情况\n";
125- strs += " virtual uint64_t GetL2Use() = 0;\n";
126- strs += " virtual uint32_t GetInitVal(uint32_t L2Size) = 0;\n";
127- strs += " // 纯虚函数,需要在子类中实现,用于判断索引 idx 处是否存在冲突\n";
128- strs += " virtual bool IsClash(uint32_t idx) = 0;\n";
129- strs += " // L2TileInput 类型的成员变量,用于存储输入数据\n";
130- strs += " L2TileInput input_;\n";
131- strs += " // 默认初始化值为 1 的 used_corenum_ 成员变量\n";
132- strs += " uint32_t used_corenum_{1};\n";
133- strs += " // 指向 uint32_t 类型的指针 blocknum_per_tile_,初始化为空指针,用于表达每个方向的输入包含多少个基本块\n";
134- strs += " uint32_t *blocknum_per_tile_{nullptr};\n";
135- strs += " // 指向 uint32_t 类型的指针 size_per_tile_,初始化为空指针,用于表达每个方向的输入的大小\n";
136- strs += " uint32_t *size_per_tile_{nullptr};\n";
137- strs += " // 指向 uint32_t 类型的指针 tilenum_,初始化为空指针,用于表达每个方向的输入的L2块的个数\n";
138- strs += " uint32_t *tilenum_{nullptr};\n";
139- strs += " // 指向 uint32_t 类型的指针 total_blocknum_,初始化为空指针,用于表达每个方向输入基本块的总数\n";
140- strs += " uint32_t *total_blocknum_{nullptr};\n";
141- strs += "\n";
142- strs += "private:\n";
143- strs += " // 私有的 CheckInput() 成员函数,返回布尔值,用于检查输入数据的有效性\n";
144- strs += " bool CheckInput();\n";
145- strs += " // 私有的 InitInput() 成员函数,用于初始化输入数据\n";
146- strs += " void InitInput();\n";
147- strs += " // 私有的 CheckSolvable() 成员函数,返回布尔值,用于检查问题是否可解\n";
148- strs += " bool CheckSolvable();\n";
149- strs += " void HandleClash(uint32_t loop_id, uint32_t *ori_val, uint32_t *best_val, uint64_t &max_l2_use);\n";
150- strs += "};\n";
151- strs += "\n";
152- return strs;
153-}
154- 
155-inline std::string GenL2CheckInput() {
156- std::string strs;
157- strs += "/**\n";
158- strs += " * 检查输入参数是否有效。\n";
159- strs += " *\n";
160- strs += " * 该函数用于检查 L2TileSolver 类的输入参数是否有效。\n";
161- strs += " * 首先检查 l2_vars 指针是否为空。如果为空,将记录一条错误消息并返回 false,表示输入无效。\n";
162- strs +=
163- " * 然后检查 size、core_num 和 l2_size 参数是否都不为零。如果其中任何一个为零,将记录一条错误消息并返回 "
164- "false,表示输入无效。\n";
165- strs +=
166- " * 最后遍历 l2_vars 指针数组中的所有 L2Var 结构。对于每个结构,检查 align、base_val 和 max_value "
167- "成员是否都不为零。如果其中任何一个为零,将记录一条错误消息并返回 false,表示输入无效。\n";
168- strs += " * 如果所有检查都通过,该函数将返回 true,表示输入有效。\n";
169- strs += " *\n";
170- strs += " * @Return bool, 表示输入参数是否有效\n";
171- strs += " */\n";
172- strs += "bool L2TileSolver::CheckInput() {\n";
173- strs += " if (input_.l2_vars == nullptr) {\n";
174- strs += " OP_LOGW(OP_NAME, \"Input l2var is null\");\n";
175- strs += " return false;\n";
176- strs += " }\n";
177- strs += " if (input_.size == 0 || input_.core_num == 0 || input_.l2_size == 0) {\n";
178- strs += " OP_LOGW(OP_NAME, \"Exist input 0, please check size, core_num and l2_size\");\n";
179- strs += " return false;\n";
180- strs += " }\n";
181- strs += " for (uint32_t i = 0; i < input_.size; i++) {\n";
182- strs += " auto var = input_.l2_vars[i];\n";
183- strs += " if (var.align == 0 || var.base_val == 0 || var.max_value == 0) {\n";
184- strs += " OP_LOGW(OP_NAME, \"Input [%u] exists 0\", i);\n";
185- strs += " return false;\n";
186- strs += " }\n";
187- strs += " }\n";
188- strs += " return true;\n";
189- strs += "}\n";
190- strs += "\n";
191- return strs;
192-}
193- 
194-inline std::string GenL2CheckSolvable() {
195- std::string strs;
196- strs += "/**\n";
197- strs += " * 检查问题是否可以解决。\n";
198- strs += " *\n";
199- strs += " * 这个函数检查当前的问题是否可以根据输入参数的设置来解决。\n";
200- strs += " * 首先,它初始化每个变量的值为 align 参数的值。\n";
201- strs += " * 然后,它调用 GetL2Use() 函数来获得所需的 L2 使用量。\n";
202- strs +=
203- " * 如果所需的 L2 使用量超过了可用的缓存大小(input_.l2_size),它将记录一条警告消息并返回 "
204- "false,表示没有解决方案。\n";
205- strs += " * 如果所需的 L2 使用量小于或等于可用的缓存大小,函数将返回 true,表示问题可以被解决。\n";
206- strs += " *\n";
207- strs += " * @Return true 如果问题可以解决,false 否则。\n";
208- strs += " */\n";
209- strs += "bool L2TileSolver::CheckSolvable() {\n";
210- strs += " for (uint32_t i = 0; i < input_.size; i++) {\n";
211- strs += " auto &var = input_.l2_vars[i];\n";
212- strs += " var.value = var.align;\n";
213- strs += " }\n";
214- strs += " if (GetL2Use() > input_.l2_size) {\n";
215- strs += " OP_LOGW(OP_NAME, \"No solution, l2 size is too small\");\n";
216- strs += " return false;\n";
217- strs += " }\n";
218- strs += " return true;\n";
219- strs += "}\n";
220- strs += "\n";
221- return strs;
222-}
223- 
224-inline std::string GetL2InitInput() {
225- std::string strs;
226- strs += "/**\n";
227- strs += " * 初始化输入数据。\n";
228- strs += " *\n";
229- strs += " * 这个函数负责初始化 L2TileSolver 对象的输入数据。\n";
230- strs +=
231- " * "
232- "它首先为输入数据结构中的每个变量计算最大值,将其向上取整为对齐值的最近倍数。这确保了每个变量的最大值是其对齐值的"
233- "倍数。\n";
234- strs += " * 然后,它找出所有变量中的最大最大值,并将每个变量的值初始化为这个最大最大值。\n";
235- strs += " *\n";
236- strs += " * 初始化过程对于准备输入数据以进行后续处理步骤(如优化或分析)至关重要。\n";
237- strs += " * 确保最大值是对齐值的倍数,可以简化数据的处理,并可能是依赖于这种属性的算法或操作所必需的。\n";
238- strs += " */\n";
239- strs += "void L2TileSolver::InitInput() {\n";
240- strs += " uint32_t init_value = 0;\n";
241- strs += " for (uint32_t i = 0; i < input_.size; i++) {\n";
242- strs += " auto &var = input_.l2_vars[i];\n";
243- strs += " var.max_value = CeilDivision(var.max_value, var.align) * var.align;\n";
244- strs += " }\n";
245- strs += "}\n";
246- strs += "\n";
247- return strs;
248-}
249- 
250-inline std::string GetHandleClash() {
251- std::string strs;
252- strs +=
253- "void L2TileSolver::HandleClash(uint32_t loop_id, uint32_t *ori_val, uint32_t *best_val, uint64_t &max_l2_use) "
254- "{\n";
255- strs += " auto max_blocknum = ori_val[loop_id];\n";
256- strs += " auto &var = input_.l2_vars[loop_id];\n";
257- strs += " for (uint32_t i = max_blocknum; i >= 1u; i--) {\n";
258- strs += " blocknum_per_tile_[loop_id] = i;\n";
259- strs += " size_per_tile_[loop_id] = blocknum_per_tile_[loop_id] * var.base_val;\n";
260- strs += " tilenum_[loop_id] = CeilDivision(var.max_value, size_per_tile_[loop_id]);\n";
261- strs += " var.value = size_per_tile_[loop_id];\n";
262- strs += " if (loop_id == input_.size-1) {\n";
263- strs += " uint32_t tmp_corenum = 1;\n";
264- strs += " for (uint32_t j = 0; j < input_.size; j++) {\n";
265- strs += " tmp_corenum *= blocknum_per_tile_[j];\n";
266- strs += " }\n";
267- strs += " used_corenum_ = std::min(input_.core_num, tmp_corenum);\n";
268- strs += " bool solved=true;\n";
269- strs += " for (uint32_t k = 0; k < input_.size; k++) {\n";
270- strs += " if (IsClash(k)) {\n";
271- strs += " solved=false;\n";
272- strs += " }\n";
273- strs += " }\n";
274- strs += " if (solved) {\n";
275- strs += " uint64_t l2_use = GetL2Use();\n";
276- strs += " if (l2_use > max_l2_use) {\n";
277- strs += " for (uint32_t l = 0; l < input_.size; l++) {\n";
278- strs += " best_val[l] = blocknum_per_tile_[l];\n";
279- strs += " }\n";
280- strs += " max_l2_use = l2_use;\n";
281- strs += " }\n";
282- strs += " return;\n";
283- strs += " }\n";
284- strs += " } else {\n";
285- strs += " HandleClash(loop_id+1, ori_val, best_val, max_l2_use);\n";
286- strs += " }\n";
287- strs += " }\n";
288- strs += "}\n";
289- strs += "\n";
290- return strs;
291-}
292- 
293-inline std::string GetL2RunAnnotation() {
294- std::string strs;
295- strs += "/**\n";
296- strs += " * 运行 L2TileSolver 算法来解决 L2 缓存分块问题。\n";
297- strs += " *\n";
298- strs += " * 这个函数是 L2TileSolver 算法的核心。它尝试根据输入参数和约束条件找到 L2 缓存的最优分块方案。\n";
299- strs += " * 函数首先通过 CheckInput() 函数检查输入参数的有效性。如果输入无效,它将记录一条错误消息并返回 false。\n";
300- strs += " * 然后通过 CheckSolvable() 函数检查问题是否有解。如果没有解,它将记录一条错误消息并返回 false。\n";
301- strs += " * 如果输入有效并且问题有解,函数将通过 InitInput() 函数初始化输入数据。\n";
302- strs += " *\n";
303- strs += " * 算法的核心是一个循环,在这个循环中,它不断地调整每个变量的值,以找到一个合适的分块方案。\n";
304- strs += " * 循环结束的条件是总内存使用量小于或等于 L2 缓存的大小。\n";
305- strs += " * 循环结束后,它将计算每个变量的每个分块的大小、总块数、分块数和占用的核心数。\n";
306- strs += " * 然后,它检查是否存在读冲突。如果存在读冲突,它将调整每个分块的大小,直到不再检测到冲突。\n";
307- strs += " * 最后,它检查每个变量的最大值是否可以放在一个分块中。\n";
308- strs += " * 如果不能,它将相应地调整分块数和每个分块的大小。\n";
309- strs += " *\n";
310- strs += " * 如果找到合适的分块方案,函数返回 true,否则返回 false。\n";
311- strs += " *\n";
312- strs += " * @Return 如果成功则返回 true,否则返回 false\n";
313- strs += " */\n";
314- return strs;
315-}
316- 
317-inline std::string GetL2Run() {
318- std::string strs;
319- strs += GetL2RunAnnotation();
320- strs += "bool L2TileSolver::Run() {\n";
321- strs += " if (!CheckInput()) {\n";
322- strs += " OP_LOGW(OP_NAME, \"Check input failed\");\n";
323- strs += " return false;\n";
324- strs += " }\n";
325- strs += " if (!CheckSolvable()) {\n";
326- strs += " OP_LOGW(OP_NAME, \"Check Solvable failed\");\n";
327- strs += " return false;\n";
328- strs += " }\n";
329- strs += " InitInput();\n";
330- strs += " uint32_t core_num = input_.core_num;\n";
331- strs += " uint32_t l2_size = input_.l2_size;\n";
332- strs += " blocknum_per_tile_ = new(std::nothrow) uint32_t[input_.size];\n";
333- strs += " size_per_tile_ = new(std::nothrow) uint32_t[input_.size];\n";
334- strs += " tilenum_ = new(std::nothrow) uint32_t[input_.size];\n";
335- strs += " total_blocknum_ = new(std::nothrow) uint32_t[input_.size];\n";
336- strs += " for (uint32_t i=0u; i < input_.size; i++) {\n";
337- strs += " auto &var = input_.l2_vars[i];\n";
338- strs += " auto val = GetInitVal(input_.l2_size);\n";
339- strs += " var.value = std::min(var.max_value, CeilDivision(val, var.align) * var.align);\n";
340- strs += " }\n";
341- strs += " // 遍历直到满足L2占用停止\n";
342- strs += " while (GetL2Use() > l2_size) {\n";
343- strs += " for (uint32_t i = 0; i < input_.size; i++) {\n";
344- strs += " auto &var = input_.l2_vars[i];\n";
345- strs += " var.value = (var.align < var.value) ? (var.value - var.align) : var.align;\n";
346- strs += " }\n";
347- strs += " }\n";
348- strs += " for (uint32_t i = 0; i < input_.size; i++) {\n";
349- strs += " auto &var = input_.l2_vars[i];\n";
350- strs += " blocknum_per_tile_[i] = CeilDivision(var.value, var.base_val);\n";
351- strs += " size_per_tile_[i] = blocknum_per_tile_[i] * var.base_val;\n";
352- strs += " tilenum_[i] = CeilDivision(var.max_value, size_per_tile_[i]);\n";
353- strs += " total_blocknum_[i] = CeilDivision(var.max_value, var.base_val);\n";
354- strs += " }\n";
355- strs += " for (uint32_t i = 0; i < input_.size; i++) {\n";
356- strs += " auto &var = input_.l2_vars[i];\n";
357- strs += " if (var.max_value <= size_per_tile_[i]) {\n";
358- strs += " tilenum_[i] = 1;\n";
359- strs += " blocknum_per_tile_[i] = CeilDivision(var.max_value, var.base_val);\n";
360- strs += " size_per_tile_[i] = blocknum_per_tile_[i] * var.base_val;\n";
361- strs += " }\n";
362- strs += " }\n";
363- strs += " return true;\n";
364- strs += "}\n";
365- return strs;
366-}
367- 
368-inline std::string GetL2SolverHead() {
369- std::string strs;
370- strs += GenCeilDivision();
371- strs += GenL2Var();
372- strs += L2TileInput();
373- strs += GenL2TileSolver();
374- return strs;
375-}
376- 
377-inline std::string GetL2SolverFunc() {
378- std::string strs;
379- strs += GenL2CheckInput();
380- strs += GenL2CheckSolvable();
381- strs += GetL2InitInput();
382- strs += GetHandleClash();
383- strs += GetL2Run();
384- return strs;
385-}
386- 
387-inline const std::string L2_SOLVER_CODE_HEAD = GetL2SolverHead();
388-inline const std::string L2_SOLVER_CODE_FUNC = GetL2SolverFunc();
389-} // namespace att
390-#endif
@@ -12,12 +12,6 @@
12 12 
13namespace att {13namespace att {
14std::string GetSolverHead(SolverType type) {14std::string GetSolverHead(SolverType type) {
15- if (type == SolverType::L0_TILE) {
16- return L0_SOLVER_CODE_HEAD;
17- }
18- if (type == SolverType::L2_TILE) {
19- return L2_SOLVER_CODE_HEAD;
20- }
21 if (type == SolverType::SEARCH_TILE) {15 if (type == SolverType::SEARCH_TILE) {
22 return GENERAL_SOLVER_CODE; // 全是inline,放在头文件16 return GENERAL_SOLVER_CODE; // 全是inline,放在头文件
23 }17 }
@@ -25,12 +19,7 @@ std::string GetSolverHead(SolverType type) {
25}19}
26 20 
27std::string GetSolverFunc(SolverType type) {21std::string GetSolverFunc(SolverType type) {
28- if (type == SolverType::L0_TILE) {22+ (void)type;
29- return L0_SOLVER_CODE_FUNC;
30- }
31- if (type == SolverType::L2_TILE) {
32- return L2_SOLVER_CODE_FUNC;
33- }
34 return "";23 return "";
35}24}
36 25 
@@ -9,8 +9,6 @@
9 */9 */
10#ifndef ATT_SOLVER_H_10#ifndef ATT_SOLVER_H_
11#define ATT_SOLVER_H_11#define ATT_SOLVER_H_
12-#include "l0_solver_code.h"
13-#include "l2_solver_code.h"
14#include "general_solver_code.h"12#include "general_solver_code.h"
15#include "axes_reorder_solver_code.h"13#include "axes_reorder_solver_code.h"
16#include "base/base_types.h"14#include "base/base_types.h"
@@ -1,245 +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-#include "l0_solver_gen.h"
11-#include "ge_common_af/ge_api_error_codes_af.h"
12-#include "ge_common/debug/log.h"
13- 
14-namespace att {
15-bool L0TileSolverGen::CheckIsInnerMost(const Expr &arg) {
16- for (const auto &innermost_arg : innermost_args_) {
17- if (innermost_arg == arg) {
18- return true;
19- }
20- }
21- return false;
22-}
23- 
24-std::string L0TileSolverGen::GenClassDef() {
25- std::stringstream ss;
26- ss << AddAnotationLine("根据tiling case id创建对应的L0TileSolver子类\n");
27- ss << "class " << tiling_case_id_ << "L0TileSolver : public L0TileSolver {\n";
28- ss << " public:\n";
29- ss << AddAnotationLine("构造函数,接受L0TileInput类型的参数 input\n", " ");
30- ss << " explicit " << tiling_case_id_ << "L0TileSolver(L0TileInput &input) : L0TileSolver(input) {};\n";
31- ss << AddAnotationLine("成员函数声明\n", " ");
32- for (auto buffer_use : buffer_use_map_) {
33- auto hardware = buffer_use.first;
34- ss << " void Set" << BaseTypeUtils::DumpHardware(hardware) << "(const uint32_t &value) { "
35- << BaseTypeUtils::DumpHardware(hardware) << "_ = value; }\n";
36- }
37- ss << " bool CheckBufferUseValid() override;\n";
38- ss << AddAnotationLine("定义private类型的成员变量\n", " ");
39- ss << " private:\n";
40- for (auto buffer_use : buffer_use_map_) {
41- auto hardware = buffer_use.first;
42- ss << " uint32_t " << BaseTypeUtils::DumpHardware(hardware) << "_;\n";
43- }
44- for (auto const_var : const_vars_) {
45- if (!IsValid(const_var.first)) {
46- continue;
47- }
48- ss << " uint32_t " << const_var.first << "{" << const_var.second << "};\n";
49- }
50- ss << "};\n";
51- return ss.str();
52-}
53- 
54-std::string L0TileSolverGen::GenSolverClassImpl() {
55- std::stringstream ss;
56- std::string strs;
57- ss << GenClassDef();
58- strs = "";
59- strs += " 对创建的L0TileSolver子类的成员函数CheckBufferUseValid进行定义,判断缓存占用是否合法\n";
60- strs += " i从0开始到待求解l0相关变量的数量,依次递增循环\n";
61- strs += " @return 如果代求解l0变量的数量为空或i>=待求解l0相关变量的数量,返回false。如果在DEBUG域内输出提示信息\n";
62- strs += " 定义uint32_t类型变量:第i个待求解l0变量的字符串名称,值为第i个待求解l0相关变量的值\n";
63- ss << AddAnotationBlock(strs, " ");
64- 
65- ss << "bool " << tiling_case_id_ << "L0TileSolver::CheckBufferUseValid() {\n";
66- for (size_t i = 0u; i < l0_args_.size(); i++) {
67- ss << " if ((input_.l0_vars == nullptr) || (" << i << " >= " << l0_args_.size() << ")) {\n";
68- ss << " OP_LOGW(OP_NAME, \"l0_vars is nullptr or overflow\");\n";
69- ss << " return false;\n"
70- << " }\n";
71- ss << " uint32_t " << l0_args_[i] << " = input_.l0_vars[" << std::to_string(i) << "].value;\n";
72- }
73- 
74- strs = "";
75- strs += " 遍历buffer_use_map_中的所有缓存占用表达式,将其转换为字符串类型,赋值给对应的设备类型\n";
76- strs += " 将每种设备类型缓存占用计算值和其设定值对比,如果缓存占用计算值大于设定值,返回false\n";
77- strs += " @return 遍历所有设备类型后,如果所有设备类型缓存占用计算值都小于设定值,返回true\n";
78- ss << AddAnotationBlock(strs, " ");
79- 
80- for (auto buffer_use : buffer_use_map_) {
81- auto hardware = buffer_use.first;
82- auto expr = buffer_use.second;
83- if ((hardware != HardwareDef::L0A) && (hardware != HardwareDef::L0B) && (hardware != HardwareDef::L0C)) {
84- continue;
85- }
86- ss << " uint32_t " << BaseTypeUtils::DumpHardware(hardware) << " = " << expr << ";\n";
87- ss << " if (" << BaseTypeUtils::DumpHardware(hardware) << " > " << BaseTypeUtils::DumpHardware(hardware)
88- << "_) {\n";
89- ss << " return false;\n"
90- << " };\n";
91- }
92- ss << " return true;\n"
93- << "}\n";
94- return ss.str();
95-}
96- 
97-af::Status L0TileSolverGen::GetLargestAlign(const Expr &arg, Expr &max_align) {
98- auto iter = father_args_map_.find(arg);
99- if (iter != father_args_map_.end()) {
100- Expr father_arg = iter->second;
101- auto arg_align_iter = arg_align_map_.find(father_arg);
102- if (arg_align_iter == arg_align_map_.end()) {
103- GELOGE(af::FAILED, "arg align map does not find arg [%s]", father_arg.Str().get());
104- max_align = af::Symbol(0U);
105- return af::FAILED;
106- }
107- Expr align = arg_align_iter->second;
108- max_align = af::sym::Max(max_align, align);
109- GetLargestAlign(father_arg, max_align);
110- }
111- return af::SUCCESS;
112-}
113- 
114-bool L0TileSolverGen::IsMulticoreArg(const Expr &arg) {
115- if (!IsValid(arg)) {
116- return false;
117- }
118- for (auto mc_arg : mc_args_) {
119- if (mc_arg == arg) {
120- return true;
121- }
122- }
123- return false;
124-}
125- 
126-bool L0TileSolverGen::IsBindMulticore(const Expr &arg) {
127- if (IsMulticoreArg(arg)) {
128- return true;
129- }
130- auto iter = father_args_map_.find(arg);
131- if (iter == father_args_map_.end()) {
132- return false;
133- }
134- return IsBindMulticore(iter->second);
135-}
136- 
137-std::string L0TileSolverGen::GenSolverFuncInvoke() {
138- std::string strs = "";
139- strs += " if (!ExecuteL0Solver(tiling_data)) {\n";
140- strs += " return false;\n";
141- strs += " }\n";
142- return strs;
143-}
144- 
145-std::string L0TileSolverGen::GenSolverInvokeDoc() const {
146- std::string strs = "";
147- strs += "定义L0求解器调用函数\n";
148- strs += "@param tiling_data 输入参数为tiling_data\n";
149- strs += "定义求解器接受的输入为L0TileInput类型结构体变量,并将结构体的成员变量赋值为对应值\n";
150- strs += "定义basem_size为L0Var类型的结构体变量,并将相关成员变量赋值为对应值\n";
151- strs += "定义basen_size为L0Var类型的结构体变量,并将相关成员变量赋值为对应值\n";
152- strs +=
153- "定义solver为case0L0TileSolver类型的结构体变量,并将此结构体变量的成员变量L1_, L2_, L0A_, L0B_, L0C_, "
154- "CORENUM_赋值为tiling_data中的设定值\n";
155- strs +=
156- "@return "
157- "执行L0TileSolver类的主要流程solver.run(),如果solver.run()的返回值不为true,则l0求解器执行失败,函数返回false\n";
158- strs += "用output指针指向solver的输出\n";
159- strs += "i从0开始到L0相关变量个数依次递增\n";
160- strs += "@return 如果output为空指针或者i>=L0相关的变量个数,则函数返回false,否则将basem_size设置为output的第i个值\n";
161- strs += "@return 若上述没有触发返回false的场景,则函数返回true\n";
162- return AddAnotationBlock(strs);
163-}
164- 
165-std::string L0TileSolverGen::GenInitTilingData() {
166- std::stringstream ss;
167- ss << " bool ExecuteL0Solver(" + type_name_ + "& tiling_data) {\n";
168- ss << " L0TileInput l0_input;\n";
169- ss << " l0_input.l0_vars = new(std::nothrow) L0Var[" << std::to_string(l0_args_.size()) << "];\n";
170- ss << " l0_input.size = " << std::to_string(l0_args_.size()) << ";\n";
171- ss << " l0_input.core_num = corenum_;\n";
172- for (size_t i = 0u; i < l0_args_.size(); i++) {
173- auto l0_arg = l0_args_[i];
174- ss << " L0Var " << l0_arg << ";\n";
175- if (arg_max_value_map_.find(l0_arg) == arg_max_value_map_.end()) {
176- GELOGE(af::FAILED, "Ori arg map does not find l0 arg [%s]", l0_arg.Str().get());
177- return kSolverGenError;
178- }
179- ss << " " << l0_arg << ".max_value = tiling_data.get_" << arg_max_value_map_[l0_arg] << "();\n";
180- std::string true_or_false = IsBindMulticore(l0_arg) ? "true" : "false";
181- ss << " " << l0_arg << ".bind_multicore = " << true_or_false << ";\n";
182- ss << " " << l0_arg << ".align = " << Str(arg_align_map_[l0_arg]) << ";\n";
183- if (arg_align_map_.find(l0_arg) == arg_align_map_.end()) {
184- GELOGE(af::FAILED, "Arg align map does not find l0 arg [%s]", l0_arg.Str().get());
185- return kSolverGenError;
186- }
187- Expr max_align = arg_align_map_[l0_arg];
188- if (max_align.IsConstExpr() && (af::SymbolicUtils::StaticCheckEq(max_align, af::Symbol(0)) == af::TriBool::kTrue)) {
189- GELOGE(af::FAILED, "l0 arg [%s] align is 0", l0_arg.Str().get());
190- return kSolverGenError;
191- }
192- GetLargestAlign(l0_arg, max_align);
193- if (max_align.IsConstExpr() && (af::SymbolicUtils::StaticCheckEq(max_align, af::Symbol(0)) == af::TriBool::kTrue)) {
194- GELOGE(af::FAILED, "Get largest align failed");
195- return kSolverGenError;
196- }
197- if (CheckIsInnerMost(l0_arg)) {
198- ss << " " << l0_arg << ".is_innermost = true;\n";
199- ss << " " << l0_arg << ".value = 256;\n";
200- } else if (max_align.IsConstExpr() &&
201- (af::SymbolicUtils::StaticCheckGt(max_align, arg_align_map_[l0_arg]) == af::TriBool::kTrue)) {
202- ss << " " << l0_arg << ".value = 64;\n";
203- } else {
204- ss << " " << l0_arg << ".value = 128;\n";
205- }
206- ss << " " << l0_arg << ".prompt_align = " << Str(max_align) << ";\n";
207- ss << " l0_input.l0_vars[" << std::to_string(i) << "] = " << l0_arg << ";\n";
208- }
209- return ss.str();
210-}
211- 
212-std::string L0TileSolverGen::GenSolverFuncImpl() {
213- std::stringstream ss;
214- ss << GenSolverInvokeDoc();
215- ss << GenInitTilingData();
216- ss << " " << tiling_case_id_ << "L0TileSolver solver(l0_input);\n";
217- for (auto buffer_use : buffer_use_map_) {
218- auto hardware = buffer_use.first;
219- auto expr = buffer_use.second;
220- ss << " solver.Set" << BaseTypeUtils::DumpHardware(hardware) << "(tiling_data.get_"
221- << BaseTypeUtils::DumpHardware(hardware) << "());\n";
222- }
223- ss << " if (!solver.Run()) {\n";
224- ss << " OP_LOGW(OP_NAME, \"l0 solver run failed\");\n";
225- ss << " return false;\n";
226- ss << " }\n";
227- ss << " uint32_t *output = solver.GetOutput();\n";
228- for (size_t i = 0u; i < l0_args_.size(); i++) {
229- auto l0_arg = l0_args_[i];
230- if (i >= l0_args_.size()) {
231- ss << " OP_LOGW(OP_NAME, \"output overflows\");\n";
232- ss << " return false;\n";
233- } else {
234- ss << " if (output == nullptr) {\n";
235- ss << " OP_LOGW(OP_NAME, \"output is nullptr\");\n";
236- ss << " return false;\n";
237- ss << " }\n";
238- }
239- ss << " tiling_data.set_" << l0_arg << "(output[" << std::to_string(i) << "]);\n";
240- }
241- ss << " return true;\n";
242- ss << " }\n";
243- return ss.str();
244-}
245-} // namespace att
@@ -1,74 +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-#ifndef ATT_L0_SOLVER_GEN_H_
11-#define ATT_L0_SOLVER_GEN_H_
12-#include <iostream>
13-#include <map>
14-#include <string>
15-#include <vector>
16-#include <sstream>
17-#include "base/base_types.h"
18-#include "common/checker.h"
19-#include "generator/solver_pass_gen/solver_gen.h"
20-#include "util/base_types_printer.h"
21- 
22-namespace att {
23-class L0TileSolverGen : public SolverGen {
24- public:
25- explicit L0TileSolverGen(const std::string &tiling_case_id, const std::string &type_name)
26- : SolverGen(tiling_case_id, type_name) {}
27- ~L0TileSolverGen() override = default;
28- std::string GenSolverClassImpl() override;
29- std::string GenSolverFuncImpl() override;
30- std::string GenSolverFuncInvoke() override;
31- void SetL0Args(const std::vector<Expr> &l0_args) {
32- l0_args_ = l0_args;
33- }
34- void SetConstVars(const ExprUintMap &const_vars) {
35- const_vars_ = const_vars;
36- }
37- void SetBufferUseAlg(const std::map<HardwareDef, Expr> &buffer_use_map) {
38- buffer_use_map_ = buffer_use_map;
39- }
40- void SetMulticoreArgs(const std::vector<Expr> &mc_args) {
41- mc_args_ = mc_args;
42- }
43- void SetFatherArgsMap(const ExprExprMap &father_args_map) {
44- father_args_map_ = father_args_map;
45- }
46- void SetArgAlignMap(const ExprExprMap &arg_align_map) {
47- arg_align_map_ = arg_align_map;
48- }
49- void SetArgtMaxValueMap(const ExprExprMap &arg_max_value_map) {
50- arg_max_value_map_ = arg_max_value_map;
51- }
52- void SetInnerMostArgs(const std::vector<Expr> innermost_args) {
53- innermost_args_ = innermost_args;
54- }
55- 
56- private:
57- af::Status GetLargestAlign(const Expr &arg, Expr &max_align);
58- bool IsBindMulticore(const Expr &arg);
59- bool IsMulticoreArg(const Expr &arg);
60- bool CheckIsInnerMost(const Expr &arg);
61- std::string GenClassDef();
62- std::string GenSolverInvokeDoc() const;
63- std::string GenInitTilingData();
64- std::vector<Expr> l0_args_;
65- std::map<HardwareDef, Expr> buffer_use_map_;
66- ExprUintMap const_vars_;
67- std::vector<Expr> mc_args_;
68- ExprExprMap father_args_map_;
69- ExprExprMap arg_align_map_;
70- ExprExprMap arg_max_value_map_;
71- std::vector<Expr> innermost_args_;
72-};
73-} // namespace att
74-#endif
@@ -1,228 +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-#include "l2_solver_gen.h"
11- 
12-namespace att {
13-bool L2TileSolverGen::IsRepeatArgs(const Expr &arg) const {
14- if (!IsValid(arg)) {
15- return false;
16- }
17- if (l2_use_.ContainVar(arg)) {
18- return true;
19- }
20- return false;
21-}
22- 
23-bool L2TileSolverGen::IsClashPossible() const {
24- std::vector<Expr> input_and_const_args;
25- input_and_const_args.insert(input_and_const_args.end(), input_args_.begin(), input_args_.end());
26- for (auto &arg : const_vars_) {
27- input_and_const_args.emplace_back(arg.first);
28- }
29- for (size_t i = 0u; i < input_and_const_args.size(); i++) {
30- bool has_repeat_arg = true;
31- for (auto &l2_arg : l2_args_) {
32- if (!IsValid(l2_arg)) {
33- return false;
34- }
35- if (input_and_const_args[i] == arg_max_value_map_.at(l2_arg)) {
36- has_repeat_arg = false;
37- }
38- }
39- if (has_repeat_arg) {
40- if (IsRepeatArgs(input_and_const_args[i])) {
41- return true;
42- }
43- }
44- }
45- return false;
46-}
47- 
48-std::string L2TileSolverGen::GetL2RelInputArg() {
49- std::string l2_rel_input_arg = "";
50- for (const auto &arg : input_args_) {
51- if (IsRepeatArgs(arg)) {
52- l2_rel_input_arg = Str(arg);
53- }
54- }
55- return l2_rel_input_arg;
56-}
57- 
58-std::string L2TileSolverGen::GenClassDef() {
59- std::stringstream ss;
60- ss << AddAnotationLine("根据tiling case id创建对应的L2TileSolver子类\n");
61- ss << "class " << tiling_case_id_ << "L2TileSolver : public L2TileSolver {\n";
62- ss << " public:\n";
63- auto l2_rel_input_arg = GetL2RelInputArg();
64- ss << AddAnotationLine("构造函数,接受L2TileInput类型的参数 input\n", " ");
65- ss << " explicit " << tiling_case_id_
66- << "L2TileSolver(L2TileInput &input, " + type_name_ + " &tiling_data) : L2TileSolver(input) {\n";
67- if (l2_rel_input_arg != "") {
68- ss << " " << l2_rel_input_arg << " = tiling_data.get_" << l2_rel_input_arg << "();\n";
69- }
70- ss << " };\n";
71- ss << AddAnotationLine("成员函数声明\n", " ");
72- ss << " uint64_t GetL2Use() override;\n";
73- ss << " bool IsClash(const uint32_t idx) override;\n";
74- ss << " uint32_t GetInitVal(uint32_t L2Size) override;\n";
75- for (auto const_var : const_vars_) {
76- ss << " uint32_t " << const_var.first << "{" << const_var.second << "};\n";
77- }
78- if (l2_rel_input_arg != "") {
79- ss << " uint64_t " << l2_rel_input_arg << ";\n";
80- }
81- ss << "};\n";
82- std::string strs = "";
83- strs += "对创建的L2TileSolver子类的成员函数GetL2Use进行定义,计算L2的可用内存大小\n";
84- strs += "用input_存储输入数据,并将输入数据对应赋值给tilem_size, tilen_size, k_size\n";
85- strs += "使用公式计算L2的可用内存大小l2_size\n";
86- strs += "@return 返回L2可用内存的大小l2_size\n";
87- ss << AddAnotationBlock(strs);
88- ss << "uint64_t " << tiling_case_id_ << "L2TileSolver::GetL2Use() {\n";
89- for (size_t i = 0u; i < l2_args_.size(); i++) {
90- ss << " uint64_t " << l2_args_[i] << " = input_.l2_vars[" << std::to_string(i) << "].value;\n";
91- }
92- ss << " uint64_t l2_size = " << l2_use_ << ";\n";
93- ss << " return l2_size;\n";
94- ss << "}\n";
95- ss << "uint32_t " << tiling_case_id_ << "L2TileSolver::GetInitVal(uint32_t L2Size) {\n";
96- if (l2_rel_input_arg != "") {
97- ss << " uint64_t kSquare = " << l2_rel_input_arg << " * " << l2_rel_input_arg << ";\n";
98- ss << " return (uint32_t)((std::sqrt(16 * kSquare + 8 * L2Size) - 4 * " << l2_rel_input_arg << ") / 4);\n";
99- } else {
100- ss << " uint32_t max_val = 0u;\n";
101- ss << " for (uint32_t i=0u; i < input_.size; i++) {\n";
102- ss << " auto &var = input_.l2_vars[i];\n";
103- ss << " max_val = var.max_value > max_val ? var.max_value : max_val;\n";
104- ss << " }\n";
105- ss << " return max_val;\n";
106- }
107- ss << "}\n";
108- return ss.str();
109-}
110- 
111-std::string L2TileSolverGen::GenSolverClassImpl() {
112- std::stringstream ss;
113- std::string strs;
114- ss << GenClassDef();
115- strs = "";
116- strs += "对创建的L2TileSolver子类的成员函数IsClash进行定义,用于检测是否存在读冲突\n";
117- strs += "@param idx 输入参数为idx,用于获取当前tile的编号\n";
118- strs += "@return 如果满足条件blocknum_per_tile_[idx] % (used_corenum_ / 2) == 0,则存在读冲突,返回true\n";
119- strs += "用公式计算尾块个数blocknum_tail\n";
120- strs += "@return 如果满足条件blocknum_tail % (used_corenum_ / 2) == 0,则存在读冲突,返回true\n";
121- strs += "@return 如果以上两个条件均不满足,说明不存在读冲突,返回false\n";
122- ss << AddAnotationBlock(strs);
123- 
124- ss << "bool " << tiling_case_id_ << "L2TileSolver::IsClash(const uint32_t idx) {\n";
125- if (!IsClashPossible()) {
126- ss << " return false;\n";
127- } else {
128- ss << " if ((used_corenum_ <= 1) || (used_corenum_ % 2 !=0)) {\n";
129- ss << " return false;\n }\n";
130- ss << " if (blocknum_per_tile_[idx] % (used_corenum_ / 2) == 0) {\n";
131- ss << " return true;\n";
132- ss << " }\n";
133- ss << " auto blocknum_tail = total_blocknum_[idx] - (tilenum_[idx] "
134- "- 1) * blocknum_per_tile_[idx];\n";
135- ss << " if (blocknum_tail % (used_corenum_ / 2) == 0) {\n";
136- ss << " return true;\n";
137- ss << " }\n";
138- ss << " return false;\n";
139- }
140- ss << "}\n";
141- return ss.str();
142-}
143- 
144-Expr L2TileSolverGen::GetRelateL0Arg(const Expr &l2_arg) {
145- Expr res;
146- auto iter = arg_max_value_map_.find(l2_arg);
147- GE_ASSERT_TRUE(iter != arg_max_value_map_.end(), "Ori arg map does not find l2 arg[%s]", l2_arg.Str().get());
148- Expr ori_l2_arg = iter->second;
149- for (auto l0_arg : l0_args_) {
150- GE_ASSERT_TRUE(arg_max_value_map_.find(l0_arg) != arg_max_value_map_.end(), "Ori arg map does not find l0 arg[%s]",
151- l0_arg.Str().get());
152- if (arg_max_value_map_[l0_arg] == ori_l2_arg) {
153- return l0_arg;
154- }
155- }
156- return res;
157-}
158- 
159-std::string L2TileSolverGen::GenSolverFuncInvoke() {
160- std::string strs = "";
161- strs += " if (!ExecuteL2Solver(tiling_data)) {\n";
162- strs += " return false;\n";
163- strs += " }\n";
164- return strs;
165-}
166- 
167-std::string L2TileSolverGen::GenSolverInvokeDoc() const {
168- std::string strs = "";
169- strs += "定义L2求解器调用函数\n";
170- strs += "@param tiling_data 输入参数为tiling_data\n";
171- strs += "定义求解器的输入为L2TileInput类型结构体变量,并将结构体的成员变量赋值为对应值\n";
172- strs += "定义tilem_size为L2Var类型的结构体变量,并将相关成员变量赋值为对应值\n";
173- strs += "定义tilen_size为L2Var类型的结构体变量,并将相关成员变量赋值为对应值\n";
174- strs += "定义l2_solver为L2TileSolver类型的结构体变量\n";
175- strs += "计算l2的可用内存大小以及是否存在读冲突\n";
176- strs +=
177- "@return "
178- "执行L2TileSolver类的主要流程l2_solver.run(),如果l2_solver.run()的返回值不为true,"
179- "则l2求解器执行失败,函数返回false\n";
180- strs += "用output指针指向l2_solver的输出\n";
181- strs += "i从0开始到L2相关变量个数依次递增\n";
182- strs += "@return 如果output为空指针或者i>=L2相关的变量个数,则函数返回false,否则将tilem_size设置为output的第i个值\n";
183- strs += "@return 若上述没有触发返回false的场景,则函数返回true\n";
184- return AddAnotationBlock(strs);
185-}
186- 
187-std::string L2TileSolverGen::GenSolverFuncImpl() {
188- std::stringstream ss;
189- ss << GenSolverInvokeDoc();
190- ss << " bool ExecuteL2Solver(" + type_name_ + "& tiling_data) {\n";
191- ss << " L2TileInput l2_input;\n";
192- ss << " l2_input.l2_vars = new(std::nothrow) L2Var[" << std::to_string(l2_args_.size()) << "];\n";
193- ss << " l2_input.size = " << std::to_string(l2_args_.size()) << ";\n";
194- ss << " l2_input.core_num = corenum_;\n";
195- ss << " l2_input.l2_size = EMPIRIC_L2_SIZE;\n";
196- for (size_t i = 0u; i < l2_args_.size(); i++) {
197- auto l2_arg = l2_args_[i];
198- ss << " L2Var " << l2_arg << ";\n";
199- ss << " " << l2_arg << ".max_value = tiling_data.get_" << arg_max_value_map_[l2_arg] << "();\n";
200- if (arg_align_map_.find(l2_arg) == arg_align_map_.end()) {
201- GELOGE(af::FAILED, "Arg align map does not find l2 arg[%s]", l2_arg.Str().get());
202- return kSolverGenError;
203- }
204- ss << " " << l2_arg << ".align = " << Str(arg_align_map_[l2_arg]) << ";\n";
205- auto relate_l0_arg = GetRelateL0Arg(l2_arg);
206- GE_ASSERT_TRUE(IsValid(relate_l0_arg), "Get relate l0 arg failed");
207- ss << " " << l2_arg << ".base_val = tiling_data.get_" << relate_l0_arg << "();\n";
208- ss << " l2_input.l2_vars[" << std::to_string(i) << "] = " << l2_arg << ";\n";
209- }
210- ss << " " << tiling_case_id_ << "L2TileSolver l2_solver(l2_input, tiling_data);\n";
211- ss << " if (!l2_solver.Run()) {\n";
212- ss << " OP_LOGW(OP_NAME, \"l2 solver run failed\");\n";
213- ss << " return false;\n";
214- ss << " }\n";
215- ss << " uint32_t *output = l2_solver.GetL2Tile();\n";
216- for (size_t i = 0u; i < l2_args_.size(); i++) {
217- auto l2_arg = l2_args_[i];
218- ss << " if ((output == nullptr) || (" << i << " >= " << l2_args_.size() << ")) {\n";
219- ss << " OP_LOGW(OP_NAME, \"l2_vars is nullptr or overflow\");\n";
220- ss << " return false;\n";
221- ss << " }\n";
222- ss << " tiling_data.set_" << l2_arg << "(output[" << std::to_string(i) << "]);\n";
223- }
224- ss << " return true;\n";
225- ss << " }\n";
226- return ss.str();
227-}
228-} // namespace att
@@ -1,69 +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-#ifndef ATT_L2_SOLVER_GEN_H_
11-#define ATT_L2_SOLVER_GEN_H_
12-#include <iostream>
13-#include <map>
14-#include <string>
15-#include <vector>
16-#include <sstream>
17-#include "generator/solver_pass_gen/solver_gen.h"
18-#include "base/base_types.h"
19-#include "util/base_types_printer.h"
20-#include "common/checker.h"
21- 
22-namespace att {
23-class L2TileSolverGen : public SolverGen {
24- public:
25- explicit L2TileSolverGen(const std::string &tiling_case_id, const std::string &type_name)
26- : SolverGen(tiling_case_id, type_name) {}
27- ~L2TileSolverGen() override = default;
28- std::string GenSolverClassImpl() override;
29- std::string GenSolverFuncImpl() override;
30- std::string GenSolverFuncInvoke() override;
31- void SetInputArgs(const std::vector<Expr> &input_args) {
32- input_args_ = input_args;
33- }
34- void SetConstVars(const ExprUintMap &const_vars) {
35- const_vars_ = const_vars;
36- }
37- void SetL0Args(const std::vector<Expr> &l0_args) {
38- l0_args_ = l0_args;
39- }
40- void SetL2Args(const std::vector<Expr> &l2_args) {
41- l2_args_ = l2_args;
42- }
43- void SetL2Use(const Expr &expr) {
44- l2_use_ = expr;
45- }
46- void SetArgtMaxValueMap(const ExprExprMap &arg_max_value_map) {
47- arg_max_value_map_ = arg_max_value_map;
48- }
49- void SetArgAlignMap(const ExprExprMap &arg_align_map) {
50- arg_align_map_ = arg_align_map;
51- }
52- 
53- private:
54- std::string GenClassDef();
55- std::string GenSolverInvokeDoc() const;
56- std::string GetL2RelInputArg();
57- Expr GetRelateL0Arg(const Expr &l2_arg);
58- bool IsRepeatArgs(const Expr &arg) const;
59- bool IsClashPossible() const;
60- std::vector<Expr> l0_args_;
61- ExprUintMap const_vars_;
62- std::vector<Expr> l2_args_;
63- std::vector<Expr> input_args_;
64- Expr l2_use_;
65- ExprExprMap arg_max_value_map_;
66- ExprExprMap arg_align_map_;
67-};
68-} // namespace att
69-#endif
@@ -14,7 +14,6 @@
14 14 
15namespace att {15namespace att {
16constexpr char kSolverGenError[] = "Solver Gen Error";16constexpr char kSolverGenError[] = "Solver Gen Error";
17-constexpr uint32_t kMaxL0VarNum = 3u;
18inline std::string GetSmoothString(std::string str) {17inline std::string GetSmoothString(std::string str) {
19 std::string ret;18 std::string ret;
20 std::string target = "Ceiling";19 std::string target = "Ceiling";
@@ -12,118 +12,6 @@
12#include "graph/symbolizer/symbolic_utils.h"12#include "graph/symbolizer/symbolic_utils.h"
13 13 
14namespace att {14namespace att {
15-bool SolverPassManager::CheckArgExist(const Expr &new_arg, const std::vector<Expr> &args) {
16- for (auto arg : args) {
17- if (IsValid(arg) && (new_arg == arg)) {
18- return true;
19- }
20- }
21- return false;
22-}
23- 
24-std::vector<Expr> SolverPassManager::GetL0Args(ArgsManager args_manager, bool is_solved = false) {
25- std::vector<Expr> l0_args;
26- if (is_solved) {
27- const std::vector<Expr> &solved_args = args_manager.GetSolvedVars();
28- for (const auto &arg : solved_args) {
29- const auto &related_hardware = args_manager.GetRelatedHardware(arg);
30- if ((std::find(related_hardware.begin(), related_hardware.end(), HardwareDef::L0A) != related_hardware.end()) ||
31- (std::find(related_hardware.begin(), related_hardware.end(), HardwareDef::L0B) != related_hardware.end()) ||
32- (std::find(related_hardware.begin(), related_hardware.end(), HardwareDef::L0C) != related_hardware.end())) {
33- l0_args.emplace_back(arg);
34- }
35- }
36- return l0_args;
37- }
38- std::vector<Expr> l0a_args = args_manager.GetSearchableVars(HardwareDef::L0A);
39- std::vector<Expr> l0b_args = args_manager.GetSearchableVars(HardwareDef::L0B);
40- std::vector<Expr> l0c_args = args_manager.GetSearchableVars(HardwareDef::L0C);
41- for (auto arg : l0c_args) {
42- if (!CheckArgExist(arg, l0_args)) {
43- l0_args.emplace_back(arg);
44- }
45- }
46- for (auto arg : l0a_args) {
47- if (!CheckArgExist(arg, l0_args)) {
48- l0_args.emplace_back(arg);
49- }
50- }
51- for (auto arg : l0b_args) {
52- if (!CheckArgExist(arg, l0_args)) {
53- l0_args.emplace_back(arg);
54- }
55- }
56- return l0_args;
57-}
58- 
59-L0TileSolverGen SolverPassManager::GenL0TileSolverGen() {
60- std::vector<Expr> l0_args = GetL0Args(args_manager_);
61- std::map<HardwareDef, Expr> buffer_use_map;
62- std::vector<Expr> mc_args;
63- ExprExprMap father_args_map, arg_max_value_map;
64- ExprExprMap arg_align_map;
65- buffer_use_map = args_manager_.GetTotalHardwareCons();
66- mc_args = args_manager_.GetSearchableVars(HardwareDef::CORENUM);
67- auto search_args = args_manager_.GetSearchableVars();
68- for (auto arg : search_args) {
69- auto father_arg = args_manager_.GetParentVars(arg);
70- if (!father_arg.empty()) {
71- father_args_map[arg] = father_arg[0];
72- }
73- Expr align = args_manager_.GetVarAlignValue(arg);
74- arg_align_map[arg] = align;
75- arg_max_value_map[arg] = args_manager_.GetMaxValue(arg);
76- }
77- for (auto it : father_args_map) {
78- auto father_arg = it.second;
79- if (arg_align_map.find(father_arg) == arg_align_map.end()) {
80- Expr align = args_manager_.GetVarAlignValue(father_arg);
81- arg_align_map[father_arg] = align;
82- }
83- }
84- std::vector<Expr> innermost_args = args_manager_.GetNodeInnerestDimSizes();
85- L0TileSolverGen solver_gen("case" + std::to_string(case_id_), tiling_data_type_);
86- solver_gen.SetL0Args(l0_args);
87- solver_gen.SetBufferUseAlg(buffer_use_map);
88- solver_gen.SetMulticoreArgs(mc_args);
89- solver_gen.SetFatherArgsMap(father_args_map);
90- solver_gen.SetArgAlignMap(arg_align_map);
91- solver_gen.SetConstVars(args_manager_.GetConstVars());
92- solver_gen.SetArgtMaxValueMap(arg_max_value_map);
93- solver_gen.SetInnerMostArgs(innermost_args);
94- return solver_gen;
95-}
96- 
97-L2TileSolverGen SolverPassManager::GenL2TileSolverGen() {
98- std::vector<Expr> l0_args = GetL0Args(args_manager_, true);
99- std::vector<Expr> l2_args = args_manager_.GetSearchableVars(HardwareDef::L2);
100- auto buffer_use_map = args_manager_.GetTotalHardwareCons();
101- std::vector<Expr> input_args = args_manager_.GetInputVars();
102- Expr l2_use = buffer_use_map[HardwareDef::L2];
103- auto search_args = args_manager_.GetSearchableVars();
104- ExprExprMap arg_align_map;
105- ExprExprMap arg_max_value_map;
106- for (auto arg : search_args) {
107- Expr align = args_manager_.GetVarAlignValue(arg);
108- arg_align_map[arg] = align;
109- arg_max_value_map[arg] = args_manager_.GetMaxValue(arg);
110- }
111- for (auto arg : l0_args) {
112- Expr align = args_manager_.GetVarAlignValue(arg);
113- arg_align_map[arg] = align;
114- arg_max_value_map[arg] = args_manager_.GetMaxValue(arg);
115- }
116- L2TileSolverGen solver_gen("case" + std::to_string(case_id_), tiling_data_type_);
117- solver_gen.SetArgAlignMap(arg_align_map);
118- solver_gen.SetL0Args(l0_args);
119- solver_gen.SetL2Args(l2_args);
120- solver_gen.SetL2Use(l2_use);
121- solver_gen.SetConstVars(args_manager_.GetConstVars());
122- solver_gen.SetArgtMaxValueMap(arg_max_value_map);
123- solver_gen.SetInputArgs(input_args);
124- return solver_gen;
125-}
126- 
127template <typename SolverGenType>15template <typename SolverGenType>
128auto SolverPassManager::GenerateSolverGen() -> SolverGenType {16auto SolverPassManager::GenerateSolverGen() -> SolverGenType {
129 ExprExprMap max_value;17 ExprExprMap max_value;
@@ -161,47 +49,13 @@ auto SolverPassManager::GenerateSolverGen() -> SolverGenType {
161 return solver_gen;49 return solver_gen;
162}50}
163 51 
164-std::pair<std::string, std::string> SolverPassManager::L0SolverPassFuncGen() {
165- std::string impl_codes, invoke_codes = "";
166- std::pair<std::string, std::string> codes;
167- std::vector<Expr> l0_args = GetL0Args(args_manager_);
168- if ((l0_args.size() == 0) || (l0_args.size() > kMaxL0VarNum)) {
169- codes = std::make_pair(impl_codes, invoke_codes);
170- return codes;
171- }
172- L0TileSolverGen solver_gen = GenL0TileSolverGen();
173- impl_codes = solver_gen.GenSolverFuncImpl();
174- invoke_codes = solver_gen.GenSolverFuncInvoke();
175- codes = std::make_pair(impl_codes, invoke_codes);
176- args_manager_.SetSolvedVars(l0_args);
177- return codes;
178-}
179- 
180-std::pair<std::string, std::string> SolverPassManager::L2SolverPassFuncGen() {
181- std::string impl_codes = "";
182- std::string invoke_codes = "";
183- std::pair<std::string, std::string> codes;
184- std::vector<Expr> l2_args = args_manager_.GetSearchableVars(HardwareDef::L2);
185- if (l2_args.size() == 0) {
186- return codes;
187- }
188- L2TileSolverGen solver_gen = GenL2TileSolverGen();
189- impl_codes = solver_gen.GenSolverFuncImpl();
190- invoke_codes = solver_gen.GenSolverFuncInvoke();
191- codes = std::make_pair(impl_codes, invoke_codes);
192- args_manager_.SetSolvedVars(l2_args);
193- return codes;
194-}
195- 
196template <typename SpecificSolverGen>52template <typename SpecificSolverGen>
197std::pair<std::string, std::string> SolverPassManager::GenerateSolverPassFunc(SpecificSolverGen solver_gen) {53std::pair<std::string, std::string> SolverPassManager::GenerateSolverPassFunc(SpecificSolverGen solver_gen) {
198 std::string impl_codes;54 std::string impl_codes;
199 std::string invoke_codes;55 std::string invoke_codes;
200 std::pair<std::string, std::string> codes;56 std::pair<std::string, std::string> codes;
201- std::vector<Expr> l0_args = GetL0Args(args_manager_);
202- std::vector<Expr> l2_args = args_manager_.GetSearchableVars(HardwareDef::L2);
203 std::vector<Expr> all_args = args_manager_.GetSearchableVars();57 std::vector<Expr> all_args = args_manager_.GetSearchableVars();
204- if (l0_args.size() + l2_args.size() >= all_args.size()) {58+ if (all_args.size() == 0) {
205 return codes;59 return codes;
206 }60 }
207 impl_codes = solver_gen.GenSolverFuncImpl();61 impl_codes = solver_gen.GenSolverFuncImpl();
@@ -212,14 +66,8 @@ std::pair<std::string, std::string> SolverPassManager::GenerateSolverPassFunc(Sp
212}66}
213 67 
214std::pair<std::string, std::string> SolverPassManager::SolverPassFuncGen(SolverType type) {68std::pair<std::string, std::string> SolverPassManager::SolverPassFuncGen(SolverType type) {
215- GE_ASSERT_TRUE((type != SolverType::L0_TILE) || (type != SolverType::L2_TILE) || (type != SolverType::SEARCH_TILE),
216- "Solver type[%u] is invalid", type);
217 std::pair<std::string, std::string> codes;69 std::pair<std::string, std::string> codes;
218- if (type == SolverType::L0_TILE) {70+ if (type == SolverType::SEARCH_TILE) {
219- codes = L0SolverPassFuncGen();
220- } else if (type == SolverType::L2_TILE) {
221- codes = L2SolverPassFuncGen();
222- } else if (type == SolverType::SEARCH_TILE) {
223 args_manager_.DoVarsReplace();71 args_manager_.DoVarsReplace();
224 codes = GenerateSolverPassFunc(GenerateSolverGen<GeneralSolverGen>());72 codes = GenerateSolverPassFunc(GenerateSolverGen<GeneralSolverGen>());
225 }73 }
@@ -241,30 +89,6 @@ std::pair<std::string, std::string> SolverPassManager::GenFuncPass() {
241 return pass_codes;89 return pass_codes;
242}90}
243 91 
244-std::pair<std::string, std::string> SolverPassManager::L0SolverDtFuncGen() {
245- std::vector<Expr> l0_args = GetL0Args(args_manager_);
246- if ((l0_args.size() == 0) || (l0_args.size() > kMaxL0VarNum)) {
247- return std::make_pair("", "");
248- } else {
249- L0TileSolverGen solver_gen("case" + std::to_string(case_id_), tiling_data_type_);
250- args_manager_.SetSolvedVars(l0_args);
251- return std::make_pair("", solver_gen.GenSolverFuncInvoke());
252- }
253-}
254- 
255-std::pair<std::string, std::string> SolverPassManager::L2SolverDtFuncGen() {
256- std::string invoke_codes;
257- std::vector<Expr> l0_args = GetL0Args(args_manager_, true);
258- std::vector<Expr> l2_args = args_manager_.GetSearchableVars(HardwareDef::L2);
259- if (l2_args.size() == 0) {
260- return std::make_pair("", "");
261- } else {
262- L2TileSolverGen solver_gen("case" + std::to_string(case_id_), tiling_data_type_);
263- args_manager_.SetSolvedVars(l2_args);
264- return std::make_pair("", solver_gen.GenSolverFuncInvoke());
265- }
266-}
267- 
268std::pair<std::string, std::string> SolverPassManager::GeneralSolverDtFuncGen() {92std::pair<std::string, std::string> SolverPassManager::GeneralSolverDtFuncGen() {
269 GeneralSolverGen solver_gen = GenerateSolverGen<GeneralSolverGen>();93 GeneralSolverGen solver_gen = GenerateSolverGen<GeneralSolverGen>();
270 std::string impl_codes = solver_gen.GenSolverDTImpl();94 std::string impl_codes = solver_gen.GenSolverDTImpl();
@@ -273,44 +97,14 @@ std::pair<std::string, std::string> SolverPassManager::GeneralSolverDtFuncGen()
273}97}
274 98 
275std::pair<std::string, std::string> SolverPassManager::SolverDtFuncGen(SolverType type) {99std::pair<std::string, std::string> SolverPassManager::SolverDtFuncGen(SolverType type) {
276- GE_ASSERT_TRUE((type != SolverType::L0_TILE) || (type != SolverType::L2_TILE) || (type != SolverType::SEARCH_TILE),
277- "Solver type[%u] is invalid", type);
278 std::pair<std::string, std::string> codes;100 std::pair<std::string, std::string> codes;
279- if (type == SolverType::L0_TILE) {101+ if (type == SolverType::SEARCH_TILE) {
280- codes = L0SolverDtFuncGen();
281- } else if (type == SolverType::L2_TILE) {
282- codes = L2SolverDtFuncGen();
283- } else if (type == SolverType::SEARCH_TILE) {
284 args_manager_.DoVarsReplace();102 args_manager_.DoVarsReplace();
285 codes = GeneralSolverDtFuncGen();103 codes = GeneralSolverDtFuncGen();
286 }104 }
287 return codes;105 return codes;
288}106}
289 107 
290-std::string SolverPassManager::L0SolverPassClassGen() {
291- std::string code;
292- std::vector<Expr> l0_args = GetL0Args(args_manager_);
293- if ((l0_args.size() == 0) || (l0_args.size() > kMaxL0VarNum)) {
294- return code;
295- }
296- L0TileSolverGen solver_gen = GenL0TileSolverGen();
297- code = solver_gen.GenSolverClassImpl();
298- args_manager_.SetSolvedVars(l0_args);
299- return code;
300-}
301- 
302-std::string SolverPassManager::L2SolverPassClassGen() {
303- std::string code;
304- std::vector<Expr> l2_args = args_manager_.GetSearchableVars(HardwareDef::L2);
305- if (l2_args.size() == 0) {
306- return code;
307- }
308- L2TileSolverGen solver_gen = GenL2TileSolverGen();
309- code = solver_gen.GenSolverClassImpl();
310- args_manager_.SetSolvedVars(l2_args);
311- return code;
312-}
313- 
314std::string SolverPassManager::GeneralSolverPassClassGen() {108std::string SolverPassManager::GeneralSolverPassClassGen() {
315 GeneralSolverGen solver_gen = GenerateSolverGen<GeneralSolverGen>();109 GeneralSolverGen solver_gen = GenerateSolverGen<GeneralSolverGen>();
316 std::string code = solver_gen.GenSolverClassImpl();110 std::string code = solver_gen.GenSolverClassImpl();
@@ -319,13 +113,7 @@ std::string SolverPassManager::GeneralSolverPassClassGen() {
319}113}
320 114 
321std::string SolverPassManager::SolverPassClassGen(SolverType type) {115std::string SolverPassManager::SolverPassClassGen(SolverType type) {
322- GE_ASSERT_TRUE((type != SolverType::L0_TILE) || (type != SolverType::L2_TILE) || (type != SolverType::SEARCH_TILE),116+ if (type == SolverType::SEARCH_TILE) {
323- "Solver type[%u] is invalid", type);
324- if (type == SolverType::L0_TILE) {
325- return L0SolverPassClassGen();
326- } else if (type == SolverType::L2_TILE) {
327- return L2SolverPassClassGen();
328- } else if (type == SolverType::SEARCH_TILE) {
329 args_manager_.DoVarsReplace();117 args_manager_.DoVarsReplace();
330 return GeneralSolverPassClassGen();118 return GeneralSolverPassClassGen();
331 }119 }
@@ -343,29 +131,9 @@ std::string SolverPassManager::GenClassPass() {
343}131}
344 132 
345bool SolverPassManager::IsNeedSolver(std::vector<ArgsManager> args_managers, SolverType type) {133bool SolverPassManager::IsNeedSolver(std::vector<ArgsManager> args_managers, SolverType type) {
346- if (type == SolverType::L0_TILE) {
347- for (auto &args_manager : args_managers) {
348- std::vector<Expr> l0_args = GetL0Args(args_manager);
349- if ((l0_args.size() > 0) && (l0_args.size() <= kMaxL0VarNum)) {
350- return true;
351- }
352- }
353- return false;
354- }
355- if (type == SolverType::L2_TILE) {
356- for (auto &args_manager : args_managers) {
357- std::vector<Expr> l2_args = args_manager.GetSearchableVars(HardwareDef::L2);
358- if (l2_args.size() > 0) {
359- return true;
360- }
361- }
362- return false;
363- }
364 if (type == SolverType::SEARCH_TILE) {134 if (type == SolverType::SEARCH_TILE) {
365 for (auto &args_manager : args_managers) {135 for (auto &args_manager : args_managers) {
366- std::vector<Expr> l0_args = GetL0Args(args_manager);136+ if (args_manager.GetSearchableVars().size() > 0) {
367- std::vector<Expr> l2_args = args_manager.GetSearchableVars(HardwareDef::L2);
368- if (l0_args.size() + l2_args.size() < args_manager.GetSearchableVars().size()) {
369 return true;137 return true;
370 }138 }
371 }139 }
@@ -18,8 +18,6 @@
18#include "generator/solver_pass/solver.h"18#include "generator/solver_pass/solver.h"
19#include "generator/solver_pass_gen/axes_reorder_solver/axes_reorder_solver_gen.h"19#include "generator/solver_pass_gen/axes_reorder_solver/axes_reorder_solver_gen.h"
20#include "generator/solver_pass_gen/general_solver/general_solver_gen.h"20#include "generator/solver_pass_gen/general_solver/general_solver_gen.h"
21-#include "generator/solver_pass_gen/l0_solver/l0_solver_gen.h"
22-#include "generator/solver_pass_gen/l2_solver/l2_solver_gen.h"
23#include "util/base_types_printer.h"21#include "util/base_types_printer.h"
24#include "autofuse_config/auto_fuse_config.h"22#include "autofuse_config/auto_fuse_config.h"
25#include "generator/solver_pass_gen/input_output_setters.h"23#include "generator/solver_pass_gen/input_output_setters.h"
@@ -76,35 +74,25 @@ class SolverPassManager : public InputOutputSettersMixin<SolverPassManager>,
76 74 
77 private:75 private:
78 // solver pass76 // solver pass
79- static bool CheckArgExist(const Expr &new_arg, const std::vector<Expr> &args);
80- static std::vector<Expr> GetL0Args(ArgsManager args_manager, bool is_solved);
81 static bool IsNeedSolver(std::vector<ArgsManager> args_managers, SolverType type);77 static bool IsNeedSolver(std::vector<ArgsManager> args_managers, SolverType type);
82 static std::string GenBaseClass(SolverType type);78 static std::string GenBaseClass(SolverType type);
83 79 
84 ExprExprMap GetInputsAlign(bool do_replace);80 ExprExprMap GetInputsAlign(bool do_replace);
85 ExprExprMap GetOriginalInputAlign() const;81 ExprExprMap GetOriginalInputAlign() const;
86 82 
87- L0TileSolverGen GenL0TileSolverGen();
88- L2TileSolverGen GenL2TileSolverGen();
89 void InitSolverGen(AxesReorderSolverGen &solver_gen);83 void InitSolverGen(AxesReorderSolverGen &solver_gen);
90 AxesReorderSolverGen GenAxesReorderGen();84 AxesReorderSolverGen GenAxesReorderGen();
91 template <typename SolverGenType>85 template <typename SolverGenType>
92 SolverGenType GenerateSolverGen();86 SolverGenType GenerateSolverGen();
93 87 
94 std::string SolverPassClassGen(SolverType type);88 std::string SolverPassClassGen(SolverType type);
95- std::string L0SolverPassClassGen();
96- std::string L2SolverPassClassGen();
97 std::string GeneralSolverPassClassGen();89 std::string GeneralSolverPassClassGen();
98 90 
99 template <typename SpecificSolverGen>91 template <typename SpecificSolverGen>
100 std::pair<std::string, std::string> GenerateSolverPassFunc(SpecificSolverGen solver_gen);92 std::pair<std::string, std::string> GenerateSolverPassFunc(SpecificSolverGen solver_gen);
101 std::pair<std::string, std::string> SolverPassFuncGen(SolverType type);93 std::pair<std::string, std::string> SolverPassFuncGen(SolverType type);
102- std::pair<std::string, std::string> L0SolverPassFuncGen();
103- std::pair<std::string, std::string> L2SolverPassFuncGen();
104 94 
105 std::pair<std::string, std::string> SolverDtFuncGen(SolverType type);95 std::pair<std::string, std::string> SolverDtFuncGen(SolverType type);
106- std::pair<std::string, std::string> L0SolverDtFuncGen();
107- std::pair<std::string, std::string> L2SolverDtFuncGen();
108 std::pair<std::string, std::string> GeneralSolverDtFuncGen();96 std::pair<std::string, std::string> GeneralSolverDtFuncGen();
109 97 
110 void AddConcatInnerDims(const Expr &arg, std::vector<Expr> &concat_inner_dims);98 void AddConcatInnerDims(const Expr &arg, std::vector<Expr> &concat_inner_dims);
@@ -258,8 +258,6 @@ file(GLOB SOURCES
258 ${ATT_DIR}/generator/solver_pass_gen/axes_reorder_solver/*.cpp258 ${ATT_DIR}/generator/solver_pass_gen/axes_reorder_solver/*.cpp
259 ${ATT_DIR}/generator/solver_pass_gen/general_solver/*.cpp259 ${ATT_DIR}/generator/solver_pass_gen/general_solver/*.cpp
260 ${ATT_DIR}/generator/solver_pass_gen/golden_solver/*.cpp260 ${ATT_DIR}/generator/solver_pass_gen/golden_solver/*.cpp
261- ${ATT_DIR}/generator/solver_pass_gen/l0_solver/*.cpp
262- ${ATT_DIR}/generator/solver_pass_gen/l2_solver/*.cpp
263 ${ATT_DIR}/generator/extra_info_gen/*.cpp261 ${ATT_DIR}/generator/extra_info_gen/*.cpp
264 ${ATT_DIR}/generator/generator_utils/*.cpp262 ${ATT_DIR}/generator/generator_utils/*.cpp
265 ${ATT_DIR}/generator/tiling_data_gen/*.cpp263 ${ATT_DIR}/generator/tiling_data_gen/*.cpp
@@ -22,8 +22,6 @@ att::Expr GetSafeOffsetDivisor(const att::Expr &expr) {
22 22 
23struct FfnExprContext {23struct FfnExprContext {
24 att::Expr expr_maxTokens;24 att::Expr expr_maxTokens;
25- att::Expr expr_basem1;
26- att::Expr expr_basem2;
27 att::Expr expr_ubm;25 att::Expr expr_ubm;
28 att::Expr expr_n1;26 att::Expr expr_n1;
29 att::Expr expr_basen1;27 att::Expr expr_basen1;
@@ -33,27 +31,17 @@ struct FfnExprContext {
33};31};
34 32 
35void BuildFfnMaxTokenAxes(att::ModelInfo &model_info, FfnExprContext &ctx, att::AttAxisPtr &maxTokens,33void BuildFfnMaxTokenAxes(att::ModelInfo &model_info, FfnExprContext &ctx, att::AttAxisPtr &maxTokens,
36- att::AttAxisPtr &basem1, att::AttAxisPtr &ubm, att::AttAxisPtr &basem2) {34+ att::AttAxisPtr &ubm) {
37 ctx.expr_maxTokens = att::CreateExpr("maxTokens");35 ctx.expr_maxTokens = att::CreateExpr("maxTokens");
38- ctx.expr_basem1 = att::CreateExpr("base_m1");
39- ctx.expr_basem2 = att::CreateExpr("base_m2");
40 ctx.expr_ubm = att::CreateExpr("ub_m");36 ctx.expr_ubm = att::CreateExpr("ub_m");
41 37 
42 att::SymVarInfoPtr sym_maxTokens = std::make_shared<att::SymVarInfo>(ctx.expr_maxTokens);38 att::SymVarInfoPtr sym_maxTokens = std::make_shared<att::SymVarInfo>(ctx.expr_maxTokens);
43- att::SymVarInfoPtr sym_basem1 = std::make_shared<att::SymVarInfo>(ctx.expr_basem1);
44- sym_basem1->align = ge::Symbol(8);
45- sym_basem1->related_scope = {att::HardwareDef::L0C};
46 att::SymVarInfoPtr sym_ubm = std::make_shared<att::SymVarInfo>(ctx.expr_ubm);39 att::SymVarInfoPtr sym_ubm = std::make_shared<att::SymVarInfo>(ctx.expr_ubm);
47 sym_ubm->align = ge::Symbol(8);40 sym_ubm->align = ge::Symbol(8);
48 sym_ubm->related_scope = {att::HardwareDef::UB};41 sym_ubm->related_scope = {att::HardwareDef::UB};
49- att::SymVarInfoPtr sym_basem2 = std::make_shared<att::SymVarInfo>(ctx.expr_basem2);
50- sym_basem2->align = ge::Symbol(8);
51- sym_basem2->related_scope = {att::HardwareDef::L0C};
52 42 
53 maxTokens = std::make_shared<att::AttAxis>();43 maxTokens = std::make_shared<att::AttAxis>();
54- basem1 = std::make_shared<att::AttAxis>();
55 ubm = std::make_shared<att::AttAxis>();44 ubm = std::make_shared<att::AttAxis>();
56- basem2 = std::make_shared<att::AttAxis>();
57 45 
58 maxTokens->name = "maxTokens";46 maxTokens->name = "maxTokens";
59 maxTokens->axis_pos = att::AxisPosition::ORIGIN;47 maxTokens->axis_pos = att::AxisPosition::ORIGIN;
@@ -62,15 +50,6 @@ void BuildFfnMaxTokenAxes(att::ModelInfo &model_info, FfnExprContext &ctx, att::
62 maxTokens->is_node_innerest_dim = false;50 maxTokens->is_node_innerest_dim = false;
63 maxTokens->size = sym_maxTokens;51 maxTokens->size = sym_maxTokens;
64 52 
65- basem1->name = "base_m1";
66- basem1->axis_pos = att::AxisPosition::INNER;
67- basem1->bind_multicore = false;
68- basem1->is_last = false;
69- basem1->is_node_innerest_dim = false;
70- basem1->size = sym_basem1;
71- basem1->orig_axis.push_back(maxTokens.get());
72- basem1->from_axis = {maxTokens.get()};
73- 
74 ubm->name = "ub_m";53 ubm->name = "ub_m";
75 ubm->axis_pos = att::AxisPosition::INNER;54 ubm->axis_pos = att::AxisPosition::INNER;
76 ubm->bind_multicore = false;55 ubm->bind_multicore = false;
@@ -78,16 +57,7 @@ void BuildFfnMaxTokenAxes(att::ModelInfo &model_info, FfnExprContext &ctx, att::
78 ubm->is_node_innerest_dim = false;57 ubm->is_node_innerest_dim = false;
79 ubm->size = sym_ubm;58 ubm->size = sym_ubm;
80 ubm->orig_axis.push_back(maxTokens.get());59 ubm->orig_axis.push_back(maxTokens.get());
81- ubm->from_axis = {basem1.get()};60+ ubm->from_axis = {maxTokens.get()};
82- 
83- basem2->name = "base_m2";
84- basem2->axis_pos = att::AxisPosition::INNER;
85- basem2->bind_multicore = false;
86- basem2->is_last = true;
87- basem2->is_node_innerest_dim = false;
88- basem2->size = sym_basem2;
89- basem2->orig_axis.push_back(maxTokens.get());
90- basem2->from_axis = {maxTokens.get()};
91}61}
92 62 
93void BuildFfnN1Axes(FfnExprContext &ctx, att::AttAxisPtr &n1, att::AttAxisPtr &basen1) {63void BuildFfnN1Axes(FfnExprContext &ctx, att::AttAxisPtr &n1, att::AttAxisPtr &basen1) {
@@ -97,7 +67,7 @@ void BuildFfnN1Axes(FfnExprContext &ctx, att::AttAxisPtr &n1, att::AttAxisPtr &b
97 att::SymVarInfoPtr sym_n1 = std::make_shared<att::SymVarInfo>(ctx.expr_n1);67 att::SymVarInfoPtr sym_n1 = std::make_shared<att::SymVarInfo>(ctx.expr_n1);
98 att::SymVarInfoPtr sym_basen1 = std::make_shared<att::SymVarInfo>(ctx.expr_basen1);68 att::SymVarInfoPtr sym_basen1 = std::make_shared<att::SymVarInfo>(ctx.expr_basen1);
99 sym_basen1->align = ge::Symbol(8);69 sym_basen1->align = ge::Symbol(8);
100- sym_basen1->related_scope = {att::HardwareDef::L0C, att::HardwareDef::UB, att::HardwareDef::BTBUF};70+ sym_basen1->related_scope = {att::HardwareDef::UB, att::HardwareDef::BTBUF};
101 71 
102 n1 = std::make_shared<att::AttAxis>();72 n1 = std::make_shared<att::AttAxis>();
103 basen1 = std::make_shared<att::AttAxis>();73 basen1 = std::make_shared<att::AttAxis>();
@@ -138,7 +108,7 @@ void BuildFfnN2Axes(FfnExprContext &ctx, att::AttAxisPtr &n2, att::AttAxisPtr &b
138 att::SymVarInfoPtr sym_n2 = std::make_shared<att::SymVarInfo>(ctx.expr_n2);108 att::SymVarInfoPtr sym_n2 = std::make_shared<att::SymVarInfo>(ctx.expr_n2);
139 att::SymVarInfoPtr sym_basen2 = std::make_shared<att::SymVarInfo>(ctx.expr_basen2);109 att::SymVarInfoPtr sym_basen2 = std::make_shared<att::SymVarInfo>(ctx.expr_basen2);
140 sym_basen2->align = ge::Symbol(8);110 sym_basen2->align = ge::Symbol(8);
141- sym_basen2->related_scope = {att::HardwareDef::L0C, att::HardwareDef::BTBUF};111+ sym_basen2->related_scope = {att::HardwareDef::BTBUF};
142 112 
143 n2 = std::make_shared<att::AttAxis>();113 n2 = std::make_shared<att::AttAxis>();
144 basen2 = std::make_shared<att::AttAxis>();114 basen2 = std::make_shared<att::AttAxis>();
@@ -160,72 +130,22 @@ void BuildFfnN2Axes(FfnExprContext &ctx, att::AttAxisPtr &n2, att::AttAxisPtr &b
160 basen2->from_axis = {n2.get()};130 basen2->from_axis = {n2.get()};
161}131}
162 132 
163-att::Expr CalcCube1Mte2(const FfnExprContext &ctx, const att::Expr &n1_cnt, const att::Expr &m1_cnt) {
164- att::Expr expr_m1n1 = ((((att::CreateExpr(0.05624f) * ctx.expr_basem1) + att::CreateExpr(0.3984f)) *
165- att::CreateExpr(6.2712e-05f) * ctx.expr_k1 * ctx.expr_basen1) +
166- (att::CreateExpr(0.0008295f) * ctx.expr_k1 * ctx.expr_basen1));
167- att::Expr weight_m1n1 = ((att::CreateExpr(0.05761f) * ctx.expr_basen1) + att::CreateExpr(0.0f));
168- att::Expr mte2_m1n1 = expr_m1n1 * weight_m1n1;
169- att::Expr expr_n1m1 = ((((att::CreateExpr(0.05940f) * ctx.expr_basen1) + att::CreateExpr(20.0944f)) *
170- att::CreateExpr(6.2712e-05f) * ctx.expr_k1 * ctx.expr_basem1) +
171- (att::CreateExpr(0.0008295f) * ctx.expr_k1 * ctx.expr_basem1));
172- att::Expr weight_n1m1 = ((att::CreateExpr(0.07543f) * ctx.expr_k1) + att::CreateExpr(0.0f));
173- att::Expr mte2_n1m1 = expr_n1m1 * weight_n1m1;
174- att::Expr weight1 = (att::CreateExpr(0.000216f) * ctx.expr_basen1) + (att::CreateExpr(0.0003614f) * ctx.expr_basem1) +
175- (att::CreateExpr(0.0005757f) * ctx.expr_k1);
176- att::Expr weight2 = (att::CreateExpr(0.0f) * ctx.expr_k1 * ctx.expr_basem1) +
177- (att::CreateExpr(0.0f) * ctx.expr_basem1 * ctx.expr_basen1) +
178- (att::CreateExpr(0.0f) * ctx.expr_k1 * ctx.expr_basen1);
179- return (mte2_m1n1 + mte2_n1m1) * (n1_cnt * m1_cnt) / (weight1 + weight2);
180-}
181- 
182-att::Expr CalcCube2Mte2(const FfnExprContext &ctx, const att::Expr &n2_cnt, const att::Expr &m2_cnt) {
183- att::Expr expr_m2n2 = ((((att::CreateExpr(0.05624f) * ctx.expr_basem2) + att::CreateExpr(0.3984f)) *
184- att::CreateExpr(6.2712e-05f) * ctx.expr_n1 * ctx.expr_basen2) +
185- (att::CreateExpr(0.0008295f) * ctx.expr_n1 * ctx.expr_basen2));
186- att::Expr weight_m2n2 = ((att::CreateExpr(0.05761f) * ctx.expr_basen2) + att::CreateExpr(0.0f));
187- att::Expr mte2_m2n2 = expr_m2n2 * weight_m2n2;
188- att::Expr expr_n2m2 = ((((att::CreateExpr(0.05940f) * ctx.expr_basen2) + att::CreateExpr(20.0944f)) *
189- att::CreateExpr(6.2712e-05f) * ctx.expr_n1 * ctx.expr_basem2) +
190- (att::CreateExpr(0.0008295f) * ctx.expr_n1 * ctx.expr_basem2));
191- att::Expr weight_n2m2 = ((att::CreateExpr(0.07543f) * ctx.expr_n1) + att::CreateExpr(0.0f));
192- att::Expr mte2_n2m2 = expr_n2m2 * weight_n2m2;
193- att::Expr weight1 = (att::CreateExpr(0.000216f) * ctx.expr_basen2) + (att::CreateExpr(0.0003614f) * ctx.expr_basem2) +
194- (att::CreateExpr(0.0005757f) * ctx.expr_n1);
195- att::Expr weight2 = (att::CreateExpr(0.0f) * ctx.expr_n1 * ctx.expr_basem2) +
196- (att::CreateExpr(0.0f) * ctx.expr_basem2 * ctx.expr_basen2) +
197- (att::CreateExpr(0.0f) * ctx.expr_n1 * ctx.expr_basen2);
198- return (mte2_m2n2 + mte2_n2m2) * (n2_cnt * m2_cnt) / (weight1 + weight2);
199-}
200- 
201void FillFfnModelInfo(att::ModelInfo &model_info, const FfnExprContext &ctx) {133void FillFfnModelInfo(att::ModelInfo &model_info, const FfnExprContext &ctx) {
202 att::Expr btbuf_occupy = af::sym::Max((att::CreateExpr(4) * ctx.expr_basen1), (att::CreateExpr(4) * ctx.expr_basen2));134 att::Expr btbuf_occupy = af::sym::Max((att::CreateExpr(4) * ctx.expr_basen1), (att::CreateExpr(4) * ctx.expr_basen2));
203- att::Expr l0c_occupy = af::sym::Max((att::CreateExpr(4) * ctx.expr_basen1 * ctx.expr_basem1),
204- (att::CreateExpr(4) * ctx.expr_basen2 * ctx.expr_basem2));
205 att::Expr ub_occupy = (att::CreateExpr(4) * ctx.expr_basen1 * ctx.expr_ubm);135 att::Expr ub_occupy = (att::CreateExpr(4) * ctx.expr_basen1 * ctx.expr_ubm);
206 model_info.hardware_cons[att::HardwareDef::BTBUF] = btbuf_occupy;136 model_info.hardware_cons[att::HardwareDef::BTBUF] = btbuf_occupy;
207- model_info.hardware_cons[att::HardwareDef::L0C] = l0c_occupy;
208 model_info.hardware_cons[att::HardwareDef::UB] = ub_occupy;137 model_info.hardware_cons[att::HardwareDef::UB] = ub_occupy;
209 138 
210- att::Expr m1_cnt = af::sym::Ceiling(ctx.expr_maxTokens / GetSafeDivisor(ctx.expr_basem1));139+ att::Expr m_cnt = af::sym::Ceiling(ctx.expr_maxTokens / GetSafeDivisor(ctx.expr_ubm));
211- att::Expr m2_cnt = af::sym::Ceiling(ctx.expr_maxTokens / GetSafeDivisor(ctx.expr_basem2));
212 att::Expr n1_cnt = af::sym::Ceiling(ctx.expr_n1 / GetSafeDivisor(ctx.expr_basen1));140 att::Expr n1_cnt = af::sym::Ceiling(ctx.expr_n1 / GetSafeDivisor(ctx.expr_basen1));
213 att::Expr n2_cnt = af::sym::Ceiling(ctx.expr_n2 / GetSafeDivisor(ctx.expr_basen2));141 att::Expr n2_cnt = af::sym::Ceiling(ctx.expr_n2 / GetSafeDivisor(ctx.expr_basen2));
214- att::Expr ubm_cnt = af::sym::Ceiling(ctx.expr_basem1 / GetSafeDivisor(ctx.expr_ubm));
215 142 
216 att::Expr vec_ub = ((att::CreateExpr(4) * ctx.expr_basen1 * ctx.expr_ubm) / GetSafeOffsetDivisor(ctx.expr_basen1) +143 att::Expr vec_ub = ((att::CreateExpr(4) * ctx.expr_basen1 * ctx.expr_ubm) / GetSafeOffsetDivisor(ctx.expr_basen1) +
217 att::CreateExpr(4));144 att::CreateExpr(4));
218- att::Expr vec_m1n1 =145+ att::Expr vec = vec_ub * (m_cnt * n1_cnt);
219- ((att::CreateExpr(8) * ctx.expr_basem1 * ctx.expr_basen1) / GetSafeOffsetDivisor(ctx.expr_basen1) +
220- att::CreateExpr(4));
221- att::Expr vec_m2n2 =
222- ((att::CreateExpr(8) * ctx.expr_basem2 * ctx.expr_basen2) / GetSafeOffsetDivisor(ctx.expr_basen2) +
223- att::CreateExpr(4));
224- att::Expr vec =
225- (vec_ub * (m1_cnt * n1_cnt * ubm_cnt)) + (vec_m1n1 * (m1_cnt * n1_cnt)) + (vec_m2n2 * (m2_cnt * n2_cnt));
226 146 
227 att::Expr mte3_ub = ((att::CreateExpr(0.01741f) * ctx.expr_basen1 * ctx.expr_ubm) + att::CreateExpr(0.22f));147 att::Expr mte3_ub = ((att::CreateExpr(0.01741f) * ctx.expr_basen1 * ctx.expr_ubm) + att::CreateExpr(0.22f));
228- att::Expr v_mte3 = mte3_ub * (m1_cnt * n1_cnt * ubm_cnt);148+ att::Expr v_mte3 = mte3_ub * (m_cnt * n1_cnt);
229 149 
230 att::Expr mte2_n1 =150 att::Expr mte2_n1 =
231 ((att::CreateExpr(5.01f) / (att::CreateExpr(27240.69f) + ctx.expr_basen1)) + att::CreateExpr(1051.66f)) *151 ((att::CreateExpr(5.01f) / (att::CreateExpr(27240.69f) + ctx.expr_basen1)) + att::CreateExpr(1051.66f)) *
@@ -234,33 +154,24 @@ void FillFfnModelInfo(att::ModelInfo &model_info, const FfnExprContext &ctx) {
234 ((att::CreateExpr(5.01f) / (att::CreateExpr(27240.69f) + ctx.expr_basen2)) + att::CreateExpr(1051.66f)) *154 ((att::CreateExpr(5.01f) / (att::CreateExpr(27240.69f) + ctx.expr_basen2)) + att::CreateExpr(1051.66f)) *
235 (ctx.expr_basen2 / att::CreateExpr(30421.24f));155 (ctx.expr_basen2 / att::CreateExpr(30421.24f));
236 att::Expr mte2_ub = (att::CreateExpr(0.007f) * ctx.expr_basen1 * ctx.expr_ubm) + att::CreateExpr(7.97f);156 att::Expr mte2_ub = (att::CreateExpr(0.007f) * ctx.expr_basen1 * ctx.expr_ubm) + att::CreateExpr(7.97f);
237- att::Expr v_mte2 = mte2_n1 * n1_cnt + mte2_n2 * n2_cnt + mte2_ub * (m1_cnt * n1_cnt * ubm_cnt);157+ att::Expr v_mte2 = mte2_n1 * n1_cnt + mte2_n2 * n2_cnt + mte2_ub * (m_cnt * n1_cnt);
238- 
239- att::Expr mte2_cube1 = CalcCube1Mte2(ctx, n1_cnt, m1_cnt);
240- att::Expr mte2_cube2 = CalcCube2Mte2(ctx, n2_cnt, m2_cnt);
241- att::Expr mte2 = mte2_cube1 + mte2_cube2;
242 158 
243 model_info.objects[att::PipeType::AIV_MTE2] = v_mte2;159 model_info.objects[att::PipeType::AIV_MTE2] = v_mte2;
244 model_info.objects[att::PipeType::AIV_MTE3] = v_mte3;160 model_info.objects[att::PipeType::AIV_MTE3] = v_mte3;
245- model_info.objects[att::PipeType::AIC_MTE2] = mte2;
246 model_info.objects[att::PipeType::AIV_VEC] = vec;161 model_info.objects[att::PipeType::AIV_VEC] = vec;
247 model_info.tiling_case_id = 0;162 model_info.tiling_case_id = 0;
248- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_basem1, ctx.expr_ubm));
249 model_info.output_size = 1;163 model_info.output_size = 1;
250}164}
251 165 
252void AppendFfnArgList(att::ModelInfo &model_info, const att::AttAxisPtr &maxTokens, const att::AttAxisPtr &basen1,166void AppendFfnArgList(att::ModelInfo &model_info, const att::AttAxisPtr &maxTokens, const att::AttAxisPtr &basen1,
253- const att::AttAxisPtr &basen2, const att::AttAxisPtr &n1, const att::AttAxisPtr &basem1,167+ const att::AttAxisPtr &basen2, const att::AttAxisPtr &n1, const att::AttAxisPtr &k1,
254- const att::AttAxisPtr &k1, const att::AttAxisPtr &n2, const att::AttAxisPtr &basem2,168+ const att::AttAxisPtr &n2, const att::AttAxisPtr &ubm) {
255- const att::AttAxisPtr &ubm) {
256 model_info.arg_list.emplace_back(maxTokens);169 model_info.arg_list.emplace_back(maxTokens);
257 model_info.arg_list.emplace_back(basen1);170 model_info.arg_list.emplace_back(basen1);
258 model_info.arg_list.emplace_back(basen2);171 model_info.arg_list.emplace_back(basen2);
259 model_info.arg_list.emplace_back(n1);172 model_info.arg_list.emplace_back(n1);
260- model_info.arg_list.emplace_back(basem1);
261 model_info.arg_list.emplace_back(k1);173 model_info.arg_list.emplace_back(k1);
262 model_info.arg_list.emplace_back(n2);174 model_info.arg_list.emplace_back(n2);
263- model_info.arg_list.emplace_back(basem2);
264 model_info.arg_list.emplace_back(ubm);175 model_info.arg_list.emplace_back(ubm);
265}176}
266} // namespace177} // namespace
@@ -271,10 +182,8 @@ ModelInfo GenFFNModelInfo() {
271 FfnExprContext ctx;182 FfnExprContext ctx;
272 183 
273 AttAxisPtr maxTokens;184 AttAxisPtr maxTokens;
274- AttAxisPtr basem1;
275 AttAxisPtr ubm;185 AttAxisPtr ubm;
276- AttAxisPtr basem2;186+ BuildFfnMaxTokenAxes(model_info, ctx, maxTokens, ubm);
277- BuildFfnMaxTokenAxes(model_info, ctx, maxTokens, basem1, ubm, basem2);
278 187 
279 AttAxisPtr n1;188 AttAxisPtr n1;
280 AttAxisPtr basen1;189 AttAxisPtr basen1;
@@ -288,7 +197,7 @@ ModelInfo GenFFNModelInfo() {
288 BuildFfnN2Axes(ctx, n2, basen2);197 BuildFfnN2Axes(ctx, n2, basen2);
289 198 
290 FillFfnModelInfo(model_info, ctx);199 FillFfnModelInfo(model_info, ctx);
291- AppendFfnArgList(model_info, maxTokens, basen1, basen2, n1, basem1, k1, n2, basem2, ubm);200+ AppendFfnArgList(model_info, maxTokens, basen1, basen2, n1, k1, n2, ubm);
292 return model_info;201 return model_info;
293}202}
294} // namespace att203} // namespace att
@@ -12,22 +12,13 @@
12#include "stub_matmul_modelinfo.h"12#include "stub_matmul_modelinfo.h"
13 13 
14namespace {14namespace {
15-att::Expr GetSafeDivisor(const att::Expr &expr) {
16- return af::sym::Max(af::sym::kSymbolOne, expr);
17-}
18- 
19struct MatmulExprContext {15struct MatmulExprContext {
20 att::Expr expr_corenum;16 att::Expr expr_corenum;
21 att::Expr expr_m;17 att::Expr expr_m;
22- att::Expr expr_tilem;
23- att::Expr expr_basem;
24 att::Expr expr_n;18 att::Expr expr_n;
25- att::Expr expr_tilen;
26- att::Expr expr_basen;
27 att::Expr expr_k;19 att::Expr expr_k;
28 att::Expr expr_stepka;20 att::Expr expr_stepka;
29 att::Expr expr_stepkb;21 att::Expr expr_stepkb;
30- att::Expr expr_basek;
31};22};
32 23 
33void BuildCoreAxis(att::ModelInfo &model_info, MatmulExprContext &ctx) {24void BuildCoreAxis(att::ModelInfo &model_info, MatmulExprContext &ctx) {
@@ -43,98 +34,38 @@ void BuildCoreAxis(att::ModelInfo &model_info, MatmulExprContext &ctx) {
43 model_info.arg_list.emplace_back(core);34 model_info.arg_list.emplace_back(core);
44}35}
45 36 
46-void BuildMAxes(MatmulExprContext &ctx, att::AttAxisPtr &m, att::AttAxisPtr &tilem, att::AttAxisPtr &basem) {37+void BuildMAxes(MatmulExprContext &ctx, att::AttAxisPtr &m) {
47 ctx.expr_m = att::CreateExpr("m_size");38 ctx.expr_m = att::CreateExpr("m_size");
48- ctx.expr_tilem = att::CreateExpr("tilem_size");
49- ctx.expr_basem = att::CreateExpr("basem_size");
50 39 
51 att::SymVarInfoPtr sym_m = std::make_shared<att::SymVarInfo>(ctx.expr_m);40 att::SymVarInfoPtr sym_m = std::make_shared<att::SymVarInfo>(ctx.expr_m);
52- att::SymVarInfoPtr sym_tilem = std::make_shared<att::SymVarInfo>(ctx.expr_tilem);
53- sym_tilem->align = ge::Symbol(16);
54- sym_tilem->related_scope = {att::HardwareDef::L2};
55- att::SymVarInfoPtr sym_basem = std::make_shared<att::SymVarInfo>(ctx.expr_basem);
56- sym_basem->align = ge::Symbol(16);
57- sym_basem->related_scope = {att::HardwareDef::L0A, att::HardwareDef::L0C, att::HardwareDef::L1};
58 41 
59 m = std::make_shared<att::AttAxis>();42 m = std::make_shared<att::AttAxis>();
60- tilem = std::make_shared<att::AttAxis>();
61- basem = std::make_shared<att::AttAxis>();
62- 
63 m->name = "m";43 m->name = "m";
64 m->axis_pos = att::AxisPosition::ORIGIN;44 m->axis_pos = att::AxisPosition::ORIGIN;
65 m->bind_multicore = false;45 m->bind_multicore = false;
66 m->is_last = false;46 m->is_last = false;
67 m->is_node_innerest_dim = false;47 m->is_node_innerest_dim = false;
68 m->size = sym_m;48 m->size = sym_m;
69- 
70- tilem->name = "tilem";
71- tilem->axis_pos = att::AxisPosition::INNER;
72- tilem->bind_multicore = false;
73- tilem->is_last = false;
74- tilem->is_node_innerest_dim = true;
75- tilem->size = sym_tilem;
76- tilem->orig_axis.push_back(m.get());
77- tilem->from_axis = {m.get()};
78- 
79- basem->name = "basem";
80- basem->axis_pos = att::AxisPosition::INNER;
81- basem->bind_multicore = false;
82- basem->is_last = true;
83- basem->is_node_innerest_dim = false;
84- basem->size = sym_basem;
85- basem->orig_axis.push_back(m.get());
86- basem->from_axis = {tilem.get()};
87}49}
88 50 
89-void BuildNAxes(MatmulExprContext &ctx, att::AttAxisPtr &n, att::AttAxisPtr &tilen, att::AttAxisPtr &basen) {51+void BuildNAxes(MatmulExprContext &ctx, att::AttAxisPtr &n) {
90 ctx.expr_n = att::CreateExpr("n_size");52 ctx.expr_n = att::CreateExpr("n_size");
91- ctx.expr_tilen = att::CreateExpr("tilen_size");
92- ctx.expr_basen = att::CreateExpr("basen_size");
93 53 
94 att::SymVarInfoPtr sym_n = std::make_shared<att::SymVarInfo>(ctx.expr_n);54 att::SymVarInfoPtr sym_n = std::make_shared<att::SymVarInfo>(ctx.expr_n);
95- att::SymVarInfoPtr sym_tilen = std::make_shared<att::SymVarInfo>(ctx.expr_tilen);
96- sym_tilen->align = ge::Symbol(16);
97- sym_tilen->related_scope = {att::HardwareDef::L2};
98- att::SymVarInfoPtr sym_basen = std::make_shared<att::SymVarInfo>(ctx.expr_basen);
99- sym_basen->align = ge::Symbol(16);
100- sym_basen->related_scope = {att::HardwareDef::L0B, att::HardwareDef::L0C, att::HardwareDef::L1};
101 55 
102 n = std::make_shared<att::AttAxis>();56 n = std::make_shared<att::AttAxis>();
103- tilen = std::make_shared<att::AttAxis>();
104- basen = std::make_shared<att::AttAxis>();
105- 
106 n->name = "n";57 n->name = "n";
107 n->axis_pos = att::AxisPosition::ORIGIN;58 n->axis_pos = att::AxisPosition::ORIGIN;
108 n->bind_multicore = false;59 n->bind_multicore = false;
109 n->is_last = false;60 n->is_last = false;
110 n->is_node_innerest_dim = false;61 n->is_node_innerest_dim = false;
111 n->size = sym_n;62 n->size = sym_n;
112- 
113- tilen->name = "tilen";
114- tilen->axis_pos = att::AxisPosition::INNER;
115- tilen->bind_multicore = false;
116- tilen->is_last = false;
117- tilen->is_node_innerest_dim = true;
118- tilen->size = sym_tilen;
119- tilen->orig_axis.push_back(n.get());
120- tilen->from_axis = {n.get()};
121- 
122- basen->name = "basen";
123- basen->axis_pos = att::AxisPosition::INNER;
124- basen->bind_multicore = false;
125- basen->is_last = true;
126- basen->is_node_innerest_dim = true;
127- basen->size = sym_basen;
128- basen->orig_axis.push_back(n.get());
129- basen->from_axis = {tilen.get()};
130}63}
131 64 
132-void BuildKAxes(MatmulExprContext &ctx, att::AttAxisPtr &k, att::AttAxisPtr &stepka, att::AttAxisPtr &stepkb,65+void BuildKAxes(MatmulExprContext &ctx, att::AttAxisPtr &k, att::AttAxisPtr &stepka, att::AttAxisPtr &stepkb) {
133- att::AttAxisPtr &basek) {
134 ctx.expr_k = att::CreateExpr("k_size");66 ctx.expr_k = att::CreateExpr("k_size");
135 ctx.expr_stepka = att::CreateExpr("stepka_size");67 ctx.expr_stepka = att::CreateExpr("stepka_size");
136 ctx.expr_stepkb = att::CreateExpr("stepkb_size");68 ctx.expr_stepkb = att::CreateExpr("stepkb_size");
137- ctx.expr_basek = att::CreateExpr("basek_size");
138 69 
139 att::SymVarInfoPtr sym_k = std::make_shared<att::SymVarInfo>(ctx.expr_k);70 att::SymVarInfoPtr sym_k = std::make_shared<att::SymVarInfo>(ctx.expr_k);
140 att::SymVarInfoPtr sym_stepka = std::make_shared<att::SymVarInfo>(ctx.expr_stepka);71 att::SymVarInfoPtr sym_stepka = std::make_shared<att::SymVarInfo>(ctx.expr_stepka);
@@ -143,14 +74,10 @@ void BuildKAxes(MatmulExprContext &ctx, att::AttAxisPtr &k, att::AttAxisPtr &ste
143 att::SymVarInfoPtr sym_stepkb = std::make_shared<att::SymVarInfo>(ctx.expr_stepkb);74 att::SymVarInfoPtr sym_stepkb = std::make_shared<att::SymVarInfo>(ctx.expr_stepkb);
144 sym_stepkb->align = ge::Symbol(16);75 sym_stepkb->align = ge::Symbol(16);
145 sym_stepkb->related_scope = {att::HardwareDef::L1};76 sym_stepkb->related_scope = {att::HardwareDef::L1};
146- att::SymVarInfoPtr sym_basek = std::make_shared<att::SymVarInfo>(ctx.expr_basek);
147- sym_basek->align = ge::Symbol(16);
148- sym_basek->related_scope = {att::HardwareDef::L0A, att::HardwareDef::L0B};
149 77 
150 k = std::make_shared<att::AttAxis>();78 k = std::make_shared<att::AttAxis>();
151 stepka = std::make_shared<att::AttAxis>();79 stepka = std::make_shared<att::AttAxis>();
152 stepkb = std::make_shared<att::AttAxis>();80 stepkb = std::make_shared<att::AttAxis>();
153- basek = std::make_shared<att::AttAxis>();
154 81 
155 k->name = "k";82 k->name = "k";
156 k->axis_pos = att::AxisPosition::ORIGIN;83 k->axis_pos = att::AxisPosition::ORIGIN;
@@ -171,117 +98,32 @@ void BuildKAxes(MatmulExprContext &ctx, att::AttAxisPtr &k, att::AttAxisPtr &ste
171 stepkb->name = "stepkb";98 stepkb->name = "stepkb";
172 stepkb->axis_pos = att::AxisPosition::INNER;99 stepkb->axis_pos = att::AxisPosition::INNER;
173 stepkb->bind_multicore = false;100 stepkb->bind_multicore = false;
174- stepkb->is_last = false;101+ stepkb->is_last = true;
175 stepkb->is_node_innerest_dim = true;102 stepkb->is_node_innerest_dim = true;
176 stepkb->size = sym_stepkb;103 stepkb->size = sym_stepkb;
177 stepkb->orig_axis.push_back(k.get());104 stepkb->orig_axis.push_back(k.get());
178 stepkb->from_axis = {stepka.get()};105 stepkb->from_axis = {stepka.get()};
179- 
180- basek->name = "basek";
181- basek->axis_pos = att::AxisPosition::INNER;
182- basek->bind_multicore = false;
183- basek->is_last = true;
184- basek->is_node_innerest_dim = false;
185- basek->size = sym_basek;
186- basek->orig_axis.push_back(k.get());
187- basek->from_axis = {stepkb.get()};
188}106}
189 107 
190-void AppendArgList(att::ModelInfo &model_info, const att::AttAxisPtr &m, const att::AttAxisPtr &tilen,108+void AppendArgList(att::ModelInfo &model_info, const att::AttAxisPtr &m, const att::AttAxisPtr &n,
191- const att::AttAxisPtr &tilem, const att::AttAxisPtr &stepka, const att::AttAxisPtr &stepkb,109+ const att::AttAxisPtr &k, const att::AttAxisPtr &stepka, const att::AttAxisPtr &stepkb) {
192- const att::AttAxisPtr &basek, const att::AttAxisPtr &basen, const att::AttAxisPtr &basem,
193- const att::AttAxisPtr &n, const att::AttAxisPtr &k) {
194 model_info.arg_list.emplace_back(m);110 model_info.arg_list.emplace_back(m);
195- model_info.arg_list.emplace_back(tilen);
196- model_info.arg_list.emplace_back(tilem);
197- model_info.arg_list.emplace_back(stepka);
198- model_info.arg_list.emplace_back(stepkb);
199- model_info.arg_list.emplace_back(basek);
200- model_info.arg_list.emplace_back(basen);
201- model_info.arg_list.emplace_back(basem);
202 model_info.arg_list.emplace_back(n);111 model_info.arg_list.emplace_back(n);
203 model_info.arg_list.emplace_back(k);112 model_info.arg_list.emplace_back(k);
113+ model_info.arg_list.emplace_back(stepka);
114+ model_info.arg_list.emplace_back(stepkb);
204}115}
205 116 
206void FillMatmulHardwareCons(att::ModelInfo &model_info, const MatmulExprContext &ctx) {117void FillMatmulHardwareCons(att::ModelInfo &model_info, const MatmulExprContext &ctx) {
207- model_info.hardware_cons[att::HardwareDef::L0A] = ctx.expr_basem * ctx.expr_basek * att::CreateExpr(4);118+ model_info.hardware_cons[att::HardwareDef::L1] = (ctx.expr_stepka * ctx.expr_stepkb * att::CreateExpr(4));
208- model_info.hardware_cons[att::HardwareDef::L0B] = ctx.expr_basek * ctx.expr_basen * att::CreateExpr(4);
209- model_info.hardware_cons[att::HardwareDef::L0C] = ctx.expr_basem * ctx.expr_basen * att::CreateExpr(4);
210- model_info.hardware_cons[att::HardwareDef::L1] =
211- (ctx.expr_stepka * ctx.expr_basem * att::CreateExpr(4)) + (ctx.expr_stepkb * ctx.expr_basen * att::CreateExpr(4));
212- model_info.hardware_cons[att::HardwareDef::L2] =
213- (ctx.expr_tilen * ctx.expr_tilem * att::CreateExpr(2)) +
214- ((ctx.expr_tilen + ctx.expr_tilem) * ctx.expr_k * att::CreateExpr(2));
215 model_info.hardware_cons[att::HardwareDef::UB] = att::CreateExpr(0L);119 model_info.hardware_cons[att::HardwareDef::UB] = att::CreateExpr(0L);
216}120}
217 121 
218-struct MatmulPerfContext {
219- att::Expr tile_cnt;
220- att::Expr base_cnt;
221- att::Expr al1_cnt;
222- att::Expr bl1_cnt;
223- att::Expr l0_cnt;
224-};
225- 
226-MatmulPerfContext CalcMatmulLoopCnts(const MatmulExprContext &ctx) {
227- MatmulPerfContext perf;
228- perf.tile_cnt = ((ctx.expr_n / GetSafeDivisor(ctx.expr_tilen)) * (ctx.expr_m / GetSafeDivisor(ctx.expr_tilem)));
229- perf.base_cnt = af::sym::Max(af::sym::kSymbolOne, (((ctx.expr_tilem / GetSafeDivisor(ctx.expr_basem)) *
230- (ctx.expr_tilen / GetSafeDivisor(ctx.expr_basen))) /
231- GetSafeDivisor(ctx.expr_corenum)));
232- perf.al1_cnt = ctx.expr_k / GetSafeDivisor(ctx.expr_stepka);
233- perf.bl1_cnt = ctx.expr_stepka / GetSafeDivisor(ctx.expr_stepkb);
234- perf.l0_cnt = ctx.expr_stepkb / GetSafeDivisor(ctx.expr_basek);
235- return perf;
236-}
237- 
238-void FillMatmulPerfObjects(att::ModelInfo &model_info, const MatmulExprContext &ctx, const MatmulPerfContext &perf) {
239- att::Expr l1_cnt = perf.al1_cnt * perf.bl1_cnt;
240- att::Expr al0_mte1 =
241- (((ctx.expr_basem * ctx.expr_basek) * att::CreateExpr(2)) / att::CreateExpr(512)) + att::CreateExpr(26);
242- att::Expr bl0_mte1 =
243- (((ctx.expr_basek * ctx.expr_basen) * att::CreateExpr(2)) / att::CreateExpr(256)) + att::CreateExpr(26);
244- att::Expr mte1 = (perf.tile_cnt * perf.base_cnt * l1_cnt * perf.l0_cnt) * (al0_mte1 + bl0_mte1);
245- std::cout << "mte1: " << mte1 << std::endl;
246- 
247- att::Expr l0_mac = af::sym::Ceiling(ctx.expr_basem / att::CreateExpr(16)) *
248- af::sym::Ceiling(ctx.expr_basek / att::CreateExpr(16)) *
249- af::sym::Ceiling(ctx.expr_basen / att::CreateExpr(16));
250- att::Expr mac = (perf.tile_cnt * perf.base_cnt * l1_cnt * perf.l0_cnt) * l0_mac;
251- std::cout << "mac: " << mac << std::endl;
252- 
253- att::Expr al1_mte2 =
254- (((ctx.expr_basem * ctx.expr_stepka) * att::CreateExpr(2)) /
255- (att::CreateExpr(32) / af::sym::Max(af::sym::kSymbolOne, (att::CreateExpr(256) / ctx.expr_stepka)))) +
256- att::CreateExpr(210);
257- att::Expr bl1_mte2 =
258- (((ctx.expr_stepkb * ctx.expr_basen) * att::CreateExpr(2)) /
259- (att::CreateExpr(32) / af::sym::Max(af::sym::kSymbolOne, (att::CreateExpr(256) / ctx.expr_basen)))) +
260- att::CreateExpr(210);
261- att::Expr mte2 = perf.tile_cnt * perf.base_cnt * perf.al1_cnt * (al1_mte2 + (perf.bl1_cnt * bl1_mte2));
262- std::cout << "mte2: " << mte2 << std::endl;
263- 
264- att::Expr base_fixpipe = ((ctx.expr_basem * ctx.expr_basen) * att::CreateExpr(4)) / att::CreateExpr(32);
265- att::Expr fixpipe = (perf.tile_cnt * perf.base_cnt) * base_fixpipe;
266- 
267- model_info.objects[att::PipeType::AIC_MAC] = mac;
268- model_info.objects[att::PipeType::AIC_MTE1] = mte1;
269- model_info.objects[att::PipeType::AIC_MTE2] = mte2;
270- model_info.objects[att::PipeType::AIC_FIXPIPE] = fixpipe;
271-}
272- 
273void FillModelInfo(att::ModelInfo &model_info, const MatmulExprContext &ctx) {122void FillModelInfo(att::ModelInfo &model_info, const MatmulExprContext &ctx) {
274 FillMatmulHardwareCons(model_info, ctx);123 FillMatmulHardwareCons(model_info, ctx);
275- MatmulPerfContext perf = CalcMatmulLoopCnts(ctx);
276- FillMatmulPerfObjects(model_info, ctx, perf);
277 124 
278 model_info.tiling_case_id = 1;125 model_info.tiling_case_id = 1;
279 model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepka, ctx.expr_stepkb));126 model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepka, ctx.expr_stepkb));
280- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepkb, ctx.expr_basek));
281- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_tilen, ctx.expr_basen));
282- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_tilem, ctx.expr_basem));
283- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_tilem - ctx.expr_m));
284- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_tilen - ctx.expr_n));
285 model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepka - ctx.expr_k));127 model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepka - ctx.expr_k));
286 model_info.container_exprs["Q1"] = (ctx.expr_m + ctx.expr_n);128 model_info.container_exprs["Q1"] = (ctx.expr_m + ctx.expr_n);
287 model_info.tensor_exprs["MATMUL_OUTPUT1"] = (ctx.expr_m + ctx.expr_n);129 model_info.tensor_exprs["MATMUL_OUTPUT1"] = (ctx.expr_m + ctx.expr_n);
@@ -296,22 +138,17 @@ ModelInfo GenMatmulModelInfo() {
296 BuildCoreAxis(model_info, ctx);138 BuildCoreAxis(model_info, ctx);
297 139 
298 AttAxisPtr m;140 AttAxisPtr m;
299- AttAxisPtr tilem;141+ BuildMAxes(ctx, m);
300- AttAxisPtr basem;
301- BuildMAxes(ctx, m, tilem, basem);
302 142 
303 AttAxisPtr n;143 AttAxisPtr n;
304- AttAxisPtr tilen;144+ BuildNAxes(ctx, n);
305- AttAxisPtr basen;
306- BuildNAxes(ctx, n, tilen, basen);
307 145 
308 AttAxisPtr k;146 AttAxisPtr k;
309 AttAxisPtr stepka;147 AttAxisPtr stepka;
310 AttAxisPtr stepkb;148 AttAxisPtr stepkb;
311- AttAxisPtr basek;149+ BuildKAxes(ctx, k, stepka, stepkb);
312- BuildKAxes(ctx, k, stepka, stepkb, basek);
313 150 
314- AppendArgList(model_info, m, tilen, tilem, stepka, stepkb, basek, basen, basem, n, k);151+ AppendArgList(model_info, m, n, k, stepka, stepkb);
315 FillModelInfo(model_info, ctx);152 FillModelInfo(model_info, ctx);
316 return model_info;153 return model_info;
317}154}
@@ -15,28 +15,19 @@
15namespace {15namespace {
16struct MatmulExprContext {16struct MatmulExprContext {
17 att::Expr expr_m;17 att::Expr expr_m;
18- att::Expr expr_tilem;
19 att::Expr expr_stepm;18 att::Expr expr_stepm;
20- att::Expr expr_basem;
21 att::Expr expr_n;19 att::Expr expr_n;
22- att::Expr expr_tilen;
23 att::Expr expr_stepn;20 att::Expr expr_stepn;
24- att::Expr expr_basen;
25 att::Expr expr_k;21 att::Expr expr_k;
26};22};
27 23 
28struct L2TileExprContext {24struct L2TileExprContext {
29 att::Expr expr_corenum;25 att::Expr expr_corenum;
30 att::Expr expr_m;26 att::Expr expr_m;
31- att::Expr expr_tilem;
32- att::Expr expr_basem;
33 att::Expr expr_n;27 att::Expr expr_n;
34- att::Expr expr_tilen;
35- att::Expr expr_basen;
36 att::Expr expr_k;28 att::Expr expr_k;
37 att::Expr expr_stepka;29 att::Expr expr_stepka;
38 att::Expr expr_stepkb;30 att::Expr expr_stepkb;
39- att::Expr expr_basek;
40};31};
41 32 
42void InitDefaultExpr(const ge::ExprType expr_type, att::Expr &default_expr, bool &is_const) {33void InitDefaultExpr(const ge::ExprType expr_type, att::Expr &default_expr, bool &is_const) {
@@ -91,52 +82,40 @@ void SetAxisInner(att::AttAxisPtr &axis, const std::string &name, const att::Sym
91void BuildCreateModelInfoMArgs(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr,82void BuildCreateModelInfoMArgs(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr,
92 MatmulExprContext &ctx) {83 MatmulExprContext &ctx) {
93 ctx.expr_m = is_const ? default_expr : att::CreateExpr("m_size");84 ctx.expr_m = is_const ? default_expr : att::CreateExpr("m_size");
94- ctx.expr_tilem = is_const ? default_expr : att::CreateExpr("tilem_size");
95 ctx.expr_stepm = is_const ? default_expr : att::CreateExpr("stepm_size");85 ctx.expr_stepm = is_const ? default_expr : att::CreateExpr("stepm_size");
96- ctx.expr_basem = att::CreateExpr("basem_size");
97 86 
98- att::SymVarInfoPtr sym_m, sym_tilem, sym_stepm, sym_basem;87+ att::SymVarInfoPtr sym_m;
88+ att::SymVarInfoPtr sym_stepm;
99 InitSymVar(sym_m, ctx.expr_m);89 InitSymVar(sym_m, ctx.expr_m);
100 sym_m->value_range.first = 1;90 sym_m->value_range.first = 1;
101 sym_m->value_range.second = 10000;91 sym_m->value_range.second = 10000;
102- InitSymVar(sym_tilem, ctx.expr_tilem, 16, {att::HardwareDef::L2});
103 InitSymVar(sym_stepm, ctx.expr_stepm, 16, {att::HardwareDef::L1, att::HardwareDef::CORENUM});92 InitSymVar(sym_stepm, ctx.expr_stepm, 16, {att::HardwareDef::L1, att::HardwareDef::CORENUM});
104- InitSymVar(sym_basem, ctx.expr_basem, 16, {att::HardwareDef::L0A, att::HardwareDef::L0C});
105 93 
106- att::AttAxisPtr m, tilem, stepm, basem;94+ att::AttAxisPtr m;
95+ att::AttAxisPtr stepm;
107 SetAxisOrigin(m, "m", sym_m);96 SetAxisOrigin(m, "m", sym_m);
108- SetAxisInner(tilem, "tilem", sym_tilem, false, false, m.get(), m.get());97+ SetAxisInner(stepm, "stepm", sym_stepm, true, true, m.get(), m.get());
109- SetAxisInner(stepm, "stepm", sym_stepm, true, false, m.get(), tilem.get());
110- SetAxisInner(basem, "basem", sym_basem, false, true, m.get(), stepm.get());
111 98 
112 model_info.arg_list.emplace_back(m);99 model_info.arg_list.emplace_back(m);
113- model_info.arg_list.emplace_back(tilem);
114 model_info.arg_list.emplace_back(stepm);100 model_info.arg_list.emplace_back(stepm);
115- model_info.arg_list.emplace_back(basem);
116}101}
117 102 
118void BuildCreateModelInfoNArgs(att::ModelInfo &model_info, MatmulExprContext &ctx) {103void BuildCreateModelInfoNArgs(att::ModelInfo &model_info, MatmulExprContext &ctx) {
119 ctx.expr_n = att::CreateExpr("n_size");104 ctx.expr_n = att::CreateExpr("n_size");
120- ctx.expr_tilen = att::CreateExpr("tilen_size");
121 ctx.expr_stepn = att::CreateExpr("stepn_size");105 ctx.expr_stepn = att::CreateExpr("stepn_size");
122- ctx.expr_basen = att::CreateExpr("basen_size");
123 106 
124- att::SymVarInfoPtr sym_n, sym_tilen, sym_stepn, sym_basen;107+ att::SymVarInfoPtr sym_n;
108+ att::SymVarInfoPtr sym_stepn;
125 InitSymVar(sym_n, ctx.expr_n);109 InitSymVar(sym_n, ctx.expr_n);
126- InitSymVar(sym_tilen, ctx.expr_tilen, 16, {att::HardwareDef::L2});
127 InitSymVar(sym_stepn, ctx.expr_stepn, 128, {att::HardwareDef::L1, att::HardwareDef::CORENUM});110 InitSymVar(sym_stepn, ctx.expr_stepn, 128, {att::HardwareDef::L1, att::HardwareDef::CORENUM});
128- InitSymVar(sym_basen, ctx.expr_basen, 16, {att::HardwareDef::L0B, att::HardwareDef::L0C});
129 111 
130- att::AttAxisPtr n, tilen, stepn, basen;112+ att::AttAxisPtr n;
113+ att::AttAxisPtr stepn;
131 SetAxisOrigin(n, "n", sym_n);114 SetAxisOrigin(n, "n", sym_n);
132- SetAxisInner(tilen, "tilen", sym_tilen, false, false, n.get(), n.get());115+ SetAxisInner(stepn, "stepn", sym_stepn, true, true, n.get(), n.get());
133- SetAxisInner(stepn, "stepn", sym_stepn, true, false, n.get(), tilen.get());
134- SetAxisInner(basen, "basen", sym_basen, false, true, n.get(), stepn.get());
135 116 
136 model_info.arg_list.emplace_back(n);117 model_info.arg_list.emplace_back(n);
137- model_info.arg_list.emplace_back(tilen);
138 model_info.arg_list.emplace_back(stepn);118 model_info.arg_list.emplace_back(stepn);
139- model_info.arg_list.emplace_back(basen);
140}119}
141 120 
142void BuildCreateModelInfoKArg(att::ModelInfo &model_info, MatmulExprContext &ctx) {121void BuildCreateModelInfoKArg(att::ModelInfo &model_info, MatmulExprContext &ctx) {
@@ -154,35 +133,20 @@ void BuildCreateModelInfoKArg(att::ModelInfo &model_info, MatmulExprContext &ctx
154}133}
155 134 
156void FillCreateModelInfo(att::ModelInfo &model_info, const MatmulExprContext &ctx) {135void FillCreateModelInfo(att::ModelInfo &model_info, const MatmulExprContext &ctx) {
157- att::Expr l0a_occupy = ctx.expr_basem * ctx.expr_k * att::CreateExpr(4);
158- att::Expr l0b_occupy = ctx.expr_k * ctx.expr_basen * att::CreateExpr(4);
159- att::Expr l0c_occupy = ctx.expr_basem * ctx.expr_basen * att::CreateExpr(4);
160 att::Expr l1_occupy =136 att::Expr l1_occupy =
161 (ctx.expr_k * ctx.expr_stepm * att::CreateExpr(4)) + (ctx.expr_k * ctx.expr_stepn * att::CreateExpr(4));137 (ctx.expr_k * ctx.expr_stepm * att::CreateExpr(4)) + (ctx.expr_k * ctx.expr_stepn * att::CreateExpr(4));
162- att::Expr l2_occupy = (ctx.expr_tilen * ctx.expr_tilem * att::CreateExpr(2)) +138+ att::Expr core_num = (ctx.expr_m / ctx.expr_stepm) * (ctx.expr_n / ctx.expr_stepn);
163- ((ctx.expr_tilen + ctx.expr_tilem) * ctx.expr_k * att::CreateExpr(2));
164- att::Expr core_num = (ctx.expr_tilem / ctx.expr_stepm) * (ctx.expr_tilen / ctx.expr_stepn);
165 139 
166- model_info.hardware_cons[att::HardwareDef::L0A] = l0a_occupy;
167- model_info.hardware_cons[att::HardwareDef::L0B] = l0b_occupy;
168- model_info.hardware_cons[att::HardwareDef::L0C] = l0c_occupy;
169 model_info.hardware_cons[att::HardwareDef::L1] = l1_occupy;140 model_info.hardware_cons[att::HardwareDef::L1] = l1_occupy;
170- model_info.hardware_cons[att::HardwareDef::L2] = l2_occupy;
171 model_info.hardware_cons[att::HardwareDef::CORENUM] = core_num;141 model_info.hardware_cons[att::HardwareDef::CORENUM] = core_num;
172 model_info.hardware_cons[att::HardwareDef::UB] = att::CreateExpr(0L);142 model_info.hardware_cons[att::HardwareDef::UB] = att::CreateExpr(0L);
173 143 
174- att::Expr mac = (ctx.expr_basem * ctx.expr_basen * ctx.expr_k) / (att::CreateExpr(16) * att::CreateExpr(256));
175 att::Expr mte =144 att::Expr mte =
176 (((ctx.expr_stepm * ctx.expr_k) / att::CreateExpr(32)) + ((ctx.expr_stepn * ctx.expr_k) / att::CreateExpr(32)));145 (((ctx.expr_stepm * ctx.expr_k) / att::CreateExpr(32)) + ((ctx.expr_stepn * ctx.expr_k) / att::CreateExpr(32)));
177- model_info.objects[att::PipeType::AIC_MAC] = mac;
178 model_info.objects[att::PipeType::AIC_MTE2] = mte;146 model_info.objects[att::PipeType::AIC_MTE2] = mte;
179 model_info.tiling_case_id = 0;147 model_info.tiling_case_id = 0;
180- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepm, ctx.expr_basem));148+ model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepm - ctx.expr_m));
181- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepn, ctx.expr_basen));149+ model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepn - ctx.expr_n));
182- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_tilem - ctx.expr_m));
183- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepm - ctx.expr_tilem));
184- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_tilen - ctx.expr_n));
185- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepn - ctx.expr_tilen));
186 model_info.output_size = 1;150 model_info.output_size = 1;
187}151}
188 152 
@@ -201,20 +165,10 @@ void BuildCoreAxis(att::ModelInfo &model_info, L2TileExprContext &ctx) {
201 165 
202void BuildL2TileMArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) {166void BuildL2TileMArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) {
203 ctx.expr_m = att::CreateExpr("m_size");167 ctx.expr_m = att::CreateExpr("m_size");
204- ctx.expr_tilem = att::CreateExpr("tilem_size");
205- ctx.expr_basem = att::CreateExpr("basem_size");
206 168 
207 att::SymVarInfoPtr sym_m = std::make_shared<att::SymVarInfo>(ctx.expr_m);169 att::SymVarInfoPtr sym_m = std::make_shared<att::SymVarInfo>(ctx.expr_m);
208- att::SymVarInfoPtr sym_tilem = std::make_shared<att::SymVarInfo>(ctx.expr_tilem);
209- sym_tilem->align = ge::Symbol(16);
210- sym_tilem->related_scope = {att::HardwareDef::L2};
211- att::SymVarInfoPtr sym_basem = std::make_shared<att::SymVarInfo>(ctx.expr_basem);
212- sym_basem->align = ge::Symbol(16);
213- sym_basem->related_scope = {att::HardwareDef::L0A, att::HardwareDef::L0C, att::HardwareDef::L1};
214 170 
215 att::AttAxisPtr m = std::make_shared<att::AttAxis>();171 att::AttAxisPtr m = std::make_shared<att::AttAxis>();
216- att::AttAxisPtr tilem = std::make_shared<att::AttAxis>();
217- att::AttAxisPtr basem = std::make_shared<att::AttAxis>();
218 172 
219 m->name = "m";173 m->name = "m";
220 m->axis_pos = att::AxisPosition::ORIGIN;174 m->axis_pos = att::AxisPosition::ORIGIN;
@@ -223,45 +177,15 @@ void BuildL2TileMArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) {
223 m->is_node_innerest_dim = false;177 m->is_node_innerest_dim = false;
224 m->size = sym_m;178 m->size = sym_m;
225 179 
226- tilem->name = "tilem";
227- tilem->axis_pos = att::AxisPosition::INNER;
228- tilem->bind_multicore = false;
229- tilem->is_last = false;
230- tilem->is_node_innerest_dim = true;
231- tilem->size = sym_tilem;
232- tilem->orig_axis.push_back(m.get());
233- tilem->from_axis = {m.get()};
234- 
235- basem->name = "basem";
236- basem->axis_pos = att::AxisPosition::INNER;
237- basem->bind_multicore = false;
238- basem->is_last = true;
239- basem->is_node_innerest_dim = false;
240- basem->size = sym_basem;
241- basem->orig_axis.push_back(m.get());
242- basem->from_axis = {tilem.get()};
243- 
244 model_info.arg_list.emplace_back(m);180 model_info.arg_list.emplace_back(m);
245- model_info.arg_list.emplace_back(tilem);
246- model_info.arg_list.emplace_back(basem);
247}181}
248 182 
249void BuildL2TileNArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) {183void BuildL2TileNArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) {
250 ctx.expr_n = att::CreateExpr("n_size");184 ctx.expr_n = att::CreateExpr("n_size");
251- ctx.expr_tilen = att::CreateExpr("tilen_size");
252- ctx.expr_basen = att::CreateExpr("basen_size");
253 185 
254 att::SymVarInfoPtr sym_n = std::make_shared<att::SymVarInfo>(ctx.expr_n);186 att::SymVarInfoPtr sym_n = std::make_shared<att::SymVarInfo>(ctx.expr_n);
255- att::SymVarInfoPtr sym_tilen = std::make_shared<att::SymVarInfo>(ctx.expr_tilen);
256- sym_tilen->align = ge::Symbol(16);
257- sym_tilen->related_scope = {att::HardwareDef::L2};
258- att::SymVarInfoPtr sym_basen = std::make_shared<att::SymVarInfo>(ctx.expr_basen);
259- sym_basen->align = ge::Symbol(16);
260- sym_basen->related_scope = {att::HardwareDef::L0B, att::HardwareDef::L0C, att::HardwareDef::L1};
261 187 
262 att::AttAxisPtr n = std::make_shared<att::AttAxis>();188 att::AttAxisPtr n = std::make_shared<att::AttAxis>();
263- att::AttAxisPtr tilen = std::make_shared<att::AttAxis>();
264- att::AttAxisPtr basen = std::make_shared<att::AttAxis>();
265 189 
266 n->name = "n";190 n->name = "n";
267 n->axis_pos = att::AxisPosition::ORIGIN;191 n->axis_pos = att::AxisPosition::ORIGIN;
@@ -270,112 +194,43 @@ void BuildL2TileNArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) {
270 n->is_node_innerest_dim = false;194 n->is_node_innerest_dim = false;
271 n->size = sym_n;195 n->size = sym_n;
272 196 
273- tilen->name = "tilen";
274- tilen->axis_pos = att::AxisPosition::INNER;
275- tilen->bind_multicore = false;
276- tilen->is_last = false;
277- tilen->is_node_innerest_dim = true;
278- tilen->size = sym_tilen;
279- tilen->orig_axis.push_back(n.get());
280- tilen->from_axis = {n.get()};
281- 
282- basen->name = "basen";
283- basen->axis_pos = att::AxisPosition::INNER;
284- basen->bind_multicore = false;
285- basen->is_last = true;
286- basen->is_node_innerest_dim = true;
287- basen->size = sym_basen;
288- basen->orig_axis.push_back(n.get());
289- basen->from_axis = {tilen.get()};
290- 
291 model_info.arg_list.emplace_back(n);197 model_info.arg_list.emplace_back(n);
292- model_info.arg_list.emplace_back(tilen);
293- model_info.arg_list.emplace_back(basen);
294}198}
295 199 
296void BuildL2TileKArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) {200void BuildL2TileKArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) {
297 ctx.expr_k = att::CreateExpr("k_size");201 ctx.expr_k = att::CreateExpr("k_size");
298 ctx.expr_stepka = att::CreateExpr("stepka_size");202 ctx.expr_stepka = att::CreateExpr("stepka_size");
299 ctx.expr_stepkb = att::CreateExpr("stepkb_size");203 ctx.expr_stepkb = att::CreateExpr("stepkb_size");
300- ctx.expr_basek = att::CreateExpr("basek_size");
301 204 
302- att::SymVarInfoPtr sym_k, sym_stepka, sym_stepkb, sym_basek;205+ att::SymVarInfoPtr sym_k;
206+ att::SymVarInfoPtr sym_stepka;
207+ att::SymVarInfoPtr sym_stepkb;
303 InitSymVar(sym_k, ctx.expr_k);208 InitSymVar(sym_k, ctx.expr_k);
304 InitSymVar(sym_stepka, ctx.expr_stepka, 256, {att::HardwareDef::L1});209 InitSymVar(sym_stepka, ctx.expr_stepka, 256, {att::HardwareDef::L1});
305 InitSymVar(sym_stepkb, ctx.expr_stepkb, 16, {att::HardwareDef::L1});210 InitSymVar(sym_stepkb, ctx.expr_stepkb, 16, {att::HardwareDef::L1});
306- InitSymVar(sym_basek, ctx.expr_basek, 16, {att::HardwareDef::L0A, att::HardwareDef::L0B});
307 211 
308- att::AttAxisPtr k, stepka, stepkb, basek;212+ att::AttAxisPtr k;
213+ att::AttAxisPtr stepka;
214+ att::AttAxisPtr stepkb;
309 SetAxisOrigin(k, "k", sym_k);215 SetAxisOrigin(k, "k", sym_k);
310 SetAxisInner(stepka, "stepka", sym_stepka, false, false, k.get(), k.get());216 SetAxisInner(stepka, "stepka", sym_stepka, false, false, k.get(), k.get());
311- SetAxisInner(stepkb, "stepkb", sym_stepkb, false, false, k.get(), stepka.get());217+ SetAxisInner(stepkb, "stepkb", sym_stepkb, false, true, k.get(), stepka.get());
312- // basek is special: is_node_innerest_dim = false
313- basek = std::make_shared<att::AttAxis>();
314- basek->name = "basek";
315- basek->axis_pos = att::AxisPosition::INNER;
316- basek->bind_multicore = false;
317- basek->is_last = true;
318- basek->is_node_innerest_dim = false;
319- basek->size = sym_basek;
320- basek->orig_axis.push_back(k.get());
321- basek->from_axis = {stepkb.get()};
322 218 
323 model_info.arg_list.emplace_back(k);219 model_info.arg_list.emplace_back(k);
324 model_info.arg_list.emplace_back(stepka);220 model_info.arg_list.emplace_back(stepka);
325 model_info.arg_list.emplace_back(stepkb);221 model_info.arg_list.emplace_back(stepkb);
326- model_info.arg_list.emplace_back(basek);
327}222}
328 223 
329void FillL2TileHardwareCons(att::ModelInfo &model_info, const L2TileExprContext &ctx) {224void FillL2TileHardwareCons(att::ModelInfo &model_info, const L2TileExprContext &ctx) {
330- att::Expr l0a_occupy = ctx.expr_basem * ctx.expr_basek * att::CreateExpr(4);225+ att::Expr l1_occupy = ctx.expr_stepka * ctx.expr_stepkb * att::CreateExpr(4);
331- att::Expr l0b_occupy = ctx.expr_basek * ctx.expr_basen * att::CreateExpr(4);
332- att::Expr l0c_occupy = ctx.expr_basem * ctx.expr_basen * att::CreateExpr(4);
333- att::Expr l1_occupy =
334- (ctx.expr_stepka * ctx.expr_basem * att::CreateExpr(4)) + (ctx.expr_stepkb * ctx.expr_basen * att::CreateExpr(4));
335- att::Expr l2_occupy = (ctx.expr_tilen * ctx.expr_tilem * att::CreateExpr(2)) +
336- ((ctx.expr_tilen + ctx.expr_tilem) * ctx.expr_k * att::CreateExpr(2));
337- model_info.hardware_cons[att::HardwareDef::L0A] = l0a_occupy;
338- model_info.hardware_cons[att::HardwareDef::L0B] = l0b_occupy;
339- model_info.hardware_cons[att::HardwareDef::L0C] = l0c_occupy;
340 model_info.hardware_cons[att::HardwareDef::L1] = l1_occupy;226 model_info.hardware_cons[att::HardwareDef::L1] = l1_occupy;
341- model_info.hardware_cons[att::HardwareDef::L2] = l2_occupy;
342 model_info.hardware_cons[att::HardwareDef::UB] = att::CreateExpr(0L);227 model_info.hardware_cons[att::HardwareDef::UB] = att::CreateExpr(0L);
343}228}
344 229 
345void FillL2TileModelInfo(att::ModelInfo &model_info, const L2TileExprContext &ctx) {230void FillL2TileModelInfo(att::ModelInfo &model_info, const L2TileExprContext &ctx) {
346 FillL2TileHardwareCons(model_info, ctx);231 FillL2TileHardwareCons(model_info, ctx);
347- att::Expr tile_cnt = ((ctx.expr_n / ctx.expr_tilen) * (ctx.expr_m / ctx.expr_tilem));
348- att::Expr base_cnt = af::sym::Max(
349- af::sym::kSymbolOne,
350- (((ctx.expr_tilem * ctx.expr_tilen) / (ctx.expr_basem * ctx.expr_basen)) / att::CreateExpr("block_dim")));
351- att::Expr al1_cnt = (ctx.expr_k / ctx.expr_stepka);
352- att::Expr bl1_cnt = (ctx.expr_stepka / ctx.expr_stepkb);
353- att::Expr l0_cnt = (ctx.expr_stepkb / ctx.expr_basek);
354- att::Expr l1_cnt = (al1_cnt * bl1_cnt);
355- att::Expr base_fixpipe_cost = ((ctx.expr_basem * ctx.expr_basen * att::CreateExpr(4)) / att::CreateExpr(32));
356- att::Expr al1_mte2 =
357- (((ctx.expr_basem * ctx.expr_stepka * att::CreateExpr(2)) /
358- (att::CreateExpr(32) / af::sym::Max(af::sym::kSymbolOne, (att::CreateExpr(256) / ctx.expr_stepka)))) +
359- att::CreateExpr(210));
360- att::Expr bl1_mte2 =
361- (((ctx.expr_basen * ctx.expr_stepkb * att::CreateExpr(2)) /
362- (att::CreateExpr(32) / af::sym::Max(af::sym::kSymbolOne, (att::CreateExpr(256) / ctx.expr_basen)))) +
363- att::CreateExpr(210));
364- att::Expr mac = (((tile_cnt * base_cnt * l1_cnt * l0_cnt)) * (ctx.expr_basem * ctx.expr_basen * ctx.expr_k) /
365- (att::CreateExpr(16) * att::CreateExpr(256)));
366- att::Expr mte2 = (tile_cnt * base_cnt * al1_cnt * (al1_mte2 + (bl1_cnt * bl1_mte2)));
367- att::Expr fixpipe = (tile_cnt * base_cnt * base_fixpipe_cost);
368- 
369- model_info.objects[att::PipeType::AIC_MAC] = mac;
370- model_info.objects[att::PipeType::AIC_MTE2] = mte2;
371- model_info.objects[att::PipeType::AIC_FIXPIPE] = fixpipe;
372 model_info.tiling_case_id = 1;232 model_info.tiling_case_id = 1;
373 model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepka, ctx.expr_stepkb));233 model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepka, ctx.expr_stepkb));
374- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_tilen, ctx.expr_basen));
375- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_tilem, ctx.expr_basem));
376- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepkb, ctx.expr_basek));
377- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_tilem - ctx.expr_m));
378- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_tilen - ctx.expr_n));
379 model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepka - ctx.expr_k));234 model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepka - ctx.expr_k));
380 model_info.container_exprs["Q1"] = (ctx.expr_m + ctx.expr_n);235 model_info.container_exprs["Q1"] = (ctx.expr_m + ctx.expr_n);
381 model_info.tensor_exprs["MATMUL_OUTPUT1"] = (ctx.expr_m + ctx.expr_n);236 model_info.tensor_exprs["MATMUL_OUTPUT1"] = (ctx.expr_m + ctx.expr_n);
@@ -13,13 +13,9 @@
13namespace {13namespace {
14struct SolverExprContext {14struct SolverExprContext {
15 att::Expr expr_m;15 att::Expr expr_m;
16- att::Expr expr_tilem;
17 att::Expr expr_stepm;16 att::Expr expr_stepm;
18- att::Expr expr_basem;
19 att::Expr expr_n;17 att::Expr expr_n;
20- att::Expr expr_tilen;
21 att::Expr expr_stepn;18 att::Expr expr_stepn;
22- att::Expr expr_basen;
23 att::Expr expr_k;19 att::Expr expr_k;
24};20};
25 21 
@@ -75,54 +71,42 @@ void SetAxisInner(att::AttAxisPtr &axis, const std::string &name, const att::Sym
75void BuildMArgList(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr,71void BuildMArgList(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr,
76 const uint32_t m_align, SolverExprContext &ctx) {72 const uint32_t m_align, SolverExprContext &ctx) {
77 ctx.expr_m = is_const ? default_expr : att::CreateExpr("m_size");73 ctx.expr_m = is_const ? default_expr : att::CreateExpr("m_size");
78- ctx.expr_tilem = is_const ? default_expr : att::CreateExpr("tilem_size");
79 ctx.expr_stepm = is_const ? default_expr : att::CreateExpr("stepm_size");74 ctx.expr_stepm = is_const ? default_expr : att::CreateExpr("stepm_size");
80- ctx.expr_basem = att::CreateExpr("basem_size");
81 75 
82- att::SymVarInfoPtr sym_m, sym_tilem, sym_stepm, sym_basem;76+ att::SymVarInfoPtr sym_m;
77+ att::SymVarInfoPtr sym_stepm;
83 InitSymVar(sym_m, ctx.expr_m);78 InitSymVar(sym_m, ctx.expr_m);
84 sym_m->value_range.first = 1;79 sym_m->value_range.first = 1;
85 sym_m->value_range.second = 10000;80 sym_m->value_range.second = 10000;
86 sym_m->align = ge::Symbol(m_align);81 sym_m->align = ge::Symbol(m_align);
87- InitSymVar(sym_tilem, ctx.expr_tilem, 16, {att::HardwareDef::L2});
88 InitSymVar(sym_stepm, ctx.expr_stepm, 16, {att::HardwareDef::L1, att::HardwareDef::CORENUM});82 InitSymVar(sym_stepm, ctx.expr_stepm, 16, {att::HardwareDef::L1, att::HardwareDef::CORENUM});
89- InitSymVar(sym_basem, ctx.expr_basem, 16, {att::HardwareDef::L0A, att::HardwareDef::L0C});
90 83 
91- att::AttAxisPtr m, tilem, stepm, basem;84+ att::AttAxisPtr m;
85+ att::AttAxisPtr stepm;
92 SetAxisOrigin(m, "m", sym_m);86 SetAxisOrigin(m, "m", sym_m);
93- SetAxisInner(tilem, "tilem", sym_tilem, false, false, m.get(), m.get());87+ SetAxisInner(stepm, "stepm", sym_stepm, true, true, m.get(), m.get());
94- SetAxisInner(stepm, "stepm", sym_stepm, true, false, m.get(), tilem.get());
95- SetAxisInner(basem, "basem", sym_basem, false, true, m.get(), stepm.get());
96 88 
97 model_info.arg_list.emplace_back(m);89 model_info.arg_list.emplace_back(m);
98- model_info.arg_list.emplace_back(tilem);
99 model_info.arg_list.emplace_back(stepm);90 model_info.arg_list.emplace_back(stepm);
100- model_info.arg_list.emplace_back(basem);
101}91}
102 92 
103void BuildNArgList(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr,93void BuildNArgList(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr,
104 SolverExprContext &ctx) {94 SolverExprContext &ctx) {
105 ctx.expr_n = is_const ? default_expr : att::CreateExpr("n_size");95 ctx.expr_n = is_const ? default_expr : att::CreateExpr("n_size");
106- ctx.expr_tilen = is_const ? default_expr : att::CreateExpr("tilen_size");
107 ctx.expr_stepn = is_const ? default_expr : att::CreateExpr("stepn_size");96 ctx.expr_stepn = is_const ? default_expr : att::CreateExpr("stepn_size");
108- ctx.expr_basen = is_const ? default_expr : att::CreateExpr("basen_size");
109 97 
110- att::SymVarInfoPtr sym_n, sym_tilen, sym_stepn, sym_basen;98+ att::SymVarInfoPtr sym_n;
99+ att::SymVarInfoPtr sym_stepn;
111 InitSymVar(sym_n, ctx.expr_n);100 InitSymVar(sym_n, ctx.expr_n);
112- InitSymVar(sym_tilen, ctx.expr_tilen, 16, {att::HardwareDef::L2});
113 InitSymVar(sym_stepn, ctx.expr_stepn, 128, {att::HardwareDef::L1, att::HardwareDef::CORENUM});101 InitSymVar(sym_stepn, ctx.expr_stepn, 128, {att::HardwareDef::L1, att::HardwareDef::CORENUM});
114- InitSymVar(sym_basen, ctx.expr_basen, 16, {att::HardwareDef::L0B, att::HardwareDef::L0C});
115 102 
116- att::AttAxisPtr n, tilen, stepn, basen;103+ att::AttAxisPtr n;
104+ att::AttAxisPtr stepn;
117 SetAxisOrigin(n, "n", sym_n);105 SetAxisOrigin(n, "n", sym_n);
118- SetAxisInner(tilen, "tilen", sym_tilen, false, false, n.get(), n.get());106+ SetAxisInner(stepn, "stepn", sym_stepn, true, true, n.get(), n.get());
119- SetAxisInner(stepn, "stepn", sym_stepn, true, false, n.get(), tilen.get());
120- SetAxisInner(basen, "basen", sym_basen, false, true, n.get(), stepn.get());
121 107 
122 model_info.arg_list.emplace_back(n);108 model_info.arg_list.emplace_back(n);
123- model_info.arg_list.emplace_back(tilen);
124 model_info.arg_list.emplace_back(stepn);109 model_info.arg_list.emplace_back(stepn);
125- model_info.arg_list.emplace_back(basen);
126}110}
127 111 
128void BuildKArg(att::ModelInfo &model_info, SolverExprContext &ctx) {112void BuildKArg(att::ModelInfo &model_info, SolverExprContext &ctx) {
@@ -140,33 +124,20 @@ void BuildKArg(att::ModelInfo &model_info, SolverExprContext &ctx) {
140}124}
141 125 
142void FillModelInfo(att::ModelInfo &model_info, const SolverExprContext &ctx) {126void FillModelInfo(att::ModelInfo &model_info, const SolverExprContext &ctx) {
143- att::Expr l0a_occupy = ctx.expr_basem * ctx.expr_k * att::CreateExpr(4);
144- att::Expr l0b_occupy = ctx.expr_k * ctx.expr_basen * att::CreateExpr(4);
145- att::Expr l0c_occupy = ctx.expr_basem * ctx.expr_basen * att::CreateExpr(4);
146 att::Expr l1_occupy =127 att::Expr l1_occupy =
147 (ctx.expr_k * ctx.expr_stepm * att::CreateExpr(4)) + (ctx.expr_k * ctx.expr_stepn * att::CreateExpr(4));128 (ctx.expr_k * ctx.expr_stepm * att::CreateExpr(4)) + (ctx.expr_k * ctx.expr_stepn * att::CreateExpr(4));
148- att::Expr l2_occupy = (ctx.expr_tilen * ctx.expr_tilem * att::CreateExpr(2)) +129+ att::Expr core_num = ((ctx.expr_m / ctx.expr_stepm) * (ctx.expr_n / ctx.expr_stepn));
149- ((ctx.expr_tilen + ctx.expr_tilem) * ctx.expr_k * att::CreateExpr(2));
150- att::Expr core_num = ((ctx.expr_tilem / ctx.expr_stepm) * (ctx.expr_tilen / ctx.expr_stepn));
151 130 
152- model_info.hardware_cons[att::HardwareDef::L0A] = l0a_occupy;
153- model_info.hardware_cons[att::HardwareDef::L0B] = l0b_occupy;
154- model_info.hardware_cons[att::HardwareDef::L0C] = l0c_occupy;
155 model_info.hardware_cons[att::HardwareDef::L1] = l1_occupy;131 model_info.hardware_cons[att::HardwareDef::L1] = l1_occupy;
156- model_info.hardware_cons[att::HardwareDef::L2] = l2_occupy;
157 model_info.hardware_cons[att::HardwareDef::UB] = ctx.expr_m * att::CreateExpr(10);132 model_info.hardware_cons[att::HardwareDef::UB] = ctx.expr_m * att::CreateExpr(10);
158 model_info.hardware_cons[att::HardwareDef::CORENUM] = core_num;133 model_info.hardware_cons[att::HardwareDef::CORENUM] = core_num;
159 134 
160- att::Expr mac = ((ctx.expr_basem * ctx.expr_basen * ctx.expr_k) / (att::CreateExpr(16) * att::CreateExpr(256)));
161 att::Expr mte =135 att::Expr mte =
162 (((ctx.expr_stepm * ctx.expr_k) / att::CreateExpr(32)) + ((ctx.expr_stepn * ctx.expr_k) / att::CreateExpr(32)));136 (((ctx.expr_stepm * ctx.expr_k) / att::CreateExpr(32)) + ((ctx.expr_stepn * ctx.expr_k) / att::CreateExpr(32)));
163- model_info.objects[att::PipeType::AIC_MAC] = mac;
164 model_info.objects[att::PipeType::AIC_MTE2] = mte;137 model_info.objects[att::PipeType::AIC_MTE2] = mte;
165 model_info.tiling_case_id = 0;138 model_info.tiling_case_id = 0;
166- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepm, ctx.expr_basem));139+ model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepm - ctx.expr_m));
167- model_info.eq_exprs[att::kFatherToChildNoTail].push_back(std::pair(ctx.expr_stepn, ctx.expr_basen));140+ model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_stepn - ctx.expr_n));
168- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_tilem - ctx.expr_stepm));
169- model_info.leq_exprs[att::kFatherToChildLarger].push_back((ctx.expr_tilen - ctx.expr_stepn));
170 model_info.container_exprs["Q1"] = (ctx.expr_m + ctx.expr_n);141 model_info.container_exprs["Q1"] = (ctx.expr_m + ctx.expr_n);
171 model_info.tensor_exprs["MATMUL_OUTPUT1"] = (ctx.expr_m + ctx.expr_n);142 model_info.tensor_exprs["MATMUL_OUTPUT1"] = (ctx.expr_m + ctx.expr_n);
172 model_info.output_size = 1;143 model_info.output_size = 1;
@@ -21,10 +21,6 @@ file(GLOB SOURCES
21 ${TC_DIR}/test_expr/*.cpp21 ${TC_DIR}/test_expr/*.cpp
22 ${TC_DIR}/solver_pass_gen/axes_reorder_solver_gen/*.cpp22 ${TC_DIR}/solver_pass_gen/axes_reorder_solver_gen/*.cpp
23 ${TC_DIR}/solver_pass_gen/general_solver_gen/*.cpp23 ${TC_DIR}/solver_pass_gen/general_solver_gen/*.cpp
24- ${TC_DIR}/solver_pass_gen/l2_solver_gen/*.cpp
25- ${TC_DIR}/solver_pass_gen/l0_solver_gen/*.cpp
26- ${TC_DIR}/solver_pass/l0_solver/*.cpp
27- ${TC_DIR}/solver_pass/l2_solver/*.cpp
28 ${TC_DIR}/solver_pass/general_solver/*.cpp24 ${TC_DIR}/solver_pass/general_solver/*.cpp
29 ${TC_DIR}/solver_pass_manager/*.cpp25 ${TC_DIR}/solver_pass_manager/*.cpp
30 ${TC_DIR}/select_model/*.cpp26 ${TC_DIR}/select_model/*.cpp
@@ -17,9 +17,6 @@ int TestCase1(uint64_t m, uint64_t n, int32_t tilingCaseId) {
17 tilingData.set_n_size(n);17 tilingData.set_n_size(n);
18 tilingData.set_block_dim(20);18 tilingData.set_block_dim(20);
19 tilingData.set_l1_size(512 * 1024);19 tilingData.set_l1_size(512 * 1024);
20- tilingData.set_l0a_size(64 * 1024);
21- tilingData.set_l0b_size(64 * 1024);
22- tilingData.set_l0c_size(128 * 1024);
23 // tilingData.z = 0;20 // tilingData.z = 0;
24 const auto status = GetTiling(tilingData, tilingCaseId);21 const auto status = GetTiling(tilingData, tilingCaseId);
25 if ((status)) {22 if ((status)) {
@@ -20,7 +20,6 @@ bool TestCase(std::vector<int64_t> shapes) {
20 tilingData.set_block_dim(48);20 tilingData.set_block_dim(48);
21 tilingData.set_ub_size(240 * 1024);21 tilingData.set_ub_size(240 * 1024);
22 tilingData.set_btbuf_size(1 * 1024);22 tilingData.set_btbuf_size(1 * 1024);
23- tilingData.set_l0c_size(128 * 1024);
24 tilingData.set_maxTokens(maxTokens);23 tilingData.set_maxTokens(maxTokens);
25 tilingData.set_N1(N1);24 tilingData.set_N1(N1);
26 tilingData.set_K1(K1);25 tilingData.set_K1(K1);
@@ -33,8 +32,6 @@ bool TestCase(std::vector<int64_t> shapes) {
33 const auto status = GetTiling(tilingData, 0u);32 const auto status = GetTiling(tilingData, 0u);
34 if ((status)) {33 if ((status)) {
35 std::cout << "ub_m" << " = " << tilingData.get_ub_m() << std::endl;34 std::cout << "ub_m" << " = " << tilingData.get_ub_m() << std::endl;
36- std::cout << "base_m1" << " = " << tilingData.get_base_m1() << std::endl;
37- std::cout << "base_m2" << " = " << tilingData.get_base_m2() << std::endl;
38 std::cout << "base_n1" << " = " << tilingData.get_base_n1() << std::endl;35 std::cout << "base_n1" << " = " << tilingData.get_base_n1() << std::endl;
39 std::cout << "base_n2" << " = " << tilingData.get_base_n2() << std::endl;36 std::cout << "base_n2" << " = " << tilingData.get_base_n2() << std::endl;
40 return true;37 return true;
@@ -17,23 +17,12 @@ int TestCase0() {
17 tilingData.set_n_size(2048);17 tilingData.set_n_size(2048);
18 tilingData.set_block_dim(20);18 tilingData.set_block_dim(20);
19 tilingData.set_l1_size(512 * 1024);19 tilingData.set_l1_size(512 * 1024);
20- tilingData.set_l0a_size(64 * 1024);
21- tilingData.set_l0b_size(64 * 1024);
22- tilingData.set_l0c_size(128 * 1024);
23 // tilingData.z = 0;20 // tilingData.z = 0;
24 const auto status = GetTiling(tilingData, 0);21 const auto status = GetTiling(tilingData, 0);
25 if ((status)) {22 if ((status)) {
26- std::cout << "basem" << " = " << tilingData.get_basem_size() << std::endl;
27- std::cout << "basen" << " = " << tilingData.get_basen_size() << std::endl;
28- std::cout << "tilem" << " = " << tilingData.get_tilem_size() << std::endl;
29- std::cout << "tilen" << " = " << tilingData.get_tilen_size() << std::endl;
30 std::cout << "stepn" << " = " << tilingData.get_stepn_size() << std::endl;23 std::cout << "stepn" << " = " << tilingData.get_stepn_size() << std::endl;
31 std::cout << "stepm" << " = " << tilingData.get_stepm_size() << std::endl;24 std::cout << "stepm" << " = " << tilingData.get_stepm_size() << std::endl;
32 std::cout << "tiling_key" << " = " << tilingData.get_tiling_key() << std::endl;25 std::cout << "tiling_key" << " = " << tilingData.get_tiling_key() << std::endl;
33- if ((tilingData.get_basem_size() != 128) || (tilingData.get_stepn_size() < 256)) {
34- std::cout << "Case0 tiling func execute failed." << std::endl;
35- return -1;
36- }
37 std::cout << "Case0 tiling func execute success." << std::endl;26 std::cout << "Case0 tiling func execute success." << std::endl;
38 return 0;27 return 0;
39 }28 }
@@ -48,25 +37,12 @@ int TestCase1() {
48 tilingData.set_k_size(2048);37 tilingData.set_k_size(2048);
49 tilingData.set_block_dim(20);38 tilingData.set_block_dim(20);
50 tilingData.set_l1_size(512 * 1024);39 tilingData.set_l1_size(512 * 1024);
51- tilingData.set_l0a_size(64 * 1024);
52- tilingData.set_l0b_size(64 * 1024);
53- tilingData.set_l0c_size(128 * 1024);
54 // tilingData.z = 0;40 // tilingData.z = 0;
55 const auto status = GetTiling(tilingData, 1);41 const auto status = GetTiling(tilingData, 1);
56 if ((status)) {42 if ((status)) {
57- std::cout << "basem" << " = " << tilingData.get_basem_size() << std::endl;
58- std::cout << "basen" << " = " << tilingData.get_basen_size() << std::endl;
59- std::cout << "tilem" << " = " << tilingData.get_tilem_size() << std::endl;
60- std::cout << "tilen" << " = " << tilingData.get_tilen_size() << std::endl;
61- std::cout << "stepn" << " = " << tilingData.get_stepn_size() << std::endl;
62- std::cout << "stepm" << " = " << tilingData.get_stepm_size() << std::endl;
63 std::cout << "stepka" << " = " << tilingData.get_stepka_size() << std::endl;43 std::cout << "stepka" << " = " << tilingData.get_stepka_size() << std::endl;
64 std::cout << "stepkb" << " = " << tilingData.get_stepkb_size() << std::endl;44 std::cout << "stepkb" << " = " << tilingData.get_stepkb_size() << std::endl;
65 std::cout << "tiling_key" << " = " << tilingData.get_tiling_key() << std::endl;45 std::cout << "tiling_key" << " = " << tilingData.get_tiling_key() << std::endl;
66- if ((tilingData.get_basen_size() != 256) || (tilingData.get_stepka_size() < 256)) {
67- std::cout << "Case1 tiling func execute failed." << std::endl;
68- return -1;
69- }
70 std::cout << "Case1 tiling func execute success." << std::endl;46 std::cout << "Case1 tiling func execute success." << std::endl;
71 return 0;47 return 0;
72 }48 }
@@ -20,11 +20,7 @@ bool TestCase(std::vector<int64_t> shapes) {
20 int64_t n = shapes[3];20 int64_t n = shapes[3];
21 MMTilingData tilingData;21 MMTilingData tilingData;
22 tilingData.set_block_dim(20);22 tilingData.set_block_dim(20);
23- tilingData.set_l2_size(128 * 1024 * 1024);
24 tilingData.set_l1_size(512 * 1024);23 tilingData.set_l1_size(512 * 1024);
25- tilingData.set_l0a_size(64 * 1024);
26- tilingData.set_l0b_size(64 * 1024);
27- tilingData.set_l0c_size(128 * 1024);
28 tilingData.set_m_size(m);24 tilingData.set_m_size(m);
29 tilingData.set_k_size(k);25 tilingData.set_k_size(k);
30 tilingData.set_n_size(n);26 tilingData.set_n_size(n);
@@ -34,13 +30,8 @@ bool TestCase(std::vector<int64_t> shapes) {
34 30 
35 const auto status = GetTiling(tilingData, 1u, nullptr);31 const auto status = GetTiling(tilingData, 1u, nullptr);
36 if ((status)) {32 if ((status)) {
37- std::cout << "tile_l2_m" << " = " << tilingData.get_tilem_size() << std::endl;
38- std::cout << "tile_l2_n" << " = " << tilingData.get_tilen_size() << std::endl;
39 std::cout << "step_ka" << " = " << tilingData.get_stepka_size() << std::endl;33 std::cout << "step_ka" << " = " << tilingData.get_stepka_size() << std::endl;
40 std::cout << "step_kb" << " = " << tilingData.get_stepkb_size() << std::endl;34 std::cout << "step_kb" << " = " << tilingData.get_stepkb_size() << std::endl;
41- std::cout << "base_k" << " = " << tilingData.get_basek_size() << std::endl;
42- std::cout << "base_m" << " = " << tilingData.get_basem_size() << std::endl;
43- std::cout << "base_n" << " = " << tilingData.get_basen_size() << std::endl;
44 return true;35 return true;
45 }36 }
46 std::cout << "mm tiling func execute failed." << std::endl;37 std::cout << "mm tiling func execute failed." << std::endl;
@@ -18,16 +18,9 @@ int TestCase1(uint64_t m, uint64_t n, uint64_t k, int32_t tilingCaseId) {
18 tilingData.set_k_size(k);18 tilingData.set_k_size(k);
19 tilingData.set_block_dim(20);19 tilingData.set_block_dim(20);
20 tilingData.set_l1_size(512 * 1024);20 tilingData.set_l1_size(512 * 1024);
21- tilingData.set_l0a_size(64 * 1024);
22- tilingData.set_l0b_size(64 * 1024);
23- tilingData.set_l0c_size(128 * 1024);
24 // tilingData.z = 0;21 // tilingData.z = 0;
25 const auto status = GetTiling(tilingData, tilingCaseId);22 const auto status = GetTiling(tilingData, tilingCaseId);
26 if ((status)) {23 if ((status)) {
27- std::cout << "basem" << " = " << tilingData.get_basem_size() << std::endl;
28- std::cout << "basen" << " = " << tilingData.get_basen_size() << std::endl;
29- std::cout << "tilem" << " = " << tilingData.get_tilem_size() << std::endl;
30- std::cout << "tilen" << " = " << tilingData.get_tilen_size() << std::endl;
31 std::cout << "stepn" << " = " << tilingData.get_stepn_size() << std::endl;24 std::cout << "stepn" << " = " << tilingData.get_stepn_size() << std::endl;
32 std::cout << "stepm" << " = " << tilingData.get_stepm_size() << std::endl;25 std::cout << "stepm" << " = " << tilingData.get_stepm_size() << std::endl;
33 std::cout << "stepka" << " = " << tilingData.get_stepka_size() << std::endl;26 std::cout << "stepka" << " = " << tilingData.get_stepka_size() << std::endl;
@@ -1,393 +0,0 @@
1-/**
2- * Copyright (c) 2026 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-#include "gtest/gtest.h"
12-#include <fstream>
13-#include "base/base_types.h"
14-#include "base/model_info.h"
15-#include "preprocess/test_stub.h"
16-#include "stub_model_info.h"
17-#include "common/util/mem_utils.h"
18-#define private public
19-#define protected public
20-#include "generator/preprocess/args_manager.h"
21-#include "generator/tiling_code_gen_impl.h"
22-#include "generator/axes_reorder_tiling_code_gen_impl.h"
23- 
24-using namespace att;
25- 
26-namespace {
27-std::set<std::string> GetVarsName(const std::vector<Expr> &args) {
28- std::set<std::string> vars_set;
29- for (const auto arg : args) {
30- DUMPS(Str(arg));
31- vars_set.insert(Str(arg));
32- }
33- return vars_set;
34-}
35- 
36-void VerifySearchArgsByHardware(ArgsManager &args_manager, HardwareDef hw, const std::set<std::string> &expected) {
37- auto args = args_manager.GetSearchableVars(hw);
38- EXPECT_EQ(GetVarsName(args), expected);
39-}
40- 
41-void VerifyBasicArgs(ArgsManager &args_manager, const std::set<std::string> &expected_search,
42- const std::set<std::string> &expected_input) {
43- auto search_args_set = GetVarsName(args_manager.GetSearchableVars());
44- EXPECT_EQ(search_args_set, expected_search);
45- auto input_args_set = GetVarsName(args_manager.GetInputVars());
46- EXPECT_EQ(input_args_set, expected_input);
47-}
48- 
49-void VerifyConstArgs(ArgsManager &args_manager, const std::map<std::string, uint32_t> &expected) {
50- auto const_args = args_manager.GetConstVars();
51- std::map<std::string, uint32_t> const_values;
52- for (const auto &const_var : const_args) {
53- const_values[Str(const_var.first)] = const_var.second;
54- }
55- EXPECT_EQ(const_values, expected);
56-}
57- 
58-void VerifyHardwareConsAndObjs(ArgsManager &args_manager, size_t expected_cons_size, size_t expected_cut_size,
59- size_t expected_obj_size) {
60- auto total_scope_cons = args_manager.GetTotalHardwareCons();
61- EXPECT_EQ(total_scope_cons.size(), expected_cons_size);
62- auto l0a_occupy = args_manager.GetUsedHardwareInfo(HardwareDef::L0A);
63- EXPECT_EQ(l0a_occupy, total_scope_cons[HardwareDef::L0A]);
64- auto l0b_occupy = args_manager.GetUsedHardwareInfo(HardwareDef::L0B);
65- EXPECT_EQ(l0b_occupy, total_scope_cons[HardwareDef::L0B]);
66- EXPECT_EQ(args_manager.GetTotalCutCons().size(), expected_cut_size);
67- EXPECT_EQ(args_manager.GetObjectFunc().size(), expected_obj_size);
68-}
69-} // namespace
70- 
71-class TestArgManager : public ::testing::Test {
72- public:
73- static void SetUpTestCase() {
74- std::cout << "Test begin." << std::endl;
75- }
76- static void TearDownTestCase() {
77- std::cout << "Test end." << std::endl;
78- }
79- 
80- void SetUp() override {
81- model_info_ = CreateModelInfo();
82- }
83- 
84- void TearDown() override {}
85- ModelInfo model_info_;
86-};
87- 
88-TEST_F(TestArgManager, test_split_args) {
89- ModelInfo model_info;
90- Expr s0 = CreateExpr("s0");
91- Expr s1 = CreateExpr("s1");
92- SymVarInfoPtr sym_m = std::make_shared<SymVarInfo>(s0 * s1);
93- AttAxisPtr z0 = std::make_shared<AttAxis>();
94- z0->axis_pos = AxisPosition::ORIGIN;
95- z0->size = sym_m;
96- model_info.arg_list.emplace_back(z0);
97- ArgsManager args_manager(model_info);
98- EXPECT_TRUE(args_manager.Process(true));
99- auto input_args = args_manager.GetInputVars();
100- std::set<std::string> target_search_args = {"s0", "s1"};
101- for (auto input_arg : input_args) {
102- EXPECT_TRUE(target_search_args.find(Str(input_arg)) != target_search_args.end());
103- }
104-}
105- 
106-TEST_F(TestArgManager, process_no_replace) {
107- ArgsManager args_manager(model_info_);
108- EXPECT_TRUE(args_manager.Process(false));
109- VerifyBasicArgs(args_manager, {"tilem_size", "stepm_size", "basem_size", "tilen_size", "stepn_size", "basen_size"},
110- {"m_size", "n_size"});
111- VerifyConstArgs(args_manager, {{"k_size", 128}});
112- 
113- VerifySearchArgsByHardware(args_manager, HardwareDef::L0A, {"basem_size"});
114- VerifySearchArgsByHardware(args_manager, HardwareDef::L0B, {"basen_size"});
115- VerifySearchArgsByHardware(args_manager, HardwareDef::L0C, {"basem_size", "basen_size"});
116- VerifySearchArgsByHardware(args_manager, HardwareDef::L1, {"stepm_size", "stepn_size"});
117- VerifySearchArgsByHardware(args_manager, HardwareDef::L2, {"tilem_size", "tilen_size"});
118- VerifySearchArgsByHardware(args_manager, HardwareDef::CORENUM, {"stepm_size", "stepn_size"});
119- 
120- EXPECT_TRUE(args_manager.GetVarsRelations().empty());
121- EXPECT_TRUE(args_manager.GetExprRelations().empty());
122- EXPECT_TRUE(args_manager.GetSolvedVars().empty());
123- 
124- EXPECT_EQ(args_manager.GetRelatedHardware(CreateExpr("basem_size")),
125- (std::vector<HardwareDef>{HardwareDef::L0A, HardwareDef::L0C}));
126- EXPECT_EQ(args_manager.GetRelatedHardware(CreateExpr("basen_size")),
127- (std::vector<HardwareDef>{HardwareDef::L0B, HardwareDef::L0C}));
128- VerifyHardwareConsAndObjs(args_manager, 7, 4, 2);
129- 
130- EXPECT_EQ(Str(args_manager.GetAncestor(CreateExpr("basem_size"))[0]), "m_size");
131- EXPECT_EQ(Str(args_manager.GetAncestor(CreateExpr("stepm_size"))[0]), "m_size");
132- EXPECT_EQ(Str(args_manager.GetAncestor(CreateExpr("tilem_size"))[0]), "m_size");
133- EXPECT_EQ(Str(args_manager.GetAncestor(CreateExpr("k_size"))[0]), "k_size");
134- 
135- EXPECT_TRUE(args_manager.GetMaxValue(CreateExpr("basem_size")) == CreateExpr("m_size"));
136- EXPECT_TRUE(args_manager.GetMaxValue(CreateExpr("stepm_size")) == CreateExpr("m_size"));
137- EXPECT_TRUE(args_manager.GetMaxValue(CreateExpr("tilem_size")) == CreateExpr("m_size"));
138- 
139- EXPECT_TRUE(args_manager.GetDefaultInitValue(CreateExpr("basem_size")) == CreateExpr("m_size"));
140- EXPECT_TRUE(args_manager.GetDefaultInitValue(CreateExpr("stepm_size")) == CreateExpr(16)); // align
141- EXPECT_TRUE(args_manager.GetDefaultInitValue(CreateExpr("tilem_size")) == CreateExpr(16)); // align
142- EXPECT_TRUE(args_manager.GetDefaultInitValue(CreateExpr("k_size")) == CreateExpr(128));
143- 
144- EXPECT_TRUE(args_manager.GetParentVars(CreateExpr("basem_size"))[0] == CreateExpr("stepm_size"));
145- DUMPS(Str(args_manager.GetParentVars(CreateExpr("basem_size"))[0]));
146- EXPECT_TRUE(args_manager.GetParentVars(CreateExpr("stepm_size"))[0] == CreateExpr("tilem_size"));
147- DUMPS(Str(args_manager.GetParentVars(CreateExpr("stepm_size"))[0]));
148- EXPECT_TRUE(args_manager.GetParentVars(CreateExpr("tilem_size"))[0] == CreateExpr("m_size"));
149- EXPECT_TRUE(args_manager.GetParentVars(CreateExpr("m_size")).empty());
150- 
151- auto innerest_dims_set = GetVarsName(args_manager.GetNodeInnerestDimSizes());
152- std::set<std::string> target_innerest_dims = {"tilem_size", "stepm_size", "basem_size",
153- "tilen_size", "stepn_size", "basen_size"};
154- EXPECT_EQ(innerest_dims_set, target_innerest_dims);
155-}
156- 
157-void VerifyReplaceAncestors(ArgsManager &args_manager) {
158- EXPECT_EQ(Str(args_manager.GetAncestor(CreateExpr("basem_size_base"))[0]), "m_size");
159- EXPECT_EQ(Str(args_manager.GetAncestor(CreateExpr("stepm_size_div_align"))[0]), "m_size");
160- EXPECT_EQ(Str(args_manager.GetAncestor(CreateExpr("tilem_size_div_align"))[0]), "m_size");
161- EXPECT_EQ(Str(args_manager.GetAncestor(CreateExpr("k_size"))[0]), "k_size");
162-}
163- 
164-void VerifyReplaceInitValues(ArgsManager &args_manager) {
165- EXPECT_TRUE(args_manager.GetDefaultInitValue(CreateExpr("basem_size_base")) ==
166- (af::sym::Log((CreateExpr("m_size") / CreateExpr(16)), CreateExpr(2))));
167- EXPECT_TRUE(args_manager.GetDefaultInitValue(CreateExpr("stepm_size_div_align")) == af::sym::kSymbolOne);
168- EXPECT_TRUE(args_manager.GetDefaultInitValue(CreateExpr("tilem_size_div_align")) == af::sym::kSymbolOne);
169- EXPECT_TRUE(args_manager.GetDefaultInitValue(CreateExpr("k_size")) == CreateExpr(128));
170-}
171- 
172-void VerifyReplaceParentVars(ArgsManager &args_manager) {
173- EXPECT_TRUE(args_manager.GetParentVars(CreateExpr("basem_size"))[0] == CreateExpr("stepm_size"));
174- EXPECT_TRUE(args_manager.GetParentVars(CreateExpr("stepm_size"))[0] == CreateExpr("tilem_size"));
175- EXPECT_TRUE(args_manager.GetParentVars(CreateExpr("tilem_size"))[0] == CreateExpr("m_size"));
176- EXPECT_TRUE(args_manager.GetParentVars(CreateExpr("m_size")).empty());
177- EXPECT_FALSE(args_manager.GetParentVars(CreateExpr("basem_size_base"))[0] == CreateExpr("stepm_size"));
178-}
179- 
180-TEST_F(TestArgManager, process_do_replace) {
181- ArgsManager args_manager(model_info_);
182- EXPECT_TRUE(args_manager.Process(true));
183- EXPECT_EQ(args_manager.GetTilingCaseId(), 0);
184- VerifyBasicArgs(args_manager,
185- {"tilem_size_div_align", "stepm_size_div_align", "basem_size_base", "tilen_size_div_align",
186- "stepn_size_div_align", "basen_size_base"},
187- {"m_size", "n_size"});
188- VerifyConstArgs(args_manager, {{"k_size", 128}});
189- 
190- VerifySearchArgsByHardware(args_manager, HardwareDef::L0A, {"basem_size_base"});
191- VerifySearchArgsByHardware(args_manager, HardwareDef::L0B, {"basen_size_base"});
192- VerifySearchArgsByHardware(args_manager, HardwareDef::L0C, {"basem_size_base", "basen_size_base"});
193- VerifySearchArgsByHardware(args_manager, HardwareDef::L1, {"stepm_size_div_align", "stepn_size_div_align"});
194- VerifySearchArgsByHardware(args_manager, HardwareDef::L2, {"tilem_size_div_align", "tilen_size_div_align"});
195- VerifySearchArgsByHardware(args_manager, HardwareDef::CORENUM, {"stepm_size_div_align", "stepn_size_div_align"});
196- 
197- EXPECT_FALSE(args_manager.GetVarsRelations().empty());
198- EXPECT_FALSE(args_manager.GetExprRelations().empty());
199- EXPECT_TRUE(args_manager.GetSolvedVars().empty());
200- EXPECT_EQ(args_manager.GetRelatedHardware(CreateExpr("basem_size_base")),
201- (std::vector<HardwareDef>{HardwareDef::L0A, HardwareDef::L0C}));
202- EXPECT_EQ(args_manager.GetRelatedHardware(CreateExpr("basen_size_base")),
203- (std::vector<HardwareDef>{HardwareDef::L0B, HardwareDef::L0C}));
204- EXPECT_EQ(args_manager.GetRelatedHardware(CreateExpr("stepn_size_div_align")),
205- (std::vector<HardwareDef>{HardwareDef::L1, HardwareDef::CORENUM}));
206- VerifyHardwareConsAndObjs(args_manager, 7, 4, 2);
207- VerifyReplaceAncestors(args_manager);
208- 
209- EXPECT_FALSE(args_manager.GetMaxValue(CreateExpr("basem_size_base")) == CreateExpr("m_size"));
210- EXPECT_FALSE(args_manager.GetMaxValue(CreateExpr("stepm_size_div_align")) == CreateExpr("m_size"));
211- EXPECT_FALSE(args_manager.GetMaxValue(CreateExpr("tilem_size_div_align")) == CreateExpr("m_size"));
212- VerifyReplaceInitValues(args_manager);
213- VerifyReplaceParentVars(args_manager);
214- 
215- auto innerest_dims_set = GetVarsName(args_manager.GetNodeInnerestDimSizes());
216- std::set<std::string> target_innerest_dims = {"tilem_size_div_align", "stepm_size_div_align", "basem_size_base",
217- "tilen_size_div_align", "stepn_size_div_align", "basen_size_base"};
218- EXPECT_EQ(innerest_dims_set, target_innerest_dims);
219-}
220- 
221-TEST_F(TestArgManager, process_replace_after_no_replace) {
222- ArgsManager args_manager(model_info_);
223- EXPECT_TRUE(args_manager.Process(false));
224- auto ori_search_args = args_manager.GetSearchableVars();
225- auto ori_search_args_set = GetVarsName(ori_search_args);
226- std::set<std::string> ori_target_search_args = {"tilem_size", "stepm_size", "basem_size",
227- "tilen_size", "stepn_size", "basen_size"};
228- EXPECT_TRUE(ori_search_args_set == ori_target_search_args);
229- 
230- auto input_args = args_manager.GetInputVars();
231- auto input_args_set = GetVarsName(input_args);
232- std::set<std::string> target_input_args = {"m_size", "n_size"};
233- EXPECT_TRUE(input_args_set == target_input_args);
234- 
235- auto const_args = args_manager.GetConstVars();
236- std::map<std::string, uint32_t> const_values;
237- for (const auto const_var : const_args) {
238- const_values[Str(const_var.first)] = const_var.second;
239- }
240- std::map<std::string, uint32_t> target_const_args = {{"k_size", 128}};
241- EXPECT_TRUE(const_values == target_const_args);
242- 
243- std::vector<Expr> sloved_vars;
244- sloved_vars.push_back(CreateExpr("basem_size"));
245- sloved_vars.push_back(CreateExpr("basen_size"));
246- args_manager.SetSolvedVars(sloved_vars);
247- EXPECT_TRUE(args_manager.DoVarsReplace() == true);
248- auto manager_solved_args = args_manager.GetSolvedVars();
249- auto manager_solved_args_set = GetVarsName(manager_solved_args);
250- std::set<std::string> target_solved_args = {"basem_size", "basen_size"};
251- EXPECT_TRUE(manager_solved_args_set == target_solved_args);
252- auto search_args = args_manager.GetSearchableVars();
253- auto search_args_set = GetVarsName(search_args);
254- std::set<std::string> target_search_args = {"tilem_size_div_align", "stepm_size_div_align", "tilen_size_div_align",
255- "stepn_size_div_align"};
256- EXPECT_TRUE(search_args_set == target_search_args);
257- auto replaced_vars = args_manager.GetVarsRelations();
258- EXPECT_FALSE(replaced_vars.empty());
259- for (auto replaced_var : replaced_vars) {
260- DUMPS("new var: " + Str(replaced_var.first) + " -- old vars: " + Str(replaced_var.second));
261- }
262- auto replaced_expr = args_manager.GetExprRelations();
263- for (auto replaced_var : replaced_expr) {
264- DUMPS("olg var: " + Str(replaced_var.first) + " -- new exprs: " + Str(replaced_var.second));
265- }
266- EXPECT_FALSE(replaced_expr.empty());
267-}
268- 
269-TEST_F(TestArgManager, process_replace_dumps) {
270- ArgsManager args_manager(model_info_);
271- EXPECT_TRUE(args_manager.Process(true));
272- auto vars_string = MakeJson(args_manager.vars_infos_);
273- std::ofstream out_file("./vars_info_st.json");
274- if (out_file.is_open()) {
275- out_file << vars_string;
276- out_file.close();
277- }
278-}
279- 
280-TEST_F(TestArgManager, vars_not_find) {
281- ArgsManager args_manager(model_info_);
282- auto test = CreateExpr("test");
283- EXPECT_TRUE(args_manager.GetAncestor(test).empty());
284- EXPECT_TRUE(args_manager.GetMaxValue(test) == 1);
285- EXPECT_TRUE(args_manager.GetMinValue(test) == 1);
286- EXPECT_TRUE(args_manager.GetDefaultInitValue(test) == 1);
287- EXPECT_TRUE(args_manager.GetVarAlignValue(test) == 0u);
288- EXPECT_TRUE(args_manager.GetVarPromptAlignValue(test) == 0u);
289-}
290- 
291-TEST_F(TestArgManager, IsConcatOuterDim) {
292- Expr var1 = CreateExpr("var1");
293- Expr var2 = CreateExpr("var2");
294- Expr undefined_var = CreateExpr("undefined_var");
295- 
296- ArgsManager manager(model_info_);
297- 
298- VarInfo info1;
299- info1.is_concat_outer_dim = true;
300- manager.vars_infos_[var1] = info1;
301- 
302- VarInfo info2;
303- info2.is_concat_outer_dim = false;
304- manager.vars_infos_[var2] = info2;
305- EXPECT_TRUE(manager.IsConcatOuterDim(var1));
306- EXPECT_FALSE(manager.IsConcatOuterDim(var2));
307- EXPECT_FALSE(manager.IsConcatOuterDim(undefined_var));
308-}
309- 
310-TEST_F(TestArgManager, IsConcatInnerDim) {
311- Expr var1 = CreateExpr("var1");
312- Expr var2 = CreateExpr("var2");
313- Expr undefined_var = CreateExpr("undefined_var");
314- 
315- ArgsManager manager(model_info_);
316- 
317- VarInfo info1;
318- info1.is_concat_inner_dim = true;
319- manager.vars_infos_[var1] = info1;
320- 
321- VarInfo info2;
322- info2.is_concat_inner_dim = false;
323- manager.vars_infos_[var2] = info2;
324- EXPECT_TRUE(manager.IsConcatInnerDim(var1));
325- EXPECT_FALSE(manager.IsConcatInnerDim(var2));
326- EXPECT_FALSE(manager.IsConcatInnerDim(undefined_var));
327-}
328- 
329-TEST_F(TestArgManager, test_small_shape_pattern) {
330- ModelInfo model_info;
331- Expr s0 = CreateExpr("s0");
332- Expr s1 = CreateExpr("s1");
333- SymVarInfoPtr sym_m = std::make_shared<SymVarInfo>(s0);
334- AttAxisPtr z0 = std::make_shared<AttAxis>();
335- z0->axis_pos = AxisPosition::ORIGIN;
336- z0->size = sym_m;
337- z0->name = "z0";
338- 
339- SymVarInfoPtr sym_n = std::make_shared<SymVarInfo>(s1);
340- AttAxisPtr z1 = std::make_shared<AttAxis>();
341- z1->axis_pos = AxisPosition::ORIGIN;
342- z1->size = sym_n;
343- z1->name = "z1";
344- 
345- model_info.arg_list.emplace_back(z0);
346- model_info.arg_list.emplace_back(z1);
347- std::map<HardwareDef, Expr> hardware_cons;
348- hardware_cons[HardwareDef::UB] = CreateExpr(1024) * s1;
349- hardware_cons[HardwareDef::CORENUM] = CreateExpr(1024) * s0;
350- model_info.hardware_cons = hardware_cons;
351- TilingCodeGenConfig config;
352- TilingModelInfo model_infos{model_info};
353- ScoreFuncs score_funcs;
354- TilingCodeGenImplPtr impl = std::shared_ptr<AxesReorderTilingCodeGenImpl>(
355- ge::MakeShared<AxesReorderTilingCodeGenImpl>("op", config, model_infos, score_funcs, false));
356- ArgsManager args_manager(model_info);
357- EXPECT_TRUE(args_manager.Process(true));
358- EXPECT_FALSE(impl->HitSmallShapePattern(args_manager));
359-}
360- 
361-TEST_F(TestArgManager, test_small_shape_pattern_case2) {
362- ModelInfo model_info;
363- Expr s0 = CreateExpr(128);
364- Expr s1 = CreateExpr("s1_size");
365- SymConstInfoPtr sym_m = std::make_shared<SymConstInfo>(s0);
366- AttAxisPtr z0 = std::make_shared<AttAxis>();
367- z0->axis_pos = AxisPosition::ORIGIN;
368- z0->size = sym_m;
369- z0->name = "z0";
370- 
371- SymVarInfoPtr sym_n = std::make_shared<SymVarInfo>(s1);
372- AttAxisPtr z0t = std::make_shared<AttAxis>();
373- z0t->axis_pos = AxisPosition::INNER;
374- z0t->size = sym_n;
375- z0t->name = "z0t";
376- z0t->orig_axis.emplace_back(z0.get());
377- 
378- model_info.arg_list.emplace_back(z0);
379- model_info.arg_list.emplace_back(z0t);
380- 
381- std::map<HardwareDef, Expr> hardware_cons;
382- hardware_cons[HardwareDef::UB] = CreateExpr(1024) * s1;
383- hardware_cons[HardwareDef::CORENUM] = CreateExpr(1024) * s0;
384- model_info.hardware_cons = hardware_cons;
385- TilingCodeGenConfig config;
386- TilingModelInfo model_infos = {model_info};
387- ScoreFuncs score_funcs;
388- TilingCodeGenImplPtr impl = std::shared_ptr<AxesReorderTilingCodeGenImpl>(
389- ge::MakeShared<AxesReorderTilingCodeGenImpl>("op", config, model_infos, score_funcs, false));
390- ArgsManager args_manager(model_info);
391- EXPECT_TRUE(args_manager.Process(true));
392- EXPECT_TRUE(impl->HitSmallShapePattern(args_manager));
393-}
@@ -1,344 +0,0 @@
1-/**
2- * Copyright (c) 2026 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-#include <symengine/functions.h>
12-#include <symengine/simplify.h>
13-#include <symengine/integer.h>
14-#include <symengine/real_double.h>
15-#include "gtest/gtest.h"
16-#include "base/base_types.h"
17-#include "solver_pass/src/l0_solver.cpp"
18-#include "test_common_utils.h"
19-using namespace att;
20- 
21-class MockStL0TileSolver : public L0TileSolver {
22- public:
23- MockStL0TileSolver() {}
24- explicit MockStL0TileSolver(L0TileInput input) : L0TileSolver(input) {}
25- void SetL0A(uint32_t value) {
26- L0A_ = value;
27- }
28- void SetL0B(uint32_t value) {
29- L0B_ = value;
30- }
31- void SetL0C(uint32_t value) {
32- L0C_ = value;
33- }
34- bool CheckBufferUseValid() override {
35- uint32_t l0A = input_.l0_vars[0].value * input_.l0_vars[2].value * 4;
36- uint32_t l0B = input_.l0_vars[2].value * input_.l0_vars[1].value * 4;
37- uint32_t l0C = input_.l0_vars[0].value * input_.l0_vars[1].value * 4;
38- if (l0A > L0A_ || l0B > L0B_ || l0C > L0C_) {
39- return false;
40- }
41- return true;
42- }
43- 
44- private:
45- uint32_t L0A_;
46- uint32_t L0B_;
47- uint32_t L0C_;
48-};
49- 
50-class case1L0TileSolver : public L0TileSolver {
51- public:
52- explicit case1L0TileSolver(L0TileInput &input) : L0TileSolver(input) {};
53- void SetL1(uint32_t &value) {
54- L1_ = value;
55- }
56- void SetL2(uint32_t &value) {
57- L2_ = value;
58- }
59- void SetL0A(uint32_t &value) {
60- L0A_ = value;
61- }
62- void SetL0B(uint32_t &value) {
63- L0B_ = value;
64- }
65- void SetL0C(uint32_t &value) {
66- L0C_ = value;
67- }
68- bool CheckBufferUseValid() override;
69- 
70- private:
71- uint32_t L1_;
72- uint32_t L2_;
73- uint32_t L0A_;
74- uint32_t L0B_;
75- uint32_t L0C_;
76-};
77-bool case1L0TileSolver::CheckBufferUseValid() {
78- uint32_t basek_size = input_.l0_vars[0].value;
79- uint32_t basem_size = input_.l0_vars[1].value;
80- uint32_t basen_size = input_.l0_vars[2].value;
81- 
82- uint32_t L0A = (4 * basek_size * basem_size);
83- if (L0A > L0A_) {
84- return false;
85- };
86- uint32_t L0B = (4 * basek_size * basen_size);
87- if (L0B > L0B_) {
88- return false;
89- };
90- uint32_t L0C = (4 * basem_size * basen_size);
91- if (L0C > L0C_) {
92- return false;
93- };
94- return true;
95-}
96- 
97-class TestL0SolverSt : public ::testing::Test {
98- public:
99- void TearDown() override {
100- // 清理测试生成的临时文件
101- autofuse::test::CleanupTestArtifacts();
102- // before the destructor).
103- }
104- MockStL0TileSolver solver;
105-};
106- 
107-TEST_F(TestL0SolverSt, TEST_CASE_01) {
108- L0Var M;
109- L0Var N;
110- L0Var K;
111- 
112- M.max_value = 32;
113- M.bind_multicore = 1;
114- M.align = 16;
115- M.prompt_align = 16;
116- 
117- N.max_value = 1024;
118- N.bind_multicore = 1;
119- N.align = 16;
120- N.prompt_align = 256;
121- 
122- K.max_value = 1024;
123- K.bind_multicore = 0;
124- K.align = 16;
125- K.prompt_align = 256;
126- 
127- L0TileInput input;
128- input.l0_vars = new L0Var[3];
129- input.l0_vars[0] = M;
130- input.l0_vars[1] = N;
131- input.l0_vars[2] = K;
132- input.size = 3;
133- input.core_num = 24;
134- solver = MockStL0TileSolver(input);
135- solver.SetL0A(64 * 1024);
136- solver.SetL0B(64 * 1024);
137- solver.SetL0C(128 * 1024);
138- EXPECT_EQ(solver.Run(), true);
139- uint32_t *output = solver.GetOutput();
140- EXPECT_NE(output, nullptr);
141- EXPECT_EQ(output[0], 16);
142- EXPECT_EQ(output[1], 256);
143- EXPECT_EQ(output[2], 64);
144- delete[] input.l0_vars;
145-}
146- 
147-TEST_F(TestL0SolverSt, TEST_CASE_02) {
148- L0Var M;
149- L0Var N;
150- L0Var K;
151- 
152- M.max_value = 1024;
153- M.bind_multicore = 1;
154- M.align = 16;
155- M.prompt_align = 16;
156- 
157- N.max_value = 1024;
158- N.bind_multicore = 1;
159- N.align = 16;
160- N.prompt_align = 256;
161- 
162- K.max_value = 1024;
163- K.bind_multicore = 0;
164- K.align = 16;
165- K.prompt_align = 256;
166- 
167- L0TileInput input;
168- input.l0_vars = new L0Var[3];
169- input.l0_vars[0] = M;
170- input.l0_vars[1] = N;
171- input.l0_vars[2] = K;
172- input.size = 3;
173- input.core_num = 24;
174- solver = MockStL0TileSolver(input);
175- solver.SetL0A(64 * 1024);
176- solver.SetL0B(64 * 1024);
177- solver.SetL0C(128 * 1024);
178- EXPECT_EQ(solver.Run(), true);
179- uint32_t *output = solver.GetOutput();
180- EXPECT_NE(output, nullptr);
181- EXPECT_EQ(output[0], 128);
182- EXPECT_EQ(output[1], 256);
183- EXPECT_EQ(output[2], 64);
184- delete[] input.l0_vars;
185-}
186- 
187-TEST_F(TestL0SolverSt, TEST_CASE_03) {
188- L0Var M;
189- L0Var N;
190- L0Var K;
191- 
192- M.max_value = 16;
193- M.bind_multicore = 1;
194- M.align = 16;
195- M.prompt_align = 16;
196- 
197- N.max_value = 16;
198- N.bind_multicore = 1;
199- N.align = 16;
200- N.prompt_align = 256;
201- 
202- K.max_value = 16;
203- K.bind_multicore = 0;
204- K.align = 16;
205- K.prompt_align = 256;
206- 
207- L0TileInput input;
208- input.l0_vars = new L0Var[3];
209- input.l0_vars[0] = M;
210- input.l0_vars[1] = N;
211- input.l0_vars[2] = K;
212- input.size = 3;
213- input.core_num = 24;
214- solver = MockStL0TileSolver(input);
215- solver.SetL0A(64 * 1024);
216- solver.SetL0B(64 * 1024);
217- solver.SetL0C(128 * 1024);
218- EXPECT_EQ(solver.Run(), true);
219- uint32_t *output = solver.GetOutput();
220- EXPECT_NE(output, nullptr);
221- EXPECT_EQ(output[0], 16);
222- EXPECT_EQ(output[1], 16);
223- EXPECT_EQ(output[2], 16);
224- delete[] input.l0_vars;
225-}
226- 
227-TEST_F(TestL0SolverSt, TEST_CASE_04) {
228- L0Var M;
229- L0Var N;
230- L0Var K;
231- 
232- M.max_value = 1;
233- M.bind_multicore = 1;
234- M.align = 16;
235- M.prompt_align = 16;
236- 
237- N.max_value = 1;
238- N.bind_multicore = 1;
239- N.align = 16;
240- N.prompt_align = 256;
241- 
242- K.max_value = 1;
243- K.bind_multicore = 0;
244- K.align = 16;
245- K.prompt_align = 256;
246- 
247- L0TileInput input;
248- input.l0_vars = new L0Var[3];
249- input.l0_vars[0] = M;
250- input.l0_vars[1] = N;
251- input.l0_vars[2] = K;
252- input.size = 3;
253- input.core_num = 24;
254- solver = MockStL0TileSolver(input);
255- solver.SetL0A(64 * 1024);
256- solver.SetL0B(64 * 1024);
257- solver.SetL0C(128 * 1024);
258- EXPECT_EQ(solver.Run(), true);
259- uint32_t *output = solver.GetOutput();
260- EXPECT_NE(output, nullptr);
261- EXPECT_EQ(output[0], 16);
262- EXPECT_EQ(output[1], 16);
263- EXPECT_EQ(output[2], 16);
264- delete[] input.l0_vars;
265-}
266- 
267-TEST_F(TestL0SolverSt, TEST_CASE_05) {
268- L0TileInput l0_input;
269- l0_input.l0_vars = new (std::nothrow) L0Var[3];
270- l0_input.size = 3;
271- l0_input.core_num = 20;
272- L0Var basek_size;
273- basek_size.max_value = 1024;
274- basek_size.bind_multicore = false;
275- basek_size.align = 16;
276- basek_size.prompt_align = 16;
277- l0_input.l0_vars[0] = basek_size;
278- L0Var basem_size;
279- basem_size.max_value = 1024;
280- basem_size.bind_multicore = false;
281- basem_size.align = 16;
282- basem_size.prompt_align = 16;
283- l0_input.l0_vars[1] = basem_size;
284- L0Var basen_size;
285- basen_size.max_value = 2048;
286- basen_size.bind_multicore = false;
287- basen_size.align = 16;
288- basen_size.prompt_align = 16;
289- l0_input.l0_vars[2] = basen_size;
290- case1L0TileSolver solver(l0_input);
291- uint32_t L0A = 64 * 1024;
292- uint32_t L0B = 64 * 1024;
293- uint32_t L0C = 128 * 1024;
294- solver.SetL0A(L0A);
295- solver.SetL0B(L0B);
296- solver.SetL0C(L0C);
297- EXPECT_EQ(solver.Run(), true);
298- uint32_t *output = solver.GetOutput();
299- EXPECT_NE(output, nullptr);
300- EXPECT_EQ(output[0], 128);
301- EXPECT_EQ(output[1], 128);
302- EXPECT_EQ(output[2], 128);
303- delete[] l0_input.l0_vars;
304-}
305- 
306-TEST_F(TestL0SolverSt, TEST_CASE_06) {
307- L0Var M;
308- L0Var N;
309- L0Var K;
310- 
311- M.max_value = 1024;
312- M.bind_multicore = 0;
313- M.align = 16;
314- M.prompt_align = 16;
315- 
316- N.max_value = 2048;
317- N.bind_multicore = 0;
318- N.align = 16;
319- N.prompt_align = 16;
320- 
321- K.max_value = 1024;
322- K.bind_multicore = 0;
323- K.align = 16;
324- K.prompt_align = 16;
325- 
326- L0TileInput input;
327- input.l0_vars = new L0Var[3];
328- input.l0_vars[0] = K;
329- input.l0_vars[1] = M;
330- input.l0_vars[2] = N;
331- input.size = 3;
332- input.core_num = 20;
333- solver = MockStL0TileSolver(input);
334- solver.SetL0A(64 * 1024);
335- solver.SetL0B(64 * 1024);
336- solver.SetL0C(128 * 1024);
337- EXPECT_EQ(solver.Run(), true);
338- uint32_t *output = solver.GetOutput();
339- EXPECT_NE(output, nullptr);
340- EXPECT_EQ(output[0], 256);
341- EXPECT_EQ(output[1], 128);
342- EXPECT_EQ(output[2], 64);
343- delete[] input.l0_vars;
344-}
@@ -1,129 +0,0 @@
1-/**
2- * Copyright (c) 2026 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-#include <symengine/functions.h>
12-#include <symengine/simplify.h>
13-#include <symengine/integer.h>
14-#include <symengine/real_double.h>
15-#include "gtest/gtest.h"
16-#include "base/base_types.h"
17-#include "solver_pass/src/l2_solver.cpp"
18-#include "test_common_utils.h"
19-using namespace att;
20- 
21-class MockStL2TileSolver : public L2TileSolver {
22- public:
23- MockStL2TileSolver() {}
24- explicit MockStL2TileSolver(L2TileInput input) : L2TileSolver(input) {};
25- uint64_t GetL2Use() override {
26- uint64_t tilel2m = input_.l2_vars[0].value;
27- uint64_t tilel2n = input_.l2_vars[1].value;
28- uint64_t k = 512;
29- uint64_t l2Use = tilel2m * k * 2 + k * tilel2n * 2 + tilel2m * tilel2n * 2;
30- return l2Use;
31- }
32- bool IsClash(uint32_t idx) override {
33- if (used_corenum_ <= 1 || used_corenum_ % 2 != 0) {
34- return false;
35- }
36- if (blocknum_per_tile_[idx] % (used_corenum_ / 2) == 0) {
37- return true;
38- }
39- auto blockNumTail = total_blocknum_[idx] - (tilenum_[idx] - 1) * blocknum_per_tile_[idx];
40- if (blocknum_per_tile_[idx] % (used_corenum_ / 2) == 0) {
41- return true;
42- }
43- return false;
44- }
45-};
46- 
47-class TestL2SolverSt : public ::testing::Test {
48- public:
49- void TearDown() override {
50- // 清理测试生成的临时文件
51- autofuse::test::CleanupTestArtifacts();
52- // before the destructor).
53- }
54-};
55- 
56-TEST_F(TestL2SolverSt, TEST_CASE_01) {
57- L2TileInput input;
58- L2Var tilem;
59- L2Var tilen;
60- tilem.max_value = 8196;
61- tilen.max_value = 8196;
62- tilem.align = 16;
63- tilen.align = 16;
64- tilem.base_val = 128;
65- tilen.base_val = 256;
66- input.l2_vars = new L2Var[2];
67- input.l2_vars[0] = tilem;
68- input.l2_vars[1] = tilen;
69- input.size = 2;
70- input.core_num = 24;
71- input.l2_size = 128 * 1024 * 1024;
72- MockStL2TileSolver l2Solver(input);
73- l2Solver.Run();
74- uint32_t *output = l2Solver.GetL2Tile();
75- EXPECT_NE(output, nullptr);
76- EXPECT_EQ(output[0], 7552);
77- EXPECT_EQ(output[1], 7680);
78- delete[] input.l2_vars;
79-}
80- 
81-TEST_F(TestL2SolverSt, TEST_CASE_02) {
82- L2TileInput input;
83- L2Var tilem;
84- L2Var tilen;
85- tilem.max_value = 16;
86- tilen.max_value = 16;
87- tilem.align = 16;
88- tilen.align = 16;
89- tilem.base_val = 128;
90- tilen.base_val = 256;
91- input.l2_vars = new L2Var[2];
92- input.l2_vars[0] = tilem;
93- input.l2_vars[1] = tilen;
94- input.size = 2;
95- input.core_num = 24;
96- input.l2_size = 128 * 1024 * 1024;
97- MockStL2TileSolver l2Solver(input);
98- l2Solver.Run();
99- uint32_t *output = l2Solver.GetL2Tile();
100- EXPECT_NE(output, nullptr);
101- EXPECT_EQ(output[0], 128);
102- EXPECT_EQ(output[1], 256);
103- delete[] input.l2_vars;
104-}
105- 
106-TEST_F(TestL2SolverSt, TEST_CASE_03) {
107- L2TileInput input;
108- L2Var tilem;
109- L2Var tilen;
110- tilem.max_value = 16;
111- tilen.max_value = 1024;
112- tilem.align = 16;
113- tilen.align = 16;
114- tilem.base_val = 128;
115- tilen.base_val = 256;
116- input.l2_vars = new L2Var[2];
117- input.l2_vars[0] = tilem;
118- input.l2_vars[1] = tilen;
119- input.size = 2;
120- input.core_num = 24;
121- input.l2_size = 128 * 1024 * 1024;
122- MockStL2TileSolver l2Solver(input);
123- l2Solver.Run();
124- uint32_t *output = l2Solver.GetL2Tile();
125- EXPECT_NE(output, nullptr);
126- EXPECT_EQ(output[0], 128);
127- EXPECT_EQ(output[1], 1024);
128- delete[] input.l2_vars;
129-}
@@ -1,568 +0,0 @@
1-/**
2- * Copyright (c) 2026 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-#include <cstdlib>
12-#include <iostream>
13-#include <memory.h>
14-#ifdef DEBUG
15-#define OP_LOGI(...)
16-#define OP_LOGW(...)
17-#define OP_LOGE(...)
18-#define OP_LOGD(...)
19-do {
20- std::cout << "[ERROR]" << log << std::endl;
21-} while (0)
22-#else
23-#define OP_LOGI(...)
24-#define OP_LOGW(...)
25-#define OP_LOGE(...)
26-#define OP_LOGD(...)
27-#endif
28- namespace att {
29- // L0Var的备选值的个数
30- static const uint32_t candidate_size = 7u;
31- // L0Var的备选值
32- static const uint32_t candidate_value[] = {16u, 32u, 64u, 128u, 256u, 512u, 1024u};
33- // 表达L0的求解值至少要满足核数的比例,可手动修改
34- static const double CORE_NUM_RATIO = 0.6f;
35- // 表达L0的求解值pad之后的值不允许超过原始值大小的倍数,可手动修改
36- static const uint32_t UPPER_BOUND_RATIO = 2u;
37- // 表达最大L0Var的个数
38- static const uint32_t MAX_L0_VAR_NUM = 3u;
39- 
40- // L0相关变量的数据结构
41- struct L0Var {
42- // 最大值,初始化为输入原始轴的大小
43- uint32_t max_value{0u};
44- // 是否绑多核
45- bool bind_multicore{false};
46- bool is_innermost{false};
47- // 对齐值
48- uint32_t align{0u};
49- // 提示当前L0Var的最佳对齐值,通常源于父轴的对齐值,
50- // 举例,stepm是basem的父轴,stepm的对齐要求是256和basem,那么basem的prompt_align就是256,
51- // 同时约束L0Var的取值必须是256对齐或者是256的因子,因为这样父轴stepm才能既满足256也满足basem对齐
52- uint32_t prompt_align{0u};
53- // L0变量的索引
54- uint32_t idx;
55- // L0变量的值
56- uint32_t value{0u};
57- };
58- 
59- // 求解器接收的输入
60- struct L0TileInput {
61- // 待求解L0变量的集合
62- L0Var *l0_vars{nullptr};
63- // 待求解L0变量的数量
64- uint32_t size;
65- // 核数
66- uint32_t core_num;
67- };
68- 
69- /**
70- * 比较两个 L0Var 类型变量的大小
71- *
72- * 这个函数用来比较两个 L0Var 类型变量的大小。它遵循特定的比较逻辑:
73- * 1. 如果 a 变量绑定到多核且 b 变量没有绑定多核,则 a 被认为大于 b,函数返回
74- * true
75- * 2. 如果 a 变量没有绑定多核且 b 变量绑定多核,则 a 被认为小于 b,函数返回
76- * false
77- * 3. 如果 a 和 b 变量都绑定多核或都未绑定多核,则比较它们的 prompt_align
78- * 属性,prompt_align 属性值大的变量被认为更大。
79- *
80- * @param a 第一个要比较的 L0Var 变量。
81- * @param b 第二个要比较的 L0Var 变量。
82- * @return 如果 a 大于 b,返回 true;如果 a 小于 b,返回 false
83- */
84- static bool L0VarCmp(L0Var a, L0Var b) {
85- if (a.bind_multicore && !b.bind_multicore) {
86- return true;
87- }
88- if (!a.bind_multicore && b.bind_multicore) {
89- return false;
90- }
91- if (a.is_innermost && !b.is_innermost) {
92- return true;
93- }
94- if (!a.is_innermost && b.is_innermost) {
95- return false;
96- }
97- return a.prompt_align > b.prompt_align;
98- }
99- /**
100- * L0 求解器类
101- */
102- class L0TileSolver {
103- public:
104- /**
105- * 构造函数
106- *
107- * @param input 一个 L0TileInput 结构体,包含了 L0 变量相关的信息
108- *
109- * 这个构造函数初始化了 L0TileSolver 对象
110- */
111- explicit L0TileSolver(L0TileInput input) : input_(input) {}
112- L0TileSolver() {};
113- /**
114- * 析构函数
115- *
116- * 当 L0TileSolver 对象被销毁时,析构函数被调用
117- * 用来释放使用 new 运算符动态分配的内存,确保没有内存泄漏
118- */
119- ~L0TileSolver() {
120- if (sortedvars_ != nullptr) {
121- delete[] sortedvars_;
122- }
123- if (output_ != nullptr) {
124- delete[] output_;
125- }
126- }
127- /**
128- * 运行求解器
129- *
130- * @return 如果求解成功,返回 true;否则返回 false
131- *
132- * 这个方法是算法的入口点,调用它会启动求解过程
133- * 成功与否取决于 CheckBufferUseValid() 方法的返回值
134- */
135- bool Run();
136- /**
137- * 获取优化结果
138- *
139- * @return 指向求解结果数据的指针
140- *
141- * 如果有求解结果,这个方法将返回一个指向结果数据的指针
142- * 结果数据的内存使用完后,应该使用 delete[] 释放内存
143- */
144- uint32_t *GetOutput() {
145- return output_;
146- }
147- 
148- protected:
149- /**
150- * 检查是否满足buffer约束
151- *
152- * @return 如果满足,返回 true;否则返回 false
153- *
154- * 这个纯虚函数要求派生类提供实现,以确保缓冲区使用是有效的
155- * 在 L0TileSolver 类中,它是一个抽象方法,需要在子类中实现
156- */
157- virtual bool CheckBufferUseValid() = 0;
158- L0TileInput input_;
159- uint32_t *output_{nullptr};
160- 
161- private:
162- /**
163- * 检查输入数据的完整性和正确性
164- *
165- * @return 如果输入数据有效,返回 true;否则返回 false
166- *
167- * 这个私有方法检查输入数据的格式和逻辑,确保它们适用求解算法
168- */
169- bool CheckInput();
170- 
171- /**
172- * 使用输入数据初始化算法所需的内部数据结构
173- *
174- * 这个方法根据输入的 L0TileInput 结构体中的数据,初始化算法所需的内部数据结构
175- * 确保 sortedvars_ 和 output_ 成员变量被正确初始化
176- */
177- void InitInput();
178- 
179- /**
180- * 检查算法运行的结果,确保它们符合预期
181- *
182- * @return 如果输出数据有效,返回 true;否则返回 false
183- *
184- * 这个方法检查运行算法后得到的结果,确保它们在逻辑上是合理的
185- */
186- bool CheckOutput();
187- 
188- /**
189- * 更新算法执行过程中的对齐设置
190- *
191- * 这个方法根据算法执行过程中的数据更新对齐提示值,确保结果按照预期的方式对齐
192- */
193- void UpdateAlign();
194- 
195- /**
196- * 为指定索引的 L0 变量获取最佳对齐值
197- *
198- * @param i 想要的变量索引值
199- * @return 最佳对齐值
200- *
201- * 这个方法计算并返回给定索引值的 L0 变量的最佳对齐值,
202- * 确保变量以最恰当的方式对齐,从而提高效率或者减少资源浪费
203- */
204- uint32_t GetBestAlign(uint32_t i) const;
205- 
206- /**
207- * 为 L0 变量找到最优值进行迭代运行
208- *
209- * @param loop_id 当前循环索引,表示正在处理的 L0 变量的位置
210- * @param best_var_value 一个指针,指向用于存储每个 L0
211- * 变量迄今为止找到的最佳值的数组
212- *
213- * 这个函数使用递归方法来遍历 L0
214- * 变量的所有可能值。对于每个值,它检查是否满足约束条件,如小于上界且满足对齐要求。如果满足这些条件,它将继续下一个循环或者递归调用自身来处理下一个
215- * L0 变量。如果是最后一个 L0
216- * 变量,它将检查当前组合是否满足优化条件,如核心数量和数据处理量。如果满足条件,它将当前组合存储为最优解。
217- *
218- * 请注意,这个函数没有返回值,而是将最优解存储在传入的 best_var_value
219- * 数组中。
220- */
221- void IterativeRun(uint32_t loop_id, uint32_t *best_var_value);
222- 
223- /**
224- * 根据 L0 变量信息和总核心数计算可以分配的最大核心数
225- *
226- * @param l0_vars 一个指向 L0Var 结构体数组的指针
227- * @param core_num 可用于分配的总核心数
228- * @return 可以分配的最大核心数
229- *
230- * 这个函数计算在给定 L0 变量信息和总核心数的情况下,可以分配的最大核心数。
231- * 它遍历输入的 L0Var
232- * 结构体数组,对于每个变量,根据其是否绑定多核心以及最大、当前和提示对齐值计算所需的块数。
233- * 通过将所有变量的块数相乘,得到总块数。
234- * 最大核心数是总块数和总核心数中的最小值,以确保核心数不会超过可用资源。
235- *
236- * 返回值表示可以分配给 L0
237- * 变量的最大核心数,这对于在多核心系统中进行资源分配是有用的。
238- */
239- int32_t MaxCoreNum(const L0Var *l0_vars, const uint32_t &core_num);
240- 
241- /**
242- * 计算所有 L0 变量值的乘积,作为mac计算量的度量
243- *
244- * @return mac计算量
245- *
246- * 这个函数计算所有 L0 变量值的乘积,结果是一个数字。
247- * 这个数字可以作为数据处理量的度量,例如在评估算法性能时。
248- * 通过不断更新 usage 变量,乘法操作确保了所有 L0 变量的影响都被计入。
249- * 最终 usage 变量中的值就是所有 L0 变量值的乘积,代表了整体的数据处理量。
250- * 返回值可以帮助了解算法处理的数据量,从而对算法的效率和扩展性有更直观的认识。
251- */
252- uint32_t GetMacUse() const;
253- /**
254- * 用于排序的 L0Var 对象数组
255- */
256- L0Var *sortedvars_{nullptr};
257- /**
258- * 最大核心数
259- */
260- int64_t max_corenum_{-1};
261- /**
262- * 最大 MAC 使用量
263- */
264- int32_t max_macuse_{-1};
265- };
266- 
267- /**
268- * 获取给定索引的 L0 变量的最佳对齐值
269- *
270- * @param i 想要的变量索引值
271- * @return 最佳对齐值
272- *
273- * 这个方法为给定索引的 L0
274- * 变量计算最佳对齐值。它考虑到变量的最大、当前和提示对齐值,以确保数据存储和访问的效率。
275- * 根据变量的原始值(ori_value),它首先确定最小和最大对齐值的范围。然后,通过在这个范围内以二的幂次方递增,它找到最大的满足条件的值。
276- * 如果没有找到这样的值,它将返回最小对齐值。如果在范围内找到了一个值,它将返回这个值的二分之一,作为最佳对齐值。
277- * 这个最佳对齐值可以用于确保数据以最有效的方式存储
278- */
279- uint32_t L0TileSolver::GetBestAlign(uint32_t i) const {
280- uint32_t ori_value = input_.l0_vars[i].max_value;
281- uint32_t min_align = input_.l0_vars[i].align;
282- uint32_t max_align = input_.l0_vars[i].prompt_align;
283- uint32_t ori_align = min_align;
284- uint32_t max_value = std::min(ori_value, max_align);
285- while (ori_align <= max_value) {
286- ori_align = ori_align << 1;
287- }
288- if (ori_align == min_align) {
289- return min_align;
290- }
291- return std::max(1u, ori_align >> 1);
292- }
293- 
294- /**
295- * 根据给定的 L0 变量信息计算可以分配的最大核心数
296- *
297- * @param l0_vars 指向 L0Var 结构数组的指针
298- * @param core_num 总核心数
299- * @return 可以分配的最大核心数
300- *
301- * 这个函数遍历 L0Var 结构数组,根据每个变量的 bind_multicore 属性以及
302- * max_value、value 和 prompt_align
303- * 的值来计算每个变量所需的块数。对于绑定多核心的变量,块数计算方式为:(max_value
304- * + max(value,prompt_align)-1)/
305- * max(value,prompt_align)。对于未绑定多核心的变量,块数为 1
306- * 总块数通过将所有变量的块数相乘得到。之后,通过比较总块数和
307- * core_num,返回两者中的最小值,作为可以分配的最大核心数。如果总块数超过了
308- * core_num,那么系统的核心数将成为瓶颈,因此需要将 core_num
309- * 设置为最大核心数。如果总块数小于等于
310- * core_num,那么总块数就是可以分配的最大核心数。
311- */
312- int32_t L0TileSolver::MaxCoreNum(const L0Var *l0_vars, const uint32_t &core_num) {
313- uint32_t total_block_size = 1u;
314- for (uint32_t i = 0u; i < input_.size; i++) {
315- auto var = l0_vars[i];
316- uint32_t block_num = var.bind_multicore ? ((var.max_value + std::max(var.value, var.prompt_align) - 1)) /
317- std::max(var.value, var.prompt_align)
318- : 1;
319- total_block_size *= block_num;
320- }
321- int64_t max_core_num = total_block_size > core_num ? core_num : total_block_size;
322- return max_core_num;
323- }
324- 
325- /**
326- * 计算所有 L0 变量值的乘积,作为数据处理量的度量
327- *
328- * @return 数据处理量
329- *
330- * 这个函数遍历 L0TileInput 结构体中的所有 L0Var 对象,计算它们的 value
331- * 属性的乘积。这个乘积代表了所有 L0
332- * 变量值的联合效应,或者说数据处理量的一个度量。 通过不断更新 usage
333- * 变量,乘法操作确保了所有 L0 变量的贡献都被包含在内。最终 usage
334- * 变量中的值就是所有 L0 变量值的乘积。
335- * 返回值可以帮助评估算法在处理给定输入数据时的效率,以及比较不同算法或优化策略的数据处理量。
336- */
337- uint32_t L0TileSolver::GetMacUse() const {
338- uint32_t usage = 1u;
339- for (uint32_t j = 0; j < input_.size; j++) {
340- usage *= input_.l0_vars[j].value;
341- }
342- return usage;
343- }
344- 
345- /**
346- * 为 L0 变量找到最优值进行迭代运行
347- *
348- * @param loop_id 当前循环索引,表示正在处理的 L0 变量的位置
349- * @param best_var_value 一个指针,指向用于存储每个 L0 变量迄今为止找到的最佳值的数组
350- *
351- * 这个函数使用递归方法来遍历 L0 变量的所有可能值。对于每个值,它检查是否满足约束条件,如小于上界且满足对齐要求。
352- * 如果满足这些条件,它将继续下一个循环或者递归调用自身来处理下一个L0 变量。
353- * 如果是最后一个 L0
354- * 变量,它将检查当前组合是否满足优化条件,如核心数量和数据处理量。如果满足条件,它将当前组合存储为最优解。
355- *
356- * 请注意,这个函数没有返回值,而是将最优解存储在传入的 best_var_value 数组中。
357- */
358- void L0TileSolver::IterativeRun(uint32_t loop_id, uint32_t * best_var_value) {
359- for (uint32_t i = 0u; i < candidate_size; i++) {
360- uint32_t candi_value = candidate_value[i];
361- const auto &l0_tile = sortedvars_[loop_id];
362- // L0Var的上限
363- uint32_t upper_bound = l0_tile.max_value * UPPER_BOUND_RATIO;
364- if (candi_value >= upper_bound) {
365- continue;
366- }
367- // 必须满足prompt_align对齐或者是prompt_align的因子
368- if ((candi_value % l0_tile.prompt_align != 0) && (l0_tile.prompt_align % candi_value != 0)) {
369- continue;
370- }
371- auto idx = l0_tile.idx;
372- input_.l0_vars[idx].value = candi_value;
373- // 终止条件为遍历到最后一个变量
374- if (loop_id == input_.size - 1) {
375- if (!CheckBufferUseValid()) {
376- break;
377- }
378- int32_t usage = GetMacUse();
379- int32_t core_num = MaxCoreNum(input_.l0_vars, input_.core_num);
380- // 最大核数如果满足核数*系数(默认0.6),则比较mac利用率即可,否则需要比较核数的使用和mac利用率
381- if (((core_num >= max_corenum_) || (core_num >= static_cast<int32_t>(input_.core_num * CORE_NUM_RATIO))) &&
382- (usage >= max_macuse_)) {
383- max_corenum_ = core_num;
384- max_macuse_ = usage;
385- for (uint32_t k = 0u; k < input_.size; k++) {
386- best_var_value[k] = input_.l0_vars[k].value;
387- }
388- }
389- } else {
390- IterativeRun(loop_id + 1, best_var_value);
391- }
392- }
393- }
394- 
395- /**
396- * 更新 L0Var 对象的对齐值
397- *
398- * 这个函数用于更新 L0Var 对象的 prompt_align
399- * 值,以确保它们在内存中按照最优方式对齐。它遍历 input_ 对象中的 l0_vars
400- * 数组,为每个 L0Var 对象计算并设置最佳的对齐值。
401- *
402- * @param 无
403- * @return
404- */
405- void L0TileSolver::UpdateAlign() {
406- for (uint32_t i = 0u; i < input_.size; i++) {
407- uint32_t best_align = GetBestAlign(i);
408- input_.l0_vars[i].prompt_align = best_align;
409- }
410- }
411- 
412- /**
413- * 检查输入数据的有效性
414- *
415- * 这个函数用来检查 L0TileSolver 类的输入数据是否有效。它验证以下几个方面:
416- * - 基础变量指针(l0_vars)是否为空。
417- * - 输入数据的大小(size)是否为0,表示没有 L0 参数需要求解。
418- * - 输入数据的大小(size)是否超过最大支持的参数数量(MAX_L0_VAR_NUM)。
419- * - 核心数量(core_num)是否为0
420- * - 对于输入数据中的每个 L0Var 对象(通过索引 i 访问),它检查几个属性:
421- * - max_value、align 和 prompt_align 是否都不等于0
422- * - align 是否不大于 prompt_align。
423- *
424- * 如果以上任何一个条件不满足,函数将通过 OP_LOG 宏记录一条错误消息,并返回
425- * false,表示输入无效。如果所有条件都满足,函数返回 true,表示输入有效。
426- *
427- * @return 如果输入数据有效,则返回 true;否则返回 false
428- */
429- bool L0TileSolver::CheckInput() {
430- if (input_.l0_vars == nullptr) {
431- OP_LOGW(OP_NAME, "Input basevar is null");
432- return false;
433- }
434- if (input_.size == 0u) {
435- OP_LOGW(OP_NAME, "Size is 0, no l0 arg to be solved");
436- return false;
437- }
438- if (input_.size > MAX_L0_VAR_NUM) {
439- OP_LOGW(OP_NAME, "L0 solver does not support more than 3 input args");
440- return false;
441- }
442- if (input_.core_num == 0) {
443- OP_LOGW(OP_NAME, "Corenum is 0");
444- return false;
445- }
446- for (uint32_t i = 0u; i < input_.size; i++) {
447- auto var = input_.l0_vars[i];
448- if ((var.max_value == 0) || (var.align == 0) || (var.prompt_align == 0)) {
449- OP_LOGW(OP_NAME, "Input [%u] exists 0", i);
450- return false;
451- }
452- if (var.align > var.prompt_align) {
453- OP_LOGW(OP_NAME, "Input [%u] align is larger than prompt align", i);
454- return false;
455- }
456- }
457- return true;
458- }
459- 
460- /**
461- * 初始化 L0Var 对象数组
462- *
463- * 这个函数用于初始化 L0TileSolver 类的 input_ 对象中的 l0_vars 数组。它遍历
464- * l0_vars 数组中的每一个元素,对于每个元素,执行以下操作:
465- * 1. 通过访问索引 i 对应的 L0Var 对象的引用 var,重置其 max_value
466- * 属性。具体重置方式是,先将 max_value 增加 align 属性值减 1,再除以 align
467- * 属性值,最后乘以 align 属性值。这样做的目的可能是为了确保 max_value 是 align
468- * 的整数倍。
469- * 2. 将当前循环的索引值 i 设置为 var 的 idx 属性。这可能是为了标记每个 L0Var
470- * 对象在数组中的位置,以便后续处理。
471- *
472- * @param 无
473- * @return
474- */
475- void L0TileSolver::InitInput() {
476- for (uint32_t i = 0u; i < input_.size; i++) {
477- auto &var = input_.l0_vars[i];
478- var.max_value = (var.max_value + var.align - 1) / var.align * var.align;
479- var.idx = i;
480- }
481- }
482- 
483- /**
484- * 检查输出数据的有效性
485- *
486- * 这个函数用于检查 L0TileSolver 类的输出数据是否有效。它首先检查 output_
487- * 指针是否为空。如果 output_ 指针为空,通过 OP_LOG 宏记录一条错误消息,并返回
488- * false,表示输出无效。
489- *
490- * 接着,函数遍历 output_ 数组中的每个元素。对于每个元素,它检查其值是否为
491- * 0。如果发现任何一个元素的值为 0,函数会通过 OP_LOG
492- * 宏记录相应的错误消息,并返回 false,表示输出数据中存在无效的元素。
493- *
494- * 如果输出数据有效,即 output_ 指针不为空且 output_ 数组中没有 0
495- * 值元素,函数返回 true
496- *
497- * @return 如果输出数据有效,则返回 true;否则返回 false
498- */
499- bool L0TileSolver::CheckOutput() {
500- if (output_ == nullptr) {
501- OP_LOGW(OP_NAME, "Output is null");
502- return false;
503- }
504- for (uint32_t i = 0u; i < input_.size; i++) {
505- if (output_[i] == 0u) {
506- OP_LOGW(OP_NAME, "Output [%u] is 0", i);
507- return false;
508- }
509- }
510- return true;
511- }
512- 
513- /**
514- * 执行 L0TileSolver 类的主要流程
515- * @return 如果所有操作成功并且输出有效,则返回 true;否则返回 false
516- */
517- bool L0TileSolver::Run() {
518- // 检查输入数据的有效性
519- if (!CheckInput()) {
520- // 如果输入检查失败,则记录一条错误日志,并返回 false
521- OP_LOGW(OP_NAME, "Check input failed");
522- return false;
523- }
524- 
525- // 初始化输入数据
526- InitInput();
527- 
528- // 更新 L0Var 对象的对齐值
529- UpdateAlign();
530- 
531- // 为排序后的变量申请内存,并初始化为 0
532- sortedvars_ = new (std::nothrow) L0Var[input_.size];
533- output_ = new (std::nothrow) uint32_t[input_.size]();
534- 
535- // 将输入数据复制到新的内存中
536- std::copy(input_.l0_vars, input_.l0_vars + input_.size, sortedvars_);
537- 
538- // 根据比较函数对变量进行排序
539- std::sort(sortedvars_, sortedvars_ + input_.size, L0VarCmp);
540- 
541- bool is_fast_mode = true;
542- for (uint32_t i = 0u; i < input_.size; i++) {
543- auto &var = input_.l0_vars[i];
544- if ((var.value == 0u) || (var.value > var.max_value * UPPER_BOUND_RATIO)) {
545- is_fast_mode = false;
546- break;
547- }
548- }
549- if (is_fast_mode && CheckBufferUseValid()) {
550- for (uint32_t k = 0u; k < input_.size; k++) {
551- output_[k] = input_.l0_vars[k].value;
552- }
553- } else {
554- // 调用 IterativeRun 函数,传递参数 0 和 output_ 数组的指针
555- IterativeRun(0u, output_);
556- }
557- 
558- // 检查输出数据的有效性
559- if (!CheckOutput()) {
560- // 如果输出检查失败,则记录一条错误日志,并返回 false
561- OP_LOGW(OP_NAME, "Check output failed");
562- return false;
563- }
564- 
565- // 如果所有操作都成功,返回 true
566- return true;
567- }
568-} // namespace att
@@ -1,342 +0,0 @@
1-/**
2- * Copyright (c) 2026 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-#include <iostream>
12-#ifdef DEBUG
13-#define OP_LOG(log)
14-do {
15- std::cout << "[ERROR]" << log << std::endl;
16-} while (0)
17-#else
18-#define OP_LOG(log)
19-#endif
20- 
21- namespace att {
22- // L2的占用经验值大小
23- const uint32_t EMPIRIC_L2_SIZE = 128 * 1024 * 1024u;
24- uint32_t CeilDivision(uint32_t a, uint32_t b) {
25- if (b == 0) {
26- return 0;
27- }
28- return uint32_t((a + b - 1) / b);
29- }
30- 
31- // 每个L2变量的信息
32- struct L2Var {
33- // 最大值,初始化为原始输入的大小
34- uint32_t max_value{0};
35- // 对气值
36- uint32_t align{0};
37- // 对应的L0基本块的大小,举例,TileL2M对应的基本块basem
38- uint32_t base_val{0};
39- // 当前变量的值
40- uint32_t value{0};
41- };
42- 
43- // 求解器的输入
44- struct L2TileInput {
45- // L2变量的集合
46- L2Var *l2_vars{nullptr};
47- // L2变量的个数
48- uint32_t size{0};
49- // 核数
50- uint32_t core_num{0};
51- // l2的大小,默认为经验值的大小
52- uint32_t l2_size{0};
53- };
54- 
55- // L2求解器的适用范围如下
56- // 举例如下,一个TileL2M * TileL2N的结果矩阵,每个小方格表示一个basem * basen的基本块
57- // TileL2M 和 TileL2N大小的结果矩阵存储在L2中,以基本块为粒度分多核,如下图所示,假设有4个核,每个核计算两个基本块
58- //
59- // tileL2N basen
60- // .-------.-------.-------.-------.
61- // | core0 | core0 | core1 | core1 | ->basem
62- // tileL2M <- '-------'-------'-------'-------'
63- // | core2 | core2 | core3 | core3 |
64- // '-------'-------'-------'-------'
65- class L2TileSolver {
66- public:
67- // 构造函数,接受 L2TileInput 类型的参数 input
68- explicit L2TileSolver(L2TileInput input) : input_(input) {};
69- // 无参构造函数
70- L2TileSolver() {}
71- // 析构函数,用于清理堆上分配的内存
72- ~L2TileSolver() {
73- // 如果 blocknum_per_tile_ 指针不为空,则释放其指向的内存
74- if (blocknum_per_tile_ != nullptr) {
75- delete[] blocknum_per_tile_;
76- }
77- // 如果 size_per_tile_ 指针不为空,则释放其指向的内存
78- if (size_per_tile_ != nullptr) {
79- delete[] size_per_tile_;
80- }
81- // 如果 tilenum_ 指针不为空,则释放其指向的内存
82- if (tilenum_ != nullptr) {
83- delete[] tilenum_;
84- }
85- // 如果 total_blocknum_ 指针不为空,则释放其指向的内存
86- if (total_blocknum_ != nullptr) {
87- delete[] total_blocknum_;
88- }
89- }
90- // Run() 成员函数,返回布尔值,可能用于指示某个操作的成功或失败
91- bool Run();
92- // GetL2Tile() 成员函数,返回 uint32_t 类型的指针
93- uint32_t *GetL2Tile() {
94- return size_per_tile_;
95- }
96- 
97- protected:
98- // 纯虚函数,需要在子类中实现,用于获取 L2 的使用情况
99- virtual uint64_t GetL2Use() = 0;
100- // 纯虚函数,需要在子类中实现,用于判断索引 idx 处是否存在冲突
101- virtual bool IsClash(uint32_t idx) = 0;
102- // L2TileInput 类型的成员变量,用于存储输入数据
103- L2TileInput input_;
104- // 默认初始化值为 1 的 used_corenum_ 成员变量
105- uint32_t used_corenum_{1};
106- // 指向 uint32_t 类型的指针 blocknum_per_tile_,初始化为空指针,用于表达每个方向的输入包含多少个基本块
107- uint32_t *blocknum_per_tile_{nullptr};
108- // 指向 uint32_t 类型的指针 size_per_tile_,初始化为空指针,用于表达每个方向的输入的大小
109- uint32_t *size_per_tile_{nullptr};
110- // 指向 uint32_t 类型的指针 tilenum_,初始化为空指针,用于表达每个方向的输入的L2块的个数
111- uint32_t *tilenum_{nullptr};
112- // 指向 uint32_t 类型的指针 total_blocknum_,初始化为空指针,用于表达每个方向输入基本块的总数
113- uint32_t *total_blocknum_{nullptr};
114- 
115- private:
116- // 私有的 CheckInput() 成员函数,返回布尔值,用于检查输入数据的有效性
117- bool CheckInput();
118- // 私有的 InitInput() 成员函数,用于初始化输入数据
119- void InitInput();
120- // 私有的 CheckSolvable() 成员函数,返回布尔值,用于检查问题是否可解
121- bool CheckSolvable();
122- void HandleClash(uint32_t loop_id, uint32_t *ori_val, uint32_t *best_val, uint64_t &max_l2_use);
123- bool CheckAllSolved();
124- void UpdateBestSolution(uint32_t *best_val, uint64_t &max_l2_use);
125- void ReduceToFitL2(uint32_t l2_size);
126- void InitTileArrays(uint32_t *best_val, uint32_t *ori_val);
127- void ApplyBestResult(uint32_t *best_val);
128- void AdjustSingleTile();
129- };
130- 
131- /**
132- * 检查输入参数是否有效。
133- *
134- * 该函数用于检查 L2TileSolver 类的输入参数是否有效。
135- * 首先检查 l2_vars 指针是否为空。如果为空,将记录一条错误消息并返回 false,表示输入无效。
136- * 然后检查 size、core_num 和 l2_size 参数是否都不为零。如果其中任何一个为零,将记录一条错误消息并返回
137- * false,表示输入无效。 最后遍历 l2_vars 指针数组中的所有 L2Var 结构。对于每个结构,检查 align、base_val 和 max_value
138- * 成员是否都不为零。如果其中任何一个为零,将记录一条错误消息并返回 false,表示输入无效。
139- * 如果所有检查都通过,该函数将返回 true,表示输入有效。
140- *
141- * @Return bool, 表示输入参数是否有效
142- */
143- bool L2TileSolver::CheckInput() {
144- if (input_.l2_vars == nullptr) {
145- OP_LOG("Input l2var is null");
146- return false;
147- }
148- if (input_.size == 0 || input_.core_num == 0 || input_.l2_size == 0) {
149- OP_LOG("Exist input 0, please check size, core_num and l2_size");
150- return false;
151- }
152- for (uint32_t i = 0; i < input_.size; i++) {
153- auto var = input_.l2_vars[i];
154- if (var.align == 0 || var.base_val == 0 || var.max_value == 0) {
155- OP_LOG("Input [" + std::to_string(i) + "] exists 0");
156- return false;
157- }
158- }
159- return true;
160- }
161- 
162- /**
163- * 检查问题是否可以解决。
164- *
165- * 这个函数检查当前的问题是否可以根据输入参数的设置来解决。
166- * 首先,它初始化每个变量的值为 align 参数的值。
167- * 然后,它调用 GetL2Use() 函数来获得所需的 L2 使用量。
168- * 如果所需的 L2 使用量超过了可用的缓存大小(input_.l2_size),它将记录一条警告消息并返回 false,表示没有解决方案。
169- * 如果所需的 L2 使用量小于或等于可用的缓存大小,函数将返回 true,表示问题可以被解决。
170- *
171- * @Return true 如果问题可以解决,false 否则。
172- */
173- bool L2TileSolver::CheckSolvable() {
174- for (uint32_t i = 0; i < input_.size; i++) {
175- auto &var = input_.l2_vars[i];
176- var.value = var.align;
177- }
178- if (GetL2Use() > input_.l2_size) {
179- OP_LOG("No solution, l2 size is too small");
180- return false;
181- }
182- return true;
183- }
184- 
185- /**
186- * 初始化输入数据。
187- *
188- * 这个函数负责初始化 L2TileSolver 对象的输入数据。
189- * 它首先为输入数据结构中的每个变量计算最大值,将其向上取整为对齐值的最近倍数。这确保了每个变量的最大值是其对齐值的倍数。
190- * 然后,它找出所有变量中的最大最大值,并将每个变量的值初始化为这个最大最大值。
191- *
192- * 初始化过程对于准备输入数据以进行后续处理步骤(如优化或分析)至关重要。
193- * 确保最大值是对齐值的倍数,可以简化数据的处理,并可能是依赖于这种属性的算法或操作所必需的。
194- */
195- void L2TileSolver::InitInput() {
196- uint32_t init_value = 0;
197- for (uint32_t i = 0; i < input_.size; i++) {
198- auto &var = input_.l2_vars[i];
199- var.max_value = CeilDivision(var.max_value, var.align) * var.align;
200- init_value = var.max_value > init_value ? var.max_value : init_value;
201- }
202- for (uint32_t i = 0; i < input_.size; i++) {
203- auto &var = input_.l2_vars[i];
204- var.value = init_value;
205- }
206- }
207- 
208- bool L2TileSolver::CheckAllSolved() {
209- for (uint32_t k = 0; k < input_.size; k++) {
210- if (IsClash(k)) {
211- return false;
212- }
213- }
214- return true;
215- }
216- 
217- void L2TileSolver::UpdateBestSolution(uint32_t * best_val, uint64_t & max_l2_use) {
218- uint32_t tmp_corenum = 1;
219- for (uint32_t j = 0; j < input_.size; j++) {
220- tmp_corenum *= blocknum_per_tile_[j];
221- }
222- used_corenum_ = std::min(input_.core_num, tmp_corenum);
223- if (!CheckAllSolved()) {
224- return;
225- }
226- uint64_t l2_use = GetL2Use();
227- if (l2_use > max_l2_use) {
228- for (uint32_t l = 0; l < input_.size; l++) {
229- best_val[l] = blocknum_per_tile_[l];
230- }
231- max_l2_use = l2_use;
232- }
233- }
234- 
235- void L2TileSolver::HandleClash(uint32_t loop_id, uint32_t * ori_val, uint32_t * best_val, uint64_t & max_l2_use) {
236- auto max_blocknum = ori_val[loop_id];
237- auto &var = input_.l2_vars[loop_id];
238- for (uint32_t i = max_blocknum; i >= 1u; i--) {
239- blocknum_per_tile_[loop_id] = i;
240- size_per_tile_[loop_id] = blocknum_per_tile_[loop_id] * var.base_val;
241- tilenum_[loop_id] = CeilDivision(var.max_value, size_per_tile_[loop_id]);
242- var.value = size_per_tile_[loop_id];
243- if (loop_id == input_.size - 1) {
244- UpdateBestSolution(best_val, max_l2_use);
245- } else {
246- HandleClash(loop_id + 1, ori_val, best_val, max_l2_use);
247- }
248- }
249- }
250- 
251- /**
252- * 运行 L2TileSolver 算法来解决 L2 缓存分块问题。
253- *
254- * 这个函数是 L2TileSolver 算法的核心。它尝试根据输入参数和约束条件找到 L2 缓存的最优分块方案。
255- * 函数首先通过 CheckInput() 函数检查输入参数的有效性。如果输入无效,它将记录一条错误消息并返回 false
256- * 然后通过 CheckSolvable() 函数检查问题是否有解。如果没有解,它将记录一条错误消息并返回 false
257- * 如果输入有效并且问题有解,函数将通过 InitInput() 函数初始化输入数据。
258- *
259- * 算法的核心是一个循环,在这个循环中,它不断地调整每个变量的值,以找到一个合适的分块方案。
260- * 循环结束的条件是总内存使用量小于或等于 L2 缓存的大小。
261- * 循环结束后,它将计算每个变量的每个分块的大小、总块数、分块数和占用的核心数。
262- * 然后,它检查是否存在读冲突。如果存在读冲突,它将调整每个分块的大小,直到不再检测到冲突。
263- * 最后,它检查每个变量的最大值是否可以放在一个分块中。
264- * 如果不能,它将相应地调整分块数和每个分块的大小。
265- *
266- * 如果找到合适的分块方案,函数返回 true,否则返回 false
267- *
268- * @Return 如果成功则返回 true,否则返回 false
269- */
270- void L2TileSolver::ReduceToFitL2(uint32_t l2_size) {
271- while (GetL2Use() > l2_size) {
272- for (uint32_t i = 0; i < input_.size; i++) {
273- auto &var = input_.l2_vars[i];
274- var.value = (var.align < var.value) ? (var.value - var.align) : var.align;
275- }
276- }
277- }
278- 
279- void L2TileSolver::InitTileArrays(uint32_t * best_val, uint32_t * ori_val) {
280- for (uint32_t i = 0; i < input_.size; i++) {
281- auto &var = input_.l2_vars[i];
282- blocknum_per_tile_[i] = CeilDivision(var.value, var.base_val);
283- size_per_tile_[i] = blocknum_per_tile_[i] * var.base_val;
284- tilenum_[i] = CeilDivision(var.max_value, size_per_tile_[i]);
285- total_blocknum_[i] = CeilDivision(var.max_value, var.base_val);
286- best_val[i] = blocknum_per_tile_[i];
287- ori_val[i] = blocknum_per_tile_[i];
288- }
289- }
290- 
291- void L2TileSolver::ApplyBestResult(uint32_t * best_val) {
292- for (uint32_t i = 0; i < input_.size; i++) {
293- auto &var = input_.l2_vars[i];
294- blocknum_per_tile_[i] = best_val[i];
295- size_per_tile_[i] = blocknum_per_tile_[i] * var.base_val;
296- }
297- }
298- 
299- void L2TileSolver::AdjustSingleTile() {
300- for (uint32_t i = 0; i < input_.size; i++) {
301- auto &var = input_.l2_vars[i];
302- if (var.max_value <= size_per_tile_[i]) {
303- tilenum_[i] = 1;
304- blocknum_per_tile_[i] = CeilDivision(var.max_value, var.base_val);
305- size_per_tile_[i] = blocknum_per_tile_[i] * var.base_val;
306- }
307- }
308- }
309- 
310- bool L2TileSolver::Run() {
311- if (!CheckInput()) {
312- OP_LOG("Check input failed");
313- return false;
314- }
315- if (!CheckSolvable()) {
316- OP_LOG("Check Solvable failed");
317- return false;
318- }
319- InitInput();
320- uint32_t l2_size = input_.l2_size;
321- blocknum_per_tile_ = new (std::nothrow) uint32_t[input_.size];
322- size_per_tile_ = new (std::nothrow) uint32_t[input_.size];
323- tilenum_ = new (std::nothrow) uint32_t[input_.size];
324- total_blocknum_ = new (std::nothrow) uint32_t[input_.size];
325- 
326- ReduceToFitL2(l2_size);
327- 
328- uint32_t *best_val = new uint32_t[input_.size];
329- uint32_t *ori_val = new uint32_t[input_.size];
330- InitTileArrays(best_val, ori_val);
331- 
332- uint64_t max_l2_use = 0u;
333- HandleClash(0, ori_val, best_val, max_l2_use);
334- 
335- ApplyBestResult(best_val);
336- delete[] best_val;
337- delete[] ori_val;
338- 
339- AdjustSingleTile();
340- return true;
341- }
342-} // namespace att
@@ -1,147 +0,0 @@
1-/**
2- * Copyright (c) 2026 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-#include "gtest/gtest.h"
12-#include "base/base_types.h"
13-#include "base/model_info.h"
14-#include "generator/solver_pass_gen/l0_solver/l0_solver_gen.h"
15-#include <symengine/functions.h>
16-#include <symengine/simplify.h>
17-#include <symengine/integer.h>
18-#include <symengine/real_double.h>
19-#include "test_common_utils.h"
20-using namespace att;
21- 
22-class TestL0SolverGen : public ::testing::Test {
23- public:
24- void TearDown() override {
25- // 清理测试生成的临时文件
26- autofuse::test::CleanupTestArtifacts();
27- // before the destructor).
28- }
29-};
30- 
31-TEST_F(TestL0SolverGen, TEST_CASE_01) {
32- L0TileSolverGen solver_gen("case0", "TilingData");
33- std::vector<Expr> ori_args;
34- std::vector<Expr> l0_args;
35- std::vector<Expr> mc_args;
36- std::map<Expr, Expr, ExprCmp> father_args_map;
37- std::map<Expr, Expr, ExprCmp> arg_align_map;
38- std::map<Expr, Expr, ExprCmp> ori_arg_map;
39- 
40- Expr m = CreateExpr("m");
41- Expr n = CreateExpr("n");
42- Expr k = CreateExpr("k");
43- Expr tilem = CreateExpr("tilem");
44- Expr tilen = CreateExpr("tilen");
45- Expr basem = CreateExpr("basem");
46- Expr basen = CreateExpr("basen");
47- Expr basek = CreateExpr("basek");
48- Expr buffer_use = ((basem * basek) + (basen * basek));
49- 
50- std::map<HardwareDef, Expr> buffer_use_map;
51- buffer_use_map[HardwareDef::L0A] = buffer_use;
52- 
53- ori_args.emplace_back(m);
54- ori_args.emplace_back(n);
55- ori_args.emplace_back(k);
56- 
57- l0_args.emplace_back(basem);
58- l0_args.emplace_back(basen);
59- l0_args.emplace_back(basek);
60- 
61- mc_args.emplace_back(tilem);
62- mc_args.emplace_back(tilen);
63- 
64- ori_arg_map[basem] = m;
65- ori_arg_map[basen] = n;
66- ori_arg_map[basek] = k;
67- 
68- father_args_map[basem] = tilem;
69- father_args_map[basen] = tilen;
70- 
71- arg_align_map[basem] = ge::Symbol(16);
72- arg_align_map[basen] = ge::Symbol(16);
73- arg_align_map[basek] = ge::Symbol(16);
74- arg_align_map[tilem] = ge::Symbol(256);
75- arg_align_map[tilen] = ge::Symbol(16);
76- 
77- solver_gen.SetMulticoreArgs(mc_args);
78- solver_gen.SetFatherArgsMap(father_args_map);
79- solver_gen.SetArgAlignMap(arg_align_map);
80- solver_gen.SetArgtMaxValueMap(ori_arg_map);
81- solver_gen.SetL0Args(l0_args);
82- solver_gen.SetBufferUseAlg(buffer_use_map);
83- std::string impl_code = solver_gen.GenSolverFuncImpl();
84- std::string invoke_code = solver_gen.GenSolverFuncInvoke();
85- // std::cout<<impl_code<<std::endl;
86- // std::cout<<invoke_code<<std::endl;
87- EXPECT_NE(impl_code, "");
88- EXPECT_NE(invoke_code, "");
89-}
90- 
91-TEST_F(TestL0SolverGen, TEST_CASE_02) {
92- L0TileSolverGen solver_gen("case1", "TilingData");
93- std::vector<Expr> ori_args;
94- std::vector<Expr> l0_args;
95- std::vector<Expr> mc_args;
96- std::map<Expr, Expr, ExprCmp> father_args_map;
97- std::map<Expr, Expr, ExprCmp> arg_align_map;
98- std::map<Expr, Expr, ExprCmp> ori_arg_map;
99- 
100- Expr m = CreateExpr("m");
101- Expr n = CreateExpr("n");
102- Expr k = CreateExpr("k");
103- Expr tilem = CreateExpr("tilem");
104- Expr tilen = CreateExpr("tilen");
105- Expr basem = CreateExpr("basem");
106- Expr basen = CreateExpr("basen");
107- Expr basek = CreateExpr("basek");
108- Expr buffer_use = ((basem * basek) + (basen * basek));
109- 
110- std::map<HardwareDef, Expr> buffer_use_map;
111- buffer_use_map[HardwareDef::L0A] = buffer_use;
112- 
113- ori_args.emplace_back(m);
114- ori_args.emplace_back(n);
115- ori_args.emplace_back(k);
116- 
117- l0_args.emplace_back(basem);
118- l0_args.emplace_back(basen);
119- l0_args.emplace_back(basek);
120- 
121- mc_args.emplace_back(tilem);
122- mc_args.emplace_back(tilen);
123- 
124- ori_arg_map[basem] = m;
125- ori_arg_map[basen] = n;
126- ori_arg_map[basek] = k;
127- 
128- father_args_map[basem] = tilem;
129- father_args_map[basen] = tilen;
130- 
131- arg_align_map[basem] = ge::Symbol(256);
132- arg_align_map[basen] = ge::Symbol(16);
133- arg_align_map[basek] = ge::Symbol(16);
134- arg_align_map[tilem] = ge::Symbol(16);
135- arg_align_map[tilen] = ge::Symbol(16);
136- 
137- solver_gen.SetMulticoreArgs(mc_args);
138- solver_gen.SetFatherArgsMap(father_args_map);
139- solver_gen.SetArgAlignMap(arg_align_map);
140- solver_gen.SetArgtMaxValueMap(ori_arg_map);
141- solver_gen.SetL0Args(l0_args);
142- solver_gen.SetBufferUseAlg(buffer_use_map);
143- std::string impl_code = solver_gen.GenSolverFuncImpl();
144- std::string invoke_code = solver_gen.GenSolverFuncInvoke();
145- EXPECT_NE(impl_code, "");
146- EXPECT_NE(invoke_code, "");
147-}
@@ -1,127 +0,0 @@
1-/**
2- * Copyright (c) 2026 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-#include "gtest/gtest.h"
12-#include "base/base_types.h"
13-#include "base/model_info.h"
14-#include "generator/solver_pass_gen/l2_solver/l2_solver_gen.h"
15-#include <symengine/functions.h>
16-#include <symengine/simplify.h>
17-#include <symengine/integer.h>
18-#include <symengine/real_double.h>
19-#include "test_common_utils.h"
20-using namespace att;
21- 
22-class TestL2SolverGen : public ::testing::Test {
23- public:
24- void TearDown() override {
25- // 清理测试生成的临时文件
26- autofuse::test::CleanupTestArtifacts();
27- // before the destructor).
28- }
29-};
30- 
31-TEST_F(TestL2SolverGen, TEST_CASE_01) {
32- Expr m = CreateExpr("m");
33- Expr n = CreateExpr("n");
34- Expr k = CreateExpr("k");
35- 
36- Expr basem = CreateExpr("basem");
37- Expr basen = CreateExpr("basen");
38- 
39- Expr tilem = CreateExpr("tilem");
40- Expr tilen = CreateExpr("tilen");
41- 
42- std::vector<att::Expr> input_args;
43- std::vector<att::Expr> l0_args;
44- std::vector<att::Expr> l2_args;
45- std::map<att::Expr, Expr, att::ExprCmp> ori_arg_map;
46- std::map<att::Expr, Expr, att::ExprCmp> arg_align_map;
47- 
48- input_args.emplace_back(m);
49- input_args.emplace_back(n);
50- input_args.emplace_back(k);
51- ori_arg_map[basem] = m;
52- ori_arg_map[basen] = n;
53- ori_arg_map[tilem] = m;
54- ori_arg_map[tilen] = n;
55- arg_align_map[tilem] = ge::Symbol(16);
56- arg_align_map[tilen] = ge::Symbol(16);
57- 
58- l0_args.emplace_back(basem);
59- l0_args.emplace_back(basen);
60- 
61- l2_args.emplace_back(tilem);
62- l2_args.emplace_back(tilen);
63- 
64- att::Expr l2_use_expr = ((((tilem * k) + (tilen * k)) + (tilem * tilen)) * CreateExpr(2));
65- 
66- att::L2TileSolverGen *solver_gen = new att::L2TileSolverGen("Case0", "TilingData");
67- solver_gen->SetArgAlignMap(arg_align_map);
68- solver_gen->SetL0Args(l0_args);
69- solver_gen->SetL2Args(l2_args);
70- solver_gen->SetL2Use(l2_use_expr);
71- solver_gen->SetArgtMaxValueMap(ori_arg_map);
72- solver_gen->SetInputArgs(input_args);
73- std::string impl_code = solver_gen->GenSolverFuncImpl();
74- std::string invoke_code = solver_gen->GenSolverFuncInvoke();
75- EXPECT_NE(impl_code, "");
76- EXPECT_NE(invoke_code, "");
77- delete solver_gen;
78-}
79- 
80-TEST_F(TestL2SolverGen, TEST_CASE_02) {
81- Expr m = CreateExpr("m");
82- Expr n = CreateExpr("n");
83- Expr k = CreateExpr("k");
84- 
85- Expr basem = CreateExpr("basem");
86- Expr basen = CreateExpr("basen");
87- 
88- Expr tilem = CreateExpr("tilem");
89- Expr tilen = CreateExpr("tilen");
90- 
91- std::vector<att::Expr> input_args;
92- std::vector<att::Expr> l0_args;
93- std::vector<att::Expr> l2_args;
94- std::map<att::Expr, Expr, att::ExprCmp> ori_arg_map;
95- std::map<att::Expr, Expr, att::ExprCmp> arg_align_map;
96- 
97- input_args.emplace_back(m);
98- input_args.emplace_back(n);
99- input_args.emplace_back(k);
100- ori_arg_map[basem] = m;
101- ori_arg_map[basen] = n;
102- ori_arg_map[tilem] = m;
103- ori_arg_map[tilen] = n;
104- arg_align_map[tilem] = ge::Symbol(16);
105- arg_align_map[tilen] = ge::Symbol(256);
106- 
107- l0_args.emplace_back(basem);
108- l0_args.emplace_back(basen);
109- 
110- l2_args.emplace_back(tilem);
111- l2_args.emplace_back(tilen);
112- 
113- att::Expr l2_use_expr = ((((tilem * k) + (tilen * k)) + (tilem * tilen)) * CreateExpr(2));
114- 
115- att::L2TileSolverGen *solver_gen = new att::L2TileSolverGen("Case0", "TilingData");
116- solver_gen->SetArgAlignMap(arg_align_map);
117- solver_gen->SetL0Args(l0_args);
118- solver_gen->SetL2Args(l2_args);
119- solver_gen->SetL2Use(l2_use_expr);
120- solver_gen->SetArgtMaxValueMap(ori_arg_map);
121- solver_gen->SetInputArgs(input_args);
122- std::string impl_code = solver_gen->GenSolverFuncImpl();
123- std::string invoke_code = solver_gen->GenSolverFuncInvoke();
124- EXPECT_NE(impl_code, "");
125- EXPECT_NE(invoke_code, "");
126- delete solver_gen;
127-}
@@ -20,15 +20,10 @@ file(GLOB SOURCES
20 # generator/preprocess20 # generator/preprocess
21 ${TOP_DIR}/tests/ut/att/testcase/preprocess/*.cpp21 ${TOP_DIR}/tests/ut/att/testcase/preprocess/*.cpp
22 # generator/solver_pass22 # generator/solver_pass
23- ${TOP_DIR}/tests/ut/att/testcase/solver_pass/l0_solver/*.cpp
24- ${TOP_DIR}/tests/ut/att/testcase/solver_pass/l2_solver/*.cpp
25 ${TOP_DIR}/tests/ut/att/testcase/solver_pass/general_solver/*.cpp23 ${TOP_DIR}/tests/ut/att/testcase/solver_pass/general_solver/*.cpp
26- ${TOP_DIR}/tests/ut/att/testcase/solver_pass/src/*.cpp
27 # generator/solver_pass_gen24 # generator/solver_pass_gen
28 ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/axes_reorder_gen/*.cpp25 ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/axes_reorder_gen/*.cpp
29 ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/general_solver_gen/*.cpp26 ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/general_solver_gen/*.cpp
30- ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/l2_solver_gen/*.cpp
31- ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/l0_solver_gen/*.cpp
32 ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/manager/*.cpp27 ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/manager/*.cpp
33 # generator/core28 # generator/core
34 ${TOP_DIR}/tests/ut/att/testcase/generator/core/*.cpp29 ${TOP_DIR}/tests/ut/att/testcase/generator/core/*.cpp
@@ -60,9 +55,6 @@ file(GLOB SOURCES
60 ${TOP_DIR}/tests/ut/att/utils/*.cpp55 ${TOP_DIR}/tests/ut/att/utils/*.cpp
61 )56 )
62 57 
63-# solver_pass/src/*.cpp 是源码镜像,从 ATT_DIR 编译,排除以避免 multiple definition
64-list(FILTER SOURCES EXCLUDE REGEX "solver_pass/src/.+\\.cpp$")
65- 
66add_executable(att_ut58add_executable(att_ut
67 ${SOURCES}59 ${SOURCES}
68 )60 )
@@ -184,21 +184,6 @@ TEST(GeneratorUT, DurationSplitSourcesIncludeDirectDependencies) {
184 ExpectSystemHeaders(solver_source, {"chrono", "memory", "new", "string"}, {});184 ExpectSystemHeaders(solver_source, {"chrono", "memory", "new", "string"}, {});
185}185}
186 186 
187-TEST(GeneratorUT, HighPerfSolverHeaderOmitsCstddefWithoutGeneralSolver) {
188- TilingModelInfo model_infos{CreateModelInfo(1U, af::ExprType::kExprConstantInteger)};
189- ASSERT_EQ(ReuseGroupUtils::InitReuseScheduleGroup({0UL, 0UL, 0UL}, model_infos), af::SUCCESS);
190- TilingCodeGenConfig config;
191- config.type = TilingImplType::HIGH_PERF;
192- config.tiling_data_type_name = "OpTestTilingData";
193- std::map<std::string, std::string> tiling_res;
194- TilingCodeGenerator generator;
195- 
196- ASSERT_EQ(generator.GenTilingCode(op_name, model_infos, config, tiling_res), af::SUCCESS);
197- const auto &solver_header = tiling_res.at(kTilingSolverHeaderIdentify);
198- EXPECT_EQ(solver_header.find("GetTemp(size_t idx)"), std::string::npos);
199- EXPECT_EQ(solver_header.find("#include <cstddef>"), std::string::npos);
200-}
201- 
202TEST(GeneratorUT, NormalGroupRegistersOnlyDirectStandardHeaders) {187TEST(GeneratorUT, NormalGroupRegistersOnlyDirectStandardHeaders) {
203 TilingModelInfo model_infos{CreateModelInfo()};188 TilingModelInfo model_infos{CreateModelInfo()};
204 ASSERT_EQ(ReuseGroupUtils::InitReuseScheduleGroup({0UL, 0UL, 0UL}, model_infos), af::SUCCESS);189 ASSERT_EQ(ReuseGroupUtils::InitReuseScheduleGroup({0UL, 0UL, 0UL}, model_infos), af::SUCCESS);
@@ -724,7 +709,13 @@ TEST(GeneratorUT, AxesReorderSolverHeaderRegistersOnlyDirectStandardHeaders) {
724 ModelInfo model_info = CreateModelInfo();709 ModelInfo model_info = CreateModelInfo();
725 size_t order = 1U;710 size_t order = 1U;
726 for (const auto &arg : model_info.arg_list) {711 for (const auto &arg : model_info.arg_list) {
727- arg->order = (arg->name == "tilem" || arg->name == "tilen") ? 0U : order++;712+ if (arg->name == "stepm" || arg->name == "stepn") {
713+ // R2 剥离后 INNER 轴仅剩 stepm/stepn;置 bind_multicore=false + 同 order 构造 equal-order 触发条件
714+ arg->bind_multicore = false;
715+ arg->order = 0U;
716+ } else {
717+ arg->order = order++;
718+ }
728 }719 }
729 TilingModelInfo model_infos{model_info};720 TilingModelInfo model_infos{model_info};
730 ASSERT_EQ(ReuseGroupUtils::InitReuseScheduleGroup({0UL, 0UL, 0UL}, model_infos), af::SUCCESS);721 ASSERT_EQ(ReuseGroupUtils::InitReuseScheduleGroup({0UL, 0UL, 0UL}, model_infos), af::SUCCESS);
@@ -1,281 +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-#include "gtest/gtest.h"
12-#include "gmock/gmock.h"
13-#include "base/base_types.h"
14-#define private public
15-#define protected public
16-#include "solver_pass/src/l0_solver.cpp"
17-#include <symengine/functions.h>
18-#include <symengine/simplify.h>
19-#include <symengine/integer.h>
20-#include <symengine/real_double.h>
21-using namespace att;
22- 
23-class MockUtL0TileSolver : public L0TileSolver {
24- public:
25- MockUtL0TileSolver() {}
26- explicit MockUtL0TileSolver(L0TileInput input) : L0TileSolver(input) {}
27- void SetL0A(uint32_t value) {
28- L0A_ = value;
29- }
30- void SetL0B(uint32_t value) {
31- L0B_ = value;
32- }
33- void SetL0C(uint32_t value) {
34- L0C_ = value;
35- }
36- bool CheckBufferUseValid() override {
37- uint32_t l0A = input_.l0_vars[0].value * input_.l0_vars[2].value * 4;
38- uint32_t l0B = input_.l0_vars[2].value * input_.l0_vars[1].value * 4;
39- uint32_t l0C = input_.l0_vars[0].value * input_.l0_vars[1].value * 4;
40- if (l0A > L0A_ || l0B > L0B_ || l0C > L0C_) {
41- return false;
42- }
43- return true;
44- }
45- 
46- private:
47- uint32_t L0A_;
48- uint32_t L0B_;
49- uint32_t L0C_;
50-};
51- 
52-class TestL0SolverUt : public ::testing::Test {
53- public:
54- static void TearDownTestCase() {
55- std::cout << "Test end." << std::endl;
56- }
57- static void SetUpTestCase() {
58- std::cout << "Test begin." << std::endl;
59- }
60- void SetUp() override {
61- // Code here will be called immediately after the constructor (right
62- // before each test).
63- }
64- 
65- void TearDown() override {
66- // Code here will be called immediately after each test (right
67- // before the destructor).
68- }
69- MockUtL0TileSolver solver;
70-};
71- 
72-TEST_F(TestL0SolverUt, TEST_CHECK_INPUT) {
73- EXPECT_EQ(solver.CheckInput(), false);
74- 
75- L0Var M;
76- L0Var N;
77- L0Var K;
78- M.max_value = 32;
79- M.bind_multicore = 1;
80- M.align = 16;
81- M.prompt_align = 16;
82- 
83- N.max_value = 1024;
84- N.bind_multicore = 1;
85- N.align = 16;
86- N.prompt_align = 256;
87- 
88- K.max_value = 1024;
89- K.bind_multicore = 0;
90- K.align = 16;
91- K.prompt_align = 256;
92- 
93- L0TileInput input;
94- input.l0_vars = new L0Var[3];
95- input.l0_vars[0] = M;
96- input.l0_vars[1] = N;
97- input.l0_vars[2] = K;
98- input.size = 3;
99- input.core_num = 24;
100- solver = MockUtL0TileSolver(input);
101- EXPECT_EQ(solver.CheckInput(), true);
102- solver.input_.core_num = 0;
103- EXPECT_EQ(solver.CheckInput(), false);
104- solver.input_.core_num = 24;
105- solver.input_.l0_vars[0].align = 0;
106- EXPECT_EQ(solver.CheckInput(), false);
107- solver.input_.l0_vars[0].align = 32;
108- EXPECT_EQ(solver.CheckInput(), false);
109- solver.input_.l0_vars[0].align = 16;
110- solver.input_.l0_vars[0].prompt_align = 0;
111- EXPECT_EQ(solver.CheckInput(), false);
112- 
113- L0TileInput input_2;
114- input_2.l0_vars = new L0Var[0];
115- input_2.size = 0;
116- input_2.core_num = 24;
117- solver = MockUtL0TileSolver(input_2);
118- EXPECT_EQ(solver.CheckInput(), false);
119- 
120- L0TileInput input_3;
121- input.l0_vars = new L0Var[4];
122- input.size = 4;
123- input.core_num = 24;
124- solver = MockUtL0TileSolver(input);
125- EXPECT_EQ(solver.CheckInput(), false);
126-}
127- 
128-TEST_F(TestL0SolverUt, TEST_MAX_CORE_NUM) {
129- L0Var M;
130- L0Var N;
131- L0Var K;
132- 
133- M.max_value = 32;
134- M.bind_multicore = 1;
135- M.align = 16;
136- M.prompt_align = 16;
137- 
138- N.max_value = 1024;
139- N.bind_multicore = 1;
140- N.align = 16;
141- N.prompt_align = 256;
142- 
143- K.max_value = 1024;
144- K.bind_multicore = 0;
145- K.align = 16;
146- K.prompt_align = 256;
147- 
148- L0TileInput input;
149- input.l0_vars = new L0Var[3];
150- input.l0_vars[0] = M;
151- input.l0_vars[1] = N;
152- input.l0_vars[2] = K;
153- input.size = 3;
154- input.core_num = 24;
155- solver = MockUtL0TileSolver(input);
156- EXPECT_EQ(solver.MaxCoreNum(input.l0_vars, input.core_num), 8);
157-}
158- 
159-TEST_F(TestL0SolverUt, TEST_GET_MAC_USE) {
160- L0Var M;
161- L0Var N;
162- L0Var K;
163- 
164- M.max_value = 32;
165- M.bind_multicore = 1;
166- M.align = 16;
167- M.prompt_align = 16;
168- 
169- N.max_value = 1024;
170- N.bind_multicore = 1;
171- N.align = 16;
172- N.prompt_align = 256;
173- 
174- K.max_value = 1024;
175- K.bind_multicore = 0;
176- K.align = 16;
177- K.prompt_align = 256;
178- 
179- L0TileInput input;
180- input.l0_vars = new L0Var[3];
181- input.l0_vars[0] = M;
182- input.l0_vars[1] = N;
183- input.l0_vars[2] = K;
184- input.l0_vars[0].value = 16;
185- input.l0_vars[1].value = 16;
186- input.l0_vars[2].value = 16;
187- input.size = 3;
188- input.core_num = 24;
189- solver = MockUtL0TileSolver(input);
190- EXPECT_EQ(solver.GetMacUse(), 4096);
191-}
192- 
193-TEST_F(TestL0SolverUt, TEST_CHECK_OUTPUT) {
194- L0Var M;
195- L0Var N;
196- L0Var K;
197- 
198- M.max_value = 32;
199- M.bind_multicore = 1;
200- M.align = 16;
201- M.prompt_align = 16;
202- 
203- N.max_value = 1024;
204- N.bind_multicore = 1;
205- N.align = 16;
206- N.prompt_align = 256;
207- 
208- K.max_value = 1024;
209- K.bind_multicore = 0;
210- K.align = 16;
211- K.prompt_align = 256;
212- 
213- L0TileInput input;
214- input.l0_vars = new L0Var[3];
215- input.l0_vars[0] = M;
216- input.l0_vars[1] = N;
217- input.l0_vars[2] = K;
218- input.size = 3;
219- input.core_num = 24;
220- solver = MockUtL0TileSolver(input);
221- EXPECT_EQ(solver.CheckOutput(), false);
222- solver.SetL0A(1);
223- solver.SetL0B(1);
224- solver.SetL0C(1);
225- solver.Run();
226- EXPECT_EQ(solver.CheckOutput(), false);
227- solver.SetL0A(64 * 1024);
228- solver.SetL0B(64 * 1024);
229- solver.SetL0C(128 * 1024);
230- solver.Run();
231- EXPECT_EQ(solver.CheckOutput(), true);
232-}
233- 
234-TEST_F(TestL0SolverUt, TEST_RUN) {
235- L0Var M;
236- L0Var N;
237- L0Var K;
238- 
239- M.max_value = 32;
240- M.bind_multicore = 1;
241- M.align = 16;
242- M.prompt_align = 16;
243- 
244- N.max_value = 1024;
245- N.bind_multicore = 1;
246- N.align = 16;
247- N.prompt_align = 256;
248- 
249- K.max_value = 1024;
250- K.bind_multicore = 0;
251- K.align = 16;
252- K.prompt_align = 256;
253- 
254- L0TileInput input;
255- input.l0_vars = new L0Var[3];
256- input.l0_vars[0] = M;
257- input.l0_vars[1] = N;
258- input.l0_vars[2] = K;
259- input.size = 3;
260- input.core_num = 24;
261- uint32_t *best_value = new uint32_t[input.size];
262- solver = MockUtL0TileSolver(input);
263- solver.SetL0A(1);
264- solver.SetL0B(1);
265- solver.SetL0C(1);
266- EXPECT_EQ(solver.Run(), false);
267- solver.SetL0A(64 * 1024);
268- solver.SetL0B(64 * 1024);
269- solver.SetL0C(128 * 1024);
270- EXPECT_EQ(solver.Run(), true);
271- uint32_t *output = solver.GetOutput();
272- EXPECT_NE(output, nullptr);
273- EXPECT_EQ(output[0], 16);
274- EXPECT_EQ(output[1], 256);
275- EXPECT_EQ(output[2], 64);
276- solver.input_.size = 0;
277- EXPECT_EQ(solver.Run(), false);
278- solver.input_.size = 3;
279- solver.input_.l0_vars[0].max_value = 0;
280- EXPECT_EQ(solver.Run(), false);
281-}
@@ -1,158 +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-#include "gtest/gtest.h"
12-#include "base/base_types.h"
13-#define private public
14-#define protected public
15-#include "solver_pass/src/l2_solver.cpp"
16-using namespace att;
17- 
18-class MockUtL2TileSolver : public L2TileSolver {
19- public:
20- MockUtL2TileSolver() {}
21- explicit MockUtL2TileSolver(L2TileInput input) : L2TileSolver(input) {};
22- uint64_t GetL2Use() override {
23- uint64_t tilel2m = input_.l2_vars[0].value;
24- uint64_t tilel2n = input_.l2_vars[1].value;
25- uint64_t k = 512;
26- uint64_t l2Use = tilel2m * k * 2 + k * tilel2n * 2 + tilel2m * tilel2n * 2;
27- return l2Use;
28- }
29- bool IsClash(uint32_t idx) override {
30- if (used_corenum_ <= 1 || used_corenum_ % 2 != 0) {
31- return false;
32- }
33- if (blocknum_per_tile_[idx] % (used_corenum_ / 2) == 0) {
34- return true;
35- }
36- auto blockNumTail = total_blocknum_[idx] - (tilenum_[idx] - 1) * blocknum_per_tile_[idx];
37- if (blocknum_per_tile_[idx] % (used_corenum_ / 2) == 0) {
38- return true;
39- }
40- return false;
41- }
42-};
43- 
44-class TestL2SolverUt : public ::testing::Test {
45- public:
46- static void TearDownTestCase() {
47- std::cout << "Test end." << std::endl;
48- }
49- static void SetUpTestCase() {
50- std::cout << "Test begin." << std::endl;
51- }
52- void SetUp() override {
53- // Code here will be called immediately after the constructor (right
54- // before each test).
55- }
56- 
57- void TearDown() override {
58- // Code here will be called immediately after each test (right
59- // before the destructor).
60- }
61-};
62- 
63-TEST_F(TestL2SolverUt, TEST_RUN) {
64- L2TileInput input;
65- L2Var tilem;
66- L2Var tilen;
67- tilem.max_value = 8196;
68- tilen.max_value = 8196;
69- tilem.align = 16;
70- tilen.align = 16;
71- tilem.base_val = 128;
72- tilen.base_val = 256;
73- input.l2_vars = new L2Var[2];
74- input.l2_vars[0] = tilem;
75- input.l2_vars[1] = tilen;
76- input.size = 2;
77- input.core_num = 24;
78- input.l2_size = 128 * 1024 * 1024;
79- MockUtL2TileSolver l2_solver(input);
80- bool status = l2_solver.Run();
81- EXPECT_EQ(status, true);
82- uint32_t *output = l2_solver.GetL2Tile();
83- EXPECT_NE(output, nullptr);
84- EXPECT_EQ(output[0], 7552);
85- EXPECT_EQ(output[1], 7680);
86-}
87- 
88-TEST_F(TestL2SolverUt, TEST_CHECK_INPUT) {
89- L2TileInput input;
90- L2Var tilem;
91- L2Var tilen;
92- tilem.max_value = 8196;
93- tilen.max_value = 8196;
94- tilem.align = 16;
95- tilen.align = 16;
96- tilem.base_val = 128;
97- tilen.base_val = 256;
98- input.l2_vars = new L2Var[2];
99- input.l2_vars[0] = tilem;
100- input.l2_vars[1] = tilen;
101- input.size = 2;
102- input.core_num = 24;
103- input.l2_size = 128 * 1024 * 1024;
104- MockUtL2TileSolver l2_solver(input);
105- EXPECT_EQ(l2_solver.CheckInput(), true);
106- l2_solver.input_.l2_vars[0].align = 0;
107- EXPECT_EQ(l2_solver.CheckInput(), false);
108- l2_solver.input_.l2_vars[0].align = 16;
109- l2_solver.input_.size = 0;
110- EXPECT_EQ(l2_solver.CheckInput(), false);
111- l2_solver.input_.size = 2;
112- 
113- MockUtL2TileSolver null_l2_solver;
114- EXPECT_EQ(null_l2_solver.CheckInput(), false);
115-}
116- 
117-TEST_F(TestL2SolverUt, TEST_CHECK_SOLVABLE) {
118- L2TileInput input;
119- L2Var tilem;
120- L2Var tilen;
121- tilem.max_value = 8196;
122- tilen.max_value = 8196;
123- tilem.align = 16;
124- tilen.align = 16;
125- tilem.base_val = 128;
126- tilen.base_val = 256;
127- input.l2_vars = new L2Var[2];
128- input.l2_vars[0] = tilem;
129- input.l2_vars[1] = tilen;
130- input.size = 2;
131- input.core_num = 24;
132- input.l2_size = 128 * 1024 * 1024;
133- MockUtL2TileSolver l2_solver(input);
134- EXPECT_EQ(l2_solver.CheckSolvable(), true);
135- l2_solver.input_.l2_size = 1;
136- EXPECT_EQ(l2_solver.CheckSolvable(), false);
137-}
138- 
139-TEST_F(TestL2SolverUt, TEST_BLOCK_NUM_PER_TILE_1) {
140- L2TileInput input;
141- L2Var tilem;
142- L2Var tilen;
143- tilem.max_value = 128;
144- tilen.max_value = 8196;
145- tilem.align = 16;
146- tilen.align = 16;
147- tilem.base_val = 128;
148- tilen.base_val = 256;
149- input.l2_vars = new L2Var[2];
150- input.l2_vars[0] = tilem;
151- input.l2_vars[1] = tilen;
152- input.size = 2;
153- input.core_num = 24;
154- input.l2_size = 128 * 1024 * 1024;
155- MockUtL2TileSolver l2_solver(input);
156- bool status = l2_solver.Run();
157- EXPECT_EQ(status, true);
158-}
@@ -1,549 +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-#include <cstdlib>
12-#include <iostream>
13-#include <memory.h>
14-#include <algorithm>
15-#ifdef DEBUG
16-#define ATT_LOG(log)
17-do {
18- std::cout << "[ERROR]" << log << std::endl;
19-} while (0)
20-#else
21-#define ATT_LOG(log)
22-#endif
23- namespace att {
24- // L0Var的备选值的个数
25- static const uint32_t candidate_size = 7u;
26- // L0Var的备选值
27- static const uint32_t candidate_value[] = {16u, 32u, 64u, 128u, 256u, 512u, 1024u};
28- // 表达L0的求解值至少要满足核数的比例,可手动修改
29- static const double CORE_NUM_RATIO = 0.6f;
30- // 表达L0的求解值pad之后的值不允许超过原始值大小的倍数,可手动修改
31- static const uint32_t UPPER_BOUND_RATIO = 2u;
32- // 表达最大L0Var的个数
33- static const uint32_t MAX_L0_VAR_NUM = 3u;
34- 
35- // L0相关变量的数据结构
36- struct L0Var {
37- // 最大值,初始化为输入原始轴的大小
38- uint32_t max_value{0u};
39- // 是否绑多核
40- bool bind_multicore{false};
41- bool is_innermost{false};
42- // 对齐值
43- uint32_t align{0u};
44- // 提示当前L0Var的最佳对齐值,通常源于父轴的对齐值,
45- // 举例,stepm是basem的父轴,stepm的对齐要求是256和basem,那么basem的prompt_align就是256,
46- // 同时约束L0Var的取值必须是256对齐或者是256的因子,因为这样父轴stepm才能既满足256也满足basem对齐
47- uint32_t prompt_align{0u};
48- // L0变量的索引
49- uint32_t idx;
50- // L0变量的值
51- uint32_t value{0u};
52- };
53- 
54- // 求解器接收的输入
55- struct L0TileInput {
56- // 待求解L0变量的集合
57- L0Var *l0_vars{nullptr};
58- // 待求解L0变量的数量
59- uint32_t size;
60- // 核数
61- uint32_t core_num;
62- };
63- 
64- /**
65- * 比较两个 L0Var 类型变量的大小
66- *
67- * 这个函数用来比较两个 L0Var 类型变量的大小。它遵循特定的比较逻辑:
68- * 1. 如果 a 变量绑定到多核且 b 变量没有绑定多核,则 a 被认为大于 b,函数返回
69- * true
70- * 2. 如果 a 变量没有绑定多核且 b 变量绑定多核,则 a 被认为小于 b,函数返回
71- * false
72- * 3. 如果 a 和 b 变量都绑定多核或都未绑定多核,则比较它们的 prompt_align
73- * 属性,prompt_align 属性值大的变量被认为更大。
74- *
75- * @param a 第一个要比较的 L0Var 变量。
76- * @param b 第二个要比较的 L0Var 变量。
77- * @return 如果 a 大于 b,返回 true;如果 a 小于 b,返回 false
78- */
79- static bool L0VarCmp(L0Var a, L0Var b) {
80- if (a.bind_multicore && !b.bind_multicore) {
81- return true;
82- }
83- if (!a.bind_multicore && b.bind_multicore) {
84- return false;
85- }
86- if (a.is_innermost && !b.is_innermost) {
87- return true;
88- }
89- if (!a.is_innermost && b.is_innermost) {
90- return false;
91- }
92- return a.prompt_align > b.prompt_align;
93- }
94- /**
95- * L0 求解器类
96- */
97- class L0TileSolver {
98- public:
99- /**
100- * 构造函数
101- *
102- * @param input 一个 L0TileInput 结构体,包含了 L0 变量相关的信息
103- *
104- * 这个构造函数初始化了 L0TileSolver 对象
105- */
106- explicit L0TileSolver(L0TileInput input) : input_(input) {}
107- L0TileSolver() {};
108- /**
109- * 析构函数
110- *
111- * 当 L0TileSolver 对象被销毁时,析构函数被调用
112- * 用来释放使用 new 运算符动态分配的内存,确保没有内存泄漏
113- */
114- ~L0TileSolver() {
115- if (sortedvars_ != nullptr) {
116- delete[] sortedvars_;
117- }
118- if (output_ != nullptr) {
119- delete[] output_;
120- }
121- }
122- /**
123- * 运行求解器
124- *
125- * @return 如果求解成功,返回 true;否则返回 false
126- *
127- * 这个方法是算法的入口点,调用它会启动求解过程
128- * 成功与否取决于 CheckBufferUseValid() 方法的返回值
129- */
130- bool Run();
131- /**
132- * 获取优化结果
133- *
134- * @return 指向求解结果数据的指针
135- *
136- * 如果有求解结果,这个方法将返回一个指向结果数据的指针
137- * 结果数据的内存使用完后,应该使用 delete[] 释放内存
138- */
139- uint32_t *GetOutput() {
140- return output_;
141- }
142- 
143- protected:
144- /**
145- * 检查是否满足buffer约束
146- *
147- * @return 如果满足,返回 true;否则返回 false
148- *
149- * 这个纯虚函数要求派生类提供实现,以确保缓冲区使用是有效的
150- * 在 L0TileSolver 类中,它是一个抽象方法,需要在子类中实现
151- */
152- virtual bool CheckBufferUseValid() = 0;
153- L0TileInput input_;
154- uint32_t *output_{nullptr};
155- 
156- private:
157- /**
158- * 检查输入数据的完整性和正确性
159- *
160- * @return 如果输入数据有效,返回 true;否则返回 false
161- *
162- * 这个私有方法检查输入数据的格式和逻辑,确保它们适用求解算法
163- */
164- bool CheckInput();
165- 
166- /**
167- * 使用输入数据初始化算法所需的内部数据结构
168- *
169- * 这个方法根据输入的 L0TileInput 结构体中的数据,初始化算法所需的内部数据结构
170- * 确保 sortedvars_ 和 output_ 成员变量被正确初始化
171- */
172- void InitInput();
173- 
174- /**
175- * 检查算法运行的结果,确保它们符合预期
176- *
177- * @return 如果输出数据有效,返回 true;否则返回 false
178- *
179- * 这个方法检查运行算法后得到的结果,确保它们在逻辑上是合理的
180- */
181- bool CheckOutput();
182- 
183- /**
184- * 更新算法执行过程中的对齐设置
185- *
186- * 这个方法根据算法执行过程中的数据更新对齐提示值,确保结果按照预期的方式对齐
187- */
188- void UpdateAlign();
189- 
190- /**
191- * 为指定索引的 L0 变量获取最佳对齐值
192- *
193- * @param i 想要的变量索引值
194- * @return 最佳对齐值
195- *
196- * 这个方法计算并返回给定索引值的 L0 变量的最佳对齐值,
197- * 确保变量以最恰当的方式对齐,从而提高效率或者减少资源浪费
198- */
199- uint32_t GetBestAlign(uint32_t i) const;
200- 
201- /**
202- * 为 L0 变量找到最优值进行迭代运行
203- *
204- * @param loop_id 当前循环索引,表示正在处理的 L0 变量的位置
205- * @param best_var_value 一个指针,指向用于存储每个 L0
206- * 变量迄今为止找到的最佳值的数组
207- *
208- * 这个函数使用递归方法来遍历 L0
209- * 变量的所有可能值。对于每个值,它检查是否满足约束条件,如小于上界且满足对齐要求。如果满足这些条件,它将继续下一个循环或者递归调用自身来处理下一个
210- * L0 变量。如果是最后一个 L0
211- * 变量,它将检查当前组合是否满足优化条件,如核心数量和数据处理量。如果满足条件,它将当前组合存储为最优解。
212- *
213- * 请注意,这个函数没有返回值,而是将最优解存储在传入的 best_var_value
214- * 数组中。
215- */
216- void IterativeRun(uint32_t loop_id, uint32_t *best_var_value);
217- 
218- /**
219- * 根据 L0 变量信息和总核心数计算可以分配的最大核心数
220- *
221- * @param l0_vars 一个指向 L0Var 结构体数组的指针
222- * @param core_num 可用于分配的总核心数
223- * @return 可以分配的最大核心数
224- *
225- * 这个函数计算在给定 L0 变量信息和总核心数的情况下,可以分配的最大核心数。
226- * 它遍历输入的 L0Var
227- * 结构体数组,对于每个变量,根据其是否绑定多核心以及最大、当前和提示对齐值计算所需的块数。
228- * 通过将所有变量的块数相乘,得到总块数。
229- * 最大核心数是总块数和总核心数中的最小值,以确保核心数不会超过可用资源。
230- *
231- * 返回值表示可以分配给 L0
232- * 变量的最大核心数,这对于在多核心系统中进行资源分配是有用的。
233- */
234- int32_t MaxCoreNum(const L0Var *l0_vars, const uint32_t &core_num);
235- 
236- /**
237- * 计算所有 L0 变量值的乘积,作为mac计算量的度量
238- *
239- * @return mac计算量
240- *
241- * 这个函数计算所有 L0 变量值的乘积,结果是一个数字。
242- * 这个数字可以作为数据处理量的度量,例如在评估算法性能时。
243- * 通过不断更新 usage 变量,乘法操作确保了所有 L0 变量的影响都被计入。
244- * 最终 usage 变量中的值就是所有 L0 变量值的乘积,代表了整体的数据处理量。
245- * 返回值可以帮助了解算法处理的数据量,从而对算法的效率和扩展性有更直观的认识。
246- */
247- uint32_t GetMacUse() const;
248- /**
249- * 用于排序的 L0Var 对象数组
250- */
251- L0Var *sortedvars_{nullptr};
252- /**
253- * 最大核心数
254- */
255- int64_t max_corenum_{-1};
256- /**
257- * 最大 MAC 使用量
258- */
259- int32_t max_macuse_{-1};
260- };
261- 
262- /**
263- * 获取给定索引的 L0 变量的最佳对齐值
264- *
265- * @param i 想要的变量索引值
266- * @return 最佳对齐值
267- *
268- * 这个方法为给定索引的 L0
269- * 变量计算最佳对齐值。它考虑到变量的最大、当前和提示对齐值,以确保数据存储和访问的效率。
270- * 根据变量的原始值(ori_value),它首先确定最小和最大对齐值的范围。然后,通过在这个范围内以二的幂次方递增,它找到最大的满足条件的值。
271- * 如果没有找到这样的值,它将返回最小对齐值。如果在范围内找到了一个值,它将返回这个值的二分之一,作为最佳对齐值。
272- * 这个最佳对齐值可以用于确保数据以最有效的方式存储
273- */
274- uint32_t L0TileSolver::GetBestAlign(uint32_t i) const {
275- uint32_t ori_value = input_.l0_vars[i].max_value;
276- uint32_t min_align = input_.l0_vars[i].align;
277- uint32_t max_align = input_.l0_vars[i].prompt_align;
278- uint32_t ori_align = min_align;
279- uint32_t max_value = std::min(ori_value, max_align);
280- while (ori_align <= max_value) {
281- ori_align = ori_align << 1;
282- }
283- if (ori_align == min_align) {
284- return min_align;
285- }
286- return std::max(1u, ori_align >> 1);
287- }
288- 
289- /**
290- * 根据给定的 L0 变量信息计算可以分配的最大核心数
291- *
292- * @param l0_vars 指向 L0Var 结构数组的指针
293- * @param core_num 总核心数
294- * @return 可以分配的最大核心数
295- *
296- * 这个函数遍历 L0Var 结构数组,根据每个变量的 bind_multicore 属性以及
297- * max_value、value 和 prompt_align
298- * 的值来计算每个变量所需的块数。对于绑定多核心的变量,块数计算方式为:(max_value
299- * + max(value,prompt_align)-1)/
300- * max(value,prompt_align)。对于未绑定多核心的变量,块数为 1
301- * 总块数通过将所有变量的块数相乘得到。之后,通过比较总块数和
302- * core_num,返回两者中的最小值,作为可以分配的最大核心数。如果总块数超过了
303- * core_num,那么系统的核心数将成为瓶颈,因此需要将 core_num
304- * 设置为最大核心数。如果总块数小于等于
305- * core_num,那么总块数就是可以分配的最大核心数。
306- */
307- int32_t L0TileSolver::MaxCoreNum(const L0Var *l0_vars, const uint32_t &core_num) {
308- uint32_t total_block_size = 1u;
309- for (uint32_t i = 0u; i < input_.size; i++) {
310- auto var = l0_vars[i];
311- uint32_t block_num = var.bind_multicore ? ((var.max_value + std::max(var.value, var.prompt_align) - 1)) /
312- std::max(var.value, var.prompt_align)
313- : 1;
314- total_block_size *= block_num;
315- }
316- int64_t max_core_num = total_block_size > core_num ? core_num : total_block_size;
317- return max_core_num;
318- }
319- 
320- /**
321- * 计算所有 L0 变量值的乘积,作为数据处理量的度量
322- *
323- * @return 数据处理量
324- *
325- * 这个函数遍历 L0TileInput 结构体中的所有 L0Var 对象,计算它们的 value
326- * 属性的乘积。这个乘积代表了所有 L0
327- * 变量值的联合效应,或者说数据处理量的一个度量。 通过不断更新 usage
328- * 变量,乘法操作确保了所有 L0 变量的贡献都被包含在内。最终 usage
329- * 变量中的值就是所有 L0 变量值的乘积。
330- * 返回值可以帮助评估算法在处理给定输入数据时的效率,以及比较不同算法或优化策略的数据处理量。
331- */
332- uint32_t L0TileSolver::GetMacUse() const {
333- uint32_t usage = 1u;
334- for (uint32_t j = 0; j < input_.size; j++) {
335- usage *= input_.l0_vars[j].value;
336- }
337- return usage;
338- }
339- 
340- /**
341- * 为 L0 变量找到最优值进行迭代运行
342- *
343- * @param loop_id 当前循环索引,表示正在处理的 L0 变量的位置
344- * @param best_var_value 一个指针,指向用于存储每个 L0 变量迄今为止找到的最佳值的数组
345- *
346- * 这个函数使用递归方法来遍历 L0 变量的所有可能值。对于每个值,它检查是否满足约束条件,如小于上界且满足对齐要求。
347- * 如果满足这些条件,它将继续下一个循环或者递归调用自身来处理下一个L0 变量。
348- * 如果是最后一个 L0
349- * 变量,它将检查当前组合是否满足优化条件,如核心数量和数据处理量。如果满足条件,它将当前组合存储为最优解。
350- *
351- * 请注意,这个函数没有返回值,而是将最优解存储在传入的 best_var_value 数组中。
352- */
353- void L0TileSolver::IterativeRun(uint32_t loop_id, uint32_t * best_var_value) {
354- for (uint32_t i = 0u; i < candidate_size; i++) {
355- uint32_t candi_value = candidate_value[i];
356- const auto &l0_tile = sortedvars_[loop_id];
357- // L0Var的上限
358- uint32_t upper_bound = l0_tile.max_value * UPPER_BOUND_RATIO;
359- if (candi_value >= upper_bound) {
360- continue;
361- }
362- // 必须满足prompt_align对齐或者是prompt_align的因子
363- if ((candi_value % l0_tile.prompt_align != 0) && (l0_tile.prompt_align % candi_value != 0)) {
364- continue;
365- }
366- auto idx = l0_tile.idx;
367- input_.l0_vars[idx].value = candi_value;
368- // 终止条件为遍历到最后一个变量
369- if (loop_id == input_.size - 1) {
370- if (!CheckBufferUseValid()) {
371- break;
372- }
373- int32_t usage = GetMacUse();
374- int32_t core_num = MaxCoreNum(input_.l0_vars, input_.core_num);
375- // 最大核数如果满足核数*系数(默认0.6),则比较mac利用率即可,否则需要比较核数的使用和mac利用率
376- if (((core_num >= max_corenum_) || (core_num >= static_cast<int32_t>(input_.core_num * CORE_NUM_RATIO))) &&
377- (usage >= max_macuse_)) {
378- max_corenum_ = core_num;
379- max_macuse_ = usage;
380- for (uint32_t k = 0u; k < input_.size; k++) {
381- best_var_value[k] = input_.l0_vars[k].value;
382- }
383- }
384- } else {
385- IterativeRun(loop_id + 1, best_var_value);
386- }
387- }
388- }
389- 
390- /**
391- * 更新 L0Var 对象的对齐值
392- *
393- * 这个函数用于更新 L0Var 对象的 prompt_align
394- * 值,以确保它们在内存中按照最优方式对齐。它遍历 input_ 对象中的 l0_vars
395- * 数组,为每个 L0Var 对象计算并设置最佳的对齐值。
396- *
397- * @param 无
398- * @return
399- */
400- void L0TileSolver::UpdateAlign() {
401- for (uint32_t i = 0u; i < input_.size; i++) {
402- uint32_t best_align = GetBestAlign(i);
403- input_.l0_vars[i].prompt_align = best_align;
404- }
405- }
406- 
407- /**
408- * 检查输入数据的有效性
409- *
410- * 这个函数用来检查 L0TileSolver 类的输入数据是否有效。它验证以下几个方面:
411- * - 基础变量指针(l0_vars)是否为空。
412- * - 输入数据的大小(size)是否为0,表示没有 L0 参数需要求解。
413- * - 输入数据的大小(size)是否超过最大支持的参数数量(MAX_L0_VAR_NUM)。
414- * - 核心数量(core_num)是否为0
415- * - 对于输入数据中的每个 L0Var 对象(通过索引 i 访问),它检查几个属性:
416- * - max_value、align 和 prompt_align 是否都不等于0
417- * - align 是否不大于 prompt_align。
418- *
419- * 如果以上任何一个条件不满足,函数将通过 ATT_LOG 宏记录一条错误消息,并返回
420- * false,表示输入无效。如果所有条件都满足,函数返回 true,表示输入有效。
421- *
422- * @return 如果输入数据有效,则返回 true;否则返回 false
423- */
424- bool L0TileSolver::CheckInput() {
425- if (input_.l0_vars == nullptr) {
426- ATT_LOG("Input basevar is null");
427- return false;
428- }
429- if (input_.size == 0u) {
430- ATT_LOG("Size is 0, no l0 arg to be solved");
431- return false;
432- }
433- if (input_.size > MAX_L0_VAR_NUM) {
434- ATT_LOG("L0 solver does not support more than 3 input args");
435- return false;
436- }
437- if (input_.core_num == 0) {
438- ATT_LOG("Corenum is 0");
439- return false;
440- }
441- for (uint32_t i = 0u; i < input_.size; i++) {
442- auto var = input_.l0_vars[i];
443- if ((var.max_value == 0) || (var.align == 0) || (var.prompt_align == 0)) {
444- ATT_LOG("Input [" + std::to_string(i) + "] exists 0");
445- return false;
446- }
447- if (var.align > var.prompt_align) {
448- ATT_LOG("Input [" + std::to_string(i) + "] align is larger than prompt align");
449- return false;
450- }
451- }
452- return true;
453- }
454- 
455- /**
456- * 初始化 L0Var 对象数组
457- *
458- * 这个函数用于初始化 L0TileSolver 类的 input_ 对象中的 l0_vars 数组。它遍历
459- * l0_vars 数组中的每一个元素,对于每个元素,执行以下操作:
460- * 1. 通过访问索引 i 对应的 L0Var 对象的引用 var,重置其 max_value
461- * 属性。具体重置方式是,先将 max_value 增加 align 属性值减 1,再除以 align
462- * 属性值,最后乘以 align 属性值。这样做的目的可能是为了确保 max_value 是 align
463- * 的整数倍。
464- * 2. 将当前循环的索引值 i 设置为 var 的 idx 属性。这可能是为了标记每个 L0Var
465- * 对象在数组中的位置,以便后续处理。
466- *
467- * @param 无
468- * @return
469- */
470- void L0TileSolver::InitInput() {
471- for (uint32_t i = 0u; i < input_.size; i++) {
472- auto &var = input_.l0_vars[i];
473- var.max_value = (var.max_value + var.align - 1) / var.align * var.align;
474- var.idx = i;
475- }
476- }
477- 
478- /**
479- * 检查输出数据的有效性
480- *
481- * 这个函数用于检查 L0TileSolver 类的输出数据是否有效。它首先检查 output_
482- * 指针是否为空。如果 output_ 指针为空,通过 ATT_LOG 宏记录一条错误消息,并返回
483- * false,表示输出无效。
484- *
485- * 接着,函数遍历 output_ 数组中的每个元素。对于每个元素,它检查其值是否为
486- * 0。如果发现任何一个元素的值为 0,函数会通过 ATT_LOG
487- * 宏记录相应的错误消息,并返回 false,表示输出数据中存在无效的元素。
488- *
489- * 如果输出数据有效,即 output_ 指针不为空且 output_ 数组中没有 0
490- * 值元素,函数返回 true
491- *
492- * @return 如果输出数据有效,则返回 true;否则返回 false
493- */
494- bool L0TileSolver::CheckOutput() {
495- if (output_ == nullptr) {
496- ATT_LOG("Output is null");
497- return false;
498- }
499- for (uint32_t i = 0u; i < input_.size; i++) {
500- if (output_[i] == 0u) {
501- ATT_LOG("Output [" + std::to_string(i) + "] is 0");
502- return false;
503- }
504- }
505- return true;
506- }
507- 
508- /**
509- * 执行 L0TileSolver 类的主要流程
510- * @return 如果所有操作成功并且输出有效,则返回 true;否则返回 false
511- */
512- bool L0TileSolver::Run() {
513- // 检查输入数据的有效性
514- if (!CheckInput()) {
515- // 如果输入检查失败,则记录一条错误日志,并返回 false
516- ATT_LOG("Check input failed");
517- return false;
518- }
519- 
520- // 初始化输入数据
521- InitInput();
522- 
523- // 更新 L0Var 对象的对齐值
524- UpdateAlign();
525- 
526- // 为排序后的变量申请内存,并初始化为 0
527- sortedvars_ = new (std::nothrow) L0Var[input_.size];
528- output_ = new (std::nothrow) uint32_t[input_.size]();
529- 
530- // 将输入数据复制到新的内存中
531- std::copy(input_.l0_vars, input_.l0_vars + input_.size, sortedvars_);
532- 
533- // 根据比较函数对变量进行排序
534- std::sort(sortedvars_, sortedvars_ + input_.size, L0VarCmp);
535- 
536- // 调用 IterativeRun 函数,传递参数 0 和 output_ 数组的指针
537- IterativeRun(0u, output_);
538- 
539- // 检查输出数据的有效性
540- if (!CheckOutput()) {
541- // 如果输出检查失败,则记录一条错误日志,并返回 false
542- ATT_LOG("Check output failed");
543- return false;
544- }
545- 
546- // 如果所有操作都成功,返回 true
547- return true;
548- }
549-} // namespace att
@@ -1,316 +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-#include <iostream>
12-#ifdef DEBUG
13-#define ATT_LOG(log)
14-do {
15- std::cout << "[ERROR]" << log << std::endl;
16-} while (0)
17-#else
18-#define ATT_LOG(log)
19-#endif
20- 
21- namespace att {
22- // L2的占用经验值大小
23- const uint32_t EMPIRIC_L2_SIZE = 128 * 1024 * 1024u;
24- uint32_t CeilDivision(uint32_t a, uint32_t b) {
25- if (b == 0) {
26- return 0;
27- }
28- return uint32_t((a + b - 1) / b);
29- }
30- 
31- // 每个L2变量的信息
32- struct L2Var {
33- // 最大值,初始化为原始输入的大小
34- uint32_t max_value{0};
35- // 对气值
36- uint32_t align{0};
37- // 对应的L0基本块的大小,举例,TileL2M对应的基本块basem
38- uint32_t base_val{0};
39- // 当前变量的值
40- uint32_t value{0};
41- };
42- 
43- // 求解器的输入
44- struct L2TileInput {
45- // L2变量的集合
46- L2Var *l2_vars{nullptr};
47- // L2变量的个数
48- uint32_t size{0};
49- // 核数
50- uint32_t core_num{0};
51- // l2的大小,默认为经验值的大小
52- uint32_t l2_size{0};
53- };
54- 
55- // L2求解器的适用范围如下
56- // 举例如下,一个TileL2M * TileL2N的结果矩阵,每个小方格表示一个basem * basen的基本块
57- // TileL2M 和 TileL2N大小的结果矩阵存储在L2中,以基本块为粒度分多核,如下图所示,假设有4个核,每个核计算两个基本块
58- //
59- // tileL2N basen
60- // .-------.-------.-------.-------.
61- // | core0 | core0 | core1 | core1 | ->basem
62- // tileL2M <- '-------'-------'-------'-------'
63- // | core2 | core2 | core3 | core3 |
64- // '-------'-------'-------'-------'
65- class L2TileSolver {
66- public:
67- // 构造函数,接受 L2TileInput 类型的参数 input
68- explicit L2TileSolver(L2TileInput input) : input_(input) {};
69- // 无参构造函数
70- L2TileSolver() {}
71- // 析构函数,用于清理堆上分配的内存
72- ~L2TileSolver() {
73- // 如果 blocknum_per_tile_ 指针不为空,则释放其指向的内存
74- if (blocknum_per_tile_ != nullptr) {
75- delete[] blocknum_per_tile_;
76- }
77- // 如果 size_per_tile_ 指针不为空,则释放其指向的内存
78- if (size_per_tile_ != nullptr) {
79- delete[] size_per_tile_;
80- }
81- // 如果 tilenum_ 指针不为空,则释放其指向的内存
82- if (tilenum_ != nullptr) {
83- delete[] tilenum_;
84- }
85- // 如果 total_blocknum_ 指针不为空,则释放其指向的内存
86- if (total_blocknum_ != nullptr) {
87- delete[] total_blocknum_;
88- }
89- }
90- // Run() 成员函数,返回布尔值,可能用于指示某个操作的成功或失败
91- bool Run();
92- // GetL2Tile() 成员函数,返回 uint32_t 类型的指针
93- uint32_t *GetL2Tile() {
94- return size_per_tile_;
95- }
96- 
97- protected:
98- // 纯虚函数,需要在子类中实现,用于获取 L2 的使用情况
99- virtual uint64_t GetL2Use() = 0;
100- // 纯虚函数,需要在子类中实现,用于判断索引 idx 处是否存在冲突
101- virtual bool IsClash(uint32_t idx) = 0;
102- // L2TileInput 类型的成员变量,用于存储输入数据
103- L2TileInput input_;
104- // 默认初始化值为 1 的 used_corenum_ 成员变量
105- uint32_t used_corenum_{1};
106- // 指向 uint32_t 类型的指针 blocknum_per_tile_,初始化为空指针,用于表达每个方向的输入包含多少个基本块
107- uint32_t *blocknum_per_tile_{nullptr};
108- // 指向 uint32_t 类型的指针 size_per_tile_,初始化为空指针,用于表达每个方向的输入的大小
109- uint32_t *size_per_tile_{nullptr};
110- // 指向 uint32_t 类型的指针 tilenum_,初始化为空指针,用于表达每个方向的输入的L2块的个数
111- uint32_t *tilenum_{nullptr};
112- // 指向 uint32_t 类型的指针 total_blocknum_,初始化为空指针,用于表达每个方向输入基本块的总数
113- uint32_t *total_blocknum_{nullptr};
114- 
115- private:
116- // 私有的 CheckInput() 成员函数,返回布尔值,用于检查输入数据的有效性
117- bool CheckInput();
118- // 私有的 InitInput() 成员函数,用于初始化输入数据
119- void InitInput();
120- // 私有的 CheckSolvable() 成员函数,返回布尔值,用于检查问题是否可解
121- bool CheckSolvable();
122- void HandleClash(uint32_t loop_id, uint32_t *ori_val, uint32_t *best_val, uint64_t &max_l2_use);
123- };
124- 
125- /**
126- * 检查输入参数是否有效。
127- *
128- * 该函数用于检查 L2TileSolver 类的输入参数是否有效。
129- * 首先检查 l2_vars 指针是否为空。如果为空,将记录一条错误消息并返回 false,表示输入无效。
130- * 然后检查 size、core_num 和 l2_size 参数是否都不为零。如果其中任何一个为零,将记录一条错误消息并返回
131- * false,表示输入无效。 最后遍历 l2_vars 指针数组中的所有 L2Var 结构。对于每个结构,检查 align、base_val 和 max_value
132- * 成员是否都不为零。如果其中任何一个为零,将记录一条错误消息并返回 false,表示输入无效。
133- * 如果所有检查都通过,该函数将返回 true,表示输入有效。
134- *
135- * @Return bool, 表示输入参数是否有效
136- */
137- bool L2TileSolver::CheckInput() {
138- if (input_.l2_vars == nullptr) {
139- ATT_LOG("Input l2var is null");
140- return false;
141- }
142- if (input_.size == 0 || input_.core_num == 0 || input_.l2_size == 0) {
143- ATT_LOG("Exist input 0, please check size, core_num and l2_size");
144- return false;
145- }
146- for (uint32_t i = 0; i < input_.size; i++) {
147- auto var = input_.l2_vars[i];
148- if (var.align == 0 || var.base_val == 0 || var.max_value == 0) {
149- ATT_LOG("Input [" + std::to_string(i) + "] exists 0");
150- return false;
151- }
152- }
153- return true;
154- }
155- 
156- /**
157- * 检查问题是否可以解决。
158- *
159- * 这个函数检查当前的问题是否可以根据输入参数的设置来解决。
160- * 首先,它初始化每个变量的值为 align 参数的值。
161- * 然后,它调用 GetL2Use() 函数来获得所需的 L2 使用量。
162- * 如果所需的 L2 使用量超过了可用的缓存大小(input_.l2_size),它将记录一条警告消息并返回 false,表示没有解决方案。
163- * 如果所需的 L2 使用量小于或等于可用的缓存大小,函数将返回 true,表示问题可以被解决。
164- *
165- * @Return true 如果问题可以解决,false 否则。
166- */
167- bool L2TileSolver::CheckSolvable() {
168- for (uint32_t i = 0; i < input_.size; i++) {
169- auto &var = input_.l2_vars[i];
170- var.value = var.align;
171- }
172- if (GetL2Use() > input_.l2_size) {
173- ATT_LOG("No solution, l2 size is too small");
174- return false;
175- }
176- return true;
177- }
178- 
179- /**
180- * 初始化输入数据。
181- *
182- * 这个函数负责初始化 L2TileSolver 对象的输入数据。
183- * 它首先为输入数据结构中的每个变量计算最大值,将其向上取整为对齐值的最近倍数。这确保了每个变量的最大值是其对齐值的倍数。
184- * 然后,它找出所有变量中的最大最大值,并将每个变量的值初始化为这个最大最大值。
185- *
186- * 初始化过程对于准备输入数据以进行后续处理步骤(如优化或分析)至关重要。
187- * 确保最大值是对齐值的倍数,可以简化数据的处理,并可能是依赖于这种属性的算法或操作所必需的。
188- */
189- void L2TileSolver::InitInput() {
190- uint32_t init_value = 0;
191- for (uint32_t i = 0; i < input_.size; i++) {
192- auto &var = input_.l2_vars[i];
193- var.max_value = CeilDivision(var.max_value, var.align) * var.align;
194- init_value = var.max_value > init_value ? var.max_value : init_value;
195- }
196- for (uint32_t i = 0; i < input_.size; i++) {
197- auto &var = input_.l2_vars[i];
198- var.value = init_value;
199- }
200- }
201- 
202- void L2TileSolver::HandleClash(uint32_t loop_id, uint32_t * ori_val, uint32_t * best_val, uint64_t & max_l2_use) {
203- auto max_blocknum = ori_val[loop_id];
204- auto &var = input_.l2_vars[loop_id];
205- for (uint32_t i = max_blocknum; i >= 1u; i--) {
206- blocknum_per_tile_[loop_id] = i;
207- size_per_tile_[loop_id] = blocknum_per_tile_[loop_id] * var.base_val;
208- tilenum_[loop_id] = CeilDivision(var.max_value, size_per_tile_[loop_id]);
209- var.value = size_per_tile_[loop_id];
210- if (loop_id == input_.size - 1) {
211- uint32_t tmp_corenum = 1;
212- for (uint32_t j = 0; j < input_.size; j++) {
213- tmp_corenum *= blocknum_per_tile_[j];
214- }
215- used_corenum_ = std::min(input_.core_num, tmp_corenum);
216- bool solved = true;
217- for (uint32_t k = 0; k < input_.size; k++) {
218- if (IsClash(k)) {
219- solved = false;
220- }
221- }
222- if (solved) {
223- uint64_t l2_use = GetL2Use();
224- if (l2_use > max_l2_use) {
225- for (uint32_t l = 0; l < input_.size; l++) {
226- best_val[l] = blocknum_per_tile_[l];
227- }
228- max_l2_use = l2_use;
229- }
230- return;
231- }
232- } else {
233- HandleClash(loop_id + 1, ori_val, best_val, max_l2_use);
234- }
235- }
236- }
237- 
238- /**
239- * 运行 L2TileSolver 算法来解决 L2 缓存分块问题。
240- *
241- * 这个函数是 L2TileSolver 算法的核心。它尝试根据输入参数和约束条件找到 L2 缓存的最优分块方案。
242- * 函数首先通过 CheckInput() 函数检查输入参数的有效性。如果输入无效,它将记录一条错误消息并返回 false
243- * 然后通过 CheckSolvable() 函数检查问题是否有解。如果没有解,它将记录一条错误消息并返回 false
244- * 如果输入有效并且问题有解,函数将通过 InitInput() 函数初始化输入数据。
245- *
246- * 算法的核心是一个循环,在这个循环中,它不断地调整每个变量的值,以找到一个合适的分块方案。
247- * 循环结束的条件是总内存使用量小于或等于 L2 缓存的大小。
248- * 循环结束后,它将计算每个变量的每个分块的大小、总块数、分块数和占用的核心数。
249- * 然后,它检查是否存在读冲突。如果存在读冲突,它将调整每个分块的大小,直到不再检测到冲突。
250- * 最后,它检查每个变量的最大值是否可以放在一个分块中。
251- * 如果不能,它将相应地调整分块数和每个分块的大小。
252- *
253- * 如果找到合适的分块方案,函数返回 true,否则返回 false
254- *
255- * @Return 如果成功则返回 true,否则返回 false
256- */
257- bool L2TileSolver::Run() {
258- if (!CheckInput()) {
259- ATT_LOG("Check input failed");
260- return false;
261- }
262- if (!CheckSolvable()) {
263- ATT_LOG("Check Solvable failed");
264- return false;
265- }
266- InitInput();
267- uint32_t core_num = input_.core_num;
268- uint32_t l2_size = input_.l2_size;
269- blocknum_per_tile_ = new (std::nothrow) uint32_t[input_.size];
270- size_per_tile_ = new (std::nothrow) uint32_t[input_.size];
271- tilenum_ = new (std::nothrow) uint32_t[input_.size];
272- total_blocknum_ = new (std::nothrow) uint32_t[input_.size];
273- 
274- // 遍历直到满足L2占用停止
275- while (GetL2Use() > l2_size) {
276- for (uint32_t i = 0; i < input_.size; i++) {
277- auto &var = input_.l2_vars[i];
278- var.value = (var.align < var.value) ? (var.value - var.align) : var.align;
279- }
280- }
281- 
282- uint32_t *best_val = new uint32_t[input_.size];
283- uint32_t *ori_val = new uint32_t[input_.size];
284- for (uint32_t i = 0; i < input_.size; i++) {
285- auto &var = input_.l2_vars[i];
286- blocknum_per_tile_[i] = CeilDivision(var.value, var.base_val);
287- size_per_tile_[i] = blocknum_per_tile_[i] * var.base_val;
288- tilenum_[i] = CeilDivision(var.max_value, size_per_tile_[i]);
289- total_blocknum_[i] = CeilDivision(var.max_value, var.base_val);
290- best_val[i] = blocknum_per_tile_[i];
291- ori_val[i] = blocknum_per_tile_[i];
292- }
293- 
294- uint64_t max_l2_use = 0u;
295- HandleClash(0, ori_val, best_val, max_l2_use);
296- 
297- for (uint32_t i = 0; i < input_.size; i++) {
298- auto &var = input_.l2_vars[i];
299- blocknum_per_tile_[i] = best_val[i];
300- size_per_tile_[i] = blocknum_per_tile_[i] * var.base_val;
301- }
302- 
303- delete[] best_val;
304- delete[] ori_val;
305- 
306- for (uint32_t i = 0; i < input_.size; i++) {
307- auto &var = input_.l2_vars[i];
308- if (var.max_value <= size_per_tile_[i]) {
309- tilenum_[i] = 1;
310- blocknum_per_tile_[i] = CeilDivision(var.max_value, var.base_val);
311- size_per_tile_[i] = blocknum_per_tile_[i] * var.base_val;
312- }
313- }
314- return true;
315- }
316-} // namespace att
@@ -1,279 +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-#include "gtest/gtest.h"
12-#include "base/base_types.h"
13-#include "base/model_info.h"
14-#define private public
15-#include "generator/solver_pass_gen/l0_solver/l0_solver_gen.h"
16-#include <symengine/functions.h>
17-#include <symengine/simplify.h>
18-#include <symengine/integer.h>
19-#include <symengine/real_double.h>
20-using namespace att;
21- 
22-class TestL0SolverGen : public ::testing::Test {
23- public:
24- static void TearDownTestCase() {
25- std::cout << "Test end." << std::endl;
26- }
27- static void SetUpTestCase() {
28- std::cout << "Test begin." << std::endl;
29- }
30- void SetUp() override {
31- // Code here will be called immediately after the constructor (right
32- // before each test).
33- }
34- 
35- void TearDown() override {
36- // Code here will be called immediately after each test (right
37- // before the destructor).
38- }
39-};
40- 
41-TEST_F(TestL0SolverGen, TEST_IS_MULTICORE_ARG) {
42- L0TileSolverGen solver_gen("case0", "TilingData");
43- std::vector<Expr> mc_args;
44- Expr null_expr;
45- Expr tilem = CreateExpr("tilem");
46- Expr tilen = CreateExpr("tilen");
47- Expr basem = CreateExpr("basem");
48- solver_gen.mc_args_.emplace_back(tilem);
49- solver_gen.mc_args_.emplace_back(tilen);
50- bool case0_is_mc = solver_gen.IsMulticoreArg(tilem);
51- bool case1_is_mc = solver_gen.IsMulticoreArg(tilen);
52- bool case2_is_mc = solver_gen.IsMulticoreArg(basem);
53- bool case3_is_mc = solver_gen.IsMulticoreArg(null_expr);
54- 
55- EXPECT_EQ(case0_is_mc, true);
56- EXPECT_EQ(case1_is_mc, true);
57- EXPECT_EQ(case2_is_mc, false);
58- EXPECT_EQ(case3_is_mc, false);
59-}
60- 
61-TEST_F(TestL0SolverGen, TEST_IS_BIND_MULTICORE) {
62- L0TileSolverGen solver_gen("case0", "TilingData");
63- std::vector<Expr> mc_args;
64- Expr tilem = CreateExpr("tilem");
65- Expr tilen = CreateExpr("tilen");
66- Expr basem = CreateExpr("basem");
67- std::map<Expr, Expr, ExprCmp> father_args_map;
68- father_args_map[basem] = tilem;
69- solver_gen.mc_args_.emplace_back(tilem);
70- solver_gen.mc_args_.emplace_back(tilen);
71- EXPECT_EQ(solver_gen.IsBindMulticore(basem), false);
72- solver_gen.SetFatherArgsMap(father_args_map);
73- EXPECT_EQ(solver_gen.IsBindMulticore(basem), true);
74-}
75- 
76-TEST_F(TestL0SolverGen, TEST_GET_LARGEST_ALIGN) {
77- L0TileSolverGen solver_gen("case0", "TilingData");
78- std::vector<Expr> mc_args;
79- Expr tilem = CreateExpr("tilem");
80- Expr tilen = CreateExpr("tilen");
81- Expr basem = CreateExpr("basem");
82- std::map<Expr, Expr, ExprCmp> father_args_map;
83- std::map<Expr, Expr, ExprCmp> arg_align_map;
84- father_args_map[basem] = tilem;
85- arg_align_map[basem] = ge::Symbol(16);
86- arg_align_map[tilem] = ge::Symbol(256);
87- arg_align_map[tilen] = ge::Symbol(16);
88- solver_gen.mc_args_.emplace_back(tilem);
89- solver_gen.mc_args_.emplace_back(tilen);
90- solver_gen.SetFatherArgsMap(father_args_map);
91- solver_gen.SetArgAlignMap(arg_align_map);
92- Expr max_align = ge::Symbol(16);
93- solver_gen.GetLargestAlign(basem, max_align);
94- EXPECT_EQ(max_align, 256);
95-}
96- 
97-TEST_F(TestL0SolverGen, TEST_GEN_SOLVER) {
98- L0TileSolverGen solver_gen("case0", "TilingData");
99- std::vector<Expr> ori_args;
100- std::vector<Expr> l0_args;
101- std::vector<Expr> mc_args;
102- ExprUintMap const_args;
103- ExprExprMap father_args_map;
104- ExprExprMap arg_align_map;
105- ExprExprMap ori_arg_map;
106- 
107- Expr m = CreateExpr("m");
108- Expr n = CreateExpr("n");
109- Expr k = CreateExpr("k");
110- Expr tilem = CreateExpr("tilem");
111- Expr tilen = CreateExpr("tilen");
112- Expr basem = CreateExpr("basem");
113- Expr basen = CreateExpr("basen");
114- Expr basek = CreateExpr("basek");
115- Expr constbl = CreateExpr("bl");
116- Expr buffer_use = ((basem * basek) + (basen * basek));
117- 
118- std::map<HardwareDef, Expr> buffer_use_map;
119- buffer_use_map[HardwareDef::L0A] = buffer_use;
120- 
121- ori_args.emplace_back(m);
122- ori_args.emplace_back(n);
123- ori_args.emplace_back(k);
124- 
125- l0_args.emplace_back(basem);
126- l0_args.emplace_back(basen);
127- l0_args.emplace_back(basek);
128- 
129- mc_args.emplace_back(tilem);
130- mc_args.emplace_back(tilen);
131- 
132- ori_arg_map[basem] = m;
133- ori_arg_map[basen] = n;
134- ori_arg_map[basek] = k;
135- 
136- father_args_map[basem] = tilem;
137- father_args_map[basen] = tilen;
138- 
139- arg_align_map[basem] = ge::Symbol(16);
140- arg_align_map[basen] = ge::Symbol(16);
141- arg_align_map[basek] = ge::Symbol(16);
142- arg_align_map[tilem] = ge::Symbol(256);
143- arg_align_map[tilen] = ge::Symbol(16);
144- 
145- const_args[constbl] = 8;
146- 
147- solver_gen.SetMulticoreArgs(mc_args);
148- solver_gen.SetFatherArgsMap(father_args_map);
149- solver_gen.SetArgAlignMap(arg_align_map);
150- solver_gen.SetArgtMaxValueMap(ori_arg_map);
151- solver_gen.SetL0Args(l0_args);
152- solver_gen.SetBufferUseAlg(buffer_use_map);
153- solver_gen.SetConstVars(const_args);
154- std::string impl_code = solver_gen.GenSolverFuncImpl();
155- std::string invoke_code = solver_gen.GenSolverFuncInvoke();
156- EXPECT_NE(impl_code, "");
157- EXPECT_NE(invoke_code, "");
158-}
159- 
160-TEST_F(TestL0SolverGen, TEST_GEN_SOLVER_ERR) {
161- L0TileSolverGen solver_gen("case0", "TilingData");
162- std::vector<Expr> ori_args;
163- std::vector<Expr> l0_args;
164- std::vector<Expr> mc_args;
165- ExprUintMap const_args;
166- ExprExprMap father_args_map;
167- ExprExprMap arg_align_map;
168- ExprExprMap ori_arg_map;
169- 
170- Expr m = CreateExpr("m");
171- Expr n = CreateExpr("n");
172- Expr k = CreateExpr("k");
173- Expr tilem = CreateExpr("tilem");
174- Expr tilen = CreateExpr("tilen");
175- Expr basem = CreateExpr("basem");
176- Expr basen = CreateExpr("basen");
177- Expr basek = CreateExpr("basek");
178- Expr constbl = CreateExpr("bl");
179- Expr buffer_use = ((basem * basek) + (basen * basek));
180- 
181- std::map<HardwareDef, Expr> buffer_use_map;
182- buffer_use_map[HardwareDef::L0A] = buffer_use;
183- 
184- ori_args.emplace_back(m);
185- ori_args.emplace_back(n);
186- ori_args.emplace_back(k);
187- 
188- l0_args.emplace_back(basem);
189- l0_args.emplace_back(basen);
190- l0_args.emplace_back(basek);
191- 
192- mc_args.emplace_back(tilem);
193- mc_args.emplace_back(tilen);
194- 
195- ori_arg_map[basem] = m;
196- ori_arg_map[basen] = n;
197- 
198- father_args_map[basem] = tilem;
199- father_args_map[basen] = tilen;
200- 
201- arg_align_map[basem] = ge::Symbol(16);
202- arg_align_map[basen] = ge::Symbol(16);
203- 
204- arg_align_map[tilem] = ge::Symbol(256);
205- arg_align_map[tilen] = ge::Symbol(16);
206- 
207- const_args[constbl] = 8;
208- 
209- solver_gen.SetMulticoreArgs(mc_args);
210- solver_gen.SetFatherArgsMap(father_args_map);
211- solver_gen.SetArgAlignMap(arg_align_map);
212- solver_gen.SetArgtMaxValueMap(ori_arg_map);
213- solver_gen.SetL0Args(l0_args);
214- solver_gen.SetBufferUseAlg(buffer_use_map);
215- solver_gen.SetConstVars(const_args);
216- std::string impl_code = solver_gen.GenInitTilingData();
217- EXPECT_EQ(impl_code, "Solver Gen Error");
218-}
219- 
220-TEST_F(TestL0SolverGen, TEST_MAX_ALIGN_ZERO) {
221- L0TileSolverGen solver_gen("case0", "TilingData");
222- std::vector<Expr> ori_args;
223- std::vector<Expr> l0_args;
224- std::vector<Expr> mc_args;
225- ExprUintMap const_args;
226- ExprExprMap father_args_map;
227- ExprExprMap arg_align_map;
228- ExprExprMap ori_arg_map;
229- 
230- Expr m = CreateExpr("m");
231- Expr n = CreateExpr("n");
232- Expr k = CreateExpr("k");
233- Expr tilem = CreateExpr("tilem");
234- Expr tilen = CreateExpr("tilen");
235- Expr basem = CreateExpr("basem");
236- Expr basen = CreateExpr("basen");
237- Expr basek = CreateExpr("basek");
238- Expr constbl = CreateExpr("bl");
239- Expr buffer_use = ((basem * basek) + (basen * basek));
240- 
241- std::map<HardwareDef, Expr> buffer_use_map;
242- buffer_use_map[HardwareDef::L0A] = buffer_use;
243- 
244- ori_args.emplace_back(m);
245- ori_args.emplace_back(n);
246- ori_args.emplace_back(k);
247- 
248- l0_args.emplace_back(basem);
249- l0_args.emplace_back(basen);
250- l0_args.emplace_back(basek);
251- 
252- mc_args.emplace_back(tilem);
253- mc_args.emplace_back(tilen);
254- 
255- ori_arg_map[basem] = m;
256- ori_arg_map[basen] = n;
257- ori_arg_map[basek] = k;
258- 
259- father_args_map[basem] = tilem;
260- father_args_map[basen] = tilen;
261- 
262- arg_align_map[basem] = ge::Symbol(0);
263- arg_align_map[basen] = ge::Symbol(16);
264- arg_align_map[basek] = ge::Symbol(16);
265- arg_align_map[tilem] = ge::Symbol(256);
266- arg_align_map[tilen] = ge::Symbol(16);
267- 
268- const_args[constbl] = 8;
269- 
270- solver_gen.SetMulticoreArgs(mc_args);
271- solver_gen.SetFatherArgsMap(father_args_map);
272- solver_gen.SetArgAlignMap(arg_align_map);
273- solver_gen.SetArgtMaxValueMap(ori_arg_map);
274- solver_gen.SetL0Args(l0_args);
275- solver_gen.SetBufferUseAlg(buffer_use_map);
276- solver_gen.SetConstVars(const_args);
277- std::string impl_code = solver_gen.GenInitTilingData();
278- EXPECT_EQ(impl_code, "Solver Gen Error");
279-}
@@ -1,231 +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-#include "gtest/gtest.h"
12-#include "base/base_types.h"
13-#include "base/model_info.h"
14-#define private public
15-#include "generator/solver_pass_gen/l2_solver/l2_solver_gen.h"
16-#include <symengine/functions.h>
17-#include <symengine/simplify.h>
18-#include <symengine/integer.h>
19-#include <symengine/real_double.h>
20-using namespace att;
21- 
22-class TestL2SolverGen : public ::testing::Test {
23- public:
24- static void TearDownTestCase() {
25- std::cout << "Test end." << std::endl;
26- }
27- static void SetUpTestCase() {
28- std::cout << "Test begin." << std::endl;
29- }
30- void SetUp() override {
31- // Code here will be called immediately after the constructor (right
32- // before each test).
33- }
34- 
35- void TearDown() override {
36- // Code here will be called immediately after each test (right
37- // before the destructor).
38- }
39-};
40- 
41-TEST_F(TestL2SolverGen, TEST_GEN_SOLVER) {
42- Expr m = CreateExpr("m");
43- Expr n = CreateExpr("n");
44- Expr k = CreateExpr("k");
45- 
46- Expr basem = CreateExpr("basem");
47- Expr basen = CreateExpr("basen");
48- 
49- Expr tilem = CreateExpr("tilem");
50- Expr tilen = CreateExpr("tilen");
51- 
52- std::vector<Expr> input_args;
53- std::vector<Expr> l0_args;
54- std::vector<Expr> l2_args;
55- std::map<Expr, Expr, ExprCmp> ori_arg_map;
56- std::map<Expr, Expr, ExprCmp> arg_align_map;
57- 
58- input_args.emplace_back(m);
59- input_args.emplace_back(n);
60- input_args.emplace_back(k);
61- ori_arg_map[basem] = m;
62- ori_arg_map[basen] = n;
63- ori_arg_map[tilem] = m;
64- ori_arg_map[tilen] = n;
65- arg_align_map[tilem] = ge::Symbol(16);
66- arg_align_map[tilen] = ge::Symbol(16);
67- 
68- l0_args.emplace_back(basem);
69- l0_args.emplace_back(basen);
70- 
71- l2_args.emplace_back(tilem);
72- l2_args.emplace_back(tilen);
73- 
74- Expr l2_use_expr = ((((tilem * k) + (tilen * k)) + (tilem * tilen)) * CreateExpr(2));
75- 
76- att::L2TileSolverGen *solver_gen = new att::L2TileSolverGen("Case0", "TilingData");
77- solver_gen->SetArgAlignMap(arg_align_map);
78- solver_gen->SetL0Args(l0_args);
79- solver_gen->SetL2Args(l2_args);
80- solver_gen->SetL2Use(l2_use_expr);
81- solver_gen->SetArgtMaxValueMap(ori_arg_map);
82- solver_gen->SetInputArgs(input_args);
83- std::string impl_code = solver_gen->GenSolverFuncImpl();
84- std::string invoke_code = solver_gen->GenSolverFuncInvoke();
85- EXPECT_NE(impl_code, "");
86- EXPECT_NE(invoke_code, "");
87-}
88- 
89-TEST_F(TestL2SolverGen, TEST_KSOLVER_GEN_ERR) {
90- Expr m = CreateExpr("m");
91- Expr n = CreateExpr("n");
92- Expr k = CreateExpr("k");
93- 
94- Expr basem = CreateExpr("basem");
95- Expr basen = CreateExpr("basen");
96- 
97- Expr tilem = CreateExpr("tilem");
98- Expr tilen = CreateExpr("tilen");
99- Expr tilek = CreateExpr("tilek");
100- 
101- std::vector<Expr> input_args;
102- std::vector<Expr> l0_args;
103- std::vector<Expr> l2_args;
104- std::map<Expr, Expr, ExprCmp> ori_arg_map;
105- std::map<Expr, Expr, ExprCmp> arg_align_map;
106- 
107- input_args.emplace_back(m);
108- input_args.emplace_back(n);
109- input_args.emplace_back(k);
110- ori_arg_map[basem] = m;
111- ori_arg_map[basen] = n;
112- ori_arg_map[tilem] = m;
113- ori_arg_map[tilen] = n;
114- arg_align_map[tilem] = ge::Symbol(16);
115- arg_align_map[tilen] = ge::Symbol(16);
116- 
117- l0_args.emplace_back(basem);
118- l0_args.emplace_back(basen);
119- 
120- l2_args.emplace_back(tilem);
121- l2_args.emplace_back(tilen);
122- l2_args.emplace_back(tilek);
123- 
124- Expr l2_use_expr = ((((tilem * k) + (tilen * k)) + (tilem * tilen)) * CreateExpr(2));
125- 
126- att::L2TileSolverGen *solver_gen = new att::L2TileSolverGen("Case0", "TilingData");
127- solver_gen->SetArgAlignMap(arg_align_map);
128- solver_gen->SetL0Args(l0_args);
129- solver_gen->SetL2Args(l2_args);
130- solver_gen->SetL2Use(l2_use_expr);
131- solver_gen->SetArgtMaxValueMap(ori_arg_map);
132- solver_gen->SetInputArgs(input_args);
133- std::string impl_code = solver_gen->GenSolverFuncImpl();
134- EXPECT_EQ(impl_code, "Solver Gen Error");
135-}
136- 
137-TEST_F(TestL2SolverGen, TEST_IS_CLASH_POSSIBLE) {
138- Expr m = CreateExpr("m");
139- Expr n = CreateExpr("n");
140- Expr k = CreateExpr("k");
141- 
142- Expr basem = CreateExpr("basem");
143- Expr basen = CreateExpr("basen");
144- 
145- Expr tilem = CreateExpr("tilem");
146- Expr tilen = CreateExpr("tilen");
147- 
148- std::vector<Expr> input_args;
149- std::vector<Expr> l0_args;
150- std::vector<Expr> l2_args;
151- std::map<Expr, Expr, ExprCmp> ori_arg_map;
152- std::map<Expr, Expr, ExprCmp> arg_align_map;
153- 
154- input_args.emplace_back(m);
155- input_args.emplace_back(n);
156- input_args.emplace_back(k);
157- ori_arg_map[basem] = m;
158- ori_arg_map[basen] = n;
159- ori_arg_map[tilem] = m;
160- ori_arg_map[tilen] = n;
161- arg_align_map[tilem] = ge::Symbol(16);
162- arg_align_map[tilen] = ge::Symbol(16);
163- 
164- l0_args.emplace_back(basem);
165- l0_args.emplace_back(basen);
166- 
167- l2_args.emplace_back(tilem);
168- l2_args.emplace_back(tilen);
169- 
170- Expr l2_use_expr = ((((tilem * k) + (tilen * k)) + (tilem * tilen)) * CreateExpr(2));
171- 
172- att::L2TileSolverGen *solver_gen = new att::L2TileSolverGen("Case0", "TilingData");
173- solver_gen->SetArgAlignMap(arg_align_map);
174- solver_gen->SetL0Args(l0_args);
175- solver_gen->SetL2Args(l2_args);
176- solver_gen->SetL2Use(l2_use_expr);
177- solver_gen->SetArgtMaxValueMap(ori_arg_map);
178- solver_gen->SetInputArgs(input_args);
179- EXPECT_EQ(solver_gen->IsClashPossible(), true);
180- // Expr mock_l2_use = Mul(tilem, tilen);
181- // solver_gen->SetL2Use(mock_l2_use);
182- // EXPECT_EQ(solver_gen->IsClashPossible(), false);
183-}
184- 
185-TEST_F(TestL2SolverGen, TEST_GET_RELATED_L0_ARGS) {
186- Expr m = CreateExpr("m");
187- Expr n = CreateExpr("n");
188- Expr k = CreateExpr("k");
189- 
190- Expr basem = CreateExpr("basem");
191- Expr basen = CreateExpr("basen");
192- 
193- Expr tilem = CreateExpr("tilem");
194- Expr tilen = CreateExpr("tilen");
195- 
196- std::vector<Expr> input_args;
197- std::vector<Expr> l0_args;
198- std::vector<Expr> l2_args;
199- std::map<Expr, Expr, ExprCmp> ori_arg_map;
200- std::map<Expr, Expr, ExprCmp> arg_align_map;
201- 
202- input_args.emplace_back(m);
203- input_args.emplace_back(n);
204- input_args.emplace_back(k);
205- ori_arg_map[basem] = m;
206- ori_arg_map[basen] = n;
207- ori_arg_map[tilem] = m;
208- ori_arg_map[tilen] = n;
209- arg_align_map[tilem] = ge::Symbol(16);
210- arg_align_map[tilen] = ge::Symbol(16);
211- 
212- l0_args.emplace_back(basem);
213- l0_args.emplace_back(basen);
214- 
215- l2_args.emplace_back(tilem);
216- l2_args.emplace_back(tilen);
217- 
218- Expr l2_use_expr = ((((tilem * k) + (tilen * k)) + (tilem * tilen)) * CreateExpr(2));
219- 
220- att::L2TileSolverGen *solver_gen = new att::L2TileSolverGen("Case0", "TilingData");
221- solver_gen->SetArgAlignMap(arg_align_map);
222- solver_gen->SetL0Args(l0_args);
223- solver_gen->SetL2Args(l2_args);
224- solver_gen->SetL2Use(l2_use_expr);
225- solver_gen->SetArgtMaxValueMap(ori_arg_map);
226- solver_gen->SetInputArgs(input_args);
227- EXPECT_EQ(solver_gen->GetRelateL0Arg(tilem), basem);
228- ori_arg_map[basem] = CreateExpr("mock");
229- solver_gen->SetArgtMaxValueMap(ori_arg_map);
230- EXPECT_FALSE(IsValid(solver_gen->GetRelateL0Arg(tilem)));
231-}
@@ -40,31 +40,6 @@ class TestSolverPassManager : public ::testing::Test {
40 }40 }
41};41};
42 42 
43-TEST_F(TestSolverPassManager, TEST_CHECK_ARG_EXIST) {
44- ModelInfo modelInfo = CreateModelInfo();
45- ArgsManager args_manager(modelInfo);
46- att::SolverPassManager manager(args_manager, {0}, "TilingData");
47- std::vector<Expr> args;
48- Expr arg0 = CreateExpr("arg0");
49- Expr arg1 = CreateExpr("arg1");
50- Expr arg2 = CreateExpr("arg2");
51- args.emplace_back(arg0);
52- args.emplace_back(arg1);
53- bool case0 = manager.CheckArgExist(arg2, args);
54- EXPECT_EQ(case0, false);
55- bool case1 = manager.CheckArgExist(arg0, args);
56- EXPECT_EQ(case1, true);
57-}
58- 
59-TEST_F(TestSolverPassManager, TEST_GET_L0_ARGS) {
60- ModelInfo modelInfo = CreateModelInfo();
61- ArgsManager args_manager(modelInfo);
62- att::SolverPassManager manager(args_manager, {0}, "TilingData");
63- args_manager.Process(false);
64- auto l0_args = manager.GetL0Args(args_manager, false);
65- EXPECT_EQ(l0_args.size(), 2);
66-}
67- 
68TEST_F(TestSolverPassManager, TEST_IS_NEED_SOLVER) {43TEST_F(TestSolverPassManager, TEST_IS_NEED_SOLVER) {
69 ModelInfo modelInfo = CreateModelInfo();44 ModelInfo modelInfo = CreateModelInfo();
70 ArgsManager args_manager(modelInfo);45 ArgsManager args_manager(modelInfo);
@@ -72,7 +47,7 @@ TEST_F(TestSolverPassManager, TEST_IS_NEED_SOLVER) {
72 std::vector<ArgsManager> args_managers;47 std::vector<ArgsManager> args_managers;
73 args_managers.emplace_back(args_manager);48 args_managers.emplace_back(args_manager);
74 att::SolverPassManager manager(args_manager, {0}, "TilingData");49 att::SolverPassManager manager(args_manager, {0}, "TilingData");
75- bool is_need = manager.IsNeedSolver(args_managers, SolverType::L0_TILE);50+ bool is_need = manager.IsNeedSolver(args_managers, SolverType::SEARCH_TILE);
76 EXPECT_EQ(is_need, true);51 EXPECT_EQ(is_need, true);
77}52}
78 53 
@@ -87,7 +62,7 @@ TEST_F(TestSolverPassManager, case0) {
87 std::string base_class_head = manager.GenCommonBaseClassesHead(args_managers);62 std::string base_class_head = manager.GenCommonBaseClassesHead(args_managers);
88 std::string base_class_func = manager.GenCommonBaseClassesFunc(args_managers);63 std::string base_class_func = manager.GenCommonBaseClassesFunc(args_managers);
89 EXPECT_NE(base_class_head, "");64 EXPECT_NE(base_class_head, "");
90- EXPECT_NE(base_class_func, "");65+ EXPECT_EQ(base_class_func, "");
91 std::string impl_code = res.first;66 std::string impl_code = res.first;
92 std::string invoke_code = res.second;67 std::string invoke_code = res.second;
93 EXPECT_NE(impl_code, "");68 EXPECT_NE(impl_code, "");
@@ -1,463 +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-#include <iostream>
12-#include "stub/stub_model_info.h"
13-namespace att {
14-ModelInfo CreateModelInfo(const ge::ExprType expr_type) {
15- ModelInfo model_info;
16- Expr default_expr;
17- bool is_const = true;
18- if (expr_type == ge::ExprType::kExprConstantRation) {
19- default_expr = ge::Symbol(8, "tmp") / ge::Symbol(3, "tmp");
20- } else if (expr_type == ge::ExprType::kExprConstantInteger) {
21- default_expr = ge::Symbol(2, "tmp");
22- } else if (expr_type == ge::ExprType::kExprVariable) {
23- is_const = false;
24- }
25- // set m
26- Expr expr_m = is_const ? default_expr : CreateExpr("m_size");
27- Expr expr_tilem = is_const ? default_expr : CreateExpr("tilem_size");
28- Expr expr_stepm = is_const ? default_expr : CreateExpr("stepm_size");
29- Expr expr_basem = CreateExpr("basem_size");
30- SymVarInfoPtr sym_m = std::make_shared<SymVarInfo>(expr_m);
31- SymVarInfoPtr sym_tilem = std::make_shared<SymVarInfo>(expr_tilem);
32- sym_tilem->align = ge::Symbol(16);
33- sym_tilem->related_scope = {HardwareDef::L2};
34- SymVarInfoPtr sym_stepm = std::make_shared<SymVarInfo>(expr_stepm);
35- sym_stepm->align = ge::Symbol(16);
36- sym_stepm->related_scope = {HardwareDef::L1, HardwareDef::CORENUM};
37- SymVarInfoPtr sym_basem = std::make_shared<SymVarInfo>(expr_basem);
38- sym_basem->align = ge::Symbol(16);
39- sym_basem->related_scope = {HardwareDef::L0A, HardwareDef::L0C};
40- AttAxisPtr m = std::make_shared<AttAxis>();
41- AttAxisPtr tilem = std::make_shared<AttAxis>();
42- AttAxisPtr stepm = std::make_shared<AttAxis>();
43- AttAxisPtr basem = std::make_shared<AttAxis>();
44- m->name = "m";
45- m->axis_pos = AxisPosition::ORIGIN;
46- m->bind_multicore = false;
47- m->is_last = false;
48- m->is_node_innerest_dim = false;
49- sym_m->value_range.first = 1;
50- sym_m->value_range.second = 10000;
51- m->size = sym_m;
52- 
53- tilem->name = "tilem";
54- tilem->axis_pos = AxisPosition::INNER;
55- tilem->bind_multicore = false;
56- tilem->is_last = false;
57- tilem->is_node_innerest_dim = true;
58- tilem->size = sym_tilem;
59- tilem->orig_axis.push_back(m.get());
60- tilem->from_axis = {m.get()};
61- 
62- stepm->name = "stepm";
63- stepm->axis_pos = AxisPosition::INNER;
64- stepm->bind_multicore = true;
65- stepm->is_last = false;
66- stepm->is_node_innerest_dim = true;
67- stepm->size = sym_stepm;
68- stepm->orig_axis.push_back(m.get());
69- stepm->from_axis = {tilem.get()};
70- 
71- basem->name = "basem";
72- basem->axis_pos = AxisPosition::INNER;
73- basem->bind_multicore = false;
74- basem->is_last = true;
75- basem->is_node_innerest_dim = true;
76- basem->size = sym_basem;
77- basem->orig_axis.push_back(m.get());
78- basem->from_axis = {stepm.get()};
79- model_info.arg_list.emplace_back(m);
80- model_info.arg_list.emplace_back(tilem);
81- model_info.arg_list.emplace_back(stepm);
82- model_info.arg_list.emplace_back(basem);
83- Optional test_optional;
84- test_optional.optional_name = "test";
85- test_optional.data_type = "int32_t";
86- test_optional.min_value = "1";
87- test_optional.max_value = "100";
88- model_info.graph_input_infos.optional_atts[1U] = test_optional;
89- 
90- // set n
91- Expr expr_n = CreateExpr("n_size");
92- Expr expr_tilen = CreateExpr("tilen_size");
93- Expr expr_stepn = CreateExpr("stepn_size");
94- Expr expr_basen = CreateExpr("basen_size");
95- SymVarInfoPtr sym_n = std::make_shared<SymVarInfo>(expr_n);
96- SymVarInfoPtr sym_tilen = std::make_shared<SymVarInfo>(expr_tilen);
97- sym_tilen->align = ge::Symbol(16);
98- sym_tilen->related_scope = {HardwareDef::L2};
99- SymVarInfoPtr sym_stepn = std::make_shared<SymVarInfo>(expr_stepn);
100- sym_stepn->align = ge::Symbol(128);
101- sym_stepn->related_scope = {HardwareDef::L1, HardwareDef::CORENUM};
102- SymVarInfoPtr sym_basen = std::make_shared<SymVarInfo>(expr_basen);
103- sym_basen->align = ge::Symbol(16);
104- sym_basen->related_scope = {HardwareDef::L0B, HardwareDef::L0C};
105- AttAxisPtr n = std::make_shared<AttAxis>();
106- AttAxisPtr tilen = std::make_shared<AttAxis>();
107- AttAxisPtr stepn = std::make_shared<AttAxis>();
108- AttAxisPtr basen = std::make_shared<AttAxis>();
109- n->name = "n";
110- n->axis_pos = AxisPosition::ORIGIN;
111- n->bind_multicore = false;
112- n->is_last = false;
113- n->is_node_innerest_dim = false;
114- n->size = sym_n;
115- 
116- tilen->name = "tilen";
117- tilen->axis_pos = AxisPosition::INNER;
118- tilen->bind_multicore = false;
119- tilen->is_last = false;
120- tilen->is_node_innerest_dim = true;
121- tilen->size = sym_tilen;
122- tilen->orig_axis.push_back(n.get());
123- tilen->from_axis = {n.get()};
124- 
125- stepn->name = "stepn";
126- stepn->axis_pos = AxisPosition::INNER;
127- stepn->bind_multicore = true;
128- stepn->is_last = false;
129- stepn->is_node_innerest_dim = true;
130- stepn->size = sym_stepn;
131- stepn->orig_axis.push_back(n.get());
132- stepn->from_axis = {tilen.get()};
133- 
134- basen->name = "basen";
135- basen->axis_pos = AxisPosition::INNER;
136- basen->bind_multicore = false;
137- basen->is_last = true;
138- basen->is_node_innerest_dim = true;
139- basen->size = sym_basen;
140- basen->orig_axis.push_back(n.get());
141- basen->from_axis = {stepn.get()};
142- 
143- model_info.arg_list.emplace_back(n);
144- model_info.arg_list.emplace_back(tilen);
145- model_info.arg_list.emplace_back(stepn);
146- model_info.arg_list.emplace_back(basen);
147- 
148- // setk
149- Expr expr_k = CreateExpr("k_size");
150- SymConstInfoPtr sym_k = std::make_shared<SymConstInfo>(expr_k);
151- sym_k->const_value = 128u;
152- AttAxisPtr k = std::make_shared<AttAxis>();
153- k->name = "k";
154- k->axis_pos = AxisPosition::ORIGIN;
155- k->bind_multicore = false;
156- k->is_last = false;
157- k->is_node_innerest_dim = false;
158- k->size = sym_k;
159- model_info.arg_list.emplace_back(k);
160- 
161- Expr l0a_occupy = expr_basem * expr_k * CreateExpr(4);
162- Expr l0b_occupy = expr_k * expr_basen * CreateExpr(4);
163- Expr l0c_occupy = expr_basem * expr_basen * CreateExpr(4);
164- Expr l1_occupy = (expr_k * expr_stepm * CreateExpr(4)) + (expr_k * expr_stepn * CreateExpr(4));
165- Expr l2_occupy = (expr_tilen * expr_tilem * CreateExpr(2)) + ((expr_tilen + expr_tilem) * expr_k * CreateExpr(2));
166- Expr core_num = (expr_tilem / expr_stepm) * (expr_tilen / expr_stepn);
167- std::map<HardwareDef, Expr> hardware_cons;
168- model_info.hardware_cons[HardwareDef::L0A] = l0a_occupy;
169- model_info.hardware_cons[HardwareDef::L0B] = l0b_occupy;
170- model_info.hardware_cons[HardwareDef::L0C] = l0c_occupy;
171- model_info.hardware_cons[HardwareDef::L1] = l1_occupy;
172- model_info.hardware_cons[HardwareDef::L2] = l2_occupy;
173- model_info.hardware_cons[HardwareDef::CORENUM] = core_num;
174- model_info.hardware_cons[HardwareDef::UB] = CreateExpr(0L);
175- 
176- Expr mac = (expr_basem * expr_basen * expr_k) / (CreateExpr(16) * CreateExpr(256));
177- Expr mte = (((expr_stepm * expr_k) / CreateExpr(32)) + ((expr_stepn * expr_k) / CreateExpr(32)));
178- model_info.objects[PipeType::AIC_MAC] = mac;
179- model_info.objects[PipeType::AIC_MTE2] = mte;
180- model_info.tiling_case_id = 0;
181- model_info.eq_exprs[kFatherToChildNoTail].push_back(std::pair(expr_stepm, expr_basem));
182- model_info.eq_exprs[kFatherToChildNoTail].push_back(std::pair(expr_stepn, expr_basen));
183- model_info.leq_exprs[kFatherToChildLarger].push_back((expr_tilem - expr_m));
184- model_info.leq_exprs[kFatherToChildLarger].push_back((expr_stepm - expr_tilem));
185- model_info.leq_exprs[kFatherToChildLarger].push_back((expr_tilen - expr_n));
186- model_info.leq_exprs[kFatherToChildLarger].push_back((expr_stepn - expr_tilen));
187- model_info.output_size = 1;
188- 
189- return model_info;
190-}
191- 
192-ModelInfo GetMatmulL2TileInfo() {
193- ModelInfo model_info;
194- Expr expr_corenum = CreateExpr("block_dim");
195- SymVarInfoPtr sym_corenum = std::make_shared<SymVarInfo>(expr_corenum);
196- AttAxisPtr core = std::make_shared<AttAxis>();
197- core->name = "corenum";
198- core->axis_pos = AxisPosition::ORIGIN;
199- core->bind_multicore = false;
200- core->is_last = false;
201- core->is_node_innerest_dim = false;
202- core->size = sym_corenum;
203- model_info.arg_list.emplace_back(core);
204- // set m
205- Expr expr_m = CreateExpr("m_size");
206- Expr expr_tilem = CreateExpr("tilem_size");
207- Expr expr_basem = CreateExpr("basem_size");
208- SymVarInfoPtr sym_m = std::make_shared<SymVarInfo>(expr_m);
209- SymVarInfoPtr sym_tilem = std::make_shared<SymVarInfo>(expr_tilem);
210- sym_tilem->align = ge::Symbol(16);
211- sym_tilem->related_scope = {HardwareDef::L2};
212- SymVarInfoPtr sym_basem = std::make_shared<SymVarInfo>(expr_basem);
213- sym_basem->align = ge::Symbol(16);
214- sym_basem->related_scope = {HardwareDef::L0A, HardwareDef::L0C, HardwareDef::L1};
215- AttAxisPtr m = std::make_shared<AttAxis>();
216- AttAxisPtr tilem = std::make_shared<AttAxis>();
217- AttAxisPtr basem = std::make_shared<AttAxis>();
218- m->name = "m";
219- m->axis_pos = AxisPosition::ORIGIN;
220- m->bind_multicore = false;
221- m->is_last = false;
222- m->is_node_innerest_dim = false;
223- m->size = sym_m;
224- 
225- tilem->name = "tilem";
226- tilem->axis_pos = AxisPosition::INNER;
227- tilem->bind_multicore = false;
228- tilem->is_last = false;
229- tilem->is_node_innerest_dim = true;
230- tilem->size = sym_tilem;
231- tilem->orig_axis.push_back(m.get());
232- tilem->from_axis = {m.get()};
233- 
234- basem->name = "basem";
235- basem->axis_pos = AxisPosition::INNER;
236- basem->bind_multicore = false;
237- basem->is_last = true;
238- basem->is_node_innerest_dim = false;
239- basem->size = sym_basem;
240- basem->orig_axis.push_back(m.get());
241- basem->from_axis = {tilem.get()};
242- model_info.arg_list.emplace_back(m);
243- model_info.arg_list.emplace_back(tilem);
244- model_info.arg_list.emplace_back(basem);
245- 
246- // set n
247- Expr expr_n = CreateExpr("n_size");
248- Expr expr_tilen = CreateExpr("tilen_size");
249- Expr expr_basen = CreateExpr("basen_size");
250- SymVarInfoPtr sym_n = std::make_shared<SymVarInfo>(expr_n);
251- SymVarInfoPtr sym_tilen = std::make_shared<SymVarInfo>(expr_tilen);
252- sym_tilen->align = ge::Symbol(16);
253- sym_tilen->related_scope = {HardwareDef::L2};
254- SymVarInfoPtr sym_basen = std::make_shared<SymVarInfo>(expr_basen);
255- sym_basen->align = ge::Symbol(16);
256- sym_basen->related_scope = {HardwareDef::L0B, HardwareDef::L0C, HardwareDef::L1};
257- AttAxisPtr n = std::make_shared<AttAxis>();
258- AttAxisPtr tilen = std::make_shared<AttAxis>();
259- AttAxisPtr basen = std::make_shared<AttAxis>();
260- n->name = "n";
261- n->axis_pos = AxisPosition::ORIGIN;
262- n->bind_multicore = false;
263- n->is_last = false;
264- n->is_node_innerest_dim = false;
265- n->size = sym_n;
266- 
267- tilen->name = "tilen";
268- tilen->axis_pos = AxisPosition::INNER;
269- tilen->bind_multicore = false;
270- tilen->is_last = false;
271- tilen->is_node_innerest_dim = true;
272- tilen->size = sym_tilen;
273- tilen->orig_axis.push_back(n.get());
274- tilen->from_axis = {n.get()};
275- 
276- basen->name = "basen";
277- basen->axis_pos = AxisPosition::INNER;
278- basen->bind_multicore = false;
279- basen->is_last = true;
280- basen->is_node_innerest_dim = true;
281- basen->size = sym_basen;
282- basen->orig_axis.push_back(n.get());
283- basen->from_axis = {tilen.get()};
284- 
285- model_info.arg_list.emplace_back(n);
286- model_info.arg_list.emplace_back(tilen);
287- model_info.arg_list.emplace_back(basen);
288- 
289- // setk
290- Expr expr_k = CreateExpr("k_size");
291- Expr expr_stepka = CreateExpr("stepka_size");
292- Expr expr_stepkb = CreateExpr("stepkb_size");
293- Expr expr_basek = CreateExpr("basek_size");
294- SymVarInfoPtr sym_k = std::make_shared<SymVarInfo>(expr_k);
295- SymVarInfoPtr sym_stepka = std::make_shared<SymVarInfo>(expr_stepka);
296- sym_stepka->align = ge::Symbol(256);
297- sym_stepka->related_scope = {HardwareDef::L1};
298- SymVarInfoPtr sym_stepkb = std::make_shared<SymVarInfo>(expr_stepkb);
299- sym_stepkb->align = ge::Symbol(16);
300- sym_stepkb->related_scope = {HardwareDef::L1};
301- SymVarInfoPtr sym_basek = std::make_shared<SymVarInfo>(expr_basek);
302- sym_basek->align = ge::Symbol(16);
303- sym_basek->related_scope = {HardwareDef::L0A, HardwareDef::L0B};
304- AttAxisPtr k = std::make_shared<AttAxis>();
305- k->name = "k";
306- k->axis_pos = AxisPosition::ORIGIN;
307- k->bind_multicore = false;
308- k->is_last = false;
309- k->is_node_innerest_dim = false;
310- k->size = sym_k;
311- 
312- AttAxisPtr stepka = std::make_shared<AttAxis>();
313- stepka->name = "stepka";
314- stepka->axis_pos = AxisPosition::INNER;
315- stepka->bind_multicore = false;
316- stepka->is_last = false;
317- stepka->is_node_innerest_dim = true;
318- stepka->size = sym_stepka;
319- stepka->orig_axis.push_back(k.get());
320- stepka->from_axis = {k.get()};
321- 
322- AttAxisPtr stepkb = std::make_shared<AttAxis>();
323- stepkb->name = "stepkb";
324- stepkb->axis_pos = AxisPosition::INNER;
325- stepkb->bind_multicore = false;
326- stepkb->is_last = false;
327- stepkb->is_node_innerest_dim = true;
328- stepkb->size = sym_stepkb;
329- stepkb->orig_axis.push_back(k.get());
330- stepkb->from_axis = {stepka.get()};
331- 
332- AttAxisPtr basek = std::make_shared<AttAxis>();
333- basek->name = "basek";
334- basek->axis_pos = AxisPosition::INNER;
335- basek->bind_multicore = false;
336- basek->is_last = true;
337- basek->is_node_innerest_dim = false;
338- basek->size = sym_basek;
339- basek->orig_axis.push_back(k.get());
340- basek->from_axis = {stepkb.get()};
341- 
342- model_info.arg_list.emplace_back(k);
343- model_info.arg_list.emplace_back(stepka);
344- model_info.arg_list.emplace_back(stepkb);
345- model_info.arg_list.emplace_back(basek);
346- 
347- Expr l0a_occupy = expr_basem * expr_basek * CreateExpr(4);
348- Expr l0b_occupy = expr_basek * expr_basen * CreateExpr(4);
349- Expr l0c_occupy = expr_basem * expr_basen * CreateExpr(4);
350- Expr l1_occupy = (expr_stepka * expr_basem * CreateExpr(4)) + (expr_stepkb * expr_basen * CreateExpr(4));
351- Expr l2_occupy = (expr_tilen * expr_tilem * CreateExpr(2)) + ((expr_tilen + expr_tilem) * expr_k * CreateExpr(2));
352- std::map<HardwareDef, Expr> hardware_cons;
353- model_info.hardware_cons[HardwareDef::L0A] = l0a_occupy;
354- model_info.hardware_cons[HardwareDef::L0B] = l0b_occupy;
355- model_info.hardware_cons[HardwareDef::L0C] = l0c_occupy;
356- model_info.hardware_cons[HardwareDef::L1] = l1_occupy;
357- model_info.hardware_cons[HardwareDef::L2] = l2_occupy;
358- model_info.hardware_cons[HardwareDef::UB] = CreateExpr(0L);
359- 
360- Expr tile_cnt = ((expr_n / expr_tilen) * (expr_m / expr_tilem));
361- Expr base_cnt = af::sym::Max(af::sym::kSymbolOne,
362- (((expr_tilem * expr_tilen) / (expr_basem * expr_basen)) / CreateExpr("block_dim")));
363- Expr al1_cnt = (expr_k / expr_stepka);
364- Expr bl1_cnt = (expr_stepka / expr_stepkb);
365- Expr l0_cnt = (expr_stepkb / expr_basek);
366- Expr l1_cnt = (al1_cnt * bl1_cnt);
367- Expr base_fixpipe_cost = ((expr_basem * expr_basen * CreateExpr(4)) / CreateExpr(32));
368- Expr al1_mte2 = (((expr_basem * expr_stepka * CreateExpr(2)) /
369- (CreateExpr(32) / af::sym::Max(af::sym::kSymbolOne, (CreateExpr(256) / expr_stepka)))) +
370- CreateExpr(210));
371- std::cout << "AL1 mte2: " << al1_mte2 << std::endl;
372- Expr bl1_mte2 = (((expr_basen * expr_stepkb * CreateExpr(2)) /
373- (CreateExpr(32) / af::sym::Max(af::sym::kSymbolOne, (CreateExpr(256) / expr_basen)))) +
374- CreateExpr(210));
375- std::cout << "BL1 mte2: " << bl1_mte2 << std::endl;
376- Expr mac = (((tile_cnt * base_cnt * l1_cnt * l0_cnt)) * (expr_basem * expr_basen * expr_k) /
377- (CreateExpr(16) * CreateExpr(256)));
378- Expr mte2 = (tile_cnt * base_cnt * al1_cnt * (al1_mte2 + (bl1_cnt * bl1_mte2)));
379- std::cout << "mte2: " << mte2 << std::endl;
380- Expr fixpipe = (tile_cnt * base_cnt * base_fixpipe_cost);
381- model_info.objects[PipeType::AIC_MAC] = mac;
382- model_info.objects[PipeType::AIC_MTE2] = mte2;
383- model_info.objects[PipeType::AIC_FIXPIPE] = fixpipe;
384- model_info.tiling_case_id = 1;
385- model_info.eq_exprs[kFatherToChildNoTail].push_back(std::pair(expr_stepka, expr_stepkb));
386- model_info.eq_exprs[kFatherToChildNoTail].push_back(std::pair(expr_tilen, expr_basen));
387- model_info.eq_exprs[kFatherToChildNoTail].push_back(std::pair(expr_tilem, expr_basem));
388- model_info.eq_exprs[kFatherToChildNoTail].push_back(std::pair(expr_stepkb, expr_basek));
389- model_info.leq_exprs[kFatherToChildLarger].push_back((expr_tilem - expr_m));
390- model_info.leq_exprs[kFatherToChildLarger].push_back((expr_tilen - expr_n));
391- model_info.leq_exprs[kFatherToChildLarger].push_back((expr_stepka - expr_k));
392- model_info.container_exprs["Q1"] = (expr_m + expr_n);
393- model_info.tensor_exprs["MATMUL_OUTPUT1"] = (expr_m + expr_n);
394- model_info.output_size = 1;
395- 
396- return model_info;
397-}
398- 
399-ModelInfo CreateCeilingModel() {
400- ModelInfo model_info;
401- Expr expr_s1 = CreateExpr("s1_size");
402- Expr expr_s2 = CreateExpr("s2_size");
403- Expr expr_s2t = CreateExpr("s2t_size");
404- Expr expr_s2T = af::sym::Ceiling(expr_s2 / expr_s2t);
405- Expr expr_s1s2T = (expr_s1 * expr_s2T);
406- Expr expr_s1s2Tb = CreateExpr("s1s2Tb_size");
407- SymVarInfoPtr sym_s1 = std::make_shared<SymVarInfo>(expr_s1);
408- SymVarInfoPtr sym_s2 = std::make_shared<SymVarInfo>(expr_s2);
409- SymVarInfoPtr sym_s2t = std::make_shared<SymVarInfo>(expr_s2t);
410- SymVarInfoPtr sym_s1s2Tb = std::make_shared<SymVarInfo>(expr_s1s2Tb);
411- 
412- AttAxisPtr s1 = std::make_shared<AttAxis>();
413- s1->name = "s1";
414- s1->axis_pos = AxisPosition::ORIGIN;
415- s1->bind_multicore = false;
416- s1->is_last = false;
417- s1->is_node_innerest_dim = false;
418- s1->size = sym_s1;
419- 
420- AttAxisPtr s2 = std::make_shared<AttAxis>();
421- s2->name = "s2";
422- s2->axis_pos = AxisPosition::ORIGIN;
423- s2->bind_multicore = false;
424- s2->is_last = false;
425- s2->is_node_innerest_dim = false;
426- s2->size = sym_s2;
427- 
428- AttAxisPtr s2t = std::make_shared<AttAxis>();
429- s2t->name = "s2t";
430- s2t->axis_pos = AxisPosition::INNER;
431- s2t->bind_multicore = false;
432- s2t->is_last = false;
433- s2t->is_node_innerest_dim = false;
434- s2t->size = sym_s2t;
435- s2t->orig_axis.push_back(s2.get());
436- s2t->from_axis = {s2.get()};
437- 
438- AttAxisPtr s1s2Tb = std::make_shared<AttAxis>();
439- s1s2Tb->name = "s1s2Tb";
440- s1s2Tb->axis_pos = AxisPosition::INNER;
441- s1s2Tb->bind_multicore = false;
442- s1s2Tb->is_last = false;
443- s1s2Tb->is_node_innerest_dim = true;
444- s1s2Tb->size = sym_s1s2Tb;
445- s1s2Tb->orig_axis.push_back(s1.get());
446- s1s2Tb->orig_axis.push_back(s2.get());
447- s1s2Tb->from_axis = {s1.get()};
448- 
449- model_info.arg_list.emplace_back(s1);
450- model_info.arg_list.emplace_back(s2);
451- model_info.arg_list.emplace_back(s2t);
452- model_info.arg_list.emplace_back(s1s2Tb);
453- 
454- Expr core_num = af::sym::Ceiling((expr_s1 * af::sym::Ceiling(expr_s2 / expr_s2t)) / expr_s1s2Tb);
455- std::map<HardwareDef, Expr> hardware_cons;
456- model_info.hardware_cons[HardwareDef::UB] = expr_s1 * CreateExpr(10);
457- model_info.hardware_cons[HardwareDef::CORENUM] = core_num;
458- 
459- model_info.output_size = 1;
460- model_info.tiling_case_id = 0;
461- return model_info;
462-}
463-} // namespace att
@@ -1,19 +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-#ifndef STUB_MODEL_INFO_H_
12-#define STUB_MODEL_INFO_H_
13-#include "base/model_info.h"
14-namespace att {
15-ModelInfo CreateModelInfo(const ge::ExprType expr_type = ge::ExprType::kExprVariable);
16-ModelInfo GetMatmulL2TileInfo();
17-ModelInfo CreateCeilingModel();
18-} // namespace att
19-#endif
@@ -20,11 +20,7 @@ bool TestCase(std::vector<int64_t> shapes) {
20 int64_t n = shapes[3];20 int64_t n = shapes[3];
21 MMTilingData tilingData;21 MMTilingData tilingData;
22 tilingData.set_block_dim(20);22 tilingData.set_block_dim(20);
23- tilingData.set_l2_size(128 * 1024 * 1024);
24 tilingData.set_l1_size(512 * 1024);23 tilingData.set_l1_size(512 * 1024);
25- tilingData.set_l0a_size(64 * 1024);
26- tilingData.set_l0b_size(64 * 1024);
27- tilingData.set_l0c_size(128 * 1024);
28 tilingData.set_m_size(m);24 tilingData.set_m_size(m);
29 tilingData.set_k_size(k);25 tilingData.set_k_size(k);
30 tilingData.set_n_size(n);26 tilingData.set_n_size(n);
@@ -34,13 +30,8 @@ bool TestCase(std::vector<int64_t> shapes) {
34 30 
35 const auto status = GetTiling(tilingData, 1u, nullptr);31 const auto status = GetTiling(tilingData, 1u, nullptr);
36 if ((status)) {32 if ((status)) {
37- std::cout << "tile_l2_m" << " = " << tilingData.get_tilem_size() << std::endl;
38- std::cout << "tile_l2_n" << " = " << tilingData.get_tilen_size() << std::endl;
39 std::cout << "step_ka" << " = " << tilingData.get_stepka_size() << std::endl;33 std::cout << "step_ka" << " = " << tilingData.get_stepka_size() << std::endl;
40 std::cout << "step_kb" << " = " << tilingData.get_stepkb_size() << std::endl;34 std::cout << "step_kb" << " = " << tilingData.get_stepkb_size() << std::endl;
41- std::cout << "base_k" << " = " << tilingData.get_basek_size() << std::endl;
42- std::cout << "base_m" << " = " << tilingData.get_basem_size() << std::endl;
43- std::cout << "base_n" << " = " << tilingData.get_basen_size() << std::endl;
44 return true;35 return true;
45 }36 }
46 std::cout << "mm tiling func execute failed." << std::endl;37 std::cout << "mm tiling func execute failed." << std::endl;