From f1044e3b9a9cb5b0db0b72a2d5ead4c62d76a278 Mon Sep 17 00:00:00 2001 From: Hariharan Seshadri Date: Mon, 19 Dec 2022 13:29:10 -0800 Subject: [PATCH] CUDA GreedySearch ProcessLogits optimization (#13823) ### Description Explore the possible re-use of the logits buffer in `GreedySearch` for cases where sequence length == 1 (Post the first decoding run, the sequence length is guaranteed to be 1). This re-use will ensure that we do not have to make copies of the logits before processing them. Currently, we make a copy of the logits even if the sequence length == 1 which is not necessary as we can directly re-use the logits buffer for the token generation step. A similar optimization exists in `BeamSearch`, but seems lacking in `GreedySearch`. Since, the logits buffer may contain padded data, we need to adjust the pieces consuming the logits buffer directly to account for any padding. A more invasive change (needs changes in a few places) will be to adjust the interfaces of `ProcessLogits()` such that it takes a reference to the logits and not a const reference as (based on my understanding) this is the only place where the logits from the decoder subgraph will ever be used and giving the `ProcessLogits()` method license to mutate/process the underlying buffer of the logits OrtValue seems reasonable (instead of making a copy and then mutating/processing them). The will also remove the ugly `const_cast`(s) seen in this change. --- .../cuda/transformers/beam_search_impl.cu | 110 ++++++++++-------- .../cuda/transformers/beam_search_impl.h | 1 + .../transformers/generation_device_helper.cc | 65 ++++++++--- 3 files changed, 108 insertions(+), 68 deletions(-) diff --git a/onnxruntime/contrib_ops/cuda/transformers/beam_search_impl.cu b/onnxruntime/contrib_ops/cuda/transformers/beam_search_impl.cu index 6bc52758c7..3b530cef6b 100644 --- a/onnxruntime/contrib_ops/cuda/transformers/beam_search_impl.cu +++ b/onnxruntime/contrib_ops/cuda/transformers/beam_search_impl.cu @@ -66,6 +66,7 @@ __global__ void LogitsProcessKernel( const int* prefix_vocab_mask, int num_beams, int vocab_size, + int padded_vocab_size, int total_elements, int demote_token_id, int32_t* sequences, @@ -75,66 +76,71 @@ __global__ void LogitsProcessKernel( int no_repeat_ngram_size) { int index = blockIdx.x * blockDim.x + threadIdx.x; if (index < total_elements) { - int batch_beam_index = index / vocab_size; - int word_id = index % vocab_size; + int batch_beam_index = index / padded_vocab_size; + int word_id = index % padded_vocab_size; - // RepetitionPenaltyLogitsProcessor - if (repetition_penalty != 1.0f) { - int32_t* current_sequence = sequences + batch_beam_index * max_sequence_length; - bool found = false; - for (int i = 0; i < current_sequence_length; i++) { - if (current_sequence[i] == word_id) { - found = true; - break; - } - } - if (found) { - float score = (float)next_token_scores[index]; - next_token_scores[index] = (T)(score < 0 ? score * repetition_penalty : score / repetition_penalty); - } - } - - // NoRepeatNGramLogitsProcessor - if (no_repeat_ngram_size > 0 && current_sequence_length >= no_repeat_ngram_size) { - int32_t* current_sequence = sequences + batch_beam_index * max_sequence_length; - bool found = false; - for (int i = no_repeat_ngram_size - 1; i < current_sequence_length; i++) { - if (current_sequence[i] == word_id) { // last token of n-gram matched - found = true; - for (int j = 0; j < no_repeat_ngram_size - 1; j++) { // match the remaining N-1 tokens - if (current_sequence[i - j - 1] != current_sequence[current_sequence_length - 1 - j]) { - found = false; - break; - } - } - if (found) { + if (word_id >= vocab_size) { + // Set any value within the padding region to the lowest value so that it isn't picked + next_token_scores[index] = cub::FpLimits::Lowest(); + } else { + // RepetitionPenaltyLogitsProcessor + if (repetition_penalty != 1.0f) { + int32_t* current_sequence = sequences + batch_beam_index * max_sequence_length; + bool found = false; + for (int i = 0; i < current_sequence_length; i++) { + if (current_sequence[i] == word_id) { + found = true; break; } } + if (found) { + float score = (float)next_token_scores[index]; + next_token_scores[index] = (T)(score < 0 ? score * repetition_penalty : score / repetition_penalty); + } } - if (found) { + // NoRepeatNGramLogitsProcessor + if (no_repeat_ngram_size > 0 && current_sequence_length >= no_repeat_ngram_size) { + int32_t* current_sequence = sequences + batch_beam_index * max_sequence_length; + bool found = false; + for (int i = no_repeat_ngram_size - 1; i < current_sequence_length; i++) { + if (current_sequence[i] == word_id) { // last token of n-gram matched + found = true; + for (int j = 0; j < no_repeat_ngram_size - 1; j++) { // match the remaining N-1 tokens + if (current_sequence[i - j - 1] != current_sequence[current_sequence_length - 1 - j]) { + found = false; + break; + } + } + if (found) { + break; + } + } + } + + if (found) { + next_token_scores[index] = cub::FpLimits::Lowest(); + return; + } + } + + // VocabMaskLogitsProcessor + if (vocab_mask != nullptr && vocab_mask[word_id] == 0) { next_token_scores[index] = cub::FpLimits::Lowest(); return; } - } - // VocabMaskLogitsProcessor - if (vocab_mask != nullptr && vocab_mask[word_id] == 0) { - next_token_scores[index] = cub::FpLimits::Lowest(); - return; - } + // PrefixVocabMaskLogitsProcessor + int batch_id = batch_beam_index / num_beams; + if (prefix_vocab_mask != nullptr && prefix_vocab_mask[batch_id * vocab_size + word_id] == 0) { + next_token_scores[index] = cub::FpLimits::Lowest(); + return; + } - // PrefixVocabMaskLogitsProcessor - int batch_id = batch_beam_index / num_beams; - if (prefix_vocab_mask != nullptr && prefix_vocab_mask[batch_id * vocab_size + word_id] == 0) { - next_token_scores[index] = cub::FpLimits::Lowest(); - return; - } - - // MinLengthLogitsProcessor - if (word_id == demote_token_id) { - next_token_scores[index] = cub::FpLimits::Lowest(); + // MinLengthLogitsProcessor + if (word_id == demote_token_id) { + next_token_scores[index] = cub::FpLimits::Lowest(); + } } } } @@ -147,6 +153,7 @@ void LaunchLogitsProcessKernel( int batch_size, int num_beams, int vocab_size, + int padded_vocab_size, int demote_token_id, int32_t* sequences, int max_sequence_length, @@ -154,7 +161,7 @@ void LaunchLogitsProcessKernel( float repetition_penalty, int no_repeat_ngram_size, cudaStream_t stream) { - int total_elements = batch_size * num_beams * vocab_size; + int total_elements = batch_size * num_beams * padded_vocab_size; constexpr int blockSize = 256; const int gridSize = (total_elements + blockSize - 1) / blockSize; LogitsProcessKernel<<>>( @@ -163,6 +170,7 @@ void LaunchLogitsProcessKernel( prefix_vocab_mask, num_beams, vocab_size, + padded_vocab_size, total_elements, demote_token_id, sequences, @@ -180,6 +188,7 @@ template void LaunchLogitsProcessKernel( int batch_size, int num_beams, int vocab_size, + int padded_vocab_size, int demote_token_id, int32_t* sequences, int max_sequence_length, @@ -195,6 +204,7 @@ template void LaunchLogitsProcessKernel( int batch_size, int num_beams, int vocab_size, + int padded_vocab_size, int demote_token_id, int32_t* sequences, int max_sequence_length, diff --git a/onnxruntime/contrib_ops/cuda/transformers/beam_search_impl.h b/onnxruntime/contrib_ops/cuda/transformers/beam_search_impl.h index b1685326a1..9b122ca797 100644 --- a/onnxruntime/contrib_ops/cuda/transformers/beam_search_impl.h +++ b/onnxruntime/contrib_ops/cuda/transformers/beam_search_impl.h @@ -32,6 +32,7 @@ void LaunchLogitsProcessKernel( int batch_size, int num_beams, int vocab_size, + int padded_vocab_size, int demote_token_id, int32_t* sequences, int max_sequence_length, diff --git a/onnxruntime/contrib_ops/cuda/transformers/generation_device_helper.cc b/onnxruntime/contrib_ops/cuda/transformers/generation_device_helper.cc index c7be91f961..ee1b90822b 100644 --- a/onnxruntime/contrib_ops/cuda/transformers/generation_device_helper.cc +++ b/onnxruntime/contrib_ops/cuda/transformers/generation_device_helper.cc @@ -366,6 +366,7 @@ Status ProcessLogits(const OrtValue& logits, // parameters->batch_size, parameters->num_beams, parameters->vocab_size, + parameters->vocab_size, (parameters->min_length > 0 && current_sequence_length < parameters->min_length) ? parameters->eos_token_id : -1, reinterpret_cast(sequences_buffer.get()), parameters->max_length, @@ -550,24 +551,36 @@ Status GreedySearchProcessLogits( // In greedy search, next_token_scores is next_token_logits. gsl::span& next_token_scores = greedy_state->next_token_scores; - // TODO(tianleiwu): use one kernel to replace a loop of memory copy. - // Move the pointer in increments of padded_vocab_size to account for any padding - // if any in the logits weight of the MatMul. - const CudaT* current_logits = logits_data + (input_length - 1) * padded_vocab_size; - for (int i = 0; i < batch_beam_size; i++) { - // We only copy what is relevant (i.e.) vocab_size as padded_vocab_size will contain - // some logits corresponding to the "padded" vocab size which we will ignore - // for token generation. - gsl::span source(reinterpret_cast(current_logits), vocab_size); - gsl::span target = next_token_scores.subspan(i * vocab_size, vocab_size); - CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(target.data(), source.data(), sizeof(T) * vocab_size, - cudaMemcpyDeviceToDevice, cuda_stream)); - current_logits += input_length * padded_vocab_size; + auto is_reuse_logits_buffer = (input_length == 1); + + // Copy over the logits data into the staging buffer, only if + // we do not plan to re-use the logits buffer directly + if (!is_reuse_logits_buffer) { + // TODO(tianleiwu): use one kernel to replace a loop of memory copy. + + // Move the pointer in increments of padded_vocab_size to account for any padding + // if any in the logits weight of the MatMul. + const CudaT* current_logits = logits_data + (input_length - 1) * padded_vocab_size; + for (int i = 0; i < batch_beam_size; i++) { + // We only copy what is relevant (i.e.) vocab_size as padded_vocab_size will contain + // some logits corresponding to the "padded" vocab size which we will ignore + // for token generation. + gsl::span source(reinterpret_cast(current_logits), vocab_size); + gsl::span target = next_token_scores.subspan(i * vocab_size, vocab_size); + CUDA_RETURN_IF_ERROR(cudaMemcpyAsync(target.data(), source.data(), sizeof(T) * vocab_size, + cudaMemcpyDeviceToDevice, cuda_stream)); + current_logits += input_length * padded_vocab_size; + } } #ifdef DEBUG_GENERATION dumper->Print("logits", logits); - dumper->Print("next_token_scores", next_token_scores.data(), batch_size, vocab_size); + if (is_reuse_logits_buffer) { + //TODO: Handle padded logits in the logits buffer before printing its contents + ORT_THROW("Dumping contents of logits buffer is not implemented yet"); + } else { + dumper->Print("next_token_scores", next_token_scores.data(), batch_size, vocab_size); + } #endif // Sequences generated by beam scorer is currently stored in CPU. @@ -583,13 +596,19 @@ Status GreedySearchProcessLogits( cudaMemcpyHostToDevice, cuda_stream)); } + // TODO(hasesh): Can we avoid the const_cast by changing the interface of + // GreedySearchProcessLogits() to take in a non-const OrtValue for logits + // as this is the only place we will ever use the logits and it may be reasonable + // to allow this method to mutate/process the logits in-place cuda::LaunchLogitsProcessKernel( - reinterpret_cast(next_token_scores.data()), + is_reuse_logits_buffer ? const_cast(logits_data) + : reinterpret_cast(next_token_scores.data()), parameters->vocab_mask.data(), step > 1 ? nullptr : parameters->prefix_vocab_mask.data(), // prefix vocab mask is applied to first step only. parameters->batch_size, parameters->num_beams, parameters->vocab_size, + is_reuse_logits_buffer ? padded_vocab_size : parameters->vocab_size, (parameters->min_length > 0 && current_sequence_length < parameters->min_length) ? parameters->eos_token_id : -1, reinterpret_cast(sequences_buffer.get()), parameters->max_length, @@ -599,20 +618,30 @@ Status GreedySearchProcessLogits( cuda_stream); #ifdef DEBUG_GENERATION - dumper->Print("next_token_scores after logits process", next_token_scores.data(), batch_size, vocab_size); + if (is_reuse_logits_buffer) { + //TODO: Handle padded logits in the logits buffer before printing its contents + ORT_THROW("Dumping contents of logits buffer is not implemented yet"); + } else { + dumper->Print("next_token_scores after logits process", next_token_scores.data(), batch_size, vocab_size); + } #endif // TODO(wy): support output_scores in greedy search ORT_UNUSED_PARAMETER(output_scores); // next_tokens = torch.argmax(scores, dim=-1) - int64_t next_token_scores_dims[] = {static_cast(batch_size), vocab_size}; + int64_t next_token_scores_dims[] = {static_cast(batch_size), + is_reuse_logits_buffer ? padded_vocab_size : vocab_size}; TensorShape next_token_scores_shape(&next_token_scores_dims[0], 2); auto element_type = DataTypeImpl::GetType(); OrtValue next_token_scores_value; + + // TODO(hasesh): Same TODO as above about avoiding the const_cast here Tensor::InitOrtValue(element_type, next_token_scores_shape, - next_token_scores.data(), + is_reuse_logits_buffer + ? const_cast(reinterpret_cast(logits_data)) + : reinterpret_cast(next_token_scores.data()), allocator->Info(), next_token_scores_value); const Tensor& input = next_token_scores_value.Get();