From ce7b12bf5d0e73202a8d958bb0f235251ac9c4c8 Mon Sep 17 00:00:00 2001 From: satyajandhyala Date: Fri, 10 Sep 2021 11:53:26 -0700 Subject: [PATCH] Added new fp16 allow/safe opcodes in PropagateCastOps (#8964) * Removed RemoveInputOutputUpDownCasts strategy in PropagatCastOps. * Added Expand, Squeeze and Unsqueeze ops to fp16 allow ops * Added onnx models for squeeze/unsqueeze tests. --- .../core/optimizer/graph_transformer_config.h | 2 - .../core/optimizer/propagate_cast_ops.cc | 175 ++---------------- .../test/optimizer/graph_transform_test.cc | 42 ++--- .../propagate_cast/gen_propagate_cast.py | 54 ++++-- .../squeeze_cast_propagation_test.onnx | Bin 0 -> 224 bytes .../unsqueeze_cast_propagation_test.onnx | Bin 0 -> 245 bytes .../python/orttraining_pybind_state.cc | 1 - ...test_ortmodule_experimental_json_config.py | 2 +- ..._ortmodule_experimental_json_config_2.json | 2 +- 9 files changed, 79 insertions(+), 199 deletions(-) create mode 100644 onnxruntime/test/testdata/transform/propagate_cast/squeeze_cast_propagation_test.onnx create mode 100644 onnxruntime/test/testdata/transform/propagate_cast/unsqueeze_cast_propagation_test.onnx diff --git a/include/onnxruntime/core/optimizer/graph_transformer_config.h b/include/onnxruntime/core/optimizer/graph_transformer_config.h index 6866433d51..ac760981b0 100644 --- a/include/onnxruntime/core/optimizer/graph_transformer_config.h +++ b/include/onnxruntime/core/optimizer/graph_transformer_config.h @@ -21,8 +21,6 @@ struct GraphTransformerConfiguration { None = 0, InsertAndReduce = 1, FloodFill = 2, /* Propagate FP16 Cast operations up and FP32 operations down */ - RemoveInputOutputUpDownCasts = 4 /* If all the floatingpoint inputs of a node are casted to FP32 and all the floatingpoint outputs - are casted to FP16. Then remove all input and output casts. */ }; using Strategy_t = std::underlying_type::type; friend constexpr Strategy operator|(const Strategy s1, const Strategy s2) { diff --git a/onnxruntime/core/optimizer/propagate_cast_ops.cc b/onnxruntime/core/optimizer/propagate_cast_ops.cc index 5bbfa1585b..4852da323d 100644 --- a/onnxruntime/core/optimizer/propagate_cast_ops.cc +++ b/onnxruntime/core/optimizer/propagate_cast_ops.cc @@ -66,8 +66,8 @@ static std::string GetName(const std::pair>& */ static std::vector> fp16_allow_ops = { /* Level 0 */ {}, - /* Level 1 */ {"Transpose", "Relu", "Reshape", "Split", "Tanh"}, - /* Level 2 */ {"BiasGelu", "Dropout", "FastGelu", "Gather", "Gelu", "LayerNormalization", "Where"}}; + /* Level 1 */ {"Expand", "Transpose", "Relu", "Reshape", "Split", "Tanh", "Squeeze", "Unsqueeze"}, + /* Level 2 */ {"Add", "BiasGelu", "Dropout", "FastGelu", "Gather", "Gelu", "LayerNormalization", "Where"}}; /* * The following two maps specify the opcode to input and opcode to output mappings to list the inputs/outputs to consider while propagating @@ -79,14 +79,20 @@ static std::unordered_map> opcode_to_input_map = { {"Gather", {0}}, {"Reshape", {0}}, {"Dropout", {0}}, + {"Expand", {0}}, {"LayerNormalization", {0, 1, 2}}, + {"Squeeze", {0}}, + {"Unsqueeze", {0}} }; static std::unordered_map> opcode_to_output_map = { {"Gather", {0}}, {"Reshape", {0}}, {"Dropout", {0}}, + {"Expand", {0}}, {"LayerNormalization", {0}}, + {"Squeeze", {0}}, + {"Unsqueeze", {0}} }; static std::unordered_set inserted_node_names; // Names of the nodes inserted @@ -386,14 +392,15 @@ static Status RemoveCastNodesChain(Graph& graph, std::vector casts, std:: * |__________| * | * _____V______ ____________ -* | Cast FP16| | Opcode 1 | -* |__________| |__________| -* | | +* | Cast | | Opcode 1 | +* |FP16->FP32| |__________| +* |__________| | * | | * | ---\ _____V______ * | ---/ | Opcode2 | * _____V______ |__________| -* | Cast FP32| +* | Cast | + |FP32->Fp16| * |__________| * | * _____V______ @@ -407,14 +414,15 @@ static Status RemoveCastNodesChain(Graph& graph, std::vector casts, std:: * |__________| * | * _____V______ ____________ -* | Cast FP32| | Opcode 1 | -* |__________| |__________| -* | | +* | Cast | | Opcode 1 | +* |FP16->FP32| |__________| +* |__________| | * | | * | ---\ _____V______ * | ---/ | Cast FP32| * _____V______ |__________| -* | Cast FP32| | +* | Cast | | +* |FP32->FP32| | * |__________| _____V______ * | | Opcode2 | * _____V______ |__________| @@ -1097,141 +1105,6 @@ static bool PropagateFP16CastsFromOutputsToInputs(Graph& graph, Node* node, return modified; } -/* -* RemoveInputOutputUpDownCasts -* When all the floatingpoint inputs on a node with ANY opcode are FP32 Cast outputs and -* all the outputs are casted to FP16, remove the input FP32 casts and output FP16 casts. -* This transformation only makes difference to opcodes not listed in FP16 allowed opcodes. -* This transformation is less aggressive than adding all such opcodes that can benefit form -* this transformation to FP16 allowed opcodes. -* -* Input0/NodeArg Input1/NodeArg -* | | -* _____V____ _____V______ -* |Cast FP32| | Cast FP32| -* |_________| |__________| -* | | -* __V______________V___ -* | Opcode | -* |(operation performed | -* | in float32) | -* |_____________________| -* | | -* _____V____ _____V______ -* |Cast FP16| | Cast FP16| -* |_________| |__________| -* | | -* V V -* -* -* Input0/NodeArg Input1/NodeArg -* | | -* __V______________V___ -* | Opcode | -* |(operation performed | -* | in float16) | -* |_____________________| -* | | -* | | -* V V -*/ -static bool RemoveInputOutputUpDownCasts(Graph& graph, Node* node, - std::deque& removed_nodes, - size_t level, - const logging::Logger& logger) { - bool modified = false; - bool has_float_outputs = false; - bool has_float_inputs = false; - bool all_float_outputs_have_casts = true; - bool all_float_inputs_have_casts = true; - std::vector input_casts; - std::vector output_casts; - std::vector& outputs = node->MutableOutputDefs(); - std::vector& inputs = node->MutableInputDefs(); - std::unordered_set require_type_change; - NodeArgToConsumerMap non_cast_consumers_map; - NodeArgToConsumerMap non_cast_producers_map; - for (auto iter = outputs.begin(); iter != outputs.end() && (level >= 2 || all_float_outputs_have_casts); ++iter) { - NodeArg* output = *iter; - if (!IsType(*output, TensorProto::FLOAT) || !IsRelevantOutput(node, output)) { - continue; - } - has_float_outputs = true; - std::vector consumers = graph.GetMutableConsumerNodes(output->Name()); - for (auto node_iter = consumers.begin(); node_iter != consumers.end() && (level >= 2 || all_float_outputs_have_casts); ++node_iter) { - Node* consumer = *node_iter; - if (nullptr != consumer && - std::find(removed_nodes.begin(), removed_nodes.end(), consumer->Index()) == removed_nodes.end()) { - if (IsCastTo(consumer, TensorProto::FLOAT16)) { - output_casts.push_back(consumer); - } else { - non_cast_consumers_map[output].push_back(consumer); - all_float_outputs_have_casts = false; - } - } - } - if (graph.IsOutput(output) && non_cast_consumers_map.find(output) == non_cast_consumers_map.end()) { - non_cast_consumers_map[output] = std::vector(); - } - if (non_cast_consumers_map.empty()) { - require_type_change.insert(output); - } - } - for (auto iter = inputs.begin(); iter != inputs.end() && (level >= 2 || all_float_inputs_have_casts); ++iter) { - NodeArg* input = *iter; - if (!IsType(*input, TensorProto::FLOAT) || !IsRelevantInput(node, input)) { - continue; - } - has_float_inputs = true; - Node* producer = graph.GetMutableProducerNode(input->Name()); - if (nullptr != producer) { - if (std::find(removed_nodes.begin(), removed_nodes.end(), producer->Index()) == removed_nodes.end()) { - if (IsCastTo(producer, TensorProto::FLOAT) && - producer->GetOutputEdgesCount() == 1 && - !graph.IsOutput(input)) { - input_casts.push_back(producer); - require_type_change.insert(input); - } else { - non_cast_producers_map[input].push_back(node); - all_float_inputs_have_casts = false; - } - } - } else if (graph_utils::IsGraphInput(graph, input)) { - non_cast_producers_map[input].push_back(node); - all_float_inputs_have_casts = false; - } - } - if (has_float_outputs && has_float_inputs && - (level >= 2 || (all_float_outputs_have_casts && all_float_inputs_have_casts)) && - (input_casts.size() > 0 && output_casts.size() > 0)) { - if (non_cast_consumers_map.size() > 0) { - InsertCastNodes(graph, non_cast_consumers_map, false, removed_nodes); - LOGS(logger, VERBOSE) << "RemoveInputOutputUpDownCasts: Inserted FP32 Cast node to " - << ConcatNames(non_cast_consumers_map, GetName); - } - for (Node* cast : input_casts) { - RemoveCastNodesChain(graph, {cast}, removed_nodes); - } - LOGS(logger, VERBOSE) << "RemoveInputOutputUpDownCasts: Removed Cast nodes " - << ConcatNames>(input_casts) - << " feeding from the same compute node " << node->Name(); - if (non_cast_producers_map.size() > 0) { - InsertCastNodes(graph, non_cast_producers_map, true, removed_nodes); - LOGS(logger, VERBOSE) << "RemoveInputOutputUpDownCasts: Inserted FP16 Cast node to " - << ConcatNames(non_cast_producers_map, GetName); - } - for (Node* cast : output_casts) { - RemoveCastNodesChain(graph, {cast}, removed_nodes); - } - LOGS(logger, VERBOSE) << "RemoveInputOutputUpDownCasts: Removed Cast nodes " - << ConcatNames>(output_casts) - << " feeding the same compute node " << node->Name(); - ChangeTypeToFP16(graph, require_type_change, false, logger); - modified = true; - } - return modified; -} - /* * CreateCast * Create a cast node based on the node_arg for the given data type. If the node_arg is a graph output is_graph_outut is set. @@ -1508,18 +1381,6 @@ Status PropagateCastOps::ApplyImpl(Graph& graph, bool& modified, int graph_level } } - // Eliminate FP32 input casts and FP16 output casts - if ((strategy_ & GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::RemoveInputOutputUpDownCasts) != - GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::None) { - for (auto node_index : node_topology_list) { - Node* node = graph.GetNode(node_index); - if (nullptr != node && - std::find(removed_nodes.begin(), removed_nodes.end(), node->Index()) == removed_nodes.end()) { - local_modified |= RemoveInputOutputUpDownCasts(graph, node, removed_nodes, level_, logger); - } - } - } - // Propagate FP32 Casts forward for (auto node_index : node_topology_list) { Node* node = graph.GetNode(node_index); diff --git a/onnxruntime/test/optimizer/graph_transform_test.cc b/onnxruntime/test/optimizer/graph_transform_test.cc index 3c773128bc..26e7d7e044 100644 --- a/onnxruntime/test/optimizer/graph_transform_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_test.cc @@ -4236,25 +4236,24 @@ TEST_F(GraphTransformationTests, FilterEnabledOptimizers) { } TEST_F(GraphTransformationTests, PropagateCastOpsTests) { + using Strategy = GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy; struct PropagateCastOpsTestSpecs { PathString model_uri; // Expected number of casts after the transformation with different stratigies and optimization levels - std::map, int> casts_count_map; + std::map, int> casts_count_map; vector allow_ops = {}; // Allowed ops for PropagateCastOps graph transformer }; - std::pair insertAndReduce0 = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::InsertAndReduce, 0); - std::pair floodFill1 = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill, 1); - std::pair floodFill2 = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill, 2); - std::pair floodFill1Plus = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill | - GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::RemoveInputOutputUpDownCasts, - 1); - std::pair floodFill2Plus = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill | - GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::RemoveInputOutputUpDownCasts, - 2); + + std::pair insertAndReduce0 = std::make_pair(Strategy::InsertAndReduce, 0); + std::pair insertAndReduce1 = std::make_pair(Strategy::InsertAndReduce, 1); + std::pair floodFill1 = std::make_pair(Strategy::FloodFill, 1); + std::pair floodFill2 = std::make_pair(Strategy::FloodFill, 2); std::vector allow_matmul = {"MatMul"}; std::vector allow_matmul_transpose = {"MatMul", "Transpose"}; std::vector allow_matmul_transpose_add = {"Add", "MatMul", "Transpose"}; const std::vector test_cases = { + {MODEL_FOLDER "propagate_cast/squeeze_cast_propagation_test.onnx", {{insertAndReduce0, 2}, {insertAndReduce1, 0}, {floodFill1, 0}, {floodFill2, 0}}}, + {MODEL_FOLDER "propagate_cast/unsqueeze_cast_propagation_test.onnx", {{insertAndReduce0, 2}, {insertAndReduce1, 0}, {floodFill1, 0}, {floodFill2, 0}}}, // Negative testcase to test that the transformer will not move cast bool to float/float16. {MODEL_FOLDER "propagate_cast/negative_test_case_bool_fp_cast.onnx", {{insertAndReduce0, 2}, {floodFill1, 2}, {floodFill2, 2}}, {"Add"}}, {MODEL_FOLDER "propagate_cast/negative_test_case_bool_fp16_cast.onnx", {{insertAndReduce0, 2}, {floodFill1, 2}, {floodFill2, 2}}, {"Add"}}, @@ -4288,16 +4287,16 @@ TEST_F(GraphTransformationTests, PropagateCastOpsTests) { {MODEL_FOLDER "propagate_cast/matmul_transpose_inputs_transpose_product_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_transpose_inputs_transpose_product_cast_inputs_cast_product.onnx", {{insertAndReduce0, 0}, {floodFill1, 0}, {floodFill2, 0}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_transpose_inputs_transpose_product_cast_product.onnx", {{insertAndReduce0, 2}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul_transpose}, - {MODEL_FOLDER "propagate_cast/matmul_add_transpose_inputs_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose}, + {MODEL_FOLDER "propagate_cast/matmul_add_transpose_inputs_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 2}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_add_transpose_inputs_cast_inputs_cast_product.onnx", {{insertAndReduce0, 0}, {floodFill1, 0}, {floodFill2, 0}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_add_transpose_inputs_cast_product.onnx", {{insertAndReduce0, 2}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul_transpose}, - {MODEL_FOLDER "propagate_cast/matmul_add_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul}, + {MODEL_FOLDER "propagate_cast/matmul_add_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 2}}, allow_matmul}, {MODEL_FOLDER "propagate_cast/matmul_add_cast_inputs_cast_product.onnx", {{insertAndReduce0, 0}, {floodFill1, 0}, {floodFill2, 0}}, allow_matmul}, {MODEL_FOLDER "propagate_cast/matmul_add_cast_product.onnx", {{insertAndReduce0, 2}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul}, - {MODEL_FOLDER "propagate_cast/matmul_add_transpose_product_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose}, + {MODEL_FOLDER "propagate_cast/matmul_add_transpose_product_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 2}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_add_transpose_product_cast_inputs_cast_product.onnx", {{insertAndReduce0, 0}, {floodFill1, 0}, {floodFill2, 0}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_add_transpose_product_cast_product.onnx", {{insertAndReduce0, 2}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul_transpose}, - {MODEL_FOLDER "propagate_cast/matmul_add_transpose_inputs_transpose_product_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose}, + {MODEL_FOLDER "propagate_cast/matmul_add_transpose_inputs_transpose_product_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 2}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_add_transpose_inputs_transpose_product_cast_inputs_cast_product.onnx", {{insertAndReduce0, 0}, {floodFill1, 0}, {floodFill2, 0}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_add_transpose_inputs_transpose_product_cast_product.onnx", {{insertAndReduce0, 2}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_add_cast_inputs_cast_product_cast_sum.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose_add}, @@ -4355,11 +4354,11 @@ TEST_F(GraphTransformationTests, PropagateCastOpsTests) { {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_before_cast_transpose_second_matmul.onnx", {{insertAndReduce0, 3}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_after_cast_transpose_second_matmul.onnx", {{insertAndReduce0, 3}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_before_cast_second_matmul.onnx", {{insertAndReduce0, 3}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul_transpose}, - {MODEL_FOLDER "propagate_cast/matmul_two_outputs_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 2}, {floodFill2, 3}}, allow_matmul}, - {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_after_cast_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 2}, {floodFill2, 3}}, allow_matmul_transpose}, - {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_before_cast_transpose_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 1}, {floodFill2, 3}}, allow_matmul_transpose}, - {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_after_cast_transpose_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 1}, {floodFill2, 3}}, allow_matmul_transpose}, - {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_before_cast_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 2}, {floodFill2, 3}}, allow_matmul_transpose}, + {MODEL_FOLDER "propagate_cast/matmul_two_outputs_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 2}, {floodFill2, 4}}, allow_matmul}, + {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_after_cast_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 2}, {floodFill2, 4}}, allow_matmul_transpose}, + {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_before_cast_transpose_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 1}, {floodFill2, 4}}, allow_matmul_transpose}, + {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_after_cast_transpose_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 1}, {floodFill2, 4}}, allow_matmul_transpose}, + {MODEL_FOLDER "propagate_cast/matmul_two_outputs_transpose_before_cast_second_matmul_add_products.onnx", {{insertAndReduce0, 5}, {floodFill1, 2}, {floodFill2, 4}}, allow_matmul_transpose}, {MODEL_FOLDER "propagate_cast/matmul_two_outputs_cast_inputs.onnx", {{insertAndReduce0, 1}, {floodFill1, 2}, {floodFill2, 2}}, allow_matmul_transpose_add}, {MODEL_FOLDER "propagate_cast/matmul_two_outputs_cast_inputs_second_matmul_add_products.onnx", {{insertAndReduce0, 2}, {floodFill1, 4}, {floodFill2, 3}}, allow_matmul_transpose_add}, {MODEL_FOLDER "propagate_cast/matmul_two_outputs_cast_inputs_second_matmul.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose_add}, @@ -4374,14 +4373,13 @@ TEST_F(GraphTransformationTests, PropagateCastOpsTests) { {MODEL_FOLDER "propagate_cast/matmul_two_outputs_cast_inputs_transpose_before_cast_second_matmul.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose_add}, {MODEL_FOLDER "propagate_cast/matmul_two_outputs_cast_inputs_transpose_before_cast_transpose.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose_add}, {MODEL_FOLDER "propagate_cast/matmul_two_outputs_cast_inputs_transpose_before_cast_transpose_second_matmul_add_products.onnx", {{insertAndReduce0, 2}, {floodFill1, 3}, {floodFill2, 2}}, allow_matmul_transpose_add}, - {MODEL_FOLDER "propagate_cast/matmul_two_outputs_cast_inputs_transpose_before_cast_transpose_second_matmul.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose_add}, - {MODEL_FOLDER "propagate_cast/matmul_cast_inputs_cast_product.onnx", {{insertAndReduce0, 3}, {floodFill1Plus, 0}, {floodFill2Plus, 0}}}}; + {MODEL_FOLDER "propagate_cast/matmul_two_outputs_cast_inputs_transpose_before_cast_transpose_second_matmul.onnx", {{insertAndReduce0, 1}, {floodFill1, 1}, {floodFill2, 1}}, allow_matmul_transpose_add}}; // Create a temporary directory, which will be deleted automatically, to save/load the transformed models. TemporaryDirectory temp_dir{ORT_TSTR("propagate_casts_test_output_dir")}; for (PropagateCastOpsTestSpecs test_case : test_cases) { for (auto scenario : test_case.casts_count_map) { - GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy strategy = scenario.first.first; + Strategy strategy = scenario.first.first; int level = scenario.first.second; int expected_casts_count = scenario.second; std::shared_ptr p_model; diff --git a/onnxruntime/test/testdata/transform/propagate_cast/gen_propagate_cast.py b/onnxruntime/test/testdata/transform/propagate_cast/gen_propagate_cast.py index 5160ac9900..8601fbda4f 100644 --- a/onnxruntime/test/testdata/transform/propagate_cast/gen_propagate_cast.py +++ b/onnxruntime/test/testdata/transform/propagate_cast/gen_propagate_cast.py @@ -3,6 +3,7 @@ from onnx import helper from onnx import TensorProto from onnx import OperatorSetIdProto import itertools +import numpy as np onnxdomain = OperatorSetIdProto() onnxdomain.version = 12 @@ -128,7 +129,7 @@ def gen_fuse_sibling_casts(model_path): type_to_string(type2), nodes, inputs, outputs, []) -def flip_type(flip, type): +def flip_type(type, flip=True): return (TensorProto.FLOAT16 if type == TensorProto.FLOAT else TensorProto.FLOAT) if flip else type @@ -203,11 +204,11 @@ def gen_propagate_cast_test_model(model_path, transpose_inputs, transpose_produc input_0, input_1 = do_transpose_inputs(input_0, input_1, nodes) if cast_inputs: input_0, input_1 = do_cast_inputs(input_0, input_1, nodes, input_type) - input_type = flip_type(True, input_type) + input_type = flip_type(input_type) else: if cast_inputs: input_0, input_1 = do_cast_inputs(input_0, input_1, nodes, input_type) - input_type = flip_type(True, input_type) + input_type = flip_type(input_type) if transpose_inputs: input_0, input_1 = do_transpose_inputs(input_0, input_1, nodes) nodes.append(helper.make_node( @@ -221,8 +222,8 @@ def gen_propagate_cast_test_model(model_path, transpose_inputs, transpose_produc product = do_transpose_product(product, nodes) if cast_product: - product = do_cast_product(product, nodes, flip_type(True, product_type)) - product_type = flip_type(True, product_type) + product = do_cast_product(product, nodes, flip_type(product_type)) + product_type = flip_type(product_type) inputs = [ helper.make_tensor_value_info( @@ -232,20 +233,19 @@ def gen_propagate_cast_test_model(model_path, transpose_inputs, transpose_produc ] if insert_add: input_2 = "input_2" - add_input_type = flip_type(cast_input2, product_type) + add_input_type = flip_type(product_type, cast_input2) inputs.append(helper.make_tensor_value_info( input_2, add_input_type, ['N', 'N'])) output = "sum" output_type = product_type if cast_input2: input_2 = do_cast_input2( - input_2, nodes, flip_type(True, add_input_type)) + input_2, nodes, flip_type(add_input_type)) nodes.append(helper.make_node( "Add", [product, input_2], [output], "Add_0")) if cast_sum: - output = do_cast_sum(output, nodes, flip_type( - True, output_type)) - output_type = flip_type(True, output_type) + output = do_cast_sum(output, nodes, flip_type(output_type)) + output_type = flip_type(output_type) else: output = product output_type = product_type @@ -273,7 +273,7 @@ def gen_matmul_two_products(model_path, transpose, transpose_before_cast, second "transpose_1_"+output_1], "Transpose_1")) output_1 = "transpose_1_"+output_1 return output_0, output_1 - input_type = flip_type(cast_inputs, TensorProto.FLOAT) + input_type = flip_type(TensorProto.FLOAT, cast_inputs) input_0 = "input_0" input_1 = "input_1" output = "product" @@ -289,7 +289,7 @@ def gen_matmul_two_products(model_path, transpose, transpose_before_cast, second "input_1", input_type, ['K', 'N']) ] if cast_inputs: - input_type = flip_type(True, input_type) + input_type = flip_type(input_type) input_0, input_1 = do_cast_inputs(input_0, input_1, nodes, input_type) cast_count +=2 output0_type = input_type @@ -317,7 +317,7 @@ def gen_matmul_two_products(model_path, transpose, transpose_before_cast, second "sum", input_type, ['M', 'N'])) if transpose > 0 and transpose_before_cast: output_0, output_1 = do_transpose(output_0, output_1, transpose, nodes) - output0_type = flip_type(True, output0_type) + output0_type = flip_type(output0_type) nodes.append(helper.make_node( "Cast", [output_0], @@ -334,7 +334,7 @@ def gen_matmul_two_products(model_path, transpose, transpose_before_cast, second "Cast_"+str(cast_count), to=TensorProto.FLOAT16)) output_1 = "cast_"+str(cast_count)+"_"+output_1 - output1_type = flip_type(True, output1_type) + output1_type = flip_type(output1_type) if transpose > 0 and not transpose_before_cast: output_0, output_1 = do_transpose(output_0, output_1, transpose, nodes) @@ -351,6 +351,7 @@ def gen_matmul_two_products(model_path, transpose, transpose_before_cast, second model_path += "_add_products" if add_products else "" save(model_path, nodes, inputs, outputs, []) + def gen_bool_to_float16_cast(model_path): X1 = helper.make_tensor_value_info('x1', TensorProto.INT64, [1, 1]) X2 = helper.make_tensor_value_info('x2', TensorProto.INT64, [1, 1]) @@ -364,6 +365,7 @@ def gen_bool_to_float16_cast(model_path): save(model_path, [less1, cast1, cast2, add1], [X1, X2, X3], [Y], []) + def gen_bool_to_float_cast(model_path): X1 = helper.make_tensor_value_info('x1', TensorProto.INT64, [1, 1]) X2 = helper.make_tensor_value_info('x2', TensorProto.INT64, [1, 1]) @@ -378,6 +380,26 @@ def gen_bool_to_float_cast(model_path): save(model_path, [less1, cast1, cast2, cast3, add1], [X1, X2, X3], [Y], []) + +def gen_one_input_one_output_test(op, model_path, axes_attribute=False): + X = helper.make_tensor_value_info('x', TensorProto.FLOAT16, [2, 2]) + output_shape = [2, 2] + if (op=="Unsqueeze"): + output_shape.append(1) + Y = helper.make_tensor_value_info('y', TensorProto.FLOAT16, output_shape) + node_inputs=[] + graph_inputs=[X] + cast1 = helper.make_node('Cast', ['x'], ['cast1'], name='cast1', to=TensorProto.FLOAT) + node_inputs.insert(0, 'cast1') + if axes_attribute: + node = helper.make_node(op, node_inputs, ['op_output'], name=op+str(1), axes=np.array([2]).astype(np.int64)) + else: + node = helper.make_node(op, node_inputs, ['op_output'], name=op+str(1)) + cast2 = helper.make_node('Cast', ['op_output'], [ + 'y'], name='cast2', to=TensorProto.FLOAT16) + save(model_path, [cast1, node, cast2], graph_inputs, [Y], []) + + for (transpose_inputs, transpose_product, cast_inputs, cast_product, insert_add, cast_sum, cast_input2) in list(itertools.product([False, True], repeat=7)): if not insert_add and (cast_sum or cast_input2): continue @@ -398,4 +420,6 @@ for (transpose, transpose_before_cast, second_matmul, add_products, cast_inputs) gen_bool_to_float16_cast("negative_test_case_bool_fp16_cast") -gen_bool_to_float_cast("negative_test_case_bool_fp_cast") \ No newline at end of file +gen_bool_to_float_cast("negative_test_case_bool_fp_cast") +gen_one_input_one_output_test("Squeeze", "squeeze_cast_propagation_test") +gen_one_input_one_output_test("Unsqueeze", "unsqueeze_cast_propagation_test", True) diff --git a/onnxruntime/test/testdata/transform/propagate_cast/squeeze_cast_propagation_test.onnx b/onnxruntime/test/testdata/transform/propagate_cast/squeeze_cast_propagation_test.onnx new file mode 100644 index 0000000000000000000000000000000000000000..26a4fda6710be4e15b3b339ebfdfc06666d50575 GIT binary patch literal 224 zcmdl7i9_DURU6($v(dR6`|pD2q#t3n4GWSP3!C2*o%qpm9Qi zAoBx?@(U8v6H8Jl7i9_DX!4G;=(&operator|)) .def("__and__", py::overload_cast