已合并
fix: 补齐 scalar-like 节点 IR 属性引用的动态 size var 收集 #2125
Jett_Woo创建于 23 天前
fix: 补齐 scalar-like 节点 IR 属性引用的动态 size var 收集 #2125
已合并
共 2 个文件变更+49-0
| @@ -42,6 +42,17 @@ void InsertFreeSymbolsIntoVarSet(const af::Expression &exp, SizeVarSet &size_var | |||
| 42 | size_vars.insert(free_symbols.begin(), free_symbols.end()); | 42 | size_vars.insert(free_symbols.begin(), free_symbols.end()); |
| 43 | } | 43 | } |
| 44 | 44 | ||
| 45 | +static void InsertIrAttrFreeSymbols(const af::NodePtr &node, const char *attr_name, SizeVarSet &size_vars) { | ||
| 46 | + auto asc_node = std::dynamic_pointer_cast<af::AscNode>(node); | ||
| 47 | + if (asc_node == nullptr || asc_node->attr.ir_attr == nullptr) { | ||
| 48 | + return; | ||
| 49 | + } | ||
| 50 | + af::Expression expr; | ||
| 51 | + if (asc_node->attr.ir_attr->GetAttrValue(attr_name, expr) == af::GRAPH_SUCCESS) { | ||
| 52 | + InsertFreeSymbolsIntoVarSet(expr, size_vars); | ||
| 53 | + } | ||
| 54 | +} | ||
| 55 | + | ||
| 45 | bool ParseKsIndex(const std::string &name, uint64_t &index) { | 56 | bool ParseKsIndex(const std::string &name, uint64_t &index) { |
| 46 | if (name.size() <= 2U || name[0] != 'k' || name[1] != 's') { | 57 | if (name.size() <= 2U || name[0] != 'k' || name[1] != 's') { |
| 47 | return false; | 58 | return false; |
| @@ -271,6 +282,18 @@ void AscGraphInfoComplete::AppendOriginalSizeVar(const af::AscGraph &graph, Size | |||
| 271 | } | 282 | } |
| 272 | auto all_nodes = graph.GetAllNodes(); | 283 | auto all_nodes = graph.GetAllNodes(); |
| 273 | for (const auto &node : all_nodes) { | 284 | for (const auto &node : all_nodes) { |
| 285 | + // scalar-like 值生产节点(IndexExpr/Arange)无输出视图,其 IR 属性表达式引用的动态符号 | ||
| 286 | + // 只能从节点属性收集;AutoScheduler 会清空 size var 后仅凭本函数重建,漏扫会使 | ||
| 287 | + // tiling data 缺字段、设备代码裸印符号(use of undeclared identifier)。 | ||
| 288 | + if (af::ops::IsOps<IndexExpr>(node)) { | ||
| 289 | + InsertIrAttrFreeSymbols(node, "expr", size_vars); | ||
| 290 | + continue; | ||
| 291 | + } | ||
| 292 | + if (af::ops::IsOps<Arange>(node)) { | ||
| 293 | + InsertIrAttrFreeSymbols(node, "base", size_vars); | ||
| 294 | + InsertIrAttrFreeSymbols(node, "step", size_vars); | ||
| 295 | + continue; | ||
| 296 | + } | ||
| 274 | if (!af::ops::IsOps<Nddma>(node) && !af::ops::IsOps<Store>(node) && !af::ops::IsOps<Load>(node) && | 297 | if (!af::ops::IsOps<Nddma>(node) && !af::ops::IsOps<Store>(node) && !af::ops::IsOps<Load>(node) && |
| 275 | !af::ops::IsOps<Gather>(node)) { | 298 | !af::ops::IsOps<Gather>(node)) { |
| 276 | continue; | 299 | continue; |
| @@ -11,7 +11,9 @@ | |||
| 11 | 11 | ||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | + | ||
| 14 | 15 | ||
| 16 | + | ||
| 15 | 17 | ||
| 16 | namespace optimize { | 18 | namespace optimize { |
| 17 | namespace { | 19 | namespace { |
| @@ -64,5 +66,29 @@ TEST(FrontendShapeVarsTest, CollectsSymbolsEmbeddedInAxisExpressions) { | |||
| 64 | EXPECT_EQ(GetNames(vars), (std::vector<std::string>{"s0"})); | 66 | EXPECT_EQ(GetNames(vars), (std::vector<std::string>{"s0"})); |
| 65 | } | 67 | } |
| 66 | 68 | ||
| 69 | +TEST(FrontendShapeVarsTest, CollectsSymbolsReferencedByScalarLikeIrAttrs) { | ||
| 70 | + af::AscGraph graph("scalar_like_shape_vars"); | ||
| 71 | + const auto ks0 = graph.CreateSizeVar("ks0"); | ||
| 72 | + graph.CreateAxis("z0", ks0); | ||
| 73 | + const auto ks1 = graph.CreateSizeVar("ks1"); | ||
| 74 | + const auto ks2 = graph.CreateSizeVar("ks2"); | ||
| 75 | + | ||
| 76 | + af::ascir_op::IndexExpr index("index", graph); | ||
| 77 | + index.ir_attr.SetExpr(ks1 + af::Symbol(2)); | ||
| 78 | + af::ascir_op::Arange arange("arange", graph); | ||
| 79 | + arange.ir_attr.SetBase(af::Symbol(0)); | ||
| 80 | + arange.ir_attr.SetStep(ks2); | ||
| 81 | + | ||
| 82 | + // AutoScheduler 会先 ClearAllSizeVar 再仅凭 AppendOriginalSizeVar 重建 size var 表, | ||
| 83 | + // IndexExpr.expr / Arange base/step 引用的动态符号必须能被重新收集, | ||
| 84 | + // 否则 tiling data 缺字段、设备代码裸印符号(use of undeclared identifier)。 | ||
| 85 | + ASSERT_EQ(ScheduleUtils::ClearAllSizeVar(graph), af::SUCCESS); | ||
| 86 | + SizeVarSet var_set; | ||
| 87 | + AscGraphInfoComplete::AppendOriginalSizeVar(graph, var_set); | ||
| 88 | + std::vector<af::Expression> vars(var_set.begin(), var_set.end()); | ||
| 89 | + ASSERT_EQ(AscGraphInfoComplete::NormalizeFrontendShapeVars(vars), af::SUCCESS); | ||
| 90 | + EXPECT_EQ(GetNames(vars), (std::vector<std::string>{"ks0", "ks1", "ks2"})); | ||
| 91 | +} | ||
| 92 | + | ||
| 67 | } // namespace | 93 | } // namespace |
| 68 | } // namespace optimize | 94 | } // namespace optimize |