This commit is contained in:
Tianlei Wu 2025-02-05 02:13:59 -08:00
parent 259bef34c7
commit e0e748cfc7

View file

@ -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<int>(j), next_token, state.eos_token_id_, next_score, next_index);
batch, batch_beam_idx, static_cast<int>(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<int>(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<BeamHypotheses> 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 <typename T>
void LaunchBeamSearchScoreCopy(gsl::span<const float> final_scores,
gsl::span<T> 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<float>::MappedType CudaT;
typedef typename ToCudaType<float>::MappedType CudaT;
FloatConvertAndCopyKernel<<<num_blocks, ThreadPerBlock, 0, stream>>>(
final_scores.data(), (CudaT*)output_scores.data(), final_scores.size());
FloatConvertAndCopyKernel<<<num_blocks, ThreadPerBlock, 0, stream>>>(
final_scores.data(), (CudaT*)output_scores.data(), final_scores.size());
}
template void LaunchBeamSearchScoreCopy(gsl::span<const float> final_scores,
@ -1552,15 +1554,14 @@ void ReorderPastStatesKernelLauncher(void* out_buffer,
template <typename T>
__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<float>::MappedType CudaT;
@ -1601,11 +1601,9 @@ void LaunchCopyCrossQKSingleDecodeStep(
num_heads,
cross_qk_layer_head_pairs,
frames,
max_length
);
max_length);
}
template <typename T>
__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<float>::MappedType CudaT;
CopyDecoderCrossQKAllStepsKernel<<<grid, block, 0, stream>>>(
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 <int ElementsPerThreads>
@ -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<float>::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<unsigned>((vocab_size + 512 * ElementsPerThreads - 1) / (512 * ElementsPerThreads));
dim3 grids(gridx, num_beams, batch_size);
ForceDecodingIdsKernel<ElementsPerThreads><<<grids, blocks, 0, stream>>>(
beam_scores, vocab_size, force_ids, id_len, step
);
beam_scores, vocab_size, force_ids, id_len, step);
}
template <typename T>
@ -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 <typename T>
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<T>::MappedType CudaT;
SaveNoSpeechProbsKernel<CudaT><<<bpg, tpb, 0, stream>>>(
(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<float>(
@ -1766,8 +1756,7 @@ template void LaunchSaveNoSpeechProbs<float>(
const int num_beams,
const int vocab_size,
const int no_speech_token_id,
cudaStream_t stream
);
cudaStream_t stream);
template void LaunchSaveNoSpeechProbs<MLFloat16>(
MLFloat16* result_no_speech_probs,
@ -1776,8 +1765,7 @@ template void LaunchSaveNoSpeechProbs<MLFloat16>(
const int num_beams,
const int vocab_size,
const int no_speech_token_id,
cudaStream_t stream
);
cudaStream_t stream);
} // namespace cuda
} // namespace contrib