diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index 12245aeb51..9661633eb9 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -823,6 +823,11 @@ class Graph { } } + /** During constant folding it may become possible to infer the shape for a node. + To avoid running a full Resolve allow an individual node to have the shape inferencing re-run. + */ + Status UpdateShapeInference(Node& node); + // Options to control Graph::Resolve. struct ResolveOptions { // Whether to override existing types with inferred types. diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index 4aad6b2ed9..16b6049189 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -1644,6 +1644,16 @@ Status Graph::InferAndVerifySubgraphTypes(const Node& node, Graph& subgraph, return Status::OK(); } +Status Graph::UpdateShapeInference(Node& node) { + // We only use this during constant folding, and we don't constant fold control flow nodes. + ORT_ENFORCE(node.GetAttributeNameToMutableSubgraphMap().empty(), + "UpdateTypeShapeInference is not intended to be used with control flow nodes containing subgraphs"); + + // Whilst the type inferencing will run again we don't allow type overrides due to using the default + // ResolveOptions settings, so essentially this can only change the shape information. + return InferAndVerifyTypeMatch(node, *node.Op(), {}); +} + // Implementation of type-inference and type-checking for a single node GSL_SUPPRESS(f .23) // spurious warning about inferred_type never being checked for null Status Graph::InferAndVerifyTypeMatch(Node& node, const OpSchema& op, const ResolveOptions& options) { diff --git a/onnxruntime/core/graph/graph_utils.cc b/onnxruntime/core/graph/graph_utils.cc index cb1cc3caa3..c546d42530 100644 --- a/onnxruntime/core/graph/graph_utils.cc +++ b/onnxruntime/core/graph/graph_utils.cc @@ -522,7 +522,8 @@ bool NodeArgIsConstant(const Graph& graph, const NodeArg& node_arg) { return IsConstantInitializer(graph, node_arg.Name(), true); } -bool AllNodeInputsAreConstant(const Graph& graph, const Node& node, InitializedTensorSet& constant_inputs) { +bool AllNodeInputsAreConstant(const Graph& graph, const Node& node, InitializedTensorSet& constant_inputs, + const std::unordered_set& excluded_initializers) { // clear so we have a known state. if we fail part way through we go back to this state. constant_inputs.clear(); @@ -537,7 +538,7 @@ bool AllNodeInputsAreConstant(const Graph& graph, const Node& node, InitializedT // because it can be overridden by the user at runtime. For constant folding to be applied, the initializer should // not appear in the graph's inputs (that is the only way to guarantee it will always be constant). const ONNX_NAMESPACE::TensorProto* initializer = GetConstantInitializer(graph, input_def->Name(), true); - if (initializer) { + if (initializer && excluded_initializers.find(input_def->Name()) == excluded_initializers.cend()) { constant_inputs.insert({input_def->Name(), initializer}); } else { constant_inputs.clear(); diff --git a/onnxruntime/core/graph/graph_utils.h b/onnxruntime/core/graph/graph_utils.h index 36303479cf..77c3888bff 100644 --- a/onnxruntime/core/graph/graph_utils.h +++ b/onnxruntime/core/graph/graph_utils.h @@ -61,9 +61,10 @@ NodeArg& AddInitializer(Graph& graph, const ONNX_NAMESPACE::TensorProto& new_ini /** Checks if the given NodeArg is constant, i.e., it appears in the graph's initializers but not in its inputs. */ bool NodeArgIsConstant(const Graph& graph, const NodeArg& node_arg); -/** Checks if the given node has only constant inputs (initializers) and if so returns them in constant_inputs as they -may come from outer scope. */ -bool AllNodeInputsAreConstant(const Graph& graph, const Node& node, InitializedTensorSet& constant_inputs); +/** Checks if the given node has only constant inputs (initializers) and no input is in excluded_initializers. +If so returns them in constant_inputs as they may come from outer scope. */ +bool AllNodeInputsAreConstant(const Graph& graph, const Node& node, InitializedTensorSet& constant_inputs, + const std::unordered_set& excluded_initializers = {}); /** Gets the name of the incoming NodeArg with the specified index for the given node. */ const std::string& GetNodeInputName(const Node& node, int index); diff --git a/onnxruntime/core/optimizer/constant_folding.cc b/onnxruntime/core/optimizer/constant_folding.cc index 629be7a0f4..acc63ddac8 100644 --- a/onnxruntime/core/optimizer/constant_folding.cc +++ b/onnxruntime/core/optimizer/constant_folding.cc @@ -11,7 +11,43 @@ using namespace onnxruntime::common; namespace onnxruntime { +// We need to handle a Shape node separately as the input doesn't need to be a constant initializer for +// Shape to be able to be constant folded. +static bool ConstantFoldShapeNode(Graph& graph, Node& node) { + auto shape = node.MutableInputDefs()[0]->Shape(); + bool is_concrete_shape = true; + std::vector dim_values; + if (shape != nullptr) { + for (int dim_index = 0; dim_index < shape->dim_size(); dim_index++) { + auto dim = shape->dim(dim_index); + if (!utils::HasDimValue(dim)) { + is_concrete_shape = false; + break; + } + dim_values.push_back(dim.dim_value()); + } + } else { + is_concrete_shape = false; + } + + if (is_concrete_shape) { + ONNX_NAMESPACE::TensorProto shape_constant; + auto* constant_arg_out = node.MutableOutputDefs()[0]; + shape_constant.set_name(constant_arg_out->Name()); + shape_constant.set_data_type(ONNX_NAMESPACE::TensorProto_DataType_INT64); + shape_constant.add_dims(dim_values.size()); + shape_constant.set_raw_data(dim_values.data(), dim_values.size() * sizeof(int64_t)); + ONNX_NAMESPACE::TensorShapeProto result_shape; + result_shape.add_dim()->set_dim_value(dim_values.size()); + constant_arg_out->SetShape(result_shape); + graph.AddInitializedTensor(shape_constant); + } + + return is_concrete_shape; // convert to constant if this is true +} + Status ConstantFolding::ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const { + bool have_updated_nodes = false; GraphViewer graph_viewer(graph); auto& order = graph_viewer.GetNodesInTopologicalOrder(); @@ -23,93 +59,113 @@ Status ConstantFolding::ApplyImpl(Graph& graph, bool& modified, int graph_level, ORT_RETURN_IF_ERROR(Recurse(*node, modified, graph_level, logger)); - InitializedTensorSet constant_inputs; - - // we currently constant fold using the CPU EP only. - // if the node is assigned to a different EP we can run it if it's an ONNX op as we have CPU based implementations - // for all ONNX ops. if it's from a different domain we can't. - // NOTE: This is in addition to the IsSupportedProvider check below which will optionally do further filtering - // on the EPs we constant fold for. - auto ep_type = node->GetExecutionProviderType(); - bool cpu_ep = ep_type == kCpuExecutionProvider; - if (!cpu_ep && node->Domain() != kOnnxDomain) { - continue; + // Updating a node may allow shape inferencing to infer output shapes of following nodes, + // so re-run the shape inferencing. use have_updated_nodes as that only applies to this Graph + // (vs. 'modified' which is passed into subgraphs and applies to the main graph and all subgraphs) + // Ignore any control flow node containing subgraphs as UpdateShapeInference is not intended to be used on it. + if (have_updated_nodes && !node->ContainsSubgraph()) { + ORT_RETURN_IF_ERROR(graph.UpdateShapeInference(*node)); } - // Check if constant folding can be applied on this node. - if (!graph_utils::IsSupportedProvider(*node, GetCompatibleExecutionProviders()) || - excluded_op_types_.find(node->OpType()) != excluded_op_types_.end() || - // constant folding does not support executing a node that includes subgraphs (control flow operators, - // such as If/Loop/Scan, fall into this category). individual nodes in the subgraph will be processed - // by the Recurse call above - node->ContainsSubgraph() || !graph_utils::AllNodeInputsAreConstant(graph, *node, constant_inputs)) { - continue; - } + bool converted_to_constant = false; + if (node->OpType().compare("Shape") == 0) { + converted_to_constant = ConstantFoldShapeNode(graph, *node); + } else { + InitializedTensorSet constant_inputs; - // override the EP while setting up OptimizerExecutionFrame::Info so that it will use the CPU kernel for Compute. - if (!cpu_ep) { - node->SetExecutionProviderType(kCpuExecutionProvider); - } - - // Create execution frame for executing constant nodes. - OptimizerExecutionFrame::Info info({node}, constant_inputs); - - // undo the EP change in case something fails prior to node removal - if (!cpu_ep) { - node->SetExecutionProviderType(ep_type); - } - - std::vector fetch_mlvalue_idxs; - for (const auto* node_out : node->OutputDefs()) { - fetch_mlvalue_idxs.push_back(info.GetMLValueIndex(node_out->Name())); - } - - OptimizerExecutionFrame frame(info, fetch_mlvalue_idxs); - - auto* kernel = info.GetKernel(node->Index()); - if (kernel == nullptr) - continue; - OpKernelContext op_kernel_context(&frame, kernel, nullptr, logger); - - ORT_RETURN_IF_ERROR(kernel->Compute(&op_kernel_context)); - - std::vector fetches; - ORT_RETURN_IF_ERROR(frame.GetOutputs(fetches)); - - // Go over all output node args and substitute them with the newly computed tensors, which will be - // added to the graph as initializers. - ORT_ENFORCE(fetches.size() == node->OutputDefs().size()); - bool unsupported_output_type = false; - for (size_t fetch_idx = 0; fetch_idx < fetches.size(); ++fetch_idx) { - OrtValue& ort_value = fetches[fetch_idx]; - - if (!ort_value.IsTensor()) { - LOGS(logger, WARNING) << "Unsupported output type of " << ort_value.Type() - << ". Can't constant fold " << node->OpType() << " node '" << node->Name() << "'"; - unsupported_output_type = true; - break; + // we currently constant fold using the CPU EP only. + // if the node is assigned to a different EP we can run it if it's an ONNX op as we have CPU based implementations + // for all ONNX ops. if it's from a different domain we can't. + // NOTE: This is in addition to the IsSupportedProvider check below which will optionally do further filtering + // on the EPs we constant fold for. + auto ep_type = node->GetExecutionProviderType(); + bool cpu_ep = ep_type == kCpuExecutionProvider; + if (!cpu_ep && node->Domain() != kOnnxDomain) { + continue; } - // Build the TensorProto that corresponds to the computed OrtValue and add it as initializer to the graph. - const auto* constant_arg_out = node->OutputDefs()[fetch_idx]; - ORT_ENFORCE(ort_value.IsTensor()); - const Tensor& out_tensor = ort_value.Get(); - ONNX_NAMESPACE::TensorProto out_tensorproto = utils::TensorToTensorProto(out_tensor, constant_arg_out->Name()); + // Check if constant folding can be applied on this node. + if (!graph_utils::IsSupportedProvider(*node, GetCompatibleExecutionProviders()) || + excluded_op_types_.find(node->OpType()) != excluded_op_types_.end() || + // constant folding does not support executing a node that includes subgraphs (control flow operators, + // such as If/Loop/Scan, fall into this category). individual nodes in the subgraph will be processed + // by the Recurse call above + node->ContainsSubgraph() || !graph_utils::AllNodeInputsAreConstant(graph, *node, constant_inputs, excluded_initializers_)) { + continue; + } - graph.AddInitializedTensor(out_tensorproto); + // override the EP while setting up OptimizerExecutionFrame::Info so that it will use the CPU kernel for Compute. + if (!cpu_ep) { + node->SetExecutionProviderType(kCpuExecutionProvider); + } + + // Create execution frame for executing constant nodes. + OptimizerExecutionFrame::Info info({node}, constant_inputs); + + // undo the EP change in case something fails prior to node removal + if (!cpu_ep) { + node->SetExecutionProviderType(ep_type); + } + + std::vector fetch_mlvalue_idxs; + for (const auto* node_out : node->OutputDefs()) { + fetch_mlvalue_idxs.push_back(info.GetMLValueIndex(node_out->Name())); + } + + OptimizerExecutionFrame frame(info, fetch_mlvalue_idxs); + + auto* kernel = info.GetKernel(node->Index()); + if (kernel == nullptr) { + continue; + } + + OpKernelContext op_kernel_context(&frame, kernel, nullptr, logger); + ORT_RETURN_IF_ERROR(kernel->Compute(&op_kernel_context)); + + std::vector fetches; + ORT_RETURN_IF_ERROR(frame.GetOutputs(fetches)); + + // Go over all output node args and substitute them with the newly computed tensors, which will be + // added to the graph as initializers. + ORT_ENFORCE(fetches.size() == node->OutputDefs().size()); + converted_to_constant = true; + for (size_t fetch_idx = 0; fetch_idx < fetches.size(); ++fetch_idx) { + OrtValue& ort_value = fetches[fetch_idx]; + + if (!ort_value.IsTensor()) { + LOGS(logger, WARNING) << "Unsupported output type of " << ort_value.Type() + << ". Can't constant fold " << node->OpType() << " node '" << node->Name() << "'"; + converted_to_constant = false; + break; + } + } + + if (converted_to_constant) { + for (size_t fetch_idx = 0; fetch_idx < fetches.size(); ++fetch_idx) { + OrtValue& ort_value = fetches[fetch_idx]; + // Build the TensorProto that corresponds to the computed OrtValue and add it as initializer to the graph. + auto* constant_arg_out = node->MutableOutputDefs()[fetch_idx]; + const Tensor& out_tensor = ort_value.Get(); + ONNX_NAMESPACE::TensorProto out_tensorproto = utils::TensorToTensorProto(out_tensor, constant_arg_out->Name()); + + ONNX_NAMESPACE::TensorShapeProto result_shape; + for (auto& dim : out_tensor.Shape().GetDims()) { + result_shape.add_dim()->set_dim_value(dim); + } + + constant_arg_out->SetShape(result_shape); + graph.AddInitializedTensor(out_tensorproto); + } + } } - if (unsupported_output_type) - continue; - - // Remove the output edges of the constant node and then remove the node itself. - graph_utils::RemoveNodeOutputEdges(graph, *node); - graph.RemoveNode(node->Index()); - - // The output nodes already have the right input arg, since we used the same name in the initializer. - // We could remove unused graph initializers here, but Graph::Resolve() will take care of it. - - modified = true; + if (converted_to_constant) { + // Remove the output edges of the constant node and then remove the node itself. + graph_utils::RemoveNodeOutputEdges(graph, *node); + graph.RemoveNode(node->Index()); + modified = true; + have_updated_nodes = true; + } } return Status::OK(); diff --git a/onnxruntime/core/optimizer/constant_folding.h b/onnxruntime/core/optimizer/constant_folding.h index 2be65a2670..c0b4118e69 100644 --- a/onnxruntime/core/optimizer/constant_folding.h +++ b/onnxruntime/core/optimizer/constant_folding.h @@ -16,8 +16,11 @@ it statically computes parts of the graph that rely only on constant initializer */ class ConstantFolding : public GraphTransformer { public: - ConstantFolding(const std::unordered_set& compatible_execution_providers = {}) noexcept - : GraphTransformer("ConstantFolding", compatible_execution_providers) {} + /** Constant folding will not be applied to nodes that have one of initializers from excluded_initializers as input. + For pre-training, the trainable weights are those initializers to be excluded. */ + ConstantFolding(const std::unordered_set& compatible_execution_providers = {}, + const std::unordered_set& excluded_initializers = {}) noexcept + : GraphTransformer("ConstantFolding", compatible_execution_providers), excluded_initializers_(excluded_initializers) {} private: /** Constant folding will not be applied to nodes whose op_type is included in this set. @@ -26,6 +29,8 @@ class ConstantFolding : public GraphTransformer { {"RandomUniform", "RandomNormal", "RandomUniformLike", "RandomNormalLike", "Multinomial"}; Status ApplyImpl(Graph& graph, bool& modified, int graph_level, const logging::Logger& logger) const override; + + const std::unordered_set excluded_initializers_; }; } // namespace onnxruntime diff --git a/onnxruntime/core/providers/cuda/tensor/where_impl.cu b/onnxruntime/core/providers/cuda/tensor/where_impl.cu index 438c4f1eac..c4c4392007 100644 --- a/onnxruntime/core/providers/cuda/tensor/where_impl.cu +++ b/onnxruntime/core/providers/cuda/tensor/where_impl.cu @@ -148,8 +148,8 @@ __global__ void _TenaryElementWiseSimple( COND_INDEX_TYPE, \ X_INDEX_TYPE, \ Y_INDEX_TYPE, \ - GridDim::maxThreadsPerBlock, \ - GridDim::maxElementsPerThread> \ + GridDim::maxThreadsPerBlock, \ + GridDim::maxElementsPerThread> \ <<>>(output_rank_or_simple_broadcast, \ cond_padded_strides, \ cond_data, \ diff --git a/onnxruntime/test/optimizer/graph_transform_test.cc b/onnxruntime/test/optimizer/graph_transform_test.cc index 52dcb637f0..d7f0fbeac2 100644 --- a/onnxruntime/test/optimizer/graph_transform_test.cc +++ b/onnxruntime/test/optimizer/graph_transform_test.cc @@ -240,6 +240,52 @@ TEST_F(GraphTransformationTests, ConstantFoldingSubgraph) { << "Constant folding should have been able to remove the Add node in both subgraphs"; } +TEST_F(GraphTransformationTests, ConstantFoldingWithShapeToInitializer) { + auto model_uri = MODEL_FOLDER "fusion/constant_folding_with_shape_to_initializer.onnx"; + std::shared_ptr model; + ASSERT_TRUE(Model::Load(model_uri, model, nullptr, *logger_).IsOK()); + Graph& graph = model->MainGraph(); + std::map op_to_count = CountOpsInGraph(graph); + ASSERT_TRUE(op_to_count["Shape"] == 2); + ASSERT_TRUE(op_to_count["MatMul"] == 2); + ASSERT_TRUE(op_to_count["Unsqueeze"] == 3); + + std::unordered_set compatible_eps; + std::unordered_set excluded_initializers; + excluded_initializers.insert("matmul_weight"); + onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; + graph_transformation_mgr.Register(onnxruntime::make_unique(compatible_eps, excluded_initializers), TransformerLevel::Level1); + + ASSERT_TRUE(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_).IsOK()); + + op_to_count = CountOpsInGraph(graph); + ASSERT_TRUE(op_to_count["Shape"] == 0); + ASSERT_TRUE(op_to_count["MatMul"] == 2); + ASSERT_TRUE(op_to_count["Unsqueeze"] == 0); +} + +TEST_F(GraphTransformationTests, ConstantFoldingWithScalarShapeToInitializer) { + auto model_uri = MODEL_FOLDER "fusion/constant_folding_with_scalar_shape_to_initializer.onnx"; + std::shared_ptr model; + ASSERT_TRUE(Model::Load(model_uri, model, nullptr, *logger_).IsOK()); + Graph& graph = model->MainGraph(); + std::map op_to_count = CountOpsInGraph(graph); + ASSERT_TRUE(op_to_count["Shape"] == 1); + ASSERT_TRUE(op_to_count["ConstantOfShape"] == 1); + ASSERT_TRUE(op_to_count["Add"] == 1); + + std::unordered_set compatible_eps; + onnxruntime::GraphTransformerManager graph_transformation_mgr{5}; + graph_transformation_mgr.Register(onnxruntime::make_unique(compatible_eps), TransformerLevel::Level1); + + ASSERT_TRUE(graph_transformation_mgr.ApplyTransformers(graph, TransformerLevel::Level1, *logger_).IsOK()); + + op_to_count = CountOpsInGraph(graph); + ASSERT_TRUE(op_to_count["Shape"] == 0); + ASSERT_TRUE(op_to_count["ConstantOfShape"] == 0); + ASSERT_TRUE(op_to_count["Add"] == 1); +} + TEST_F(GraphTransformationTests, ShapeToInitializer) { auto model_uri = MODEL_FOLDER "shape-add.onnx"; std::shared_ptr model; diff --git a/onnxruntime/test/testdata/transform/fusion/constant_folding_with_scalar_shape_to_initializer.onnx b/onnxruntime/test/testdata/transform/fusion/constant_folding_with_scalar_shape_to_initializer.onnx new file mode 100644 index 0000000000..ee1583e836 Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/constant_folding_with_scalar_shape_to_initializer.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/constant_folding_with_shape_to_initializer.onnx b/onnxruntime/test/testdata/transform/fusion/constant_folding_with_shape_to_initializer.onnx new file mode 100644 index 0000000000..072723f4cd Binary files /dev/null and b/onnxruntime/test/testdata/transform/fusion/constant_folding_with_shape_to_initializer.onnx differ diff --git a/onnxruntime/test/testdata/transform/fusion/constant_folding_with_shape_to_initializer.py b/onnxruntime/test/testdata/transform/fusion/constant_folding_with_shape_to_initializer.py new file mode 100644 index 0000000000..ab1b9d00b4 --- /dev/null +++ b/onnxruntime/test/testdata/transform/fusion/constant_folding_with_shape_to_initializer.py @@ -0,0 +1,80 @@ +import onnx +from onnx import helper +from onnx import TensorProto, GraphProto, OperatorSetIdProto +from onnx import numpy_helper +import numpy as np + +X = helper.make_tensor_value_info('input', TensorProto.FLOAT, [2, 4, 8]) +Y = helper.make_tensor_value_info('output', TensorProto.FLOAT, [2, 4, 16]) + +matmul_weight_vals = (0.01 * np.arange(2 * 4 * 4, dtype=np.float32)).reshape((2, 4, 4)) +matmul_weight_initializer = numpy_helper.from_array(matmul_weight_vals, 'matmul_weight') +gather_constant_zero = numpy_helper.from_array(np.int64(0), name='gather_constant_zero') +gather_constant_one = numpy_helper.from_array(np.int64(1), name='gather_constant_one') +div_constant_two = numpy_helper.from_array(np.int64(2), name='div_constant_two') +unsqueeze_constant_16 = numpy_helper.from_array(np.int64(16), name='unsqueeze_constant_16') + +shape1 = helper.make_node('Shape', ['input'], ['shape1'], name='shape1') +constant_of_shape = helper.make_node('ConstantOfShape', ['shape1'], ['constant_of_shape'], name='constant_of_shape') +transpose = helper.make_node('Transpose', ['constant_of_shape'], ['transpose'], name='transpose', perm=[0,2,1]) +matmul1 = helper.make_node('MatMul', ['transpose', matmul_weight_initializer.name], ['matmul1'], name='matmul1') +matmul2 = helper.make_node('MatMul', ['matmul1', 'input'], ['matmul2'], name='matmul2') +shape2 = helper.make_node('Shape', ['matmul2'], ['shape2'], name='shape2') +gather1 = helper.make_node('Gather', ['shape2', gather_constant_zero.name], ['gather1'], name='gather1', axis=0) +gather2 = helper.make_node('Gather', ['shape2', gather_constant_one.name], ['gather2'], name='gather2', axis=0) +div = helper.make_node('Div', ['gather2', div_constant_two.name], ['div'], name='div') +unsqueeze1 = helper.make_node('Unsqueeze', ['gather1'], ['unsqueeze1'], name='unsqueeze1', axes=[0]) +unsqueeze2 = helper.make_node('Unsqueeze', ['div'], ['unsqueeze2'], name='unsqueeze2', axes=[0]) +unsqueeze3 = helper.make_node('Unsqueeze', [unsqueeze_constant_16.name], ['unsqueeze3'], name='unsqueeze3', axes=[0]) +concat = helper.make_node('Concat', ['unsqueeze1', 'unsqueeze2', 'unsqueeze3'], ['concat'], name='concat', axis=0) +reshape = helper.make_node('Reshape', ['matmul2', 'concat'], ['output'], name='reshape') + +# Create the graph (GraphProto) +graph_def = helper.make_graph( + [shape1, constant_of_shape, transpose, matmul1, matmul2, shape2, gather1, gather2, div, unsqueeze1, unsqueeze2, unsqueeze3, concat, reshape], + 'constant_folding_with_shape_to_initializer_model', + [X], + [Y], + [matmul_weight_initializer, gather_constant_zero, gather_constant_one, div_constant_two, unsqueeze_constant_16] +) + +opsets = [] +onnxdomain = OperatorSetIdProto() +onnxdomain.version = 12 +onnxdomain.domain = "" # The empty string ("") or absence of this field implies the operator set that is defined as part of the ONNX specification. +opsets.append(onnxdomain) + +msdomain = OperatorSetIdProto() +msdomain.version = 1 +msdomain.domain = 'com.microsoft' + +opsets.append(msdomain) +kwargs={} +kwargs['opset_imports'] = opsets + +# Create the model (ModelProto) +model_def = helper.make_model(graph_def, producer_name='onnx-example', **kwargs) +onnx.save(model_def, 'constant_folding_with_shape_to_initializer.onnx') + + + +X = helper.make_tensor_value_info('input', TensorProto.FLOAT, [1]) +Y = helper.make_tensor_value_info('output', TensorProto.FLOAT, [1]) + +squeeze = helper.make_node('Squeeze', ['input'], ['squeeze'], name='squeeze', axes=[0]) +shape = helper.make_node('Shape', ['squeeze'], ['shape'], name='shape') +constant_of_shape = helper.make_node('ConstantOfShape', ['shape'], ['constant_of_shape'], name='constant_of_shape') +add = helper.make_node('Add', ['squeeze', 'constant_of_shape'], ['add'], name='add') +unsqueeze = helper.make_node('Unsqueeze', ['add'], ['output'], name='unsqueeze', axes=[0]) + +# Create the graph (GraphProto) +graph_def = helper.make_graph( + [squeeze, shape, constant_of_shape, add, unsqueeze], + 'constant_folding_with_scalar_shape_to_initializer_model', + [X], + [Y] +) + +# Create the model (ModelProto) +model_def = helper.make_model(graph_def, producer_name='onnx-example', **kwargs) +onnx.save(model_def, 'constant_folding_with_scalar_shape_to_initializer.onnx') diff --git a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc index d7b76764aa..796a2746a4 100644 --- a/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc +++ b/orttraining/orttraining/core/optimizer/graph_transformer_utils.cc @@ -39,6 +39,7 @@ namespace onnxruntime { namespace transformer_utils { std::vector> GeneratePreTrainingTransformers(TransformerLevel level, + const std::unordered_set& weights_to_train, const std::vector& transformers_and_rules_to_enable) { std::vector> transformers; std::unique_ptr rule_transformer = nullptr; @@ -62,6 +63,7 @@ std::vector> GeneratePreTrainingTransformers(T transformers.emplace_back(onnxruntime::make_unique(compatible_eps)); transformers.emplace_back(onnxruntime::make_unique(compatible_eps)); transformers.emplace_back(onnxruntime::make_unique(compatible_eps)); + transformers.emplace_back(onnxruntime::make_unique(compatible_eps, weights_to_train)); auto horizontal_parallel_size = training::DistributedRunContext::GroupSize(training::WorkerGroupType::HorizontalParallel); if (horizontal_parallel_size > 1) { LOGS_DEFAULT(WARNING) << horizontal_parallel_size << "-way horizontal model parallel is enabled"; @@ -87,6 +89,9 @@ std::vector> GeneratePreTrainingTransformers(T // if the custom list to enable transformers\rules is empty then return the default generated transformers and rules // otherwise generate a filtered list based on the provided custom list. + // Note that some rule-based transformers are depending on some custom transformers, + // e.g., ExpandElimination and CastElimination are depending on ConstantFolding to fold the constant first, + // so we should always push the rule-based transformer to the end, this is expecially important when transformation step is 1. if (transformers_and_rules_to_enable.empty()) { if (rule_transformer != nullptr) { transformers.emplace_back(std::move(rule_transformer)); @@ -94,10 +99,6 @@ std::vector> GeneratePreTrainingTransformers(T return transformers; } std::vector> filtered_list; - // If the rule-based transformer is not empty, it should be included in the custom transformer list below. - if (rule_transformer != nullptr) { - filtered_list.emplace_back(std::move(rule_transformer)); - } // pick custom transformers enabled for this session for (const auto& t_name : transformers_and_rules_to_enable) { std::for_each(transformers.begin(), transformers.end(), @@ -107,6 +108,10 @@ std::vector> GeneratePreTrainingTransformers(T } }); } + // If the rule-based transformer is not empty, it should be included in the custom transformer list below. + if (rule_transformer != nullptr) { + filtered_list.emplace_back(std::move(rule_transformer)); + } return filtered_list; } diff --git a/orttraining/orttraining/core/optimizer/graph_transformer_utils.h b/orttraining/orttraining/core/optimizer/graph_transformer_utils.h index 921cdd0767..34fce66875 100644 --- a/orttraining/orttraining/core/optimizer/graph_transformer_utils.h +++ b/orttraining/orttraining/core/optimizer/graph_transformer_utils.h @@ -14,6 +14,7 @@ namespace transformer_utils { /** Generates all pre-training transformers for this level. */ std::vector> GeneratePreTrainingTransformers(TransformerLevel level, + const std::unordered_set& weights_to_train, const std::vector& rules_and_transformers_to_enable = {}); /** Generates all predefined (both rule-based and non-rule-based) transformers for this level. diff --git a/orttraining/orttraining/core/session/training_session.cc b/orttraining/orttraining/core/session/training_session.cc index cdfbe7cb9a..029cd16d6d 100644 --- a/orttraining/orttraining/core/session/training_session.cc +++ b/orttraining/orttraining/core/session/training_session.cc @@ -130,7 +130,18 @@ Status TrainingSession::ConfigureForTraining( config.distributed_config.horizontal_parallel_size, config.distributed_config.pipeline_parallel_size}); - ORT_RETURN_IF_ERROR(ApplyTransformationsToMainGraph()); + // We need to get trainable weights to prevent constant folding from them. This works well if trainable weights are passed from config. + // For case we use GetTrainableModelInitializers to get trainable weights such as C++ frontend, it may get more initializers + // than trainable weights here as it's before transformers. So the constant folding may miss some nodes we actually can fold. + std::unordered_set excluded_initializers = + !config.weight_names_to_train.empty() + ? config.weight_names_to_train + : GetTrainableModelInitializers(config.immutable_weights); + for (const auto& weight_name_to_not_train : config.weight_names_to_not_train) { + excluded_initializers.erase(weight_name_to_not_train); + } + + ORT_RETURN_IF_ERROR(ApplyTransformationsToMainGraph(excluded_initializers)); is_mixed_precision_enabled_ = config.mixed_precision_config.has_value(); @@ -424,9 +435,9 @@ static Status AddGradientAccumulationNodes(Graph& graph, return GraphAugmenter::AugmentGraph(graph, graph_defs); } -Status TrainingSession::ApplyTransformationsToMainGraph() { +Status TrainingSession::ApplyTransformationsToMainGraph(const std::unordered_set& weights_to_train) { GraphTransformerManager graph_transformation_mgr{1}; - AddPreTrainingTransformers(graph_transformation_mgr); + AddPreTrainingTransformers(graph_transformation_mgr, weights_to_train); // apply transformers Graph& graph = model_->MainGraph(); @@ -438,11 +449,12 @@ Status TrainingSession::ApplyTransformationsToMainGraph() { // Registers all the pre transformers with transformer manager void TrainingSession::AddPreTrainingTransformers(GraphTransformerManager& transformer_manager, + const std::unordered_set& weights_to_train, TransformerLevel graph_optimization_level, const std::vector& custom_list) { auto add_transformers = [&](TransformerLevel level) { // Generate and register transformers for level - auto transformers_to_register = transformer_utils::GeneratePreTrainingTransformers(level, custom_list); + auto transformers_to_register = transformer_utils::GeneratePreTrainingTransformers(level, weights_to_train, custom_list); for (auto& entry : transformers_to_register) { transformer_manager.Register(std::move(entry), level); } diff --git a/orttraining/orttraining/core/session/training_session.h b/orttraining/orttraining/core/session/training_session.h index c9d5ab081e..93314e2f55 100644 --- a/orttraining/orttraining/core/session/training_session.h +++ b/orttraining/orttraining/core/session/training_session.h @@ -335,10 +335,11 @@ class TrainingSession : public InferenceSession { std::string& backward_waited_output_name, std::string& backward_recorded_output_name); - common::Status ApplyTransformationsToMainGraph(); + common::Status ApplyTransformationsToMainGraph(const std::unordered_set& weights_to_train); /** configure initial transformers for training */ void AddPreTrainingTransformers(GraphTransformerManager& transformer_manager, + const std::unordered_set& weights_to_train, TransformerLevel graph_optimization_level = TransformerLevel::MaxLevel, const std::vector& custom_list = {});