diff --git a/docs/ContribOperators.md b/docs/ContribOperators.md index 0d782a96a4..7f4a960bfd 100644 --- a/docs/ContribOperators.md +++ b/docs/ContribOperators.md @@ -401,6 +401,8 @@ This version of the operator has been available since version 1 of the 'com.micr
The subgraph for initialization of encoder and decoder. It will be called once before decoder subgraph.
eos_token_id : int (required)
The id of the end-of-sequence token
+
init_decoder : graph
+
The subgraph for the first decoding run. It will be called once before `decoder` subgraph. This is relevant only for the GPT2 model. If this attribute is missing, the `decoder` subgraph will be used for all decoding runs
model_type : int
model type: 0 for GPT-2; 1 for encoder decoder like T5
no_repeat_ngram_size : int
@@ -1700,9 +1702,11 @@ This version of the operator has been available since version 1 of the 'com.micr
decoder_start_token_id : int
The id of the token that indicates decoding starts.
encoder : graph
-
The subgraph for initialization of encoder and decoder. It will be called once before decoder subgraph.
+
The subgraph for initialization of encoder and decoder. It will be called once before `decoder` subgraph.
eos_token_id : int (required)
The id of the end-of-sequence token
+
init_decoder : graph
+
The subgraph for the first decoding run. It will be called once before `decoder` subgraph. This is relevant only for the GPT2 model. If this attribute is missing, the `decoder` subgraph will be used for all decoding runs
model_type : int
model type: 0 for decoder only like GPT-2; 1 for encoder decoder like Bart
no_repeat_ngram_size : int
diff --git a/onnxruntime/contrib_ops/cpu/transformers/beam_search.cc b/onnxruntime/contrib_ops/cpu/transformers/beam_search.cc index d4132675ed..b9b93a45d9 100644 --- a/onnxruntime/contrib_ops/cpu/transformers/beam_search.cc +++ b/onnxruntime/contrib_ops/cpu/transformers/beam_search.cc @@ -34,6 +34,7 @@ #include "contrib_ops/cpu/transformers/beam_search_scorer.h" #include "contrib_ops/cpu/transformers/beam_search_impl_gpt.h" #include "contrib_ops/cpu/transformers/beam_search_impl_t5.h" +#include "contrib_ops/cpu/transformers/greedy_search_impl_gpt.h" using namespace ONNX_NAMESPACE; using namespace onnxruntime::common; @@ -65,11 +66,20 @@ void BeamSearch::Init(const OpKernelInfo& info) { ONNX_NAMESPACE::GraphProto proto; if (parameters_.model_type != IBeamSearchParameters::kModelTypeGpt) { + // Make sure the encoder sub-graph attribute is present for the T5 model. ORT_ENFORCE(info.GetAttr("encoder", &proto).IsOK()); } - // Make sure the decoder attribute was present even though we don't need it here. + if (parameters_.model_type == IBeamSearchParameters::kModelTypeGpt) { + // Check if the init_decoder sub-graph attribute is present for the GPT2 model. + if (info.GetAttr("init_decoder", &proto).IsOK()) { + has_init_decoder_ = true; + } + } + + // Make sure the decoder sub-graph attribute is present for all model types. ORT_ENFORCE(info.GetAttr("decoder", &proto).IsOK()); + ORT_IGNORE_RETURN_VALUE(proto); } @@ -80,14 +90,30 @@ Status BeamSearch::SetupSubgraphExecutionInfo(const SessionState& session_state, if (parameters_.model_type == IBeamSearchParameters::kModelTypeGpt) { if (attribute_name == "decoder") { ORT_ENFORCE(gpt_subgraph_ == nullptr, "SetupSubgraphExecutionInfo should only be called once for each subgraph."); - gpt_subgraph_ = std::make_unique(node, attribute_name, subgraph_session_state.GetGraphViewer()); - ORT_RETURN_IF_ERROR(gpt_subgraph_->Setup(session_state, subgraph_session_state)); + auto res = gpt_details::CreateGptSubgraphAndUpdateParameters(node, session_state, attribute_name, + subgraph_session_state, parameters_); + + auto status = res.first; + if (!status.IsOK()) { + return status; + } + + gpt_subgraph_ = std::move(res.second); decoder_feeds_fetches_manager_ = gpt_subgraph_->GetFeedsFetchesManager(); - parameters_.SetSubgraphParameters(gpt_subgraph_->vocab_size, - gpt_subgraph_->num_heads, - gpt_subgraph_->head_size, - gpt_subgraph_->num_layers); + } else if (attribute_name == "init_decoder") { + ORT_ENFORCE(init_run_gpt_subgraph_ == nullptr, "SetupSubgraphExecutionInfo should only be called once for each subgraph."); + auto res = gpt_details::CreateGptSubgraphAndUpdateParameters(node, session_state, attribute_name, + subgraph_session_state, parameters_); + + auto status = res.first; + if (!status.IsOK()) { + return status; + } + + init_run_gpt_subgraph_ = std::move(res.second); + init_run_decoder_feeds_fetches_manager_ = init_run_gpt_subgraph_->GetFeedsFetchesManager(); } + } else if (parameters_.model_type == IBeamSearchParameters::kModelTypeT5) { if (attribute_name == "encoder") { ORT_ENFORCE(t5_encoder_subgraph_ == nullptr, @@ -130,6 +156,12 @@ Status BeamSearch::Compute(OpKernelContext* ctx) const { ORT_ENFORCE(decoder_session_state, "Subgraph SessionState was not found for 'decoder' attribute."); ORT_ENFORCE(decoder_feeds_fetches_manager_, "CreateFeedsFetchesManager must be called prior to execution of graph."); + auto* init_run_decoder_session_state = ctx_internal->SubgraphSessionState("init_decoder"); + if (has_init_decoder_) { + ORT_ENFORCE(init_run_decoder_session_state, "Subgraph SessionState was not found for 'decoder' attribute."); + ORT_ENFORCE(init_run_decoder_feeds_fetches_manager_, "CreateFeedsFetchesManager must be called prior to execution of graph."); + } + concurrency::ThreadPool* thread_pool = ctx->GetOperatorThreadPool(); // Make a copy of parameters since we will update it based on inputs later @@ -138,7 +170,12 @@ Status BeamSearch::Compute(OpKernelContext* ctx) const { if (parameters_.model_type == IBeamSearchParameters::kModelTypeGpt) { if (!gpt_subgraph_->IsOutputFloat16()) { // Output float32 BeamSearchGpt impl{ - *ctx_internal, *decoder_session_state, *gpt_subgraph_, thread_pool, cuda_stream_, dumper_, parameters, + *ctx_internal, + has_init_decoder_ ? init_run_decoder_session_state : nullptr, + has_init_decoder_ ? init_run_gpt_subgraph_.get() : nullptr, + *decoder_session_state, + *gpt_subgraph_, + thread_pool, cuda_stream_, dumper_, parameters, GenerationCpuDeviceHelper::CreateGptInputs, add_to_feeds_func_ ? add_to_feeds_func_ : GenerationCpuDeviceHelper::AddToFeeds, topk_func_ ? topk_func_ : GenerationCpuDeviceHelper::TopK, @@ -149,10 +186,15 @@ Status BeamSearch::Compute(OpKernelContext* ctx) const { update_gpt_feeds_func_ ? update_gpt_feeds_func_ : GenerationCpuDeviceHelper::UpdateGptFeeds}; ORT_RETURN_IF_ERROR(impl.Initialize()); - return impl.Execute(*decoder_feeds_fetches_manager_); + return impl.Execute(init_run_decoder_feeds_fetches_manager_, *decoder_feeds_fetches_manager_); } else { // Output float16 BeamSearchGpt impl{ - *ctx_internal, *decoder_session_state, *gpt_subgraph_, thread_pool, cuda_stream_, dumper_, parameters, + *ctx_internal, + has_init_decoder_ ? init_run_decoder_session_state : nullptr, + has_init_decoder_ ? init_run_gpt_subgraph_.get() : nullptr, + *decoder_session_state, + *gpt_subgraph_, + thread_pool, cuda_stream_, dumper_, parameters, GenerationCpuDeviceHelper::CreateGptInputs, add_to_feeds_func_ ? add_to_feeds_func_ : GenerationCpuDeviceHelper::AddToFeeds, topk_func_ ? topk_func_ : GenerationCpuDeviceHelper::TopK, @@ -163,7 +205,7 @@ Status BeamSearch::Compute(OpKernelContext* ctx) const { update_gpt_feeds_fp16_func_}; ORT_RETURN_IF_ERROR(impl.Initialize()); - return impl.Execute(*decoder_feeds_fetches_manager_); + return impl.Execute(init_run_decoder_feeds_fetches_manager_, *decoder_feeds_fetches_manager_); } } diff --git a/onnxruntime/contrib_ops/cpu/transformers/beam_search.h b/onnxruntime/contrib_ops/cpu/transformers/beam_search.h index e957efc3df..c33dfa162b 100644 --- a/onnxruntime/contrib_ops/cpu/transformers/beam_search.h +++ b/onnxruntime/contrib_ops/cpu/transformers/beam_search.h @@ -119,17 +119,34 @@ class BeamSearch : public IControlFlowKernel { //------------------------------------------------------------ // Subgraph and FeedsFetchesManager re-used for each subgraph execution. //------------------------------------------------------------ + + // Relevant only for GPT2 + // The init_run_gpt_subgraph_ (if the `init_decoder` attribute is present) will be + // used for the first decoding run and the gpt_subgraph_ will be used + // for subsequent runs. + // If the `init_decoder` attribute is missing, the `gpt_subgraph_` will be + // used for all decoding runs. + std::unique_ptr init_run_gpt_subgraph_; std::unique_ptr gpt_subgraph_; + + // Relevant only for T5 + // Same concept as above. + // The encoder will be used for the first run and the decoder will + // be used for subsequent runs. std::unique_ptr t5_encoder_subgraph_; std::unique_ptr t5_decoder_subgraph_; + FeedsFetchesManager* encoder_feeds_fetches_manager_; FeedsFetchesManager* decoder_feeds_fetches_manager_; + FeedsFetchesManager* init_run_decoder_feeds_fetches_manager_; void* cuda_stream_; IConsoleDumper* dumper_; BeamSearchParameters parameters_; + + bool has_init_decoder_ = false; }; } // namespace transformers diff --git a/onnxruntime/contrib_ops/cpu/transformers/beam_search_impl_gpt.h b/onnxruntime/contrib_ops/cpu/transformers/beam_search_impl_gpt.h index 50606dfc8f..afac73f3a7 100644 --- a/onnxruntime/contrib_ops/cpu/transformers/beam_search_impl_gpt.h +++ b/onnxruntime/contrib_ops/cpu/transformers/beam_search_impl_gpt.h @@ -17,6 +17,8 @@ template class BeamSearchGpt : public BeamSearchBase { public: BeamSearchGpt(OpKernelContextInternal& context, + const SessionState* init_run_decoder_session_state, + GptSubgraph* init_run_gpt_subgraph, const SessionState& decoder_session_state, GptSubgraph& gpt_subgraph, concurrency::ThreadPool* thread_pool, @@ -34,6 +36,8 @@ class BeamSearchGpt : public BeamSearchBase { : BeamSearchBase(context, decoder_session_state, thread_pool, cuda_stream, cuda_dumper, params, topk_func, process_logits_func, device_copy_func, device_copy_int32_func), + init_run_decoder_session_state_(init_run_decoder_session_state), + init_run_gpt_subgraph_(init_run_gpt_subgraph), gpt_subgraph_(gpt_subgraph), create_inputs_func_(create_inputs_func), add_to_feeds_func_(add_to_feeds_func), @@ -43,7 +47,8 @@ class BeamSearchGpt : public BeamSearchBase { // Execute beam search in iterations util stopping criteria is reached. // In each iteration, GPT subgraph is called, and next token for each sequence is generated. - Status Execute(const FeedsFetchesManager& feeds_fetches_manager); + Status Execute(const FeedsFetchesManager* init_run_feeds_fetches_manager, + const FeedsFetchesManager& feeds_fetches_manager); private: // Prepare the inputs for first inference of subgraph @@ -62,6 +67,8 @@ class BeamSearchGpt : public BeamSearchBase { gsl::span beam_next_tokens, gsl::span beam_indices); + const SessionState* init_run_decoder_session_state_ = nullptr; + GptSubgraph* init_run_gpt_subgraph_ = nullptr; GptSubgraph& gpt_subgraph_; // Device specific functions @@ -79,6 +86,21 @@ Status BeamSearchGpt::CreateInitialFeeds(gsl::span& sequence_lengths const OrtValue* input_ids_value = this->context_.GetInputOrtValue(0); const Tensor& input_ids = input_ids_value->Get(); const OrtValue* attn_mask_value = this->context_.GetInputOrtValue(9); + + if (init_run_gpt_subgraph_ != nullptr) { + return init_run_gpt_subgraph_->CreateInitialFeeds(input_ids, + this->implicit_inputs_, + this->parameters_->num_beams, + this->parameters_->pad_token_id, + sequence_lengths, + expanded_input_ids, + attn_mask_value, + feeds, + this->create_inputs_func_, + this->add_to_feeds_func_, + buffer); + } + return gpt_subgraph_.CreateInitialFeeds(input_ids, this->implicit_inputs_, this->parameters_->num_beams, @@ -116,7 +138,8 @@ Status BeamSearchGpt::UpdateFeeds( } template -Status BeamSearchGpt::Execute(const FeedsFetchesManager& feeds_fetches_manager) { +Status BeamSearchGpt::Execute(const FeedsFetchesManager* init_run_feeds_fetches_manager, + const FeedsFetchesManager& feeds_fetches_manager) { auto status = Status::OK(); const BeamSearchParameters* parameters = this->parameters_; int64_t sequences_dims[] = {parameters->batch_size, parameters->num_return_sequences, parameters->max_length}; @@ -206,7 +229,6 @@ Status BeamSearchGpt::Execute(const FeedsFetchesManager& feeds_fetches_manage int current_length = parameters->sequence_length; int iteration_counter = 0; while (current_length < parameters->max_length) { - iteration_counter++; #ifdef DEBUG_GENERATION auto cur_len = std::to_string(current_length); dumper->Print("***CurrentLength", cur_len, true); @@ -221,14 +243,27 @@ Status BeamSearchGpt::Execute(const FeedsFetchesManager& feeds_fetches_manage } #endif - status = utils::ExecuteSubgraph(this->decoder_session_state_, - feeds_fetches_manager, - feeds, - fetches, - {}, - ExecutionMode::ORT_SEQUENTIAL, - this->context_.GetTerminateFlag(), - this->context_.Logger()); + // For the first iteration use the init_run_decoder subgraph (if present) + if (iteration_counter++ == 0 && + init_run_decoder_session_state_ != nullptr) { + status = utils::ExecuteSubgraph(*init_run_decoder_session_state_, + *init_run_feeds_fetches_manager, + feeds, + fetches, + {}, + ExecutionMode::ORT_SEQUENTIAL, + this->context_.GetTerminateFlag(), + this->context_.Logger()); + } else { + status = utils::ExecuteSubgraph(this->decoder_session_state_, + feeds_fetches_manager, + feeds, + fetches, + {}, + ExecutionMode::ORT_SEQUENTIAL, + this->context_.GetTerminateFlag(), + this->context_.Logger()); + } ORT_RETURN_IF_ERROR(status); diff --git a/onnxruntime/contrib_ops/cpu/transformers/greedy_search.cc b/onnxruntime/contrib_ops/cpu/transformers/greedy_search.cc index 0299912cab..cda1a6c1ce 100644 --- a/onnxruntime/contrib_ops/cpu/transformers/greedy_search.cc +++ b/onnxruntime/contrib_ops/cpu/transformers/greedy_search.cc @@ -12,6 +12,7 @@ #include #include #include +#include #include "core/common/safeint.h" #include "core/providers/cpu/math/top_k.h" #include "core/providers/cpu/tensor/utils.h" @@ -53,20 +54,50 @@ REGISTER_KERNEL_TYPED(float) namespace transformers { +namespace gpt_details { +std::pair> CreateGptSubgraphAndUpdateParameters( + const Node& node, + const SessionState& session_state, + const std::string& attribute_name, + const SessionState& subgraph_session_state, + /*out*/ BeamSearchParameters& parameters) { + auto gpt_subgraph = std::make_unique(node, attribute_name, subgraph_session_state.GetGraphViewer()); + auto status = gpt_subgraph->Setup(session_state, subgraph_session_state); + if (!status.IsOK()) { + return std::make_pair(status, std::move(gpt_subgraph)); + } + + parameters.SetSubgraphParameters(gpt_subgraph->vocab_size, + gpt_subgraph->num_heads, + gpt_subgraph->head_size, + gpt_subgraph->num_layers); + + return std::make_pair(status, std::move(gpt_subgraph)); +} +} // namespace gpt_details + void GreedySearch::Init(const OpKernelInfo& info) { parameters_.ParseFromAttributes(info); - // Check model_type 0 (GPT-2) and 1 (encoder-decoder like T5) - ORT_ENFORCE(parameters_.model_type == 0 || parameters_.model_type == 1); + // Model_type could be either 0 (GPT-2) or 1 (encoder-decoder like T5) + ORT_ENFORCE(parameters_.model_type == IBeamSearchParameters::kModelTypeGpt || + parameters_.model_type == IBeamSearchParameters::kModelTypeT5); - // Make sure the decoder attribute was present even though we don't need it here. ONNX_NAMESPACE::GraphProto proto; - if (parameters_.model_type != 0) { + if (parameters_.model_type != IBeamSearchParameters::kModelTypeGpt) { + // Make sure the encoder sub-graph attribute is present for the T5 model. ORT_ENFORCE(info.GetAttr("encoder", &proto).IsOK()); } + if (parameters_.model_type == IBeamSearchParameters::kModelTypeGpt) { + // Check if the init_decoder sub-graph attribute is present for the GPT2 model. + if (info.GetAttr("init_decoder", &proto).IsOK()) { + has_init_decoder_ = true; + } + } + + // Make sure the decoder sub-graph attribute is present for all model types. ORT_ENFORCE(info.GetAttr("decoder", &proto).IsOK()); - ORT_IGNORE_RETURN_VALUE(proto); } Status GreedySearch::SetupSubgraphExecutionInfo(const SessionState& session_state, @@ -75,16 +106,31 @@ Status GreedySearch::SetupSubgraphExecutionInfo(const SessionState& session_stat const auto& node = Node(); if (parameters_.model_type == IBeamSearchParameters::kModelTypeGpt) { // GPT-2 if (attribute_name == "decoder") { - ORT_ENFORCE(gpt_subgraph_ == nullptr, - "SetupSubgraphExecutionInfo should only be called once for each subgraph."); - gpt_subgraph_ = std::make_unique(node, attribute_name, subgraph_session_state.GetGraphViewer()); - ORT_RETURN_IF_ERROR(gpt_subgraph_->Setup(session_state, subgraph_session_state)); + ORT_ENFORCE(gpt_subgraph_ == nullptr, "SetupSubgraphExecutionInfo should only be called once for each subgraph."); + auto res = gpt_details::CreateGptSubgraphAndUpdateParameters(node, session_state, attribute_name, + subgraph_session_state, parameters_); + + auto status = res.first; + if (!status.IsOK()) { + return status; + } + + gpt_subgraph_ = std::move(res.second); decoder_feeds_fetches_manager_ = gpt_subgraph_->GetFeedsFetchesManager(); - parameters_.SetSubgraphParameters(gpt_subgraph_->vocab_size, - gpt_subgraph_->num_heads, - gpt_subgraph_->head_size, - gpt_subgraph_->num_layers); + } else if (attribute_name == "init_decoder") { + ORT_ENFORCE(init_run_gpt_subgraph_ == nullptr, "SetupSubgraphExecutionInfo should only be called once for each subgraph."); + auto res = gpt_details::CreateGptSubgraphAndUpdateParameters(node, session_state, attribute_name, + subgraph_session_state, parameters_); + + auto status = res.first; + if (!status.IsOK()) { + return status; + } + + init_run_gpt_subgraph_ = std::move(res.second); + init_run_decoder_feeds_fetches_manager_ = init_run_gpt_subgraph_->GetFeedsFetchesManager(); } + } else if (parameters_.model_type == IBeamSearchParameters::kModelTypeT5) { // encoder-decoder like T5 ORT_THROW("Not Implemented"); // if (attribute_name == "encoder") { @@ -122,6 +168,12 @@ Status GreedySearch::Compute(OpKernelContext* ctx) const { ORT_ENFORCE(decoder_session_state, "Subgraph SessionState was not found for 'decoder' attribute."); ORT_ENFORCE(decoder_feeds_fetches_manager_, "CreateFeedsFetchesManager must be called prior to execution of graph."); + auto* init_run_decoder_session_state = ctx_internal->SubgraphSessionState("init_decoder"); + if (has_init_decoder_) { + ORT_ENFORCE(init_run_decoder_session_state, "Subgraph SessionState was not found for 'decoder' attribute."); + ORT_ENFORCE(init_run_decoder_feeds_fetches_manager_, "CreateFeedsFetchesManager must be called prior to execution of graph."); + } + concurrency::ThreadPool* thread_pool = ctx->GetOperatorThreadPool(); // make a copy since we will update the parameters based on inputs later @@ -132,6 +184,8 @@ Status GreedySearch::Compute(OpKernelContext* ctx) const { if (!gpt_subgraph_->IsOutputFloat16()) { GreedySearchGpt impl{ *ctx_internal, + has_init_decoder_ ? init_run_decoder_session_state : nullptr, + has_init_decoder_ ? init_run_gpt_subgraph_.get() : nullptr, *decoder_session_state, *gpt_subgraph_, thread_pool, @@ -147,10 +201,12 @@ Status GreedySearch::Compute(OpKernelContext* ctx) const { update_gpt_feeds_func_ ? update_gpt_feeds_func_ : GenerationCpuDeviceHelper::UpdateGptFeeds}; ORT_RETURN_IF_ERROR(impl.Initialize()); - return impl.Execute(*decoder_feeds_fetches_manager_); + return impl.Execute(init_run_decoder_feeds_fetches_manager_, *decoder_feeds_fetches_manager_); } else { GreedySearchGpt impl{ *ctx_internal, + has_init_decoder_ ? init_run_decoder_session_state : nullptr, + has_init_decoder_ ? init_run_gpt_subgraph_.get() : nullptr, *decoder_session_state, *gpt_subgraph_, thread_pool, @@ -166,7 +222,7 @@ Status GreedySearch::Compute(OpKernelContext* ctx) const { update_gpt_feeds_fp16_func_}; ORT_RETURN_IF_ERROR(impl.Initialize()); - return impl.Execute(*decoder_feeds_fetches_manager_); + return impl.Execute(init_run_decoder_feeds_fetches_manager_, *decoder_feeds_fetches_manager_); } } diff --git a/onnxruntime/contrib_ops/cpu/transformers/greedy_search.h b/onnxruntime/contrib_ops/cpu/transformers/greedy_search.h index 281834d9a9..d56f0b5485 100644 --- a/onnxruntime/contrib_ops/cpu/transformers/greedy_search.h +++ b/onnxruntime/contrib_ops/cpu/transformers/greedy_search.h @@ -91,17 +91,34 @@ class GreedySearch : public IControlFlowKernel { //------------------------------------------------------------ // Subgraph and FeedsFetchesManager re-used for each subgraph execution. //------------------------------------------------------------ + + // Relevant only for GPT2 + // The init_run_gpt_subgraph_ (if the `init_decoder` attribute is present) will be + // used for the first decoding run and the gpt_subgraph_ will be used + // for subsequent runs. + // If the `init_decoder` attribute is missing, the `gpt_subgraph_` will be + // used for all decoding runs. + std::unique_ptr init_run_gpt_subgraph_; std::unique_ptr gpt_subgraph_; + + // Relevant only for T5 + // Same concept as above. + // The encoder will be used for the first run and the decoder will + // be used for subsequent runs. // std::unique_ptr t5_encoder_subgraph_; // std::unique_ptr t5_decoder_subgraph_; + // FeedsFetchesManager* encoder_feeds_fetches_manager_; FeedsFetchesManager* decoder_feeds_fetches_manager_; + FeedsFetchesManager* init_run_decoder_feeds_fetches_manager_; void* cuda_stream_; IConsoleDumper* dumper_; GreedySearchParameters parameters_; + + bool has_init_decoder_ = false; }; } // namespace transformers diff --git a/onnxruntime/contrib_ops/cpu/transformers/greedy_search_impl_gpt.h b/onnxruntime/contrib_ops/cpu/transformers/greedy_search_impl_gpt.h index cbbcf29672..bff151bb54 100644 --- a/onnxruntime/contrib_ops/cpu/transformers/greedy_search_impl_gpt.h +++ b/onnxruntime/contrib_ops/cpu/transformers/greedy_search_impl_gpt.h @@ -13,11 +13,23 @@ namespace contrib { namespace transformers { -// Beam search implementation for GPT-2 model. +namespace gpt_details { +// Some common helpers that can be shared around +std::pair> CreateGptSubgraphAndUpdateParameters( + const Node& node, + const SessionState& session_state, + const std::string& attribute_name, + const SessionState& subgraph_session_state, + /*out*/ BeamSearchParameters& parameters); +} // namespace gpt_details + +// Greedy search implementation for GPT-2 model. template class GreedySearchGpt : public GreedySearchBase { public: GreedySearchGpt(OpKernelContextInternal& context, + const SessionState* init_run_decoder_session_state, + GptSubgraph* init_run_gpt_subgraph, const SessionState& decoder_session_state, GptSubgraph& gpt_subgraph, concurrency::ThreadPool* thread_pool, @@ -40,6 +52,8 @@ class GreedySearchGpt : public GreedySearchBase { topk_func, process_logits_func, device_copy_func), + init_run_decoder_session_state_(init_run_decoder_session_state), + init_run_gpt_subgraph_(init_run_gpt_subgraph), gpt_subgraph_(gpt_subgraph), create_inputs_func_(create_inputs_func), add_to_feeds_func_(add_to_feeds_func), @@ -49,7 +63,8 @@ class GreedySearchGpt : public GreedySearchBase { // Execute beam search in iterations util stopping criteria is reached. // In each iteration, GPT subgraph is called, and next token for each sequence is generated. - Status Execute(const FeedsFetchesManager& feeds_fetches_manager); + Status Execute(const FeedsFetchesManager* init_run_feeds_fetches_manager, + const FeedsFetchesManager& feeds_fetches_manager); private: // Prepare the inputs for first inference of subgraph @@ -67,6 +82,8 @@ class GreedySearchGpt : public GreedySearchBase { bool increase_position, gsl::span next_tokens); + const SessionState* init_run_decoder_session_state_ = nullptr; + GptSubgraph* init_run_gpt_subgraph_ = nullptr; GptSubgraph& gpt_subgraph_; // Device specific functions @@ -84,6 +101,21 @@ Status GreedySearchGpt::CreateInitialFeeds(gsl::span& sequence_lengt const OrtValue* input_ids_value = this->context_.GetInputOrtValue(0); const Tensor& input_ids = input_ids_value->Get(); const OrtValue* attn_mask_value = this->context_.GetInputOrtValue(6); + + if (init_run_gpt_subgraph_ != nullptr) { + return init_run_gpt_subgraph_->CreateInitialFeeds(input_ids, + this->implicit_inputs_, + this->parameters_->num_beams, + this->parameters_->pad_token_id, + sequence_lengths, + expanded_input_ids, + attn_mask_value, + feeds, + this->create_inputs_func_, + this->add_to_feeds_func_, + buffer); + } + return gpt_subgraph_.CreateInitialFeeds(input_ids, this->implicit_inputs_, this->parameters_->num_beams, @@ -121,7 +153,8 @@ Status GreedySearchGpt::UpdateFeeds( } template -Status GreedySearchGpt::Execute(const FeedsFetchesManager& feeds_fetches_manager) { +Status GreedySearchGpt::Execute(const FeedsFetchesManager* init_run_feeds_fetches_manager, + const FeedsFetchesManager& feeds_fetches_manager) { auto status = Status::OK(); const GreedySearchParameters* parameters = this->parameters_; @@ -173,7 +206,6 @@ Status GreedySearchGpt::Execute(const FeedsFetchesManager& feeds_fetches_mana int current_length = parameters->sequence_length; int iteration_counter = 0; while (current_length < parameters->max_length) { - iteration_counter++; #ifdef DEBUG_GENERATION auto cur_len = std::to_string(current_length); dumper->Print("***CurrentLength", cur_len, true); @@ -182,14 +214,27 @@ Status GreedySearchGpt::Execute(const FeedsFetchesManager& feeds_fetches_mana dumper->Print("attention_mask", feeds[2]); #endif - status = utils::ExecuteSubgraph(this->decoder_session_state_, - feeds_fetches_manager, - feeds, - fetches, - {}, - ExecutionMode::ORT_SEQUENTIAL, - this->context_.GetTerminateFlag(), - this->context_.Logger()); + // For the first iteration use the init_run_decoder subgraph (if present) + if (iteration_counter++ == 0 && + init_run_decoder_session_state_ != nullptr) { + status = utils::ExecuteSubgraph(*init_run_decoder_session_state_, + *init_run_feeds_fetches_manager, + feeds, + fetches, + {}, + ExecutionMode::ORT_SEQUENTIAL, + this->context_.GetTerminateFlag(), + this->context_.Logger()); + } else { + status = utils::ExecuteSubgraph(this->decoder_session_state_, + feeds_fetches_manager, + feeds, + fetches, + {}, + ExecutionMode::ORT_SEQUENTIAL, + this->context_.GetTerminateFlag(), + this->context_.Logger()); + } ORT_RETURN_IF_ERROR(status); @@ -230,8 +275,8 @@ Status GreedySearchGpt::Execute(const FeedsFetchesManager& feeds_fetches_mana gsl::span output = output_sequences->MutableDataAsSpan(); for (int batch_id = 0; batch_id < parameters->batch_size; ++batch_id) { auto batch_output = output.subspan( - static_cast(batch_id) * parameters->max_length, - parameters->max_length); + static_cast(batch_id) * parameters->max_length, + parameters->max_length); gsl::span sequence_source = greedy_state.sequences.GetSequence(batch_id); gsl::copy(sequence_source, batch_output); } diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index 5d577ebb3b..03a18f1247 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -136,7 +136,7 @@ void convTransposeWithDynamicPadsShapeInference(InferenceContext& ctx) { *final_output_shape->add_dim() = ctx.getInputType(1)->tensor_type().shape().dim(1) * group; // channels should be the second dim of second input multiply - // group. + // group. int size_of_output; if (output_shape_presented) { @@ -1044,6 +1044,10 @@ ONNX_MS_OPERATOR_SET_SCHEMA(BeamSearch, 1, .Attr("early_stopping", "early stop or not", AttributeProto::INT, static_cast(0)) .Attr("model_type", "model type: 0 for GPT-2; 1 for encoder decoder like T5", AttributeProto::INT, static_cast(0)) .Attr("encoder", "The subgraph for initialization of encoder and decoder. It will be called once before decoder subgraph.", AttributeProto::GRAPH, OPTIONAL_VALUE) + .Attr("init_decoder", + "The subgraph for the first decoding run. It will be called once before `decoder` subgraph. " + "This is relevant only for the GPT2 model. If this attribute is missing, the `decoder` subgraph will be used for all decoding runs", + AttributeProto::GRAPH, OPTIONAL_VALUE) .Attr("decoder", "Decoder subgraph to execute in a loop.", AttributeProto::GRAPH) .Attr("vocab_size", "Size of the vocabulary. " @@ -1085,7 +1089,11 @@ ONNX_MS_OPERATOR_SET_SCHEMA(GreedySearch, 1, .Attr("decoder_start_token_id", "The id of the token that indicates decoding starts.", AttributeProto::INT, static_cast(-1)) .Attr("no_repeat_ngram_size", "no repeat ngrams size", AttributeProto::INT, static_cast(0)) .Attr("model_type", "model type: 0 for decoder only like GPT-2; 1 for encoder decoder like Bart", AttributeProto::INT, static_cast(0)) - .Attr("encoder", "The subgraph for initialization of encoder and decoder. It will be called once before decoder subgraph.", AttributeProto::GRAPH, OPTIONAL_VALUE) + .Attr("encoder", "The subgraph for initialization of encoder and decoder. It will be called once before `decoder` subgraph.", AttributeProto::GRAPH, OPTIONAL_VALUE) + .Attr("init_decoder", + "The subgraph for the first decoding run. It will be called once before `decoder` subgraph. " + "This is relevant only for the GPT2 model. If this attribute is missing, the `decoder` subgraph will be used for all decoding runs", + AttributeProto::GRAPH, OPTIONAL_VALUE) .Attr("decoder", "Decoder subgraph to execute in a loop.", AttributeProto::GRAPH) .Attr("vocab_size", "Size of the vocabulary. " diff --git a/onnxruntime/python/tools/transformers/convert_generation.py b/onnxruntime/python/tools/transformers/convert_generation.py index dcc2657966..8163b0cace 100644 --- a/onnxruntime/python/tools/transformers/convert_generation.py +++ b/onnxruntime/python/tools/transformers/convert_generation.py @@ -189,7 +189,7 @@ def parse_arguments(argv: Optional[List[str]] = None) -> argparse.Namespace: action="store_true", help="Have separate decoder subgraphs for initial and remaining runs. This allows for optimizations based on sequence lengths in each subgraph", ) - output_group.set_defaults(separate_gpt2_decoder_for_init_run=False) + output_group.set_defaults(separate_gpt2_decoder_for_init_run=True) output_group.add_argument( "-i", @@ -1210,6 +1210,7 @@ def convert_generation_model(args: argparse.Namespace, generation_type: Generati decoder_model = onnx.load_model(args.decoder_onnx, load_external_data=True) decoder_model.graph.name = f"{args.model_type} decoder" + gpt2_init_decoder_model = None if args.model_type == "gpt2": verify_gpt2_subgraph(decoder_model.graph, args.precision) @@ -1340,8 +1341,6 @@ def convert_generation_model(args: argparse.Namespace, generation_type: Generati ) else: if gpt2_init_decoder_generated: - gpt2_init_decoder_model = onnx.load_model(gpt2_init_decoder_onnx_path, load_external_data=True) - # Move shared initializers (shared between init decoder and decoder models) to the main # graph and remove them from these models if not args.disable_shared_initializers: diff --git a/onnxruntime/test/contrib_ops/beam_search_test.cc b/onnxruntime/test/contrib_ops/beam_search_test.cc index 798bbb7e8b..2e709d3235 100644 --- a/onnxruntime/test/contrib_ops/beam_search_test.cc +++ b/onnxruntime/test/contrib_ops/beam_search_test.cc @@ -73,6 +73,10 @@ TEST(BeamSearchTest, GptBeamSearchFp32) { Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(session_options, 0)); #endif + // The ONNX model is generated like the following: + // python convert_generation.py --model_type gpt2 -m hf-internal-testing/tiny-random-gpt2 + // --output tiny_gpt2_beamsearch_fp16.onnx --use_gpu --max_length 20 + // (with separate_gpt2_decoder_for_init_run set to False as it is now set to True by default) Ort::Session session(*ort_env, ORT_TSTR("testdata/transformers/tiny_gpt2_beamsearch.onnx"), session_options); auto ort_outputs = session.Run(Ort::RunOptions{}, input_names, ort_inputs.data(), ort_inputs.size(), output_names, 1); @@ -156,6 +160,7 @@ TEST(BeamSearchTest, GptBeamSearchFp16) { // The ONNX model is generated like the following: // python convert_generation.py --model_type gpt2 -m hf-internal-testing/tiny-random-gpt2 // --output tiny_gpt2_beamsearch_fp16.onnx -p fp16 --use_gpu --max_length 20 + // (with separate_gpt2_decoder_for_init_run set to False as it is now set to True by default) Ort::Session session(*ort_env, ORT_TSTR("testdata/transformers/tiny_gpt2_beamsearch_fp16.onnx"), session_options); auto ort_outputs = session.Run(Ort::RunOptions{}, input_names, ort_inputs.data(), ort_inputs.size(), @@ -175,6 +180,91 @@ TEST(BeamSearchTest, GptBeamSearchFp16) { } } +TEST(BeamSearchTest, GptBeamSearchWithInitDecoderFp16) { + std::vector input_ids_shape{3, 12}; + std::vector input_ids{ + 0, 0, 0, 0, 0, 52, 195, 731, 321, 301, 734, 620, + 41, 554, 74, 622, 206, 222, 75, 223, 221, 198, 224, 572, + 0, 0, 0, 52, 328, 219, 328, 206, 288, 227, 896, 328}; + + std::vector parameter_shape{1}; + std::vector max_length{20}; + std::vector min_length{1}; + std::vector num_beams{4}; + std::vector num_return_sequences{1}; + std::vector length_penalty{1.0f}; + std::vector repetition_penalty{1.0f}; + + std::vector expected_output_shape{input_ids_shape[0], num_return_sequences[0], max_length[0]}; + + std::vector expected_output{ + 0, 0, 0, 0, 0, 52, 195, 731, 321, 301, 734, 620, 131, 131, 131, 181, 638, 638, 638, 638, + 41, 554, 74, 622, 206, 222, 75, 223, 221, 198, 224, 572, 292, 292, 292, 292, 292, 292, 292, 292, + 0, 0, 0, 52, 328, 219, 328, 206, 288, 227, 896, 328, 328, 669, 669, 669, 669, 669, 669, 669}; + + Ort::MemoryInfo info("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault); + auto input_ids_tensor = Ort::Value::CreateTensor( + info, input_ids.data(), input_ids.size(), input_ids_shape.data(), input_ids_shape.size()); + + auto max_length_tensor = Ort::Value::CreateTensor( + info, max_length.data(), max_length.size(), parameter_shape.data(), parameter_shape.size()); + + auto min_length_tensor = Ort::Value::CreateTensor( + info, min_length.data(), min_length.size(), parameter_shape.data(), parameter_shape.size()); + + auto num_beams_tensor = Ort::Value::CreateTensor( + info, num_beams.data(), num_beams.size(), parameter_shape.data(), parameter_shape.size()); + + auto num_return_sequences_tensor = Ort::Value::CreateTensor( + info, num_return_sequences.data(), num_return_sequences.size(), parameter_shape.data(), parameter_shape.size()); + + auto length_penalty_tensor = Ort::Value::CreateTensor( + info, length_penalty.data(), length_penalty.size(), parameter_shape.data(), parameter_shape.size()); + + auto repetition_penalty_tensor = Ort::Value::CreateTensor( + info, repetition_penalty.data(), repetition_penalty.size(), parameter_shape.data(), parameter_shape.size()); + + std::vector ort_inputs; + ort_inputs.push_back(std::move(input_ids_tensor)); + ort_inputs.push_back(std::move(max_length_tensor)); + ort_inputs.push_back(std::move(min_length_tensor)); + ort_inputs.push_back(std::move(num_beams_tensor)); + ort_inputs.push_back(std::move(num_return_sequences_tensor)); + ort_inputs.push_back(std::move(length_penalty_tensor)); + ort_inputs.push_back(std::move(repetition_penalty_tensor)); + const char* input_names[] = {"input_ids", "max_length", "min_length", "num_beams", "num_return_sequences", + "length_penalty", "repetition_penalty"}; + const char* const output_names[] = {"sequences"}; + + constexpr int min_cuda_architecture = 530; + if (HasCudaEnvironment(min_cuda_architecture)) { + Ort::SessionOptions session_options; +#ifdef USE_CUDA + Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(session_options, 0)); +#endif + + // The ONNX model is generated like the following: + // python convert_generation.py --model_type gpt2 -m hf-internal-testing/tiny-random-gpt2 + // --output tiny_gpt2_beamsearch_with_init_decoder_fp16.onnx -p fp16 --use_gpu --max_length 20 + // (with separate_gpt2_decoder_for_init_run set to True as is the default option) + Ort::Session session(*ort_env, ORT_TSTR("testdata/transformers/tiny_gpt2_beamsearch_with_init_decoder_fp16.onnx"), session_options); + + auto ort_outputs = session.Run(Ort::RunOptions{}, input_names, ort_inputs.data(), ort_inputs.size(), + output_names, 1); + + ASSERT_EQ(ort_outputs.size(), 1U); + const auto& sequences = ort_outputs[0]; + ASSERT_TRUE(sequences.IsTensor()); + + auto result_ts = sequences.GetTensorTypeAndShapeInfo(); + ASSERT_EQ(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32, result_ts.GetElementType()); + + ASSERT_EQ(expected_output_shape, result_ts.GetShape()); + const auto* result_vals = sequences.GetTensorData(); + auto result_span = gsl::make_span(result_vals, expected_output.size()); + ASSERT_TRUE(std::equal(expected_output.cbegin(), expected_output.cend(), result_span.begin(), result_span.end())); + } +} TEST(BeamSearchTest, GptBeamSearchFp16_VocabPadded) { std::vector input_ids_shape{3, 12}; std::vector input_ids{ @@ -259,5 +349,6 @@ TEST(BeamSearchTest, GptBeamSearchFp16_VocabPadded) { ASSERT_TRUE(std::equal(expected_output.cbegin(), expected_output.cend(), result_span.begin(), result_span.end())); } } + } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/contrib_ops/greedy_search_test.cc b/onnxruntime/test/contrib_ops/greedy_search_test.cc index 4a46d96266..f5259c1391 100644 --- a/onnxruntime/test/contrib_ops/greedy_search_test.cc +++ b/onnxruntime/test/contrib_ops/greedy_search_test.cc @@ -58,8 +58,9 @@ TEST(GreedySearchTest, GptGreedySearchFp16_VocabPadded) { #endif // The following model was obtained by padding the vocabulary size in testdata/transformers/tiny_gpt2_beamsearch_fp16.onnx - // (by making beam_size == 1) from 1000 to 1600 (just for illustrative and testing purposes) to see if the beam search - // implementation can handle such a scenario + // (by making beam_size == 1) from 1000 to 1600 (just for illustrative and testing purposes) to see if the greedy search + // implementation can handle such a scenario. + // Check beam_search_test.cc to see how tiny_gpt2_beamsearch_fp16.onnx was generated. Ort::Session session(*ort_env, ORT_TSTR("testdata/transformers/tiny_gpt2_greedysearch_fp16_padded_vocab.onnx"), session_options); auto ort_outputs = session.Run(Ort::RunOptions{}, input_names, ort_inputs.data(), ort_inputs.size(), @@ -78,5 +79,69 @@ TEST(GreedySearchTest, GptGreedySearchFp16_VocabPadded) { ASSERT_TRUE(std::equal(expected_output.cbegin(), expected_output.cend(), result_span.begin(), result_span.end())); } } + +TEST(GreedySearchTest, GptGreedySearchFp32) { + std::vector input_ids_shape{2, 4}; + std::vector input_ids{ + 0, 0, 0, 52, 0, 0, 195, 731}; + + std::vector parameter_shape{1}; + std::vector max_length{10}; + std::vector min_length{1}; + std::vector repetition_penalty{1.0f}; + + std::vector expected_output_shape{input_ids_shape[0], max_length[0]}; + + std::vector expected_output{ + 0, 0, 0, 52, 204, 204, 204, 204, 204, 204, + 0, 0, 195, 731, 731, 114, 114, 114, 114, 114}; + + Ort::MemoryInfo info("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault); + auto input_ids_tensor = Ort::Value::CreateTensor( + info, input_ids.data(), input_ids.size(), input_ids_shape.data(), input_ids_shape.size()); + + auto max_length_tensor = Ort::Value::CreateTensor( + info, max_length.data(), max_length.size(), parameter_shape.data(), parameter_shape.size()); + + auto min_length_tensor = Ort::Value::CreateTensor( + info, min_length.data(), min_length.size(), parameter_shape.data(), parameter_shape.size()); + + auto repetition_penalty_tensor = Ort::Value::CreateTensor( + info, repetition_penalty.data(), repetition_penalty.size(), parameter_shape.data(), parameter_shape.size()); + + std::vector ort_inputs; + ort_inputs.push_back(std::move(input_ids_tensor)); + ort_inputs.push_back(std::move(max_length_tensor)); + ort_inputs.push_back(std::move(min_length_tensor)); + ort_inputs.push_back(std::move(repetition_penalty_tensor)); + const char* input_names[] = {"input_ids", "max_length", "min_length", "repetition_penalty"}; + const char* const output_names[] = {"sequences"}; + + constexpr int min_cuda_architecture = 530; + if (HasCudaEnvironment(min_cuda_architecture)) { + Ort::SessionOptions session_options; +#ifdef USE_CUDA + Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(session_options, 0)); +#endif + + Ort::Session session(*ort_env, ORT_TSTR("testdata/transformers/tiny_gpt2_greedysearch_with_init_decoder.onnx"), session_options); + + auto ort_outputs = session.Run(Ort::RunOptions{}, input_names, ort_inputs.data(), ort_inputs.size(), + output_names, 1); + + ASSERT_EQ(ort_outputs.size(), 1U); + const auto& sequences = ort_outputs[0]; + ASSERT_TRUE(sequences.IsTensor()); + + auto result_ts = sequences.GetTensorTypeAndShapeInfo(); + ASSERT_EQ(ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32, result_ts.GetElementType()); + + ASSERT_EQ(expected_output_shape, result_ts.GetShape()); + const auto* result_vals = sequences.GetTensorData(); + auto result_span = gsl::make_span(result_vals, expected_output.size()); + ASSERT_TRUE(std::equal(expected_output.cbegin(), expected_output.cend(), result_span.begin(), result_span.end())); + } +} + } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/testdata/transformers/tiny_gpt2_beamsearch_with_init_decoder_fp16.onnx b/onnxruntime/test/testdata/transformers/tiny_gpt2_beamsearch_with_init_decoder_fp16.onnx new file mode 100644 index 0000000000..3b419dd3d4 Binary files /dev/null and b/onnxruntime/test/testdata/transformers/tiny_gpt2_beamsearch_with_init_decoder_fp16.onnx differ diff --git a/onnxruntime/test/testdata/transformers/tiny_gpt2_greedysearch_with_init_decoder.onnx b/onnxruntime/test/testdata/transformers/tiny_gpt2_greedysearch_with_init_decoder.onnx new file mode 100644 index 0000000000..770a99576d Binary files /dev/null and b/onnxruntime/test/testdata/transformers/tiny_gpt2_greedysearch_with_init_decoder.onnx differ