已合并
fix: split shared data for matmul vector partition #1430
ling-DT创建于 7月23日
fix: split shared data for matmul vector partition #1430
已合并
共 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 | + | ||
| 39 | void ConstructJustMatMul(af::AscGraph &graph) { | 109 | void 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 | ||
| 110 | void ConstructMatMulAndAdd(af::AscGraph &graph) { | 180 | void 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; | ||
| 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 | ||
| 200 | void ConstructJustMatMulBias(af::AscGraph &graph) { | 277 | void 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 | + | ||
| 312 | TEST_F(CubeScheduleCaseGeneratorTest, Test_MatMul_Bias_Store) { | 420 | TEST_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 | 11 | ||
| 12 | 12 | ||
| 13 | + | ||
| 13 | 14 | ||
| 14 | 15 | ||
| 15 | 16 | ||
| @@ -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 | + | ||
| 180 | Status GetPrioritySequence(const af::AscGraph &graph, std::unordered_set<af::Node *> &priority_sequences, | 297 | Status 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 | 11 | ||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | + | ||
| 15 | + | ||
| 14 | 16 | ||
| 15 | 17 | ||
| 16 | 18 | ||
| @@ -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 optimize | 40 | } // namespace optimize |
🟡 Medium Priority
在
CountDataNodeByNameAndIndex函数(测试辅助函数)的第 206 行,代码直接对node->attr.ir_attr调用->DownCastTo<>()而不先检查ir_attr是否为 null。ir_attr的类型是std::unique_ptr<AscIrAttrDefBase>(定义于ascendc_ir_def.h:370),默认初始化为nullptr。对 nullunique_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。