From e0e748cfc720ed8f2da18c2326715cea9cdc4835 Mon Sep 17 00:00:00 2001 From: Tianlei Wu Date: Wed, 5 Feb 2025 02:13:59 -0800 Subject: [PATCH] update 2 --- .../cuda/transformers/generation_cuda_impl.cu | 134 ++++++++---------- 1 file changed, 61 insertions(+), 73 deletions(-) diff --git a/onnxruntime/contrib_ops/cuda/transformers/generation_cuda_impl.cu b/onnxruntime/contrib_ops/cuda/transformers/generation_cuda_impl.cu index e7722bdf34..989c5e264c 100644 --- a/onnxruntime/contrib_ops/cuda/transformers/generation_cuda_impl.cu +++ b/onnxruntime/contrib_ops/cuda/transformers/generation_cuda_impl.cu @@ -1,7 +1,6 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. - // cub.cuh includes device/dispatch_radix_sort.cuh which has assignment in conditional expressions #if defined(_MSC_VER) #pragma warning(push) @@ -336,8 +335,8 @@ __device__ void BeamHypotheses::Output( int top_k, int max_length, int pad_token_id, - int32_t* sequences, // buffer of shape (num_return_sequences, max_length) - T* sequences_scores) // buffer of shape (num_return_sequences) or empty + int32_t* sequences, // buffer of shape (num_return_sequences, max_length) + T* sequences_scores) // buffer of shape (num_return_sequences) or empty { // Copy the top_k beams into the sequences for (int index = 0; index < top_k; index++) { @@ -394,7 +393,7 @@ __global__ void BeamSearchScorer_Process(BeamScorerState& state_cpu, #ifdef DEBUG_GENERATION printf("\nbatch=%d batch_beam_idx=%d j=%d next_token=%d eos_token_id=%d next_score=%f next_index=%d\n", - batch, batch_beam_idx, static_cast(j), next_token, state.eos_token_id_, next_score, next_index); + batch, batch_beam_idx, static_cast(j), next_token, state.eos_token_id_, next_score, next_index); #endif // Add to generated hypotheses if end of sentence. if ((state.eos_token_id_ >= 0) && (next_token == state.eos_token_id_)) { @@ -429,6 +428,9 @@ __global__ void BeamSearchScorer_Process(BeamScorerState& state_cpu, #endif if (beam_hyp.beams_used_ == state.num_beams_) { bool is_done = state.early_stopping_; +#ifdef DEBUG_GENERATION + printf("\n --- state.early_stopping_=%d (batch %d)\n", static_cast(state.early_stopping_), batch); +#endif if (is_done) { #ifdef DEBUG_GENERATION printf("\n --- state.early_stopping_ (batch %d)\n", batch); @@ -445,7 +447,7 @@ __global__ void BeamSearchScorer_Process(BeamScorerState& state_cpu, if (is_done) { beam_hyp.done_ = true; - if (atomicAdd(&state.not_done_count_, -1) == 1) + if (atomicAdd(&(state.not_done_count_), -1) == 1) state_cpu.not_done_count_ = 0; // Update the CPU side #ifdef DEBUG_GENERATION printf("\n --- BeamSearchScorer_Process updated cpu state for batch %d\n", batch); @@ -470,12 +472,12 @@ __global__ void DumpBeamScorerState(BeamScorerState* state) { state->Print(false); } -void DumpBeamScorerStates(BeamScorerState& state_cpu, BeamScorerState* state, cudaStream_t stream){ - state_cpu.Print(true); +void DumpBeamScorerStates(BeamScorerState& state_cpu, BeamScorerState* state, cudaStream_t stream) { + state_cpu.Print(true); - cudaDeviceSynchronize(); - DumpBeamScorerState<<<1, 1, 0, stream>>>(state); - cudaDeviceSynchronize(); + cudaDeviceSynchronize(); + DumpBeamScorerState<<<1, 1, 0, stream>>>(state); + cudaDeviceSynchronize(); } __global__ void DumpBeamSearchScorer(const BeamHypotheses& beam_hyp) { @@ -483,13 +485,13 @@ __global__ void DumpBeamSearchScorer(const BeamHypotheses& beam_hyp) { } void DumpBeamHypotheses(gsl::span beam_hyps, cudaStream_t stream) { - printf("\n BeamHypotheses of size %zu: \n", beam_hyps.size()); - for (size_t i = 0; i < beam_hyps.size(); i++) { - printf("\n [%zu]:\n", i); - cudaDeviceSynchronize(); - DumpBeamSearchScorer<<<1, 1, 0, stream>>>(beam_hyps[i]); - cudaDeviceSynchronize(); - } + printf("\n BeamHypotheses of size %zu: \n", beam_hyps.size()); + for (size_t i = 0; i < beam_hyps.size(); i++) { + printf("\n [%zu]:\n", i); + cudaDeviceSynchronize(); + DumpBeamSearchScorer<<<1, 1, 0, stream>>>(beam_hyps[i]); + cudaDeviceSynchronize(); + } } void LaunchBeamSearchScorer_Process(BeamScorerState& state_cpu, @@ -708,14 +710,14 @@ template void LaunchBeamSearchScoreCopy(gsl::span final_scores, gsl::span output_scores, cudaStream_t stream) { - ORT_ENFORCE(final_scores.size() == output_scores.size()); - constexpr unsigned ThreadPerBlock = 256; - unsigned num_blocks = (unsigned)((final_scores.size() + (ThreadPerBlock - 1))/ ThreadPerBlock); + ORT_ENFORCE(final_scores.size() == output_scores.size()); + constexpr unsigned ThreadPerBlock = 256; + unsigned num_blocks = (unsigned)((final_scores.size() + (ThreadPerBlock - 1)) / ThreadPerBlock); - typedef typename ToCudaType::MappedType CudaT; + typedef typename ToCudaType::MappedType CudaT; - FloatConvertAndCopyKernel<<>>( - final_scores.data(), (CudaT*)output_scores.data(), final_scores.size()); + FloatConvertAndCopyKernel<<>>( + final_scores.data(), (CudaT*)output_scores.data(), final_scores.size()); } template void LaunchBeamSearchScoreCopy(gsl::span final_scores, @@ -1552,15 +1554,14 @@ void ReorderPastStatesKernelLauncher(void* out_buffer, template __global__ void CopyCrossQKSingleDecodeStepKernel( - T* target, // shape [batchxbeam, layer_head_pair_count, max_length, frame] + T* target, // shape [batchxbeam, layer_head_pair_count, max_length, frame] T** qk_layer_pointers, int token_index, int num_layers, int num_heads, const int* cross_qk_layer_head_pairs, int frames, - int max_length -) { + int max_length) { const int pair = blockIdx.x; const int layer_head_pair_count = gridDim.x; const int bbm = blockIdx.y; @@ -1572,7 +1573,7 @@ __global__ void CopyCrossQKSingleDecodeStepKernel( T* src = qk_layer_pointers[layer] + ((int64_t)bbm * num_heads + head) * frames; for (int tid = threadIdx.x; tid < frames; tid += blockDim.x) { - target[tid] = src[tid]; // use vectorized read write in future if needed + target[tid] = src[tid]; // use vectorized read write in future if needed } } @@ -1587,8 +1588,7 @@ void LaunchCopyCrossQKSingleDecodeStep( int cross_qk_layer_head_pair_count, const int* cross_qk_layer_head_pairs, int frames, - int max_length -) { + int max_length) { dim3 block(512); dim3 grid(cross_qk_layer_head_pair_count, batchxbeam); typedef typename ToCudaType::MappedType CudaT; @@ -1601,11 +1601,9 @@ void LaunchCopyCrossQKSingleDecodeStep( num_heads, cross_qk_layer_head_pairs, frames, - max_length - ); + max_length); } - template __global__ void CopyDecoderCrossQKAllStepsKernel( int context_decoding_len, @@ -1613,11 +1611,10 @@ __global__ void CopyDecoderCrossQKAllStepsKernel( int num_return_sequences, int max_length, int frames_of_k, - const T* cross_qk_buffer_data, // [batch, num_beams, layer_head_pair_count, max_length, frames] - T* cross_qk_output, // [batch, num_return_sequences, layer_head_pair_count, total_decoding_length, frames] - const int* cache_indir_data, // [batch, num_beams, max_length] - const int32_t* beam_indices -) { + const T* cross_qk_buffer_data, // [batch, num_beams, layer_head_pair_count, max_length, frames] + T* cross_qk_output, // [batch, num_return_sequences, layer_head_pair_count, total_decoding_length, frames] + const int* cache_indir_data, // [batch, num_beams, max_length] + const int32_t* beam_indices) { const int pair = blockIdx.y; const int layer_head_pair_count = gridDim.y; const int total_decoding_length = gridDim.x; @@ -1630,15 +1627,15 @@ __global__ void CopyDecoderCrossQKAllStepsKernel( const int src_beam = beam_indices[batch * num_beams + ret_seq_id] % num_beams; const int64_t offset_in_cache = ((int64_t)batch * num_beams + src_beam) * max_length + token_decoding_index + context_decoding_len; - int bm_mapped = ((num_beams <= 1) ? 0: ((token_decoding_index == total_decoding_length - 1) ? ret_seq_id : cache_indir_data[offset_in_cache])); + int bm_mapped = ((num_beams <= 1) ? 0 : ((token_decoding_index == total_decoding_length - 1) ? ret_seq_id : cache_indir_data[offset_in_cache])); int bi_src = batch * num_beams + bm_mapped; - T* target = cross_qk_output + - (((int64_t)br * layer_head_pair_count + (int64_t)pair) * total_decoding_length + token_decoding_index) * frames_of_k; + T* target = cross_qk_output + + (((int64_t)br * layer_head_pair_count + (int64_t)pair) * total_decoding_length + token_decoding_index) * frames_of_k; const T* src = cross_qk_buffer_data + - ((int64_t)bi_src * layer_head_pair_count * max_length + (int64_t)pair * max_length + token_decoding_index) * frames_of_k; + ((int64_t)bi_src * layer_head_pair_count * max_length + (int64_t)pair * max_length + token_decoding_index) * frames_of_k; for (int tid = threadIdx.x; tid < frames_of_k; tid += blockDim.x) { - target[tid] = src[tid]; // use vectorized read write in future if needed + target[tid] = src[tid]; // use vectorized read write in future if needed } } @@ -1656,8 +1653,7 @@ void LaunchFinalizeCrossQK( float* cross_qk_output, int num_return_sequences, const int* cache_indir_data, - const int32_t* beam_indices -) { + const int32_t* beam_indices) { int64_t br = (int64_t)batch_size * num_return_sequences; ORT_ENFORCE(br < 65536L && cross_qk_layer_head_pair_count < 65536); const int total_decoding_length = iteration_number - 1; @@ -1666,15 +1662,15 @@ void LaunchFinalizeCrossQK( typedef typename ToCudaType::MappedType CudaT; CopyDecoderCrossQKAllStepsKernel<<>>( - context_decoding_len, - num_beams, - num_return_sequences, - max_length, - frames_of_k, - (const CudaT*)cross_qk_buffer_data, - (CudaT*)cross_qk_output, - cache_indir_data, - beam_indices); + context_decoding_len, + num_beams, + num_return_sequences, + max_length, + frames_of_k, + (const CudaT*)cross_qk_buffer_data, + (CudaT*)cross_qk_output, + cache_indir_data, + beam_indices); } template @@ -1683,12 +1679,11 @@ __global__ void ForceDecodingIdsKernel( const int vocab_size, const int32_t* force_ids, int id_len, - int step -) { + int step) { const int num_beams = gridDim.y; const int beam = blockIdx.y; const int batch = blockIdx.z; - beam_scores += (((int64_t)batch * num_beams + beam)* vocab_size); // move to (batch, beam) + beam_scores += (((int64_t)batch * num_beams + beam) * vocab_size); // move to (batch, beam) const int32_t id_wanted = force_ids[((int64_t)batch * id_len) + step]; if (id_wanted < 0 || id_wanted >= vocab_size) return; @@ -1696,7 +1691,7 @@ __global__ void ForceDecodingIdsKernel( const int32_t block_start_id = blockIdx.x * elements_per_block; int32_t token_id = block_start_id + (int)threadIdx.x; - #pragma unroll +#pragma unroll for (int elem = 0; elem < ElementsPerThreads; elem++) { if (token_id < vocab_size) { beam_scores[token_id] = ((token_id == id_wanted) ? 0.0f : cub::FpLimits::Lowest()); @@ -1705,7 +1700,6 @@ __global__ void ForceDecodingIdsKernel( } } - void LaunchForceDecodingIds( float* beam_scores, const int batch_size, @@ -1714,15 +1708,13 @@ void LaunchForceDecodingIds( const int32_t* force_ids, int id_len, int step, - cudaStream_t stream -) { + cudaStream_t stream) { dim3 blocks(512); constexpr int ElementsPerThreads = 4; unsigned gridx = static_cast((vocab_size + 512 * ElementsPerThreads - 1) / (512 * ElementsPerThreads)); dim3 grids(gridx, num_beams, batch_size); ForceDecodingIdsKernel<<>>( - beam_scores, vocab_size, force_ids, id_len, step - ); + beam_scores, vocab_size, force_ids, id_len, step); } template @@ -1732,8 +1724,7 @@ __global__ void SaveNoSpeechProbsKernel( const int batch_size, const int num_beams, const int vocab_size, - const int no_speech_token_id -) { + const int no_speech_token_id) { int b = blockIdx.x * blockDim.x + threadIdx.x; if (b < batch_size) { int64_t src_offset = b * num_beams * vocab_size + no_speech_token_id; @@ -1743,20 +1734,19 @@ __global__ void SaveNoSpeechProbsKernel( template void LaunchSaveNoSpeechProbs( - T* result_no_speech_probs, /* [batch]*/ - const float* probs, /* [batch, num_beams, vocab_size]*/ + T* result_no_speech_probs, /* [batch]*/ + const float* probs, /* [batch, num_beams, vocab_size]*/ const int batch_size, const int num_beams, const int vocab_size, const int no_speech_token_id, - cudaStream_t stream -) { + cudaStream_t stream) { int tpb = 256; int bpg = (batch_size + 255) / 256; typedef typename ToCudaType::MappedType CudaT; SaveNoSpeechProbsKernel<<>>( - (CudaT*)result_no_speech_probs, probs, batch_size, num_beams, vocab_size, no_speech_token_id); + (CudaT*)result_no_speech_probs, probs, batch_size, num_beams, vocab_size, no_speech_token_id); } template void LaunchSaveNoSpeechProbs( @@ -1766,8 +1756,7 @@ template void LaunchSaveNoSpeechProbs( const int num_beams, const int vocab_size, const int no_speech_token_id, - cudaStream_t stream -); + cudaStream_t stream); template void LaunchSaveNoSpeechProbs( MLFloat16* result_no_speech_probs, @@ -1776,8 +1765,7 @@ template void LaunchSaveNoSpeechProbs( const int num_beams, const int vocab_size, const int no_speech_token_id, - cudaStream_t stream -); + cudaStream_t stream); } // namespace cuda } // namespace contrib