remove const cast from conv related fusion.

This commit is contained in:
linkerzhang 2018-11-29 17:45:44 -08:00
parent 728c3c078e
commit 35f94cac27
4 changed files with 11 additions and 7 deletions

View file

@ -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; }

View file

@ -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;
}
}
}

View file

@ -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;
}
}
}

View file

@ -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;
}
}
}