mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
remove const cast from conv related fusion.
This commit is contained in:
parent
728c3c078e
commit
35f94cac27
4 changed files with 11 additions and 7 deletions
|
|
@ -130,6 +130,11 @@ class Node {
|
|||
return definitions_.input_defs;
|
||||
}
|
||||
|
||||
/** Gets a modifiable collection of the Node's input definitions. */
|
||||
std::vector<NodeArg*>& MutableOutputDefs() noexcept {
|
||||
return definitions_.output_defs;
|
||||
}
|
||||
|
||||
/** Gets the count of arguments for each of the Node's explicit inputs. */
|
||||
const std::vector<int>& InputArgCount() const noexcept { return definitions_.input_arg_count; }
|
||||
|
||||
|
|
|
|||
|
|
@ -93,17 +93,16 @@ Status ConvAddFusion::Apply(onnxruntime::Graph& graph, bool& modified) const {
|
|||
|
||||
// Replace the input of the node following add node
|
||||
const NodeArg* add_output_def = add_node.OutputDefs()[0];
|
||||
const NodeArg* conv_output_def = conv_node.OutputDefs()[0];
|
||||
NodeArg* conv_output_def = conv_node.MutableOutputDefs()[0];
|
||||
for (auto it = add_node.OutputNodesBegin(); it != add_node.OutputNodesEnd(); ++it) {
|
||||
auto output_node = graph.GetNode((*it).Index());
|
||||
if (!output_node) {
|
||||
return Status(ONNXRUNTIME, INVALID_ARGUMENT);
|
||||
}
|
||||
|
||||
auto& input_defs = output_node->MutableInputDefs();
|
||||
for (auto& def : input_defs) {
|
||||
if (def == add_output_def) {
|
||||
def = const_cast<NodeArg*>(conv_output_def);
|
||||
def = conv_output_def;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -155,7 +155,7 @@ Status ConvBNFusion::Apply(onnxruntime::Graph& graph, bool& modified) const {
|
|||
|
||||
// Replace the input of the nodes following batch normalization node
|
||||
const NodeArg* bn_output_def = bn_node.OutputDefs()[0];
|
||||
const NodeArg* conv_output_def = conv_node.OutputDefs()[0];
|
||||
NodeArg* conv_output_def = conv_node.MutableOutputDefs()[0];
|
||||
for (auto it = bn_node.OutputNodesBegin(); it != bn_node.OutputNodesEnd(); ++it) {
|
||||
auto output_node = graph.GetNode((*it).Index());
|
||||
if (!output_node) {
|
||||
|
|
@ -165,7 +165,7 @@ Status ConvBNFusion::Apply(onnxruntime::Graph& graph, bool& modified) const {
|
|||
auto& input_defs = output_node->MutableInputDefs();
|
||||
for (auto& def : input_defs) {
|
||||
if (def == bn_output_def) {
|
||||
def = const_cast<NodeArg*>(conv_output_def);
|
||||
def = conv_output_def;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -94,7 +94,7 @@ Status ConvMulFusion::Apply(onnxruntime::Graph& graph, bool& modified) const {
|
|||
|
||||
// Replace the input of the node following mul node
|
||||
const NodeArg* mul_output_def = mul_node.OutputDefs()[0];
|
||||
const NodeArg* conv_output_def = conv_node.OutputDefs()[0];
|
||||
NodeArg* conv_output_def = conv_node.MutableOutputDefs()[0];
|
||||
for (auto it = mul_node.OutputNodesBegin(); it != mul_node.OutputNodesEnd(); ++it) {
|
||||
auto output_node = graph.GetNode((*it).Index());
|
||||
if (!output_node) {
|
||||
|
|
@ -104,7 +104,7 @@ Status ConvMulFusion::Apply(onnxruntime::Graph& graph, bool& modified) const {
|
|||
auto& input_defs = output_node->MutableInputDefs();
|
||||
for (auto& def : input_defs) {
|
||||
if (def == mul_output_def) {
|
||||
def = const_cast<NodeArg*>(conv_output_def);
|
||||
def = conv_output_def;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Reference in a new issue