已合并
fix: 补齐 scalar-like 节点 IR 属性引用的动态 size var 收集 #2125
fix: 补齐 scalar-like 节点 IR 属性引用的动态 size var 收集 #2125
已合并
Jett_Woo创建于 23 天前
共 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+ 
45bool ParseKsIndex(const std::string &name, uint64_t &index) {56bool 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#include <gtest/gtest.h>11#include <gtest/gtest.h>
12 12 
13#include "ascgraph_info_complete.h"13#include "ascgraph_info_complete.h"
14+#include "ascir_ops.h"
14#include "graph/symbolizer/symbolic_utils.h"15#include "graph/symbolizer/symbolic_utils.h"
16+#include "schedule_utils.h"
15 17 
16namespace optimize {18namespace optimize {
17namespace {19namespace {
@@ -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} // namespace93} // namespace
68} // namespace optimize94} // namespace optimize