diff --git a/orttraining/orttraining/core/framework/module_gradient_graph_builder.cc b/orttraining/orttraining/core/framework/module_gradient_graph_builder.cc index 1e76789901..af5a4fb215 100644 --- a/orttraining/orttraining/core/framework/module_gradient_graph_builder.cc +++ b/orttraining/orttraining/core/framework/module_gradient_graph_builder.cc @@ -168,6 +168,11 @@ Status ModuleGradientGraphBuilder::BuildAndSplit(std::istream& model_istream, gradient_graph.Resolve(); + // Run the transformers again mainly for backward part. + for (int i = static_cast(TransformerLevel::Level1); i <= static_cast(TransformerLevel::MaxLevel); i++) { + ORT_RETURN_IF_ERROR(graph_transformation_mgr.ApplyTransformers(gradient_graph, static_cast(i), *logger_)); + } + // Create two copies of gradient model for forward and backward models respectively. auto gradient_model_proto = model_->ToProto(); ORT_RETURN_IF_ERROR(Model::Load(gradient_model_proto, forward_model_, nullptr, *logger_));