mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-28 20:11:22 +00:00
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.
This commit is contained in:
parent
31af88c0bc
commit
ce7b12bf5d
9 changed files with 79 additions and 199 deletions
|
|
@ -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<Strategy>::type;
|
||||
friend constexpr Strategy operator|(const Strategy s1, const Strategy s2) {
|
||||
|
|
|
|||
|
|
@ -66,8 +66,8 @@ static std::string GetName(const std::pair<const NodeArg*, std::vector<Node*>>&
|
|||
*/
|
||||
static std::vector<std::unordered_set<std::string>> 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<std::string, std::vector<int>> opcode_to_input_map = {
|
|||
{"Gather", {0}},
|
||||
{"Reshape", {0}},
|
||||
{"Dropout", {0}},
|
||||
{"Expand", {0}},
|
||||
{"LayerNormalization", {0, 1, 2}},
|
||||
{"Squeeze", {0}},
|
||||
{"Unsqueeze", {0}}
|
||||
};
|
||||
|
||||
static std::unordered_map<std::string, std::vector<int>> opcode_to_output_map = {
|
||||
{"Gather", {0}},
|
||||
{"Reshape", {0}},
|
||||
{"Dropout", {0}},
|
||||
{"Expand", {0}},
|
||||
{"LayerNormalization", {0}},
|
||||
{"Squeeze", {0}},
|
||||
{"Unsqueeze", {0}}
|
||||
};
|
||||
|
||||
static std::unordered_set<std::string> inserted_node_names; // Names of the nodes inserted
|
||||
|
|
@ -386,14 +392,15 @@ static Status RemoveCastNodesChain(Graph& graph, std::vector<Node*> 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<Node*> 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<NodeIndex>& 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<Node*> input_casts;
|
||||
std::vector<Node*> output_casts;
|
||||
std::vector<NodeArg*>& outputs = node->MutableOutputDefs();
|
||||
std::vector<NodeArg*>& inputs = node->MutableInputDefs();
|
||||
std::unordered_set<NodeArg*> 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<Node*> 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<Node*>();
|
||||
}
|
||||
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<NodeArgToConsumerMap>(non_cast_consumers_map, GetName);
|
||||
}
|
||||
for (Node* cast : input_casts) {
|
||||
RemoveCastNodesChain(graph, {cast}, removed_nodes);
|
||||
}
|
||||
LOGS(logger, VERBOSE) << "RemoveInputOutputUpDownCasts: Removed Cast nodes "
|
||||
<< ConcatNames<std::vector<Node*>>(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<NodeArgToConsumerMap>(non_cast_producers_map, GetName);
|
||||
}
|
||||
for (Node* cast : output_casts) {
|
||||
RemoveCastNodesChain(graph, {cast}, removed_nodes);
|
||||
}
|
||||
LOGS(logger, VERBOSE) << "RemoveInputOutputUpDownCasts: Removed Cast nodes "
|
||||
<< ConcatNames<std::vector<Node*>>(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);
|
||||
|
|
|
|||
|
|
@ -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<std::pair<GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy, int>, int> casts_count_map;
|
||||
std::map<std::pair<Strategy, int>, int> casts_count_map;
|
||||
vector<std::string> allow_ops = {}; // Allowed ops for PropagateCastOps graph transformer
|
||||
};
|
||||
std::pair<GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy, int> insertAndReduce0 = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::InsertAndReduce, 0);
|
||||
std::pair<GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy, int> floodFill1 = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill, 1);
|
||||
std::pair<GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy, int> floodFill2 = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill, 2);
|
||||
std::pair<GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy, int> floodFill1Plus = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill |
|
||||
GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::RemoveInputOutputUpDownCasts,
|
||||
1);
|
||||
std::pair<GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy, int> floodFill2Plus = std::make_pair(GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill |
|
||||
GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::RemoveInputOutputUpDownCasts,
|
||||
2);
|
||||
|
||||
std::pair<Strategy, int> insertAndReduce0 = std::make_pair(Strategy::InsertAndReduce, 0);
|
||||
std::pair<Strategy, int> insertAndReduce1 = std::make_pair(Strategy::InsertAndReduce, 1);
|
||||
std::pair<Strategy, int> floodFill1 = std::make_pair(Strategy::FloodFill, 1);
|
||||
std::pair<Strategy, int> floodFill2 = std::make_pair(Strategy::FloodFill, 2);
|
||||
std::vector<std::string> allow_matmul = {"MatMul"};
|
||||
std::vector<std::string> allow_matmul_transpose = {"MatMul", "Transpose"};
|
||||
std::vector<std::string> allow_matmul_transpose_add = {"Add", "MatMul", "Transpose"};
|
||||
const std::vector<PropagateCastOpsTestSpecs> 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<Model> p_model;
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
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)
|
||||
|
|
|
|||
BIN
onnxruntime/test/testdata/transform/propagate_cast/squeeze_cast_propagation_test.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/propagate_cast/squeeze_cast_propagation_test.onnx
vendored
Normal file
Binary file not shown.
BIN
onnxruntime/test/testdata/transform/propagate_cast/unsqueeze_cast_propagation_test.onnx
vendored
Normal file
BIN
onnxruntime/test/testdata/transform/propagate_cast/unsqueeze_cast_propagation_test.onnx
vendored
Normal file
Binary file not shown.
|
|
@ -623,7 +623,6 @@ void addObjectMethodsForTraining(py::module& m, ExecutionProviderRegistrationFn
|
|||
.value("NONE", GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::None)
|
||||
.value("INSERT_AND_REDUCE", GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::InsertAndReduce)
|
||||
.value("FLOOD_FILL", GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::FloodFill)
|
||||
.value("REMOVE_INPUT_OUTPUT_UP_DOWN_CASTS", GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy::RemoveInputOutputUpDownCasts)
|
||||
.def("__or__", py::overload_cast<GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy,
|
||||
GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy>(&operator|))
|
||||
.def("__and__", py::overload_cast<GraphTransformerConfiguration::PropagateCastOpsConfiguration::Strategy,
|
||||
|
|
|
|||
|
|
@ -88,7 +88,7 @@ def test_load_config_from_json_2():
|
|||
ort_model_attributes = model._torch_module._execution_manager(training_mode)
|
||||
|
||||
# test propagate cast ops
|
||||
assert ort_model_attributes._propagate_cast_ops_strategy == C.PropagateCastOpsStrategy.REMOVE_INPUT_OUTPUT_UP_DOWN_CASTS
|
||||
assert ort_model_attributes._propagate_cast_ops_strategy == C.PropagateCastOpsStrategy.INSERT_AND_REDUCE
|
||||
assert ort_model_attributes._propagate_cast_ops_level == 5
|
||||
assert ort_model_attributes._propagate_cast_ops_allow == ["XYZ", "PQR"]
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
{
|
||||
"PropagateCastOps":
|
||||
{
|
||||
"Strategy": "REMOVE_INPUT_OUTPUT_UP_DOWN_CASTS",
|
||||
"Strategy": "INSERT_AND_REDUCE",
|
||||
"Level": 5,
|
||||
"Allow": ["XYZ", "PQR"]
|
||||
},
|
||||
|
|
|
|||
Loading…
Reference in a new issue