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:
satyajandhyala 2021-09-10 11:53:26 -07:00 committed by GitHub
parent 31af88c0bc
commit ce7b12bf5d
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
9 changed files with 79 additions and 199 deletions

View file

@ -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) {

View file

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

View file

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

View file

@ -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)

View file

@ -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,

View file

@ -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"]

View file

@ -1,7 +1,7 @@
{
"PropagateCastOps":
{
"Strategy": "REMOVE_INPUT_OUTPUT_UP_DOWN_CASTS",
"Strategy": "INSERT_AND_REDUCE",
"Level": 5,
"Allow": ["XYZ", "PQR"]
},