已合并
refactor: 面向VV融合的求解器,删除L0/L2求解器 #1684
gcw_NLOuEjCz创建于 22 天前
refactor: 面向VV融合的求解器,删除L0/L2求解器 #1684
已合并
共 42 个文件变更+75-7108
| @@ -38,7 +38,7 @@ using ge::SymbolicUtils; | |||
| 38 | 38 | ||
| 39 | namespace att { | 39 | namespace att { |
| 40 | using Expr = af::Expression; | 40 | using 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 | ||
| 43 | enum class HardwareDef { | 43 | enum 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 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 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 | - | ||
| @@ -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 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 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 | - | ||
| @@ -12,12 +12,6 @@ | |||
| 12 | 12 | ||
| 13 | namespace att { | 13 | namespace att { |
| 14 | std::string GetSolverHead(SolverType type) { | 14 | std::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 | ||
| 27 | std::string GetSolverFunc(SolverType type) { | 21 | std::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 | 10 | ||
| 11 | 11 | ||
| 12 | - | ||
| 13 | - | ||
| 14 | 12 | ||
| 15 | 13 | ||
| 16 | 14 | ||
| @@ -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 | - | ||
| 11 | - | ||
| 12 | - | ||
| 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 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 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 | - | ||
| @@ -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 | - | ||
| 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 | - | ||
| 11 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 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 | - | ||
| @@ -14,7 +14,6 @@ | |||
| 14 | 14 | ||
| 15 | namespace att { | 15 | namespace att { |
| 16 | constexpr char kSolverGenError[] = "Solver Gen Error"; | 16 | constexpr char kSolverGenError[] = "Solver Gen Error"; |
| 17 | -constexpr uint32_t kMaxL0VarNum = 3u; | ||
| 18 | inline std::string GetSmoothString(std::string str) { | 17 | inline 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 | 12 | ||
| 13 | 13 | ||
| 14 | namespace att { | 14 | namespace 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 | - | ||
| 127 | template <typename SolverGenType> | 15 | template <typename SolverGenType> |
| 128 | auto SolverPassManager::GenerateSolverGen() -> SolverGenType { | 16 | auto 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 | - | ||
| 196 | template <typename SpecificSolverGen> | 52 | template <typename SpecificSolverGen> |
| 197 | std::pair<std::string, std::string> SolverPassManager::GenerateSolverPassFunc(SpecificSolverGen solver_gen) { | 53 | std::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 | ||
| 214 | std::pair<std::string, std::string> SolverPassManager::SolverPassFuncGen(SolverType type) { | 68 | std::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 | - | ||
| 268 | std::pair<std::string, std::string> SolverPassManager::GeneralSolverDtFuncGen() { | 92 | std::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 | ||
| 275 | std::pair<std::string, std::string> SolverPassManager::SolverDtFuncGen(SolverType type) { | 99 | std::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 | - | ||
| 314 | std::string SolverPassManager::GeneralSolverPassClassGen() { | 108 | std::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 | ||
| 321 | std::string SolverPassManager::SolverPassClassGen(SolverType type) { | 115 | std::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 | ||
| 345 | bool SolverPassManager::IsNeedSolver(std::vector<ArgsManager> args_managers, SolverType type) { | 133 | bool 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 | 18 | ||
| 19 | 19 | ||
| 20 | 20 | ||
| 21 | - | ||
| 22 | - | ||
| 23 | 21 | ||
| 24 | 22 | ||
| 25 | 23 | ||
| @@ -76,35 +74,25 @@ class SolverPassManager : public InputOutputSettersMixin<SolverPassManager>, | |||
| 76 | 74 | ||
| 77 | private: | 75 | private: |
| 78 | // solver pass | 76 | // 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/*.cpp | 258 | ${ATT_DIR}/generator/solver_pass_gen/axes_reorder_solver/*.cpp |
| 259 | ${ATT_DIR}/generator/solver_pass_gen/general_solver/*.cpp | 259 | ${ATT_DIR}/generator/solver_pass_gen/general_solver/*.cpp |
| 260 | ${ATT_DIR}/generator/solver_pass_gen/golden_solver/*.cpp | 260 | ${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/*.cpp | 261 | ${ATT_DIR}/generator/extra_info_gen/*.cpp |
| 264 | ${ATT_DIR}/generator/generator_utils/*.cpp | 262 | ${ATT_DIR}/generator/generator_utils/*.cpp |
| 265 | ${ATT_DIR}/generator/tiling_data_gen/*.cpp | 263 | ${ATT_DIR}/generator/tiling_data_gen/*.cpp |
| @@ -22,8 +22,6 @@ att::Expr GetSafeOffsetDivisor(const att::Expr &expr) { | |||
| 22 | 22 | ||
| 23 | struct FfnExprContext { | 23 | struct 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 | ||
| 35 | void BuildFfnMaxTokenAxes(att::ModelInfo &model_info, FfnExprContext &ctx, att::AttAxisPtr &maxTokens, | 33 | void 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 | ||
| 93 | void BuildFfnN1Axes(FfnExprContext &ctx, att::AttAxisPtr &n1, att::AttAxisPtr &basen1) { | 63 | void 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 | - | ||
| 201 | void FillFfnModelInfo(att::ModelInfo &model_info, const FfnExprContext &ctx) { | 133 | void 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 | ||
| 252 | void AppendFfnArgList(att::ModelInfo &model_info, const att::AttAxisPtr &maxTokens, const att::AttAxisPtr &basen1, | 166 | void 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 | } // namespace | 177 | } // 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 att | 203 | } // namespace att |
| @@ -12,22 +12,13 @@ | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | namespace { | 14 | namespace { |
| 15 | -att::Expr GetSafeDivisor(const att::Expr &expr) { | ||
| 16 | - return af::sym::Max(af::sym::kSymbolOne, expr); | ||
| 17 | -} | ||
| 18 | - | ||
| 19 | struct MatmulExprContext { | 15 | struct 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 | ||
| 33 | void BuildCoreAxis(att::ModelInfo &model_info, MatmulExprContext &ctx) { | 24 | void 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 | ||
| 206 | void FillMatmulHardwareCons(att::ModelInfo &model_info, const MatmulExprContext &ctx) { | 117 | void 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 | - | ||
| 273 | void FillModelInfo(att::ModelInfo &model_info, const MatmulExprContext &ctx) { | 122 | void 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 @@ | |||
| 15 | namespace { | 15 | namespace { |
| 16 | struct MatmulExprContext { | 16 | struct 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 | ||
| 28 | struct L2TileExprContext { | 24 | struct 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 | ||
| 42 | void InitDefaultExpr(const ge::ExprType expr_type, att::Expr &default_expr, bool &is_const) { | 33 | void 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 | |||
| 91 | void BuildCreateModelInfoMArgs(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr, | 82 | void 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 | ||
| 118 | void BuildCreateModelInfoNArgs(att::ModelInfo &model_info, MatmulExprContext &ctx) { | 103 | void 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 | ||
| 142 | void BuildCreateModelInfoKArg(att::ModelInfo &model_info, MatmulExprContext &ctx) { | 121 | void BuildCreateModelInfoKArg(att::ModelInfo &model_info, MatmulExprContext &ctx) { |
| @@ -154,35 +133,20 @@ void BuildCreateModelInfoKArg(att::ModelInfo &model_info, MatmulExprContext &ctx | |||
| 154 | } | 133 | } |
| 155 | 134 | ||
| 156 | void FillCreateModelInfo(att::ModelInfo &model_info, const MatmulExprContext &ctx) { | 135 | void 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 | ||
| 202 | void BuildL2TileMArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) { | 166 | void 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 | ||
| 249 | void BuildL2TileNArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) { | 183 | void 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 | ||
| 296 | void BuildL2TileKArgs(att::ModelInfo &model_info, L2TileExprContext &ctx) { | 200 | void 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 | ||
| 329 | void FillL2TileHardwareCons(att::ModelInfo &model_info, const L2TileExprContext &ctx) { | 224 | void 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 | ||
| 345 | void FillL2TileModelInfo(att::ModelInfo &model_info, const L2TileExprContext &ctx) { | 230 | void 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 @@ | |||
| 13 | namespace { | 13 | namespace { |
| 14 | struct SolverExprContext { | 14 | struct 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 | |||
| 75 | void BuildMArgList(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr, | 71 | void 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 | ||
| 103 | void BuildNArgList(att::ModelInfo &model_info, const bool is_const, const att::Expr &default_expr, | 93 | void 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 | ||
| 128 | void BuildKArg(att::ModelInfo &model_info, SolverExprContext &ctx) { | 112 | void BuildKArg(att::ModelInfo &model_info, SolverExprContext &ctx) { |
| @@ -140,33 +124,20 @@ void BuildKArg(att::ModelInfo &model_info, SolverExprContext &ctx) { | |||
| 140 | } | 124 | } |
| 141 | 125 | ||
| 142 | void FillModelInfo(att::ModelInfo &model_info, const SolverExprContext &ctx) { | 126 | void 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/*.cpp | 21 | ${TC_DIR}/test_expr/*.cpp |
| 22 | ${TC_DIR}/solver_pass_gen/axes_reorder_solver_gen/*.cpp | 22 | ${TC_DIR}/solver_pass_gen/axes_reorder_solver_gen/*.cpp |
| 23 | ${TC_DIR}/solver_pass_gen/general_solver_gen/*.cpp | 23 | ${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/*.cpp | 24 | ${TC_DIR}/solver_pass/general_solver/*.cpp |
| 29 | ${TC_DIR}/solver_pass_manager/*.cpp | 25 | ${TC_DIR}/solver_pass_manager/*.cpp |
| 30 | ${TC_DIR}/select_model/*.cpp | 26 | ${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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 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 | -} | ||
Dautofuse/tests/st/att/testcase/source_mirror/generator/solver_pass/l0_solver/test_l0_solver.cpp+0-344
| @@ -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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 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 | -} | ||
Dautofuse/tests/st/att/testcase/source_mirror/generator/solver_pass/l2_solver/test_l2_solver.cpp+0-129
| @@ -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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | -do { | ||
| 20 | - std::cout << "[ERROR]" << log << std::endl; | ||
| 21 | -} while (0) | ||
| 22 | - | ||
| 23 | - | ||
| 24 | - | ||
| 25 | - | ||
| 26 | - | ||
| 27 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | -do { | ||
| 15 | - std::cout << "[ERROR]" << log << std::endl; | ||
| 16 | -} while (0) | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 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/preprocess | 20 | # generator/preprocess |
| 21 | ${TOP_DIR}/tests/ut/att/testcase/preprocess/*.cpp | 21 | ${TOP_DIR}/tests/ut/att/testcase/preprocess/*.cpp |
| 22 | # generator/solver_pass | 22 | # 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/*.cpp | 23 | ${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_gen | 24 | # generator/solver_pass_gen |
| 28 | ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/axes_reorder_gen/*.cpp | 25 | ${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/*.cpp | 26 | ${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/*.cpp | 27 | ${TOP_DIR}/tests/ut/att/testcase/solver_pass_gen/manager/*.cpp |
| 33 | # generator/core | 28 | # generator/core |
| 34 | ${TOP_DIR}/tests/ut/att/testcase/generator/core/*.cpp | 29 | ${TOP_DIR}/tests/ut/att/testcase/generator/core/*.cpp |
| @@ -60,9 +55,6 @@ file(GLOB SOURCES | |||
| 60 | ${TOP_DIR}/tests/ut/att/utils/*.cpp | 55 | ${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 | - | ||
| 66 | add_executable(att_ut | 58 | add_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 | - | ||
| 202 | TEST(GeneratorUT, NormalGroupRegistersOnlyDirectStandardHeaders) { | 187 | TEST(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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 20 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | -do { | ||
| 18 | - std::cout << "[ERROR]" << log << std::endl; | ||
| 19 | -} while (0) | ||
| 20 | - | ||
| 21 | - | ||
| 22 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | -do { | ||
| 15 | - std::cout << "[ERROR]" << log << std::endl; | ||
| 16 | -} while (0) | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | - | ||
| 15 | - | ||
| 16 | - | ||
| 17 | - | ||
| 18 | - | ||
| 19 | - | ||
| 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 | - | ||
| 68 | TEST_F(TestSolverPassManager, TEST_IS_NEED_SOLVER) { | 43 | TEST_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 | - | ||
| 12 | - | ||
| 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 | - | ||
| 12 | - | ||
| 13 | - | ||
| 14 | -namespace att { | ||
| 15 | -ModelInfo CreateModelInfo(const ge::ExprType expr_type = ge::ExprType::kExprVariable); | ||
| 16 | -ModelInfo GetMatmulL2TileInfo(); | ||
| 17 | -ModelInfo CreateCeilingModel(); | ||
| 18 | -} // namespace att | ||
| 19 | - | ||
| @@ -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; |