已合并
fix: split shared data for matmul vector partition #1430
fix: split shared data for matmul vector partition #1430
已合并
ling-DT创建于 7月23日
共 3 个文件变更+302-71
@@ -36,6 +36,76 @@ class CubeScheduleCaseGeneratorTest : public ::testing::Test {
36 }36 }
37};37};
38 38 
39+struct TwoDimGraphVars {
40+ af::Expression s0;
41+ af::Expression s1;
42+ int64_t z0_id;
43+ int64_t z1_id;
44+};
45+ 
46+TwoDimGraphVars CreateTwoDimGraphVars(af::AscGraph &graph, int64_t dim0, int64_t dim1) {
47+ auto s0 = graph.CreateSizeVar(dim0);
48+ auto s1 = graph.CreateSizeVar(dim1);
49+ auto z0 = graph.CreateAxis("z0", s0);
50+ auto z1 = graph.CreateAxis("z1", s1);
51+ return {s0, s1, z0.id, z1.id};
52+}
53+ 
54+template <typename Op>
55+void SetSchedAxis2D(Op &op, const TwoDimGraphVars &vars) {
56+ op.attr.sched.axis = {vars.z0_id, vars.z1_id};
57+}
58+ 
59+template <typename Tensor>
60+void SetTensor2D(Tensor &tensor, ge::DataType dtype, const TwoDimGraphVars &vars,
61+ const std::vector<af::Expression> &strides, const std::vector<af::Expression> &repeats) {
62+ tensor.dtype = dtype;
63+ *tensor.axis = {vars.z0_id, vars.z1_id};
64+ *tensor.strides = strides;
65+ *tensor.repeats = repeats;
66+}
67+ 
68+void InitData2D(Data &data, ge::DataType dtype, const TwoDimGraphVars &vars, const std::vector<af::Expression> &strides,
69+ const std::vector<af::Expression> &repeats, int64_t index) {
70+ SetSchedAxis2D(data, vars);
71+ SetTensor2D(data.y, dtype, vars, strides, repeats);
72+ data.attr.api.compute_type = af::ComputeType::kComputeInvalid;
73+ data.ir_attr.SetIndex(index);
74+}
75+ 
76+template <typename Tensor>
77+void InitLoad2D(Load &load, const Tensor &input, ge::DataType dtype, const TwoDimGraphVars &vars,
78+ const std::vector<af::Expression> &strides, const std::vector<af::Expression> &repeats) {
79+ load.x = input;
80+ SetSchedAxis2D(load, vars);
81+ SetTensor2D(load.y, dtype, vars, strides, repeats);
82+}
83+ 
84+template <typename Op, typename Lhs, typename Rhs>
85+void InitBinary2D(Op &op, const Lhs &lhs, const Rhs &rhs, const TwoDimGraphVars &vars) {
86+ SetSchedAxis2D(op, vars);
87+ op.x1 = lhs;
88+ op.x2 = rhs;
89+ SetTensor2D(op.y, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1});
90+}
91+ 
92+template <typename Lhs, typename Rhs>
93+void InitMatMul2D(MatMul &matmul, const Lhs &lhs, const Rhs &rhs, const TwoDimGraphVars &vars) {
94+ InitBinary2D(matmul, lhs, rhs, vars);
95+ matmul.attr.api.compute_type = af::ComputeType::kComputeCube;
96+}
97+ 
98+template <typename Tensor>
99+void InitStoreOutput2D(Store &store_op, Output &output_op, const Tensor &input, const TwoDimGraphVars &vars) {
100+ SetSchedAxis2D(store_op, vars);
101+ store_op.x = input;
102+ SetTensor2D(store_op.y, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1});
103+ 
104+ output_op.x = store_op.y;
105+ output_op.y.dtype = af::DT_FLOAT;
106+ output_op.ir_attr.SetIndex(0);
107+}
108+ 
39void ConstructJustMatMul(af::AscGraph &graph) {109void ConstructJustMatMul(af::AscGraph &graph) {
40 auto s0 = graph.CreateSizeVar(64);110 auto s0 = graph.CreateSizeVar(64);
41 auto s1 = graph.CreateSizeVar(64);111 auto s1 = graph.CreateSizeVar(64);
@@ -108,93 +178,100 @@ void ConstructJustMatMul(af::AscGraph &graph) {
108}178}
109 179 
110void ConstructMatMulAndAdd(af::AscGraph &graph) {180void ConstructMatMulAndAdd(af::AscGraph &graph) {
111- auto s0 = graph.CreateSizeVar(64);181+ const auto vars = CreateTwoDimGraphVars(graph, 64, 64);
112- auto s1 = graph.CreateSizeVar(64);
113- auto z0 = graph.CreateAxis("z0", s0);
114- auto z1 = graph.CreateAxis("z1", s1);
115 182 
116 Data data0("data0", graph);183 Data data0("data0", graph);
117- data0.attr.sched.axis = {z0.id, z1.id};
118- data0.y.dtype = af::DT_FLOAT;
119- *data0.y.axis = {z0.id, z1.id};
120- data0.attr.api.compute_type = af::ComputeType::kComputeInvalid;
121- *data0.y.strides = {s1, af::ops::One};
122- *data0.y.repeats = {s0, s1};
123- data0.ir_attr.SetIndex(0);
124- 
125 Load load0("load0");184 Load load0("load0");
126- load0.attr.sched.axis = {z0.id, z1.id};185+ InitData2D(data0, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1}, 0);
127- load0.x = data0.y;186+ InitLoad2D(load0, data0.y, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1});
128- *load0.y.axis = {z0.id, z1.id};
129- load0.y.dtype = af::DT_FLOAT;
130- *load0.y.strides = {s1, af::ops::One};
131- *load0.y.repeats = {s0, s1};
132 187 
133 Data data1("data1", graph);188 Data data1("data1", graph);
134- data1.y.dtype = af::DT_FLOAT;
135- data1.attr.sched.axis = {z0.id, z1.id};
136- *data1.y.axis = {z0.id, z1.id};
137- data1.attr.api.compute_type = af::ComputeType::kComputeInvalid;
138- *data1.y.repeats = {One, One};
139- *data1.y.strides = {Zero, Zero};
140- data1.ir_attr.SetIndex(1);
141- 
142 Load load1("load1");189 Load load1("load1");
143- load1.x = data1.y;190+ InitData2D(data1, af::DT_FLOAT, vars, {Zero, Zero}, {One, One}, 1);
144- load1.attr.sched.axis = {z0.id, z1.id};191+ InitLoad2D(load1, data1.y, af::DT_FLOAT, vars, {Zero, Zero}, {One, One});
145- load1.y.dtype = af::DT_FLOAT;
146- *load1.y.axis = {z0.id, z1.id};
147- *load1.y.strides = {Zero, Zero};
148- *load1.y.repeats = {One, One};
149 192 
150 Data data2("data2", graph);193 Data data2("data2", graph);
151- data2.y.dtype = af::DT_FLOAT;
152- data2.attr.sched.axis = {z0.id, z1.id};
153- *data2.y.axis = {z0.id, z1.id};
154- data2.attr.api.compute_type = af::ComputeType::kComputeInvalid;
155- *data2.y.repeats = {One, One};
156- *data2.y.strides = {Zero, Zero};
157- data2.ir_attr.SetIndex(1);
158- 
159 Load load2("load2");194 Load load2("load2");
160- load2.x = data2.y;195+ InitData2D(data2, af::DT_FLOAT, vars, {Zero, Zero}, {One, One}, 1);
161- load2.attr.sched.axis = {z0.id, z1.id};196+ InitLoad2D(load2, data2.y, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1});
162- load2.y.dtype = af::DT_FLOAT;
163- *load2.y.axis = {z0.id, z1.id};
164- *load2.y.strides = {s1, af::ops::One};
165- *load2.y.repeats = {s0, s1};
166 197 
167 MatMul matmul("matmul");198 MatMul matmul("matmul");
168- matmul.attr.sched.axis = {z0.id, z1.id};199+ InitMatMul2D(matmul, load0.y, load1.y, vars);
169- matmul.x1 = load0.y;
170- matmul.x2 = load1.y;
171- matmul.attr.api.compute_type = af::ComputeType::kComputeCube;
172- matmul.y.dtype = af::DT_FLOAT;
173- *matmul.y.axis = {z0.id, z1.id};
174- *matmul.y.repeats = {s0, s1};
175- *matmul.y.strides = {s1, af::ops::One};
176 200 
177 af::ascir_op::Add add_op("add");201 af::ascir_op::Add add_op("add");
178- add_op.attr.sched.axis = {z0.id, z1.id};202+ InitBinary2D(add_op, matmul.y, load2.y, vars);
179- add_op.x1 = matmul.y;
180- add_op.x2 = load2.y;
181- add_op.y.dtype = af::DT_FLOAT;
182- *add_op.y.axis = {z0.id, z1.id};
183- *add_op.y.strides = {s1, af::ops::One};
184- *add_op.y.repeats = {s0, s1};
185 203 
186 Store store_op("store");204 Store store_op("store");
187- store_op.attr.sched.axis = {z0.id, z1.id};
188- store_op.x = add_op.y;
189- *store_op.y.axis = {z0.id, z1.id};
190- store_op.y.dtype = af::DT_FLOAT;
191- *store_op.y.strides = {s1, af::ops::One};
192- *store_op.y.repeats = {s0, s1};
193- 
194 Output output_op("output");205 Output output_op("output");
195- output_op.x = store_op.y;206+ InitStoreOutput2D(store_op, output_op, add_op.y, vars);
196- output_op.y.dtype = af::DT_FLOAT;207+}
197- output_op.ir_attr.SetIndex(0);208+ 
209+size_t CountDataNodeByNameAndIndex(const ascir::ImplGraph &graph, const std::string &name, int64_t index) {
210+ size_t count = 0UL;
211+ for (const auto &node : graph.GetAllNodes()) {
212+ if (!af::ops::IsOps<af::ascir_op::Data>(node) || node->GetName() != name) {
213+ continue;
214+ }
215+ if (node->attr.ir_attr == nullptr) {
216+ continue;
217+ }
218+ auto ir_attr = node->attr.ir_attr->DownCastTo<af::AscDataIrAttrDef>();
219+ if (ir_attr == nullptr) {
220+ continue;
atomgit-bot
atomgit-botatomgit-bot7月23日

🟡 Medium Priority

在 CountDataNodeByNameAndIndex 函数(测试辅助函数)的第 206 行,代码直接对 node->attr.ir_attr 调用 ->DownCastTo<>() 而不先检查 ir_attr 是否为 null。

ir_attr 的类型是 std::unique_ptr<AscIrAttrDefBase>(定义于 ascendc_ir_def.h:370),默认初始化为 nullptr。对 null unique_ptr 调用 operator->() 属于未定义行为(通常会导致崩溃)。

虽然当前测试图中所有 Data 节点都通过 dataX.ir_attr.SetIndex(...) 设置了 ir_attr,但该函数作为通用辅助函数可能被复用于其他图场景。函数内部的第 207 行 if (ir_attr == nullptr) 空值检查也表明作者预期了可能为 null 的情况,但空指针解引用已经在第 206 行发生了。

证据链:node->attr.ir_attr(unique_ptr,可为 null)→ 第 206 行无保护地调用 ->DownCastTo<>() → 若 ir_attr 为空则 UB/崩溃 → 第 207 行的空值检查永远不会被触发。

修复方案:在解引用前先检查 ir_attr 是否为 null。

likedislike
不准确?
221+ }
222+ int64_t data_index = -1;
223+ if (ir_attr->GetIndex(data_index) == af::SUCCESS && data_index == index) {
224+ ++count;
225+ }
226+ }
227+ return count;
228+}
229+ 
230+size_t CountLoadsFromData(const ascir::ImplGraph &graph, const std::string &data_name) {
231+ size_t load_count = 0UL;
232+ for (const auto &node : graph.GetAllNodes()) {
233+ if (!af::ops::IsOps<af::ascir_op::Data>(node) || node->GetName() != data_name) {
234+ continue;
235+ }
236+ for (const auto &out_node : node->GetOutDataNodes()) {
237+ auto load_asc_node = std::dynamic_pointer_cast<af::AscNode>(out_node);
238+ if (load_asc_node != nullptr && ScheduleUtils::IsLoad(load_asc_node)) {
239+ ++load_count;
240+ }
241+ }
242+ }
243+ return load_count;
244+}
245+ 
246+void ConstructMatMulAddSharedInput(af::AscGraph &graph) {
247+ const auto vars = CreateTwoDimGraphVars(graph, 64, 64);
248+ 
249+ Data data0("data0", graph);
250+ Load load0("load0");
251+ Load load2("load2");
252+ Load load3("load3");
253+ InitData2D(data0, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1}, 0);
254+ InitLoad2D(load0, data0.y, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1});
255+ InitLoad2D(load2, data0.y, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1});
256+ InitLoad2D(load3, data0.y, af::DT_FLOAT, vars, {vars.s1, af::ops::One}, {vars.s0, vars.s1});
257+ 
258+ Data data1("data1", graph);
259+ Load load1("load1");
260+ InitData2D(data1, af::DT_FLOAT, vars, {af::ops::Zero, af::ops::Zero}, {af::ops::One, af::ops::One}, 1);
261+ InitLoad2D(load1, data1.y, af::DT_FLOAT, vars, {af::ops::Zero, af::ops::Zero}, {af::ops::One, af::ops::One});
262+ 
263+ MatMul matmul("matmul");
264+ InitMatMul2D(matmul, load0.y, load1.y, vars);
265+ 
266+ af::ascir_op::Add add_op("add");
267+ InitBinary2D(add_op, matmul.y, load2.y, vars);
268+ 
269+ af::ascir_op::Mul mul_op("mul");
270+ InitBinary2D(mul_op, add_op.y, load3.y, vars);
271+ 
272+ Store store_op("store");
273+ Output output_op("output");
274+ InitStoreOutput2D(store_op, output_op, mul_op.y, vars);
198}275}
199 276 
200void ConstructJustMatMulBias(af::AscGraph &graph) {277void ConstructJustMatMulBias(af::AscGraph &graph) {
@@ -309,6 +386,37 @@ TEST_F(CubeScheduleCaseGeneratorTest, Test_MatMul_Add_Store) {
309 ASSERT_EQ(tasks[1].grouped_graphs.size(), 2UL);386 ASSERT_EQ(tasks[1].grouped_graphs.size(), 2UL);
310}387}
311 388 
389+TEST_F(CubeScheduleCaseGeneratorTest, Test_MatMul_Add_Shared_Data_Split) {
390+ af::AscGraph graph("matmul_add_shared_data_split");
391+ ConstructMatMulAddSharedInput(graph);
392+ std::vector<ScheduleTask> tasks;
393+ optimize::CubeFusionCaseGenerator generator;
394+ OptimizerOptions options;
395+ ASSERT_EQ(generator.GeneratorTask(graph, tasks, options), af::SUCCESS);
396+ ASSERT_EQ(tasks.size(), 2UL);
397+ ASSERT_EQ(tasks[0].grouped_graphs.size(), 2UL);
398+ ASSERT_EQ(tasks[1].grouped_graphs.size(), 2UL);
399+ 
400+ size_t data0_graph_count = 0UL;
401+ size_t data0_cube_graph_count = 0UL;
402+ size_t data0_vector_graph_load_count = 0UL;
403+ for (const auto &grouped_graph : tasks[0].grouped_graphs) {
404+ const auto count = CountDataNodeByNameAndIndex(grouped_graph, "data0", 0);
405+ if (count == 0UL) {
406+ continue;
407+ }
408+ ++data0_graph_count;
409+ if (ScheduleUtils::HasComputeType(grouped_graph, af::ComputeType::kComputeCube)) {
410+ ++data0_cube_graph_count;
411+ } else {
412+ data0_vector_graph_load_count = CountLoadsFromData(grouped_graph, "data0");
413+ }
414+ }
415+ EXPECT_EQ(data0_graph_count, 2UL);
416+ EXPECT_EQ(data0_cube_graph_count, 1UL);
417+ EXPECT_EQ(data0_vector_graph_load_count, 2UL);
418+}
419+ 
312TEST_F(CubeScheduleCaseGeneratorTest, Test_MatMul_Bias_Store) {420TEST_F(CubeScheduleCaseGeneratorTest, Test_MatMul_Bias_Store) {
313 af::AscGraph graph("matmul_bias_store");421 af::AscGraph graph("matmul_bias_store");
314 ConstructJustMatMulBias(graph);422 ConstructJustMatMulBias(graph);
@@ -10,6 +10,7 @@
10 10 
11#include "task_generator/cube_schedule_case_generator.h"11#include "task_generator/cube_schedule_case_generator.h"
12#include <queue>12#include <queue>
13+#include <unordered_set>
13#include "graph/ascendc_ir/utils/asc_graph_utils.h"14#include "graph/ascendc_ir/utils/asc_graph_utils.h"
14#include "graph/symbolizer/symbolic_utils.h"15#include "graph/symbolizer/symbolic_utils.h"
15#include "graph/utils/graph_utils.h"16#include "graph/utils/graph_utils.h"
@@ -177,6 +178,122 @@ bool IsCubeFixpip(const af::AscNodePtr &cube_node) {
177 return true;178 return true;
178}179}
179 180 
181+bool IsReachableToCube(const af::AscNodePtr &start_node) {
182+ GE_ASSERT_NOTNULL(start_node);
183+ std::queue<af::AscNodePtr> node_queue;
184+ std::unordered_set<const af::Node *> visited;
185+ node_queue.push(start_node);
186+ visited.insert(start_node.get());
187+ while (!node_queue.empty()) {
188+ auto current_node = node_queue.front();
189+ node_queue.pop();
190+ GE_ASSERT_NOTNULL(current_node);
191+ if (ScheduleUtils::IsCube(current_node)) {
192+ return true;
193+ }
194+ for (const auto &out_node : current_node->GetOutDataNodes()) {
195+ auto out_asc_node = std::dynamic_pointer_cast<af::AscNode>(out_node);
196+ if (out_asc_node == nullptr || visited.find(out_asc_node.get()) != visited.end()) {
197+ continue;
198+ }
199+ visited.insert(out_asc_node.get());
200+ node_queue.push(out_asc_node);
201+ }
202+ }
203+ return false;
204+}
205+ 
206+std::string GetUniqueSplitDataName(af::AscGraph &graph, const std::string &base_name) {
207+ std::string split_name = base_name + "_shared_split_data";
208+ size_t suffix = 0UL;
209+ while (graph.FindNode(split_name.c_str()) != nullptr) {
210+ split_name = base_name + "_shared_split_data_" + std::to_string(++suffix);
211+ }
212+ return split_name;
213+}
214+ 
215+Status CopyDataNodeForSharedSplit(af::AscGraph &graph, const af::AscNodePtr &src_data_node,
216+ const std::string &split_name, af::AscNodePtr &split_data_node) {
217+ GE_ASSERT_NOTNULL(src_data_node);
218+ af::ascir_op::Data data(split_name.c_str());
219+ data.attr = src_data_node->attr;
220+ data.y.dtype = src_data_node->outputs[0].attr.dtype;
221+ *data.y.axis = src_data_node->outputs[0].attr.axis;
222+ *data.y.repeats = src_data_node->outputs[0].attr.repeats;
223+ *data.y.strides = src_data_node->outputs[0].attr.strides;
224+ split_data_node = graph.AddNode(data);
225+ GE_ASSERT_NOTNULL(split_data_node);
226+ split_data_node->attr = src_data_node->attr;
227+ split_data_node->outputs[0].attr = src_data_node->outputs[0].attr;
228+ return af::SUCCESS;
229+}
230+ 
231+Status SplitSharedDataForCubeAndVector(ascir::ImplGraph &graph,
232+ std::vector<std::pair<std::string, std::string>> &renamed_data_nodes) {
233+ auto all_nodes = graph.GetAllNodes();
234+ for (const auto &node : all_nodes) {
235+ GE_ASSERT_NOTNULL(node);
236+ if (!af::ops::IsOps<af::ascir_op::Data>(node)) {
237+ continue;
238+ }
239+ auto data_node = std::dynamic_pointer_cast<af::AscNode>(node);
240+ GE_ASSERT_NOTNULL(data_node);
241+ std::vector<af::AscNodePtr> cube_load_nodes;
242+ std::vector<af::AscNodePtr> vector_load_nodes;
243+ for (const auto &out_node : data_node->GetOutDataNodes()) {
244+ auto load_node = std::dynamic_pointer_cast<af::AscNode>(out_node);
245+ if (load_node == nullptr || !ScheduleUtils::IsLoad(load_node)) {
246+ continue;
247+ }
248+ if (IsReachableToCube(load_node)) {
249+ cube_load_nodes.emplace_back(load_node);
250+ } else {
251+ vector_load_nodes.emplace_back(load_node);
252+ }
253+ }
254+ if (cube_load_nodes.empty() || vector_load_nodes.empty()) {
255+ continue;
256+ }
257+ 
258+ const auto original_name = data_node->GetName();
259+ const auto split_name = GetUniqueSplitDataName(graph, original_name);
260+ af::AscNodePtr split_data_node;
261+ GE_CHK_STATUS_RET(CopyDataNodeForSharedSplit(graph, data_node, split_name, split_data_node));
262+ for (const auto &load_node : vector_load_nodes) {
263+ GE_ASSERT_NOTNULL(load_node);
264+ auto in_anchor = load_node->GetInDataAnchor(0);
265+ GE_CHECK_NOTNULL(in_anchor);
266+ auto src_anchor = data_node->GetOutDataAnchor(0);
267+ GE_CHECK_NOTNULL(src_anchor);
268+ auto new_src_anchor = split_data_node->GetOutDataAnchor(0);
269+ GE_CHECK_NOTNULL(new_src_anchor);
270+ GE_CHK_STATUS_RET(af::GraphUtils::ReplaceEdgeSrc(src_anchor, in_anchor, new_src_anchor));
271+ }
272+ renamed_data_nodes.emplace_back(split_name, original_name);
273+ }
274+ return af::SUCCESS;
275+}
276+ 
277+Status RestoreSplitDataNames(std::vector<ascir::ImplGraph> &grouped_graphs,
278+ const std::vector<std::pair<std::string, std::string>> &renamed_data_nodes) {
279+ for (auto &grouped_graph : grouped_graphs) {
280+ for (const auto &name_pair : renamed_data_nodes) {
281+ const auto &split_name = name_pair.first;
282+ const auto &original_name = name_pair.second;
283+ auto split_node = grouped_graph.FindNode(split_name.c_str());
284+ if (split_node == nullptr) {
285+ continue;
286+ }
287+ auto original_node = grouped_graph.FindNode(original_name.c_str());
288+ GE_ASSERT_TRUE(original_node == nullptr, "Graph %s already has Data node %s, cannot restore split Data name.",
289+ grouped_graph.GetName().c_str(), original_name.c_str());
290+ GE_CHECK_NOTNULL(split_node->GetOpDesc());
291+ split_node->GetOpDesc()->SetName(original_name.c_str());
292+ }
293+ }
294+ return af::SUCCESS;
295+}
296+ 
180Status GetPrioritySequence(const af::AscGraph &graph, std::unordered_set<af::Node *> &priority_sequences,297Status GetPrioritySequence(const af::AscGraph &graph, std::unordered_set<af::Node *> &priority_sequences,
181 std::unordered_set<af::Node *> &store_sequences) {298 std::unordered_set<af::Node *> &store_sequences) {
182 std::unordered_set<const af::Node *> visited;299 std::unordered_set<const af::Node *> visited;
@@ -369,6 +486,8 @@ Status CubeFusionCaseGenerator::GenerateGeneralCase(ascir::HintGraph &graph, std
369 }486 }
370 ascir::utils::DumpGraph(graph, "before_partition");487 ascir::utils::DumpGraph(graph, "before_partition");
371 ascir::utils::DumpGraph(optimize_graph, "after_partition");488 ascir::utils::DumpGraph(optimize_graph, "after_partition");
489+ split_data_names_.clear();
490+ GE_ASSERT_SUCCESS(SplitSharedDataForCubeAndVector(optimize_graph, split_data_names_));
372 GE_ASSERT_SUCCESS(FillInputDataAttrForCvGraph(optimize_graph));491 GE_ASSERT_SUCCESS(FillInputDataAttrForCvGraph(optimize_graph));
373 graphs.emplace_back(optimize_graph);492 graphs.emplace_back(optimize_graph);
374 return af::GRAPH_SUCCESS;493 return af::GRAPH_SUCCESS;
@@ -522,6 +641,7 @@ Status CubeFusionCaseGenerator::GeneratorTask(ascir::HintGraph &optimize_graph,
522 ScheduleTask task{graph, {}, score_funcs[i], {}, ReduceTemplateType::kDefault, ascir::CubeTemplateType::kFixpip};641 ScheduleTask task{graph, {}, score_funcs[i], {}, ReduceTemplateType::kDefault, ascir::CubeTemplateType::kFixpip};
523 GE_CHK_STATUS_RET(ScheduleGroupGraphPartitioner::PartitionByConnectivity(graph, task.grouped_graphs, node_order_),642 GE_CHK_STATUS_RET(ScheduleGroupGraphPartitioner::PartitionByConnectivity(graph, task.grouped_graphs, node_order_),
524 "Failed to partition graph");643 "Failed to partition graph");
644+ GE_CHK_STATUS_RET(RestoreSplitDataNames(task.grouped_graphs, split_data_names_), "Restore split Data names failed");
525 if (task.grouped_graphs.size() > 1U) {645 if (task.grouped_graphs.size() > 1U) {
526 task.cube_type = ascir::CubeTemplateType::kCommon;646 task.cube_type = ascir::CubeTemplateType::kCommon;
527 MoveCubeGraphsToEnd(task.grouped_graphs);647 MoveCubeGraphsToEnd(task.grouped_graphs);
@@ -11,6 +11,8 @@
11#ifndef ASCGEN_DEV_OPTIMIZE_TASK_GENERATOR_CUBE_SCHEDULE_CASE_GENERATOR_H_11#ifndef ASCGEN_DEV_OPTIMIZE_TASK_GENERATOR_CUBE_SCHEDULE_CASE_GENERATOR_H_
12#define ASCGEN_DEV_OPTIMIZE_TASK_GENERATOR_CUBE_SCHEDULE_CASE_GENERATOR_H_12#define ASCGEN_DEV_OPTIMIZE_TASK_GENERATOR_CUBE_SCHEDULE_CASE_GENERATOR_H_
13 13 
14+#include <utility>
15+ 
14#include "ascir_ops.h"16#include "ascir_ops.h"
15#include "ascir/meta/ascir.h"17#include "ascir/meta/ascir.h"
16#include "common/ascgen_log.h"18#include "common/ascgen_log.h"
@@ -32,6 +34,7 @@ class CubeFusionCaseGenerator : public FusionCaseGenerator {
32 static Status SwapCastBrcAndGenNddma(const af::AscNodePtr &node_cast, const af::AscNodePtr &node_load,34 static Status SwapCastBrcAndGenNddma(const af::AscNodePtr &node_cast, const af::AscNodePtr &node_load,
33 af::AscGraph &new_case);35 af::AscGraph &new_case);
34 std::vector<af::AscNodePtr> node_order_{};36 std::vector<af::AscNodePtr> node_order_{};
37+ std::vector<std::pair<std::string, std::string>> split_data_names_{};
35 bool partition_ = false;38 bool partition_ = false;
36};39};
37} // namespace optimize40} // namespace optimize