mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-28 20:11:22 +00:00
merge conflicts.
This commit is contained in:
parent
1326eb3230
commit
319a071a6e
5 changed files with 5 additions and 8 deletions
2
cmake/external/SafeInt/safeint
vendored
2
cmake/external/SafeInt/safeint
vendored
|
|
@ -1 +1 @@
|
|||
Subproject commit a104e0cf23be4fe848f7ef1f3e8996fe429b06bb
|
||||
Subproject commit 0ce2bb56e1ce46264333d75507bf99f2e4c07e3b
|
||||
2
cmake/external/flatbuffers
vendored
2
cmake/external/flatbuffers
vendored
|
|
@ -1 +1 @@
|
|||
Subproject commit 6df40a2471737b27271bdd9b900ab5f3aec746c7
|
||||
Subproject commit db2aa9b4eca0edc94aea9e2059c21ce8786fc762
|
||||
2
cmake/external/onnx-tensorrt
vendored
2
cmake/external/onnx-tensorrt
vendored
|
|
@ -1 +1 @@
|
|||
Subproject commit a3a4e38b2dfa7a62b6dcae33c0d1678b3bb5ef2a
|
||||
Subproject commit fbebb144744b3be8defecd8478d74940056df305
|
||||
|
|
@ -135,15 +135,12 @@ Status OptimizerGraphBuilder::AddGradientScalingNodes(
|
|||
std::vector<ArgDef>& gradient_argdefs, // update argdefs in place
|
||||
std::vector<ArgDef>& 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<float>(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);
|
||||
|
||||
|
|
|
|||
|
|
@ -88,7 +88,7 @@ class OptimizerGraphBuilder {
|
|||
std::vector<ArgDef>& gradient_argdefs, // update argdefs in place
|
||||
std::vector<ArgDef>& 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,
|
||||
|
|
|
|||
Loading…
Reference in a new issue