Line data Source code
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 : #include "graph/passes/variable_optimize/variable_op_pass.h"
12 : #include <cinttypes>
13 : #include <string>
14 : #include <vector>
15 :
16 : #include "formats/formats.h"
17 : #include "formats/utils/formats_trans_utils.h"
18 : #include "common/datatype_transfer/datatype_transfer.h"
19 : #include "common/checker.h"
20 : #include "graph/ge_context.h"
21 : #include "graph/graph.h"
22 : #include "graph/manager/graph_var_manager.h"
23 : #include "graph/utils/graph_utils.h"
24 : #include "graph/utils/tensor_utils.h"
25 : #include "graph/utils/type_utils.h"
26 : #include "common/plugin/ge_make_unique_util.h"
27 : #include "graph/utils/op_type_utils.h"
28 :
29 : namespace ge {
30 : namespace {
31 : const int32_t kTransOpOutIndex = 0;
32 : const std::unordered_set<std::string> kDataUnChangedNodeType = {RESHAPE, REFORMAT, SQUEEZEV2, UNSQUEEZEV2};
33 :
34 : std::string GetKey(Format format, DataType type, const std::vector<int64_t> &dims) {
35 : std::stringstream key;
36 : key << static_cast<int32_t>(format) << '-';
37 : key << static_cast<int32_t>(type) << '-';
38 : for (auto dim : dims) {
39 : key << dim << '-';
40 : }
41 : return key.str();
42 : }
43 :
44 : Status ByPassTransNode(NodePtr &trans_node, NodePtr &ref_node) {
45 : GE_CHECK_NOTNULL(trans_node);
46 : GE_CHECK_NOTNULL(ref_node);
47 : GELOGD("Begin to bypass trans node %s", trans_node->GetName().c_str());
48 : auto ret = GraphUtils::CopyInCtrlEdges(trans_node, ref_node);
49 : if (ret != GRAPH_SUCCESS) {
50 : REPORT_INNER_ERR_MSG("E19999", "Copy in control edge from node:%s(%s) to node:%s(%s) failed",
51 : trans_node->GetName().c_str(), trans_node->GetType().c_str(), ref_node->GetName().c_str(),
52 : ref_node->GetType().c_str());
53 : GELOGE(INTERNAL_ERROR, "[Copy][InCtrlEdges] from node:%s(%s) to node:%s(%s) failed", trans_node->GetName().c_str(),
54 : trans_node->GetType().c_str(), ref_node->GetName().c_str(), ref_node->GetType().c_str());
55 : return INTERNAL_ERROR;
56 : }
57 : auto ref_in_anchor = ref_node->GetInDataAnchor(0);
58 : if (ref_in_anchor == nullptr) {
59 : REPORT_INNER_ERR_MSG("E19999", "Node:%s(%s) has no input anchor, check invalid", ref_node->GetName().c_str(),
60 : ref_node->GetType().c_str());
61 : GELOGE(INTERNAL_ERROR, "[Get][InDataAnchor] failed, The variable ref node %s does not have an input anchor",
62 : ref_node->GetName().c_str());
63 : return INTERNAL_ERROR;
64 : }
65 : ref_in_anchor->UnlinkAll();
66 : auto trans_in_anchor = trans_node->GetInDataAnchor(0);
67 : if (trans_in_anchor == nullptr) {
68 : REPORT_INNER_ERR_MSG("E19999", "Node:%s(%s) has no input anchor, check invalid", trans_node->GetName().c_str(),
69 : trans_node->GetType().c_str());
70 : GELOGE(INTERNAL_ERROR, "[Get][InDataAnchor] failed, Node:%s(%s) has no input anchor", trans_node->GetName().c_str(),
71 : trans_node->GetType().c_str());
72 : return INTERNAL_ERROR;
73 : }
74 : auto prev_trans_node_out_anchor = trans_in_anchor->GetPeerOutAnchor();
75 : if (prev_trans_node_out_anchor == nullptr) {
76 : GELOGW(
77 : "The trans node %s does not have an input, so the ref node %s does"
78 : " not have any inputs after bypass",
79 : trans_node->GetName().c_str(), trans_node->GetName().c_str());
80 : } else {
81 : ret = GraphUtils::AddEdge(prev_trans_node_out_anchor, ref_in_anchor);
82 : if (ret != GRAPH_SUCCESS) {
83 : REPORT_INNER_ERR_MSG("E19999", "Add edge between op:%s(%s)(index:%d) and op:%s(%s)(index:0) failed",
84 : prev_trans_node_out_anchor->GetOwnerNode()->GetName().c_str(),
85 : prev_trans_node_out_anchor->GetOwnerNode()->GetType().c_str(),
86 : prev_trans_node_out_anchor->GetIdx(), ref_node->GetName().c_str(),
87 : ref_node->GetType().c_str());
88 : GELOGE(INTERNAL_ERROR, "[Add][Edge] between op:%s(%s)(index:%d) and op:%s(%s)(index:0) failed",
89 : prev_trans_node_out_anchor->GetOwnerNode()->GetName().c_str(),
90 : prev_trans_node_out_anchor->GetOwnerNode()->GetType().c_str(), prev_trans_node_out_anchor->GetIdx(),
91 : ref_node->GetName().c_str(), ref_node->GetType().c_str());
92 : return INTERNAL_ERROR;
93 : }
94 : }
95 : return SUCCESS;
96 : }
97 :
98 : bool IsTransSupport(const TransNodeInfo &trans_info) {
99 : if (trans_info.output.GetShape().IsUnknownShape()) {
100 : return false;
101 : }
102 : if (kDataUnChangedNodeType.count(trans_info.node_type) > 0U) {
103 : return true;
104 : } else if (trans_info.node_type == TRANSDATA || trans_info.node_type == TRANSPOSED) {
105 : const Format src_primary_format =
106 : static_cast<Format>(GetPrimaryFormat(static_cast<int32_t>(trans_info.input.GetFormat())));
107 : const Format dst_primary_format =
108 : static_cast<Format>(GetPrimaryFormat(static_cast<int32_t>(trans_info.output.GetFormat())));
109 : const Format src_sub_format = static_cast<Format>(GetSubFormat(static_cast<int32_t>(trans_info.input.GetFormat())));
110 : const Format dst_sub_format =
111 : static_cast<Format>(GetSubFormat(static_cast<int32_t>(trans_info.output.GetFormat())));
112 : const int64_t src_c0_format = GetC0Value(static_cast<int32_t>(trans_info.input.GetFormat()));
113 : const int64_t dts_c0_format = GetC0Value(static_cast<int32_t>(trans_info.output.GetFormat()));
114 : formats::TransArgs args{nullptr,
115 : trans_info.input.GetFormat(),
116 : trans_info.output.GetFormat(),
117 : src_primary_format,
118 : dst_primary_format,
119 : src_sub_format,
120 : dst_sub_format,
121 : src_c0_format,
122 : dts_c0_format,
123 : trans_info.input.GetShape().GetDims(),
124 : trans_info.output.GetShape().GetDims(),
125 : trans_info.input.GetDataType()};
126 : return formats::IsTransFormatSupport(args);
127 : } else if (trans_info.node_type == CAST) {
128 : formats::CastArgs datatype_args{nullptr, static_cast<size_t>(trans_info.input.GetShape().GetShapeSize()),
129 : trans_info.input.GetDataType(), trans_info.output.GetDataType()};
130 : return formats::IsTransDataTypeSupport(datatype_args);
131 : } else {
132 : return false;
133 : }
134 : }
135 : } // namespace
136 :
137 : Status VariableOpPass::Run(ge::ComputeGraphPtr graph) {
138 : if (graph == nullptr) {
139 : REPORT_INNER_ERR_MSG("E19999", "Param graph is nullptr, check invalid");
140 : GELOGE(INTERNAL_ERROR, "[Check][Param] Failed to run variable op pass, null graph");
141 : return INTERNAL_ERROR;
142 : }
143 : GE_ASSERT_NOTNULL(VarManager::Instance(graph->GetSessionID()));
144 : // In the multi-batch training scenario, multiple branches use and update the same variable weight,
145 : // so it cannot fuse variables and conversion operators based on format.
146 : GE_CHECK_NOTNULL(VarManager::Instance(graph->GetSessionID()));
147 : bool no_need_fusion = (graph->GetParentGraph() != nullptr) ||
148 : (VarManager::Instance(graph->GetSessionID())->HasSharedVarMemBetweenBatch());
149 : if (no_need_fusion) {
150 : return SUCCESS;
151 : }
152 :
153 : auto graph_id = graph->GetGraphID();
154 : GELOGD("Begin to run variable op pass on graph %s, session %" PRIu64 ", graph id %u", graph->GetName().c_str(),
155 : GetContext().SessionId(), graph_id);
156 :
157 : if (var_accelerate_ctrl_ == nullptr) {
158 : REPORT_INNER_ERR_MSG("E19999", "The variable accelerate control is nullptr, check invalid");
159 : GELOGE(INTERNAL_ERROR, "[Check][Param] Failed to run var op pass, the variable accelerate control is null");
160 : return INTERNAL_ERROR;
161 : }
162 :
163 : GELOGD("Begin to generate ref map for variable and refs, graph name:%s.", graph->GetName().c_str());
164 : if (RenewVarDesc(graph) != SUCCESS) {
165 : GELOGE(INTERNAL_ERROR, "[Renew][VarDesc] on graph:%s failed", graph->GetName().c_str());
166 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
167 : }
168 :
169 : if (GenerateVariableVariableRefMap(graph) != SUCCESS) {
170 : GELOGE(INTERNAL_ERROR, "[Generate][VariableMap] for graph:%s failed", graph->GetName().c_str());
171 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
172 : }
173 :
174 : GELOGD("Begin to fusion variables and trans nodes");
175 : for (auto &var_to_refs : var_and_var_ref_map_) {
176 : GE_CHECK_NOTNULL(var_accelerate_ctrl_);
177 : if (!var_accelerate_ctrl_->IsVarPermitToChangeFormats(var_to_refs.first->var_name)) {
178 : GELOGD("The var %s does not permit to change formats, skip it", var_to_refs.first->var_name.c_str());
179 : continue;
180 : }
181 :
182 : VarTransRoad fusion_road;
183 : auto ret = FusionIfNeed(var_to_refs.first, fusion_road);
184 : if (ret != SUCCESS) {
185 : GELOGE(FAILED, "[Call][FusionIfNeed] for node:%s failed", var_to_refs.first->var_name.c_str());
186 : return ret;
187 : }
188 :
189 : if (fusion_road.empty()) {
190 : GELOGD("No need to fusion variable and trans op for var %s", var_to_refs.first->var_name.c_str());
191 : continue;
192 : }
193 :
194 : auto start_iter = fusion_road.begin();
195 : auto end_iter = fusion_road.rbegin();
196 : GELOGI(
197 : "Trans variable data for %s from format %s to %s, shape %s to %s "
198 : "data-type %s to %s, path len %zu success",
199 : var_to_refs.first->var_name.c_str(), TypeUtils::FormatToSerialString(start_iter->input.GetFormat()).c_str(),
200 : TypeUtils::FormatToSerialString(end_iter->output.GetFormat()).c_str(),
201 : formats::ShapeToString(start_iter->input.GetShape().GetDims()).c_str(),
202 : formats::ShapeToString(end_iter->output.GetShape().GetDims()).c_str(),
203 : TypeUtils::DataTypeToSerialString(start_iter->input.GetDataType()).c_str(),
204 : TypeUtils::DataTypeToSerialString(end_iter->output.GetDataType()).c_str(), fusion_road.size());
205 :
206 : ret = VarManager::Instance(graph->GetSessionID())->SetTransRoad(var_to_refs.first->var_name, fusion_road);
207 : if (ret != SUCCESS) {
208 : REPORT_INNER_ERR_MSG("E19999", "Set Trans road for node:%s(Variable) failed, session_id:%" PRIu64 "",
209 : var_to_refs.first->var_name.c_str(), graph->GetSessionID());
210 : GELOGE(INTERNAL_ERROR, "[Set][TransRoad] for node:%s(Variable) failed, session_id:%" PRIu64 "",
211 : var_to_refs.first->var_name.c_str(), graph->GetSessionID());
212 : return INTERNAL_ERROR;
213 : }
214 : ret = VarManager::Instance(graph->GetSessionID())->SetChangedGraphId(var_to_refs.first->var_name, graph_id);
215 : if (ret != SUCCESS) {
216 : REPORT_INNER_ERR_MSG("E19999", "Update graph_id:%u for node:%s(Variable) failed, session_id:%" PRIu64 "",
217 : graph_id, var_to_refs.first->var_name.c_str(), graph->GetSessionID());
218 : GELOGE(INTERNAL_ERROR, "[Update][GraphId] %u for node:%s(Variable) failed, session_id:%" PRIu64 "", graph_id,
219 : var_to_refs.first->var_name.c_str(), graph->GetSessionID());
220 : return INTERNAL_ERROR;
221 : }
222 : var_accelerate_ctrl_->SetStateChanged(var_to_refs.first->var_name);
223 :
224 : GELOGD("Begin to update format info for var %s.", var_to_refs.first->var_name.c_str());
225 : if (UpdateIOFormatInfo(end_iter->output, var_to_refs.first) != SUCCESS) {
226 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
227 : }
228 :
229 : for (const auto &node : var_to_refs.first->var_nodes) {
230 : // renew var desc if the trans_road is all reshape or reformat
231 : ret = RenewVarDesc(graph->GetSessionID(), node, fusion_road);
232 : if (ret != SUCCESS) {
233 : GELOGE(FAILED, "[Renew][VarDesc] for var[%s] failed!", node->GetName().c_str());
234 : return FAILED;
235 : }
236 : }
237 : }
238 :
239 : return SUCCESS;
240 : }
241 :
242 : Status VariableOpPass::DealFusion(const SameVarPtr &same_vars) {
243 : for (const auto &var_node : same_vars->var_nodes) {
244 : GE_CHECK_NOTNULL(var_node);
245 : GELOGD("Begin to fusion var %s with trans", var_node->GetName().c_str());
246 : auto graph = var_node->GetOwnerComputeGraph();
247 : for (auto &trans_node : var_node->GetOutDataNodes()) {
248 : GELOGD("Remove node %s type %s when fusion with variable %s", trans_node->GetName().c_str(),
249 : trans_node->GetType().c_str(), var_node->GetName().c_str());
250 :
251 : if (GraphUtils::IsolateNode(trans_node, {0}) != SUCCESS) {
252 : REPORT_INNER_ERR_MSG("E19999", "Isolate node:%s(%s) failed", trans_node->GetName().c_str(),
253 : trans_node->GetType().c_str());
254 : GELOGE(GE_GRAPH_VARIABLE_OP_PASS_FAILED, "[Isolate][Node] %s(%s) failed", trans_node->GetName().c_str(),
255 : trans_node->GetType().c_str());
256 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
257 : }
258 :
259 : if (GraphUtils::RemoveNodeWithoutRelink(graph, trans_node) != SUCCESS) {
260 : REPORT_INNER_ERR_MSG("E19999", "Remove node:%s(%s) without relink in graph:%s failed",
261 : trans_node->GetName().c_str(), trans_node->GetType().c_str(), graph->GetName().c_str());
262 : GELOGE(GE_GRAPH_VARIABLE_OP_PASS_FAILED, "[Remove][Node] %s(%s) without relink in graph:%s failed",
263 : trans_node->GetName().c_str(), trans_node->GetType().c_str(), graph->GetName().c_str());
264 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
265 : }
266 : }
267 : }
268 :
269 : for (auto ref_node : GetRefVars(same_vars)) {
270 : GE_CHECK_NOTNULL(ref_node);
271 : for (auto &trans_node : ref_node->GetInDataNodes()) {
272 : GELOGD("Remove node %s type %s when fusion with variable %s", trans_node->GetName().c_str(),
273 : trans_node->GetType().c_str(), same_vars->var_name.c_str());
274 : auto graph = trans_node->GetOwnerComputeGraph();
275 : if (trans_node->GetOutDataNodes().size() > 1) {
276 : GELOGD(
277 : "The trans node %s type %s connecting with var-ref %s has more"
278 : " than one output data nodes, unlink the edge between them",
279 : trans_node->GetName().c_str(), trans_node->GetType().c_str(), ref_node->GetName().c_str());
280 : if (ByPassTransNode(trans_node, ref_node) != SUCCESS) {
281 : GELOGE(INTERNAL_ERROR, "[ByPass][TransNode] %s to ref %s failed", trans_node->GetName().c_str(),
282 : ref_node->GetName().c_str());
283 : return INTERNAL_ERROR;
284 : }
285 : } else {
286 : GELOGD(
287 : "The trans node %s type %s connecting with var-ref %s has only"
288 : " one output data nodes, isolate and remove it.",
289 : trans_node->GetName().c_str(), trans_node->GetType().c_str(), ref_node->GetName().c_str());
290 : if (GraphUtils::IsolateNode(trans_node, {0}) != SUCCESS) {
291 : REPORT_INNER_ERR_MSG("E19999", "Isolate node:%s(%s) failed", trans_node->GetName().c_str(),
292 : trans_node->GetType().c_str());
293 : GELOGE(GE_GRAPH_VARIABLE_OP_PASS_FAILED, "[Isolate][Node] %s(%s) failed", trans_node->GetName().c_str(),
294 : trans_node->GetType().c_str());
295 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
296 : }
297 : if (GraphUtils::RemoveNodeWithoutRelink(graph, trans_node) != SUCCESS) {
298 : REPORT_INNER_ERR_MSG("E19999", "Remove node:%s(%s) without relink in graph:%s failed",
299 : trans_node->GetName().c_str(), trans_node->GetType().c_str(), graph->GetName().c_str());
300 : GELOGE(GE_GRAPH_VARIABLE_OP_PASS_FAILED, "[Remove][Node] %s(%s) without relink in graph:%s failed",
301 : trans_node->GetName().c_str(), trans_node->GetType().c_str(), graph->GetName().c_str());
302 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
303 : }
304 : }
305 : }
306 : }
307 :
308 : return SUCCESS;
309 : }
310 :
311 : Status VariableOpPass::CheckSameAndTransOp(const SameVarPtr &same_vars, bool &is_matched,
312 : VarTransRoad &fusion_road) const {
313 : std::set<std::string> data_type_and_formats;
314 : std::string trans_op_type;
315 : ge::NodePtr out_node;
316 : ge::GeTensorDesc output_desc;
317 : for (const auto &var_node : same_vars->var_nodes) {
318 : GE_CHECK_NOTNULL(var_node);
319 : for (const auto &out_node_and_anchor : var_node->GetOutDataNodesAndAnchors()) {
320 : auto in_anchor = out_node_and_anchor.second;
321 : GE_CHECK_NOTNULL(in_anchor);
322 : out_node = out_node_and_anchor.first;
323 : GE_CHECK_NOTNULL(out_node);
324 : auto trans_op_desc = out_node->GetOpDesc();
325 : GE_CHECK_NOTNULL(trans_op_desc);
326 : trans_op_type = trans_op_desc->GetType();
327 :
328 : GELOGD("current node type is %s.", trans_op_type.c_str());
329 : int32_t data_index = TransOpUtil::GetTransOpDataIndex(trans_op_type);
330 : if (data_index < 0) {
331 : GELOGD("Variables only can be fusion with trans_op, the next op is %s type %s", out_node->GetName().c_str(),
332 : out_node->GetType().c_str());
333 : return SUCCESS;
334 : }
335 : if (data_index != in_anchor->GetIdx()) {
336 : GELOGD(
337 : "Variables only can be fusion with trans nodes, the next node %s"
338 : " type %s index %d does not trans anything(correct index %d)",
339 : out_node->GetName().c_str(), out_node->GetType().c_str(), in_anchor->GetIdx(), data_index);
340 : return SUCCESS;
341 : }
342 :
343 : output_desc = trans_op_desc->GetOutputDesc(kTransOpOutIndex);
344 :
345 : auto trans_op_format = output_desc.GetFormat();
346 : auto trans_op_data_type = output_desc.GetDataType();
347 : auto shape = output_desc.GetShape().GetDims();
348 : auto datatype_and_format = GetKey(trans_op_format, trans_op_data_type, shape);
349 : data_type_and_formats.insert(datatype_and_format);
350 : }
351 : }
352 :
353 : if (data_type_and_formats.empty()) {
354 : return SUCCESS;
355 : }
356 :
357 : if (data_type_and_formats.size() > 1UL) {
358 : std::stringstream type_and_formats_stream;
359 : bool first_time = true;
360 : for (const auto &data_type_and_format : data_type_and_formats) {
361 : if (first_time) {
362 : first_time = false;
363 : } else {
364 : type_and_formats_stream << "|";
365 : }
366 : type_and_formats_stream << data_type_and_format;
367 : }
368 :
369 : GELOGW(
370 : "trans_op type size for var Node(%s) is over 1, Currently not"
371 : " supported, dataTypeAndFormats is %s.",
372 : same_vars->var_name.c_str(), type_and_formats_stream.str().c_str());
373 : return SUCCESS;
374 : }
375 :
376 : GE_ASSERT_NOTNULL(out_node);
377 : int32_t tran_in_index = TransOpUtil::GetTransOpDataIndex(out_node->GetType());
378 : auto out_op_desc = out_node->GetOpDesc();
379 : GE_CHECK_NOTNULL(out_op_desc);
380 : TransNodeInfo trans_node_info;
381 : trans_node_info.node_type = out_node->GetType();
382 : trans_node_info.input = out_op_desc->GetInputDesc(tran_in_index);
383 : trans_node_info.output = out_op_desc->GetOutputDesc(kTransOpOutIndex);
384 :
385 : if (!IsTransSupport(trans_node_info)) {
386 : GELOGD("The trans node %s does not support, skip the variable accelerating", trans_node_info.node_type.c_str());
387 : return SUCCESS;
388 : }
389 :
390 : is_matched = true;
391 : fusion_road.emplace_back(trans_node_info);
392 :
393 : return SUCCESS;
394 : }
395 :
396 : Status VariableOpPass::CheckVariableRefLegally(const SameVarPtr &same_vars, bool &is_var_ref_legally) {
397 : is_var_ref_legally = true;
398 : auto var_ref_nodes = GetRefVars(same_vars);
399 : GELOGD("var name %s, ref var count %zu.", same_vars->var_name.c_str(), var_ref_nodes.size());
400 : for (const auto &var_node : same_vars->var_nodes) {
401 : GE_CHECK_NOTNULL(var_node);
402 : for (const auto &var_ref_node : var_ref_nodes) {
403 : if (CheckVarAndVarRefAreAlike(var_node, var_ref_node, is_var_ref_legally) != SUCCESS) {
404 : GELOGE(FAILED, "[Call][CheckVarAndVarRefAreAlike] for node:%s failed", var_node->GetName().c_str());
405 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
406 : }
407 :
408 : GELOGD("var name %s, is_var_ref_legally is %d", same_vars->var_name.c_str(), is_var_ref_legally);
409 :
410 : if (!is_var_ref_legally) {
411 : return SUCCESS;
412 : }
413 : }
414 : }
415 : return SUCCESS;
416 : }
417 :
418 : Status VariableOpPass::UpdateVarAndRefOutputFormatInfo(const GeTensorDesc &final_output, const ge::NodePtr &node,
419 : const SameVarPtr &same_vars) {
420 : if (node == nullptr || node->GetOpDesc() == nullptr) {
421 : REPORT_INNER_ERR_MSG("E19999", "Param node or its op_desc is nullptr, check invalid");
422 : GELOGE(FAILED, "[Check][Param] node or its opdesc is nullptr");
423 : return FAILED;
424 : }
425 : const Format &format = final_output.GetFormat();
426 : const DataType &data_type = final_output.GetDataType();
427 : const GeShape &shape = final_output.GetShape();
428 : GELOGD("last ref is (%s, %s, %" PRIu64 "), var_ref_name is %s.", TypeUtils::DataTypeToSerialString(data_type).c_str(),
429 : TypeUtils::FormatToSerialString(format).c_str(), shape.GetDims().size(), node->GetName().c_str());
430 :
431 : auto node_desc = node->GetOpDesc()->GetOutputDesc(0);
432 : CopyVariableFormatDataTypeAndShape(final_output, node_desc);
433 : if (node->GetOpDesc()->UpdateOutputDesc(0, node_desc) != GRAPH_SUCCESS) {
434 : REPORT_INNER_ERR_MSG("E19999", "Update output:0 desc in op:%s(%s) failed.", node->GetName().c_str(),
435 : node->GetType().c_str());
436 : GELOGE(FAILED, "[Update][OutputDesc] in op:%s(%s) failed, index:0", node->GetName().c_str(),
437 : node->GetType().c_str());
438 : return FAILED;
439 : }
440 : GELOGD("node ref is (%s, %s, %" PRIu64 "), var_ref_name is %s.",
441 : TypeUtils::DataTypeToSerialString(node->GetOpDesc()->GetOutputDesc(0).GetDataType()).c_str(),
442 : TypeUtils::FormatToSerialString(node->GetOpDesc()->GetOutputDesc(0).GetFormat()).c_str(),
443 : node->GetOpDesc()->GetOutputDesc(0).GetShape().GetDims().size(), node->GetName().c_str());
444 :
445 : auto iterator = var_and_var_ref_map_.find(same_vars);
446 : if (iterator == var_and_var_ref_map_.end()) {
447 : auto graph = node->GetOwnerComputeGraph();
448 : if (GenerateVariableVariableRefMap(graph) != SUCCESS) {
449 : GELOGE(INTERNAL_ERROR, "[Generate][VariableMap] for graph:%s failed", graph->GetName().c_str());
450 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
451 : }
452 : }
453 :
454 : iterator = var_and_var_ref_map_.find(same_vars);
455 : if (iterator != var_and_var_ref_map_.end()) {
456 : for (const auto &var_ref_node : iterator->second) {
457 : const auto var_ref_node_description = var_ref_node->GetOpDesc();
458 : GE_CHECK_NOTNULL(var_ref_node_description);
459 :
460 : GELOGD("var_ref_node before is (%s, %s, %zu), var_ref_name is %s.",
461 : TypeUtils::DataTypeToSerialString(data_type).c_str(), TypeUtils::FormatToSerialString(format).c_str(),
462 : shape.GetDims().size(), var_ref_node->GetName().c_str());
463 : if (var_ref_node_description->UpdateOutputDesc(0U, node_desc) != GRAPH_SUCCESS) {
464 : GELOGW("UpdateOutputDesc fail.");
465 : }
466 : if (var_ref_node_description->UpdateInputDesc(0U, node_desc) != GRAPH_SUCCESS) {
467 : GELOGW("UpdateInputDesc fail.");
468 : }
469 : const auto &input_desc = var_ref_node_description->MutableInputDesc(0U);
470 : const auto &output_desc = var_ref_node_description->MutableOutputDesc(0U);
471 : GE_CHECK_NOTNULL(input_desc);
472 : GE_CHECK_NOTNULL(output_desc);
473 : GELOGD("var_ref_node ref is (%s, %s, %zu), var_ref_name is %s.",
474 : TypeUtils::DataTypeToSerialString(input_desc->GetDataType()).c_str(),
475 : TypeUtils::FormatToSerialString(input_desc->GetFormat()).c_str(), output_desc->GetShape().GetDims().size(),
476 : var_ref_node->GetName().c_str());
477 : }
478 : }
479 :
480 : return SUCCESS;
481 : }
482 :
483 : Status VariableOpPass::GenerateVariableVariableRefMap(const ComputeGraphPtr &compute_graph) {
484 : std::map<std::string, std::set<NodePtr>> names_to_var;
485 : std::map<std::string, std::set<NodePtr>> names_to_refs;
486 : GE_CHECK_NOTNULL(compute_graph);
487 : for (const auto &node : compute_graph->GetAllNodes()) {
488 : if (!OpTypeUtils::IsVariableNode(node->GetType())) {
489 : continue;
490 : }
491 : GE_CHECK_NOTNULL(node->GetOpDesc());
492 : std::string ref_var_name;
493 : if (ge::AttrUtils::GetStr(node->GetOpDesc(), REF_VAR_SRC_VAR_NAME, ref_var_name)) {
494 : names_to_refs[ref_var_name].insert(node);
495 : } else {
496 : names_to_var[node->GetName()].insert(node);
497 : }
498 : }
499 :
500 : for (const auto &name_to_var : names_to_var) {
501 : SameVarPtr same_vars = MakeShared<SameVariable>();
502 : GE_CHECK_NOTNULL(same_vars);
503 : same_vars->var_name = name_to_var.first;
504 : same_vars->var_nodes = name_to_var.second;
505 : var_and_var_ref_map_[same_vars] = names_to_refs[name_to_var.first];
506 : }
507 : return SUCCESS;
508 : }
509 :
510 : Status VariableOpPass::CheckVarAndVarRefAreAlike(const NodePtr &var_node, const NodePtr &var_ref_node,
511 : bool &is_var_and_variable_ref_are_alike) const {
512 : GE_CHECK_NOTNULL(var_node);
513 : GE_CHECK_NOTNULL(var_ref_node);
514 : GELOGD("var_node GetOutDataNodes. name is %s.", var_node->GetName().c_str());
515 : const auto &var_node_trans_nodes = var_node->GetOutDataNodes();
516 : GELOGD("var_node_trans_nodes size is %zu.", var_node_trans_nodes.size());
517 : GELOGD("var_ref_node GetOutDataNodes. name is %s.", var_ref_node->GetName().c_str());
518 : const auto &var_ref_node_trans_nodes = var_ref_node->GetInDataNodes();
519 : GELOGD("var_ref_node_trans_nodes size is %zu.", var_ref_node_trans_nodes.size());
520 :
521 : if (var_ref_node_trans_nodes.size() > 1) {
522 : REPORT_INNER_ERR_MSG("E19999", "In data node num:%zu of node:%s(%s) bigger than 1, check invalid",
523 : var_ref_node_trans_nodes.size(), var_ref_node->GetName().c_str(),
524 : var_ref_node->GetType().c_str());
525 :
526 : GELOGE(GE_GRAPH_VARIABLE_OP_PASS_FAILED, "[Check][Param] In data node num:%zu of node:%s(%s) bigger than 1.",
527 : var_ref_node_trans_nodes.size(), var_ref_node->GetName().c_str(), var_ref_node->GetType().c_str());
528 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
529 : }
530 :
531 : const auto &var_node_trans_node = var_node_trans_nodes.at(0);
532 : const auto &var_ref_node_trans_node = var_ref_node_trans_nodes.at(0);
533 :
534 : if (CheckTransNodeAreInverse(var_node_trans_node, var_ref_node_trans_node, is_var_and_variable_ref_are_alike) !=
535 : SUCCESS) {
536 : GELOGE(FAILED, "[Call][CheckTransNodeAreInverse] failed");
537 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
538 : }
539 :
540 : return SUCCESS;
541 : }
542 :
543 : Status VariableOpPass::CheckTransNodeAreInverse(const NodePtr &node_a, const NodePtr &node_b, bool &is_same) const {
544 : GELOGD("In CheckTransNodeAreInverse.");
545 : GE_CHECK_NOTNULL(node_a);
546 : GE_CHECK_NOTNULL(node_b);
547 : const auto &node_a_op_desc = node_a->GetOpDesc();
548 : const auto &node_b_op_desc = node_b->GetOpDesc();
549 : GE_CHECK_NOTNULL(node_a_op_desc);
550 : GE_CHECK_NOTNULL(node_b_op_desc);
551 : const auto &node_a_out_op_desc = node_a_op_desc->MutableOutputDesc(0);
552 : const auto &node_a_in_op_desc = node_a_op_desc->MutableInputDesc(0);
553 : GE_CHECK_NOTNULL(node_a_out_op_desc);
554 : GE_CHECK_NOTNULL(node_a_in_op_desc);
555 :
556 : const auto &node_b_out_op_desc = node_b_op_desc->MutableOutputDesc(0);
557 : const auto &node_b_in_op_desc = node_b_op_desc->MutableInputDesc(0);
558 : GE_CHECK_NOTNULL(node_b_out_op_desc);
559 : GE_CHECK_NOTNULL(node_b_in_op_desc);
560 :
561 : is_same = IsOpDescSame(node_a_out_op_desc, node_b_in_op_desc) && IsOpDescSame(node_b_out_op_desc, node_a_in_op_desc);
562 :
563 : return SUCCESS;
564 : }
565 :
566 : bool VariableOpPass::IsOpDescSame(const GeTensorDescPtr &op_desc_a, const GeTensorDescPtr &op_desc_b) const {
567 : const auto &format_a = op_desc_a->GetFormat();
568 : const auto &type_a = op_desc_a->GetDataType();
569 : const auto &shape_a = op_desc_a->GetShape();
570 :
571 : const auto &format_b = op_desc_b->GetFormat();
572 : const auto &type_b = op_desc_b->GetDataType();
573 : const auto &shape_b = op_desc_b->GetShape();
574 :
575 : const auto &dims_a = shape_a.GetDims();
576 : const auto &dims_b = shape_b.GetDims();
577 : GELOGD("(format, data type, shape) = (%s, %s, %zu) (%s, %s, %zu)", TypeUtils::FormatToSerialString(format_a).c_str(),
578 : TypeUtils::DataTypeToSerialString(type_a).c_str(), dims_a.size(),
579 : TypeUtils::FormatToSerialString(format_b).c_str(), TypeUtils::DataTypeToSerialString(type_b).c_str(),
580 : dims_b.size());
581 : return (format_a == format_b) && (type_a == type_b) && (dims_a == dims_b);
582 : }
583 :
584 : void VariableOpPass::CopyVariableFormatDataTypeAndShape(const GeTensorDesc &src_tensor_desc,
585 : GeTensorDesc &dst_tensor_desc) const {
586 : dst_tensor_desc.SetOriginShape(src_tensor_desc.GetOriginShape());
587 : dst_tensor_desc.SetShape(src_tensor_desc.GetShape());
588 : dst_tensor_desc.SetFormat(src_tensor_desc.GetFormat());
589 : dst_tensor_desc.SetDataType(src_tensor_desc.GetDataType());
590 : }
591 :
592 : Status VariableOpPass::CheckIfCouldBeOptimized(const SameVarPtr &same_vars, bool &flag, VarTransRoad &fusion_road) {
593 : bool is_matched = false;
594 : auto ret = CheckSameAndTransOp(same_vars, is_matched, fusion_road);
595 : if (ret != SUCCESS) {
596 : GELOGE(FAILED, "[Call][CheckSameAndTransOp] failed, node:%s", same_vars->var_name.c_str());
597 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
598 : }
599 : if (!is_matched) {
600 : flag = false;
601 : return SUCCESS;
602 : }
603 :
604 : bool is_var_ref_legally = false;
605 : ret = CheckVariableRefLegally(same_vars, is_var_ref_legally);
606 : if (ret != SUCCESS) {
607 : GELOGE(FAILED, "[Call][CheckVariableRefLegally] failed, node:%s", same_vars->var_name.c_str());
608 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
609 : }
610 : GELOGD("is_var_ref_legally is %d.", is_var_ref_legally);
611 : if (!is_var_ref_legally) {
612 : GELOGI("variable ref connection are illegally");
613 : flag = false;
614 : fusion_road.clear();
615 : return SUCCESS;
616 : }
617 :
618 : flag = true;
619 : GELOGD("node %s, is_matched = %d is_var_ref_legally = %d, flag = %d", same_vars->var_name.c_str(), is_matched,
620 : is_var_ref_legally, flag);
621 :
622 : return SUCCESS;
623 : }
624 :
625 : Status VariableOpPass::FusionIfNeed(const SameVarPtr &same_vars, VarTransRoad &fusion_road) {
626 : bool can_fusion = false;
627 : bool first_time = true;
628 : while (true) {
629 : auto ret = CheckIfCouldBeOptimized(same_vars, can_fusion, fusion_road);
630 : if (ret != SUCCESS) {
631 : GELOGE(FAILED, "[Call][CheckIfCouldBeOptimized] failed");
632 : return ret;
633 : }
634 : if (!can_fusion) {
635 : break;
636 : }
637 : // In original graph where the variable output format and the output trans node input format are not
638 : // continuous, you need to insert ReFormat on trans_road. And only need to make this judgment in the first loop,
639 : // the trans node will be deleted in each loop, and the output format of the variable will
640 : // not change, so there will be scenarios where variables and trans node formats are not continuous in next loop.
641 : if (first_time) {
642 : const NodePtr var_node = *(same_vars->var_nodes.cbegin());
643 : auto var_op_desc = var_node->GetOpDesc();
644 : GE_CHECK_NOTNULL(var_op_desc);
645 : ge::GeTensorDesc var_output_desc = var_op_desc->GetOutputDesc(0);
646 : if (var_output_desc.GetFormat() != (fusion_road.begin()->input).GetFormat()) {
647 : TransNodeInfo additional_trans_node_info;
648 : additional_trans_node_info.node_type = REFORMAT;
649 : additional_trans_node_info.input = var_output_desc;
650 : additional_trans_node_info.output = fusion_road.begin()->input;
651 : fusion_road.insert(fusion_road.cbegin(), additional_trans_node_info);
652 : }
653 : first_time = false;
654 : }
655 : ret = DealFusion(same_vars);
656 : if (ret != SUCCESS) {
657 : GELOGE(FAILED, "[Call][DealFusion] failed");
658 : return ret;
659 : }
660 : }
661 : return SUCCESS;
662 : }
663 :
664 : Status VariableOpPass::UpdateIOFormatInfo(const GeTensorDesc &final_output, const SameVarPtr &same_vars) {
665 : for (const auto &need_set_node : same_vars->var_nodes) {
666 : auto ret = UpdateVarAndRefOutputFormatInfo(final_output, need_set_node, same_vars);
667 : if (ret != SUCCESS) {
668 : GELOGE(FAILED, "[Call][UpdateVarAndRefOutputFormatInfo] failed");
669 : return GE_GRAPH_VARIABLE_OP_PASS_FAILED;
670 : }
671 : }
672 : return SUCCESS;
673 : }
674 :
675 : Status VariableOpPass::RenewVarDesc(const ge::ComputeGraphPtr &graph) const {
676 : GE_CHECK_NOTNULL(graph);
677 : GE_ASSERT_NOTNULL(ge::VarManager::Instance(graph->GetSessionID()));
678 : // renew var manager desc
679 : auto graph_id = graph->GetGraphID();
680 : Status ret = SUCCESS;
681 : for (auto &node : graph->GetAllNodes()) {
682 : if (OpTypeUtils::IsVariableNode(node->GetType())) {
683 : auto var_manager = ge::VarManager::Instance(graph->GetSessionID());
684 : GE_CHECK_NOTNULL(var_manager);
685 : if (!var_manager->IsVarExist(node->GetName())) {
686 : GELOGD("var manager does not exist var node[%s]", node->GetName().c_str());
687 : continue;
688 : }
689 : GELOGD("var manager exist var node[%s], graph name[%s]", node->GetName().c_str(), graph->GetName().c_str());
690 : GE_CHECK_NOTNULL(node->GetOpDesc());
691 : ret = var_manager->RenewCurVarDesc(node->GetName(), node->GetOpDesc());
692 : if (ret != SUCCESS) {
693 : REPORT_INNER_ERR_MSG("E19999", "Renew descriptor for node:%s(%s) failed, session_id:%" PRIu64 "",
694 : node->GetName().c_str(), node->GetType().c_str(), graph->GetSessionID());
695 : GELOGE(FAILED, "[Renew][Descriptor] for node:%s(%s) failed, session_id:%" PRIu64 "", node->GetName().c_str(),
696 : node->GetType().c_str(), graph->GetSessionID());
697 : return FAILED;
698 : }
699 :
700 : ret = var_manager->RecordStagedVarDesc(graph_id, node->GetName(), node->GetOpDesc()->GetOutputDesc(0U));
701 : if (ret != SUCCESS) {
702 : REPORT_INNER_ERR_MSG("E19999", "Record staged descriptor for node:%s(%s) failed, session_id:%" PRIu64 "",
703 : node->GetName().c_str(), node->GetType().c_str(), graph->GetSessionID());
704 : GELOGE(FAILED, "[Record][Descriptor] for node:%s(%s) failed, session_id:%" PRIu64 "", node->GetName().c_str(),
705 : node->GetType().c_str(), graph->GetSessionID());
706 : return FAILED;
707 : }
708 : }
709 : }
710 : return SUCCESS;
711 : }
712 :
713 : Status VariableOpPass::RenewVarDesc(uint64_t session_id, const NodePtr &node, const VarTransRoad &fusion_road) const {
714 : // renew var desc if the trans_road is all reshape or reformat
715 : for (const auto &road : fusion_road) {
716 : if (kDataUnChangedNodeType.count(road.node_type) == 0) {
717 : return SUCCESS;
718 : }
719 : }
720 : GE_ASSERT_NOTNULL(ge::VarManager::Instance(session_id));
721 : if (!ge::VarManager::Instance(session_id)->IsVarExist(node->GetName())) {
722 : GELOGD("var manager does not exist var node[%s]", node->GetName().c_str());
723 : return SUCCESS;
724 : }
725 : GELOGD("var manager exist var node[%s]", node->GetName().c_str());
726 : GE_CHECK_NOTNULL(node->GetOpDesc());
727 : Status ret = ge::VarManager::Instance(session_id)->RenewCurVarDesc(node->GetName(), node->GetOpDesc());
728 : if (ret != SUCCESS) {
729 : REPORT_INNER_ERR_MSG("E19999", "Renew descriptor for node:%s(%s) failed, session_id:%" PRIu64 "",
730 : node->GetName().c_str(), node->GetType().c_str(), session_id);
731 : GELOGE(FAILED, "[Renew][Descriptor] for node:%s(%s) failed, session_id:%" PRIu64 "", node->GetName().c_str(),
732 : node->GetType().c_str(), session_id);
733 : return FAILED;
734 : }
735 :
736 : return SUCCESS;
737 : }
738 :
739 31 : std::vector<NodePtr> VariableOpPass::GetRefVars(const SameVarPtr &same_vars) const {
740 : std::vector<NodePtr> nodes;
741 : auto iter = var_and_var_ref_map_.find(same_vars);
742 : if (iter != var_and_var_ref_map_.end()) {
743 : for (const auto &node : iter->second) {
744 : if ((node->GetOpDesc() != nullptr) && AttrUtils::HasAttr(node->GetOpDesc(), REF_VAR_SRC_VAR_NAME)) {
745 : nodes.emplace_back(node);
746 : }
747 : }
748 : }
749 : return nodes;
750 : }
751 :
752 : REG_PASS_OPTION("VariableOpPass").LEVELS(OoLevel::kO3);
753 : } // namespace ge
|