From 41acd8c543207a9976d06d08fae0b4e350f88c1c Mon Sep 17 00:00:00 2001 From: pengwa Date: Tue, 9 Apr 2024 09:24:48 +0800 Subject: [PATCH] Support more ops for recompute (#20234) ### Support more ops for recompute To cover Mistral model, and support padding elimination ops. ### Motivation and Context --- .../memory_optimizer/recompute_analysis.cc | 80 ++++++++++++++----- 1 file changed, 59 insertions(+), 21 deletions(-) diff --git a/orttraining/orttraining/core/optimizer/memory_optimizer/recompute_analysis.cc b/orttraining/orttraining/core/optimizer/memory_optimizer/recompute_analysis.cc index 37ac1c4950..be2f1387fb 100644 --- a/orttraining/orttraining/core/optimizer/memory_optimizer/recompute_analysis.cc +++ b/orttraining/orttraining/core/optimizer/memory_optimizer/recompute_analysis.cc @@ -144,6 +144,20 @@ const InlinedHashMap& GetAllowedRecompu {20, {0}}, }, }, + { + utils::GetFullQualifiedOpName("Cos", kOnnxDomain), + { + {7, {}}, + }, + }, + { + utils::GetFullQualifiedOpName("CumSum", kOnnxDomain), + { + // The axis input is trivial + {11, {1}}, + {14, {1}}, + }, + }, { utils::GetFullQualifiedOpName("Dropout", kOnnxDomain), { @@ -162,27 +176,6 @@ const InlinedHashMap& GetAllowedRecompu {14, {}}, }, }, - { - utils::GetFullQualifiedOpName("Expand", kOnnxDomain), - { - {8, {1}}, // Ignore the shape. - {13, {1}}, - }, - }, - { - utils::GetFullQualifiedOpName("Cos", kOnnxDomain), - { - {7, {}}, - }, - }, - { - utils::GetFullQualifiedOpName("CumSum", kOnnxDomain), - { - // The axis input is trivial - {11, {1}}, - {14, {1}}, - }, - }, { utils::GetFullQualifiedOpName("Einsum", kOnnxDomain), { @@ -199,12 +192,25 @@ const InlinedHashMap& GetAllowedRecompu {19, {}}, }, }, + { + utils::GetFullQualifiedOpName("Expand", kOnnxDomain), + { + {8, {1}}, // Ignore the shape. + {13, {1}}, + }, + }, { utils::GetFullQualifiedOpName("FastGelu", kMSDomain), { {1, {}}, }, }, + { + utils::GetFullQualifiedOpName("FlattenAndUnpad", kMSDomain), + { + {1, {1}}, // ignore the indices + }, + }, { utils::GetFullQualifiedOpName("Gather", kOnnxDomain), { @@ -225,6 +231,17 @@ const InlinedHashMap& GetAllowedRecompu {1, {}}, }, }, + { + utils::GetFullQualifiedOpName("Gemm", kOnnxDomain), + { + {1, {}}, + {6, {}}, + {7, {}}, + {9, {}}, + {11, {}}, + {13, {}}, + }, + }, { utils::GetFullQualifiedOpName("Less", kOnnxDomain), { @@ -244,6 +261,27 @@ const InlinedHashMap& GetAllowedRecompu {14, {}}, }, }, + { + utils::GetFullQualifiedOpName("Neg", kOnnxDomain), + { + {1, {}}, + {6, {}}, + {13, {}}, + }, + }, + { + utils::GetFullQualifiedOpName("NonZero", kOnnxDomain), + { + {9, {}}, + {13, {}}, + }, + }, + { + utils::GetFullQualifiedOpName("PadAndUnflatten", kMSDomain), + { + {1, {1, 2}}, // ignore the indices and unflatten_dims + }, + }, { utils::GetFullQualifiedOpName("Range", kOnnxDomain), {