diff --git a/cmake/external/SafeInt/safeint b/cmake/external/SafeInt/safeint index a104e0cf23..0ce2bb56e1 160000 --- a/cmake/external/SafeInt/safeint +++ b/cmake/external/SafeInt/safeint @@ -1 +1 @@ -Subproject commit a104e0cf23be4fe848f7ef1f3e8996fe429b06bb +Subproject commit 0ce2bb56e1ce46264333d75507bf99f2e4c07e3b diff --git a/cmake/external/flatbuffers b/cmake/external/flatbuffers index 6df40a2471..db2aa9b4ec 160000 --- a/cmake/external/flatbuffers +++ b/cmake/external/flatbuffers @@ -1 +1 @@ -Subproject commit 6df40a2471737b27271bdd9b900ab5f3aec746c7 +Subproject commit db2aa9b4eca0edc94aea9e2059c21ce8786fc762 diff --git a/cmake/external/onnx-tensorrt b/cmake/external/onnx-tensorrt index a3a4e38b2d..fbebb14474 160000 --- a/cmake/external/onnx-tensorrt +++ b/cmake/external/onnx-tensorrt @@ -1 +1 @@ -Subproject commit a3a4e38b2dfa7a62b6dcae33c0d1678b3bb5ef2a +Subproject commit fbebb144744b3be8defecd8478d74940056df305 diff --git a/orttraining/orttraining/core/graph/optimizer_graph_builder.cc b/orttraining/orttraining/core/graph/optimizer_graph_builder.cc index ad37d15bef..a39a4d8913 100644 --- a/orttraining/orttraining/core/graph/optimizer_graph_builder.cc +++ b/orttraining/orttraining/core/graph/optimizer_graph_builder.cc @@ -135,15 +135,12 @@ Status OptimizerGraphBuilder::AddGradientScalingNodes( std::vector& gradient_argdefs, // update argdefs in place std::vector& output_gradient_argdef, // update argdef in place GraphAugmenter::GraphDefs& graph_defs, - const bool allreduce_in_fp16) { + ONNX_NAMESPACE::TensorProto_DataType target_type) { ArgDef pre_allreduce_scale(nodearg_name_generator("pre_allreduce_scale"), graph_defs.CreateTypeProto({}, ONNX_NAMESPACE::TensorProto_DataType_FLOAT)); graph_defs.AddInitializers({CreateTensorProto(pre_allreduce_scale.name, scale, {})}); - auto target_type = allreduce_in_fp16 ? ONNX_NAMESPACE::TensorProto_DataType_FLOAT16 - : ONNX_NAMESPACE::TensorProto_DataType_FLOAT; - TypeProto* fused_gradient_type_proto = graph_defs.CreateTypeProto(); fused_gradient_type_proto->mutable_tensor_type()->set_elem_type(target_type); diff --git a/orttraining/orttraining/core/graph/optimizer_graph_builder.h b/orttraining/orttraining/core/graph/optimizer_graph_builder.h index 2965800ca5..f3fea149d6 100644 --- a/orttraining/orttraining/core/graph/optimizer_graph_builder.h +++ b/orttraining/orttraining/core/graph/optimizer_graph_builder.h @@ -88,7 +88,7 @@ class OptimizerGraphBuilder { std::vector& gradient_argdefs, // update argdefs in place std::vector& output_gradient_argdef, // update argdef in place GraphAugmenter::GraphDefs& graph_defs, - const bool allreduce_in_fp16); + ONNX_NAMESPACE::TensorProto_DataType target_type); Status AddGradientNorm( const NodeArgNameGeneratorFn& nodearg_name_generator,