diff --git a/src/beam_search_scorer.cpp b/src/beam_search_scorer.cpp index b9cbceffef..2b846636f6 100644 --- a/src/beam_search_scorer.cpp +++ b/src/beam_search_scorer.cpp @@ -122,7 +122,7 @@ void BeamSearchScorer::Process(Sequences& sequences, int const batch_beam_idx = static_cast(batch * num_beams_) + next_index; // Add to generated hypotheses if end of sentence. - if ((eos_token_id_ >= 0) && (next_token == eos_token_id_)) { + if (contains(eos_token_id_, next_token)) { bool const is_beam_token_worse_than_top_num_beams = (j >= num_beams_); if (is_beam_token_worse_than_top_num_beams) { continue; diff --git a/src/beam_search_scorer.h b/src/beam_search_scorer.h index 2fa44767fb..1d42124453 100644 --- a/src/beam_search_scorer.h +++ b/src/beam_search_scorer.h @@ -53,7 +53,7 @@ struct BeamSearchScorer { int num_beams_; int max_length_; int pad_token_id_; - int eos_token_id_; + std::vector eos_token_id_; bool early_stopping_; int not_done_count_; // When zero, every batch entry is done (starts at batch_size_) diff --git a/src/config.cpp b/src/config.cpp index 5f7f40bab7..6ab5a4cd73 100644 --- a/src/config.cpp +++ b/src/config.cpp @@ -31,6 +31,17 @@ struct NamedStrings_Element : JSON::Element { std::vector& v_; }; +struct Int_Array_Element : JSON::Element { + explicit Int_Array_Element(std::vector& v) : v_{v} {} + + void OnValue(std::string_view name, JSON::Value value) override { + v_.emplace_back(static_cast(JSON::Get(value))); + } + + private: + std::vector& v_; +}; + struct ProviderOptionsObject_Element : JSON::Element { explicit ProviderOptionsObject_Element(std::vector& v) : v_{v} {} @@ -509,29 +520,6 @@ struct Speech_Element : JSON::Element { SpeechOutputs_Element outputs_{v_.outputs}; }; -struct Eos_Array_Element : JSON::Element { - explicit Eos_Array_Element(Config::Model& v) : v_{v} {} - - void OnValue(std::string_view name, JSON::Value value) override { - v_.eos_token_ids.push_back(static_cast(JSON::Get(value))); - } - - void OnComplete(bool empty) override { - if (v_.eos_token_ids.empty()) - return; // Empty array, nothign to do - - // Copy the first eos_token_id into the eos_token_id value, it will be our primary eos token - v_.eos_token_id = v_.eos_token_ids.front(); - - // If the array is just one value, clear the array and just act like a single value was set - if (v_.eos_token_ids.size() == 1) - v_.eos_token_ids.clear(); - } - - private: - Config::Model& v_; -}; - struct EmbeddingInputs_Element : JSON::Element { explicit EmbeddingInputs_Element(Config::Model::Embedding::Inputs& v) : v_{v} {} @@ -602,7 +590,7 @@ struct Model_Element : JSON::Element { } else if (name == "pad_token_id") { v_.pad_token_id = static_cast(JSON::Get(value)); } else if (name == "eos_token_id") { - v_.eos_token_id = static_cast(JSON::Get(value)); + v_.eos_token_id.assign(1, static_cast(JSON::Get(value))); } else if (name == "bos_token_id") { v_.bos_token_id = static_cast(JSON::Get(value)); } else if (name == "decoder_start_token_id") { @@ -615,7 +603,7 @@ struct Model_Element : JSON::Element { Element& OnArray(std::string_view name) override { if (name == "eos_token_id") - return eos_token_ids_; + return eos_token_id_; throw JSON::unknown_value_error{}; } @@ -642,7 +630,7 @@ struct Model_Element : JSON::Element { Config::Model& v_; EncoderDecoderInit_Element encoder_decoder_init_{v_.encoder_decoder_init}; Decoder_Element decoder_{v_.decoder}; - Eos_Array_Element eos_token_ids_{v_}; + Int_Array_Element eos_token_id_{v_.eos_token_id}; Vision_Element vision_{v_.vision}; Embedding_Element embedding_{v_.embedding}; Speech_Element speech_{v_.speech}; @@ -715,11 +703,8 @@ void ClearProviders(Config& config) { } void SetProviderOption(Config& config, std::string_view provider_name, std::string_view option_name, std::string_view option_value) { - if (std::find(config.model.decoder.session_options.providers.begin(), - config.model.decoder.session_options.providers.end(), provider_name) == - config.model.decoder.session_options.providers.end()) { + if (!contains(config.model.decoder.session_options.providers, provider_name)) config.model.decoder.session_options.providers.push_back(std::string(provider_name)); - } std::ostringstream json; json << R"({")" << provider_name << R"(":{)"; @@ -830,6 +815,10 @@ Config::Config(const fs::path& path, std::string_view json_overlay) : config_pat if (search.max_length == 0) search.max_length = model.context_length; + // If no eos_token_id was set, set it to the pad token id + if (model.eos_token_id.empty()) + model.eos_token_id.push_back(model.pad_token_id); + for (const auto& provider_option : model.decoder.session_options.provider_options) { model.decoder.session_options.providers.push_back(provider_option.name); } diff --git a/src/config.h b/src/config.h index cf406b3709..89831bb11c 100644 --- a/src/config.h +++ b/src/config.h @@ -81,12 +81,11 @@ struct Config { struct Model { std::string type; - int pad_token_id{}; // The id of the padding token. - int eos_token_id{}; // The id of the end-of-stream token. - std::vector eos_token_ids; // If eos_token_id is passed as an array, this is where the values go (eos_token_id gets set to the first entry in the array) - int bos_token_id{}; // The id of the beginning-of-stream token. - int sep_token_id{}; // The id of the separation token. - int decoder_start_token_id{}; // If an encoder-decoder model starts decoding with a different token than bos, the id of that token. + int pad_token_id{}; // The id of the padding token. + std::vector eos_token_id; // The end-of-stream tokens (when set as a single value it is converted to a vector with one value). + int bos_token_id{}; // The id of the beginning-of-stream token. + int sep_token_id{}; // The id of the separation token. + int decoder_start_token_id{}; // If an encoder-decoder model starts decoding with a different token than bos, the id of that token. int vocab_size{}; int context_length{}; diff --git a/src/cuda/beam_search_scorer_cuda.cpp b/src/cuda/beam_search_scorer_cuda.cpp index a59bc80bc4..61fc86accb 100644 --- a/src/cuda/beam_search_scorer_cuda.cpp +++ b/src/cuda/beam_search_scorer_cuda.cpp @@ -10,14 +10,14 @@ namespace Generators { -BeamSearchScorer_Cuda::BeamSearchScorer_Cuda(const GeneratorParams& parameters) - : stream_{GetStream()} { +BeamSearchScorer_Cuda::BeamSearchScorer_Cuda(const GeneratorParams& parameters, std::span cuda_eos_tokens) + : stream_{GetStream()}, + eos_tokens_{cuda_eos_tokens} { state_cpu_ = CudaMallocHostArray(1); state_cpu_->batch_size_ = static_cast(parameters.search.batch_size); state_cpu_->num_beams_ = static_cast(parameters.search.num_beams); state_cpu_->max_length_ = static_cast(parameters.search.max_length); state_cpu_->pad_token_id_ = parameters.config.model.pad_token_id; - state_cpu_->eos_token_id_ = parameters.config.model.eos_token_id; state_cpu_->early_stopping_ = parameters.search.early_stopping; state_cpu_->not_done_count_ = parameters.search.batch_size; state_cpu_->hypothesis_buffer_used_ = 0; @@ -51,6 +51,7 @@ void BeamSearchScorer_Cuda::Process(Sequences& sequences, std::span next_indices) { cuda::LaunchBeamSearchScorer_Process(*state_cpu_, *state_gpu_, + eos_tokens_, sequences.GetSequences().Span(), sequences.GetSequenceLength(), beam_hyps_, diff --git a/src/cuda/beam_search_scorer_cuda.cu b/src/cuda/beam_search_scorer_cuda.cu index 3add21bc1c..d7247d55e4 100644 --- a/src/cuda/beam_search_scorer_cuda.cu +++ b/src/cuda/beam_search_scorer_cuda.cu @@ -76,6 +76,8 @@ __device__ bool BeamHypotheses::CanImprove(float best_sum_logprobs, int current_ __global__ void BeamSearchScorer_Process(BeamScorerState& state_cpu, BeamScorerState& state, + const int32_t* eos_token_ids, + const int eos_token_count, const int32_t* sequences_buffer, int sequence_length, BeamHypotheses* beam_hyps_, @@ -105,7 +107,14 @@ __global__ void BeamSearchScorer_Process(BeamScorerState& state_cpu, int batch_beam_idx = batch_start + next_index; // Add to generated hypotheses if end of sentence. - if ((state.eos_token_id_ >= 0) && (next_token == state.eos_token_id_)) { + bool is_eos_token = false; + for (unsigned eos_index = 0; eos_index < eos_token_count; eos_index++) { + if (next_token == eos_token_ids[eos_index]) { + is_eos_token = true; + break; + } + } + if (is_eos_token) { bool is_beam_token_worse_than_top_num_beams = (j >= state.num_beams_); if (is_beam_token_worse_than_top_num_beams) { continue; @@ -152,6 +161,7 @@ __global__ void BeamSearchScorer_Process(BeamScorerState& state_cpu, void LaunchBeamSearchScorer_Process(BeamScorerState& state_cpu, BeamScorerState& state, + std::span eos_token_ids, std::span sequences, int sequence_length, std::span beam_hyps, @@ -165,6 +175,8 @@ void LaunchBeamSearchScorer_Process(BeamScorerState& state_cpu, cudaStream_t stream) { BeamSearchScorer_Process<<<1, state_cpu.batch_size_, 0, stream>>>(state_cpu, state, + eos_token_ids.data(), + static_cast(eos_token_ids.size()), sequences.data(), sequence_length, beam_hyps.data(), diff --git a/src/cuda/beam_search_scorer_cuda.cuh b/src/cuda/beam_search_scorer_cuda.cuh index 7a8834b694..dbee2baae8 100644 --- a/src/cuda/beam_search_scorer_cuda.cuh +++ b/src/cuda/beam_search_scorer_cuda.cuh @@ -32,7 +32,6 @@ struct BeamScorerState { int num_beams_; int max_length_; int pad_token_id_; - int eos_token_id_; bool early_stopping_; int not_done_count_; // When zero, every batch entry is done (starts at batch_size_) @@ -43,6 +42,7 @@ void LaunchInitializeBeamHypotheses(std::span beam_hyps, float l void LaunchBeamSearchScorer_Process(BeamScorerState& state_cpu, BeamScorerState& state, + std::span eos_token_ids, std::span sequences, int sequence_length, std::span beam_hyps_, diff --git a/src/cuda/beam_search_scorer_cuda.h b/src/cuda/beam_search_scorer_cuda.h index 7ec2084852..95ce56069e 100644 --- a/src/cuda/beam_search_scorer_cuda.h +++ b/src/cuda/beam_search_scorer_cuda.h @@ -4,7 +4,7 @@ namespace Generators { struct BeamSearchScorer_Cuda { - BeamSearchScorer_Cuda(const GeneratorParams& parameters); + BeamSearchScorer_Cuda(const GeneratorParams& parameters, std::span cuda_eos_tokens); void Process(Sequences& sequences, std::span next_scores, @@ -26,7 +26,9 @@ struct BeamSearchScorer_Cuda { mutable cuda_event_holder event_process_complete_; cuda_host_unique_ptr state_cpu_; cuda_unique_ptr state_gpu_; + cudaStream_t stream_; + std::span eos_tokens_; DeviceSpan next_beam_scores_; DeviceSpan next_beam_tokens_; diff --git a/src/cuda/interface.cpp b/src/cuda/interface.cpp index f3fcee3e60..ed9f5e795a 100644 --- a/src/cuda/interface.cpp +++ b/src/cuda/interface.cpp @@ -143,10 +143,6 @@ struct CudaInterfaceImpl final : DeviceInterface { return true; } - void LaunchHandleEOSArray(float* batch_logits, int batch_beam_size, int vocab_size, const int32_t* eos_token_ids, int eos_token_ids_count) override { - cuda::LaunchHandleEOSArray(batch_logits, batch_beam_size, vocab_size, eos_token_ids, eos_token_ids_count, GetStream()); - } - void UpdateCacheIndirectionKernelLauncher(int32_t* tgt_indir_cache, const int32_t* src_indir_cache, const int32_t* beam_ids, int batch_size, int beam_width, int input_seq_length, int max_seq_length, int current_length) override { cuda::UpdateCacheIndirectionKernelLauncher(tgt_indir_cache, src_indir_cache, beam_ids, batch_size, beam_width, input_seq_length, max_seq_length, current_length, GetStream()); } diff --git a/src/cuda/kernels.h b/src/cuda/kernels.h index 9a60ba5112..25fe6600c6 100644 --- a/src/cuda/kernels.h +++ b/src/cuda/kernels.h @@ -11,8 +11,6 @@ void Launch_UpdatePositionIds(T* positions, int batch_beam_size, int total_lengt template void Launch_UpdateAttentionMask(T* mask_data, T* old_data, int batch_beam_size, int new_kv_length, int total_length, int max_length, bool update_only, cudaStream_t stream); -void LaunchHandleEOSArray(float* batch_logits, int batch_beam_size, int vocab_size, const int32_t* eos_token_ids, int eos_token_ids_count, cudaStream_t stream); - void LaunchFp16ToFp32(const uint16_t* fp16, float* fp32, int count, cudaStream_t stream); void LaunchFp32ToFp16(const float* fp32, uint16_t* fp16, int count, cudaStream_t stream); void LaunchInt32ToInt64(const int32_t* src, int64_t* dst, int count, cudaStream_t stream); diff --git a/src/cuda/model_kernels.cu b/src/cuda/model_kernels.cu index 00dd4fde9c..a22499ee18 100644 --- a/src/cuda/model_kernels.cu +++ b/src/cuda/model_kernels.cu @@ -82,25 +82,6 @@ void Launch_UpdateAttentionMask(T* next_mask_data, T* mask_data, int batch_beam_ template void Launch_UpdateAttentionMask(int32_t* next_mask_data, int32_t* mask_data, int batch_beam_size, int new_kv_length, int total_length, int max_length, bool update_only, cudaStream_t stream); template void Launch_UpdateAttentionMask(int64_t* next_mask_data, int64_t* mask_data, int batch_beam_size, int new_kv_length, int total_length, int max_length, bool update_only, cudaStream_t stream); -__global__ void HandleEOSArray(float* batch_logits, int batch_beam_size, int vocab_size, const int32_t* eos_token_ids, int eos_token_ids_count) { - int index = blockIdx.x * blockDim.x + threadIdx.x; - if (index >= batch_beam_size) - return; - - float* logits = batch_logits + index * vocab_size; - float max = std::numeric_limits::lowest(); - for (int i = 0; i < eos_token_ids_count; i++) { - max = std::max(max, logits[eos_token_ids[i]]); - logits[eos_token_ids[i]] = std::numeric_limits::lowest(); // Set all EOS token options to never happen (the first will get the max of all) - } - - logits[eos_token_ids[0]] = max; // Set the score of the primary EOS token to the highest of any of the EOS tokens -} - -void LaunchHandleEOSArray(float* batch_logits, int batch_beam_size, int vocab_size, const int32_t* eos_token_ids, int eos_token_ids_count, cudaStream_t stream) { - HandleEOSArray<<<(batch_beam_size + 255) / 256, 256, 0, stream>>>(batch_logits, batch_beam_size, vocab_size, eos_token_ids, eos_token_ids_count); -} - __global__ void ConvertFp16ToFp32(const half* src, float* dst, int count) { int idx = threadIdx.x + blockIdx.x * blockDim.x; if (idx < count) diff --git a/src/cuda/search_cuda.cpp b/src/cuda/search_cuda.cpp index 5bfbc6ac8e..d67d24930a 100644 --- a/src/cuda/search_cuda.cpp +++ b/src/cuda/search_cuda.cpp @@ -24,8 +24,12 @@ Search_Cuda::Search_Cuda(const GeneratorParams& params) auto batch_beam_size = params.BatchBeamSize(); sequence_lengths_ = params.p_device->Allocate(batch_beam_size); - eos_meet_buffer_ = CudaMallocArray(batch_beam_size, &eos_meet_); - cudaMemsetAsync(eos_meet_.data(), 0, eos_meet_.size_bytes(), GetStream()); + eos_seen_buffer_ = CudaMallocArray(batch_beam_size, &eos_seen_); + cudaMemsetAsync(eos_seen_.data(), 0, eos_seen_.size_bytes(), GetStream()); + + eos_token_ids_ = params.p_device->Allocate(params.config.model.eos_token_id.size()); + copy(std::span{params.config.model.eos_token_id}, eos_token_ids_.CpuSpan()); + eos_token_ids_.CopyCpuToDevice(); done_cpu_ = CudaMallocHostArray(1); *done_cpu_ = false; @@ -49,7 +53,7 @@ BeamSearch_Cuda::BeamSearch_Cuda(const GeneratorParams& params) : Search_Cuda{params} { assert(params_->search.num_beams > 1); // If 1, use GreedySearch auto batch_beam_size = params_->BatchBeamSize(); - beam_scorer_ = std::make_unique(*params_); + beam_scorer_ = std::make_unique(*params_, eos_token_ids_.Span()); topk_next_tokens_ = CudaMallocArray(2 * batch_beam_size); topk_next_indices_ = CudaMallocArray(2 * batch_beam_size); @@ -153,9 +157,8 @@ void GreedySearch_Cuda::SampleTopKTopP(int k, float p, float temperature) { params_->search.batch_size, k, p, temperature); // Check for EOS - assert(next_tokens_.size() == eos_meet_.size()); - // Don't replace EOS with pad for batch_size == 1 for continuous decoding mode - cuda::Launch_CheckForEOSAndPad(next_tokens_.data(), static_cast(next_tokens_.size()), eos_meet_.data(), params_->config.model.eos_token_id, params_->search.batch_size > 1 ? params_->config.model.pad_token_id : params_->config.model.eos_token_id, done_cpu_.get(), GetStream()); + assert(next_tokens_.size() == eos_seen_.size()); + cuda::Launch_CheckForEOSAndPad(next_tokens_.data(), static_cast(next_tokens_.size()), eos_seen_.data(), eos_token_ids_.Span().data(), static_cast(eos_token_ids_.Span().size()), params_->config.model.pad_token_id, done_cpu_.get(), GetStream()); // Append tokens cuda::Launch_AppendNextTokensToSequences(next_tokens_buffer_.Span(), sequences_.GetSequences().Span(), params_->BatchBeamSize(), sequences_.GetSequenceLength(), sequences_.max_length_, GetStream()); @@ -210,7 +213,7 @@ std::span Search_Cuda::GetScores() { // Set user input tokens (batch_beam_size, sequence_length) void GreedySearch_Cuda::AppendTokens(DeviceSpan& next_tokens) { - cudaMemsetAsync(eos_meet_.data(), 0, eos_meet_.size_bytes(), GetStream()); + cudaMemsetAsync(eos_seen_.data(), 0, eos_seen_.size_bytes(), GetStream()); *done_cpu_ = false; auto next_tokens_gpu = next_tokens.Span(); @@ -224,7 +227,7 @@ void GreedySearch_Cuda::AppendTokens(DeviceSpan& next_tokens) { return; } - cudaMemsetAsync(eos_meet_.data(), 0, eos_meet_.size_bytes(), GetStream()); + cudaMemsetAsync(eos_seen_.data(), 0, eos_seen_.size_bytes(), GetStream()); *done_cpu_ = false; } @@ -237,7 +240,7 @@ void BeamSearch_Cuda::AppendTokens(DeviceSpan& next_tokens) { } void GreedySearch_Cuda::RewindTo(size_t index) { - cudaMemsetAsync(eos_meet_.data(), 0, eos_meet_.size_bytes(), GetStream()); + cudaMemsetAsync(eos_seen_.data(), 0, eos_seen_.size_bytes(), GetStream()); *done_cpu_ = false; if (index > 0) cuda::Launch_GetLastTokens(next_tokens_.data(), sequences_.GetSequences().Span().data(), static_cast(params_->BatchBeamSize()), static_cast(index), sequences_.max_length_, GetStream()); @@ -250,7 +253,8 @@ void Search_Cuda::ApplyMinLength(int min_length) { if (sequences_.GetSequenceLength() >= min_length) return; - cuda::LaunchSetScoreProcessor(GetScores().data(), params_->BatchBeamSize(), params_->config.model.vocab_size, params_->config.model.eos_token_id, std::numeric_limits::lowest(), GetStream()); + for (auto eos_token_id : params_->config.model.eos_token_id) + cuda::LaunchSetScoreProcessor(GetScores().data(), params_->BatchBeamSize(), params_->config.model.vocab_size, eos_token_id, std::numeric_limits::lowest(), GetStream()); } void Search_Cuda::ApplyRepetitionPenalty(float penalty) { diff --git a/src/cuda/search_cuda.cu b/src/cuda/search_cuda.cu index fcb21200d2..8474ae53ed 100644 --- a/src/cuda/search_cuda.cu +++ b/src/cuda/search_cuda.cu @@ -74,21 +74,29 @@ struct ArgMaxDataImpl : ArgMaxData { cuda_unique_ptr> argmaxen_owner_; }; -__global__ void CheckForEOSAndPad(int32_t* next_tokens, int next_tokens_count, bool* eos_meet, int eos_token_id, int pad_token_id, bool* done_cpu) { - // Look for EOS tokens, if seen set EOS flag and replace with pad token - for (size_t batch_id = 0; batch_id < next_tokens_count; ++batch_id) { - if (next_tokens[batch_id] == eos_token_id || eos_meet[batch_id] == true) { - eos_meet[batch_id] = true; +__global__ void CheckForEOSAndPad(int32_t* next_tokens, int next_tokens_count, bool* eos_seen, const int* eos_token_ids, int eos_token_count, int pad_token_id, bool* done_cpu) { + for (int batch_id = 0; batch_id < next_tokens_count; ++batch_id) { + // If EOS already met, pad + if (eos_seen[batch_id]) { next_tokens[batch_id] = pad_token_id; + continue; + } + + // Look for EOS + for (int eos_id = 0; eos_id < eos_token_count; ++eos_id) { + if (next_tokens[batch_id] == eos_token_ids[eos_id]) { + eos_seen[batch_id] = true; + break; + } } } // When all batches are finished, stop earlier to avoid wasting computation. // TODO: Merge this with the above so we don't have to double scan. Just keep track of 'batches left' { - size_t batch_id = 0; + int batch_id = 0; while (batch_id < next_tokens_count) { - if (eos_meet[batch_id] == false) { + if (eos_seen[batch_id] == false) { break; } ++batch_id; @@ -100,8 +108,8 @@ __global__ void CheckForEOSAndPad(int32_t* next_tokens, int next_tokens_count, b } } -void Launch_CheckForEOSAndPad(int32_t* next_tokens, int next_tokens_count, bool* eos_meet, int eos_token_id, int pad_token_id, bool* done_cpu, cudaStream_t stream) { - CheckForEOSAndPad<<<1, 1, 0, stream>>>(next_tokens, next_tokens_count, eos_meet, eos_token_id, pad_token_id, done_cpu); +void Launch_CheckForEOSAndPad(int32_t* next_tokens, int next_tokens_count, bool* eos_seen, const int *eos_token_ids, int eos_token_count, int pad_token_id, bool* done_cpu, cudaStream_t stream) { + CheckForEOSAndPad<<<1, 1, 0, stream>>>(next_tokens, next_tokens_count, eos_seen, eos_token_ids, eos_token_count, pad_token_id, done_cpu); } __global__ void AddProbsKernel(float* log_probs, diff --git a/src/cuda/search_cuda.cuh b/src/cuda/search_cuda.cuh index a0a07fb9fc..3aa5aa0fc5 100644 --- a/src/cuda/search_cuda.cuh +++ b/src/cuda/search_cuda.cuh @@ -9,7 +9,8 @@ struct ArgMaxData { virtual ~ArgMaxData() = default; }; -void Launch_CheckForEOSAndPad(int32_t* next_tokens, int next_tokens_count, bool* eos_meet, int eos_token_id, int pad_token_id, bool* done_cpu, cudaStream_t stream); +void Launch_CheckForEOSAndPad(int32_t* next_tokens, int next_tokens_count, bool* eos_seen, const int32_t* eos_token_ids, int eos_token_count, bool* done_cpu, cudaStream_t stream); +void Launch_CheckForEOSAndPad(int32_t* next_tokens, int next_tokens_count, bool* eos_seen, const int32_t* eos_token_ids, int eos_token_count, int pad_token_id, bool* done_cpu, cudaStream_t stream); void Launch_ExpandInputSequences(const std::span input_sequences, std::span sequences, int batch_size, int beam_size, int max_length, cudaStream_t stream); void Launch_AppendNextTokensToSequences(std::span next_tokens, std::span sequences, int batch_beam_size, int past_length, int max_length, cudaStream_t stream); void Launch_GetLastTokens(int32_t* next_tokens, const int32_t* sequences, int batch_beam_size, int sequence_length, int max_length, cudaStream_t stream); diff --git a/src/cuda/search_cuda.h b/src/cuda/search_cuda.h index acdd0525fa..e21904a9a4 100644 --- a/src/cuda/search_cuda.h +++ b/src/cuda/search_cuda.h @@ -31,8 +31,9 @@ struct Search_Cuda : Search { DeviceSpan sequence_lengths_; // shape (beam_size*batch_size) - gpu_span eos_meet_; // shape (beam_size*batch_size) - cuda_unique_ptr eos_meet_buffer_; + gpu_span eos_seen_; // shape (beam_size*batch_size) + cuda_unique_ptr eos_seen_buffer_; + DeviceSpan eos_token_ids_; gpu_span next_tokens_; // shape (beam_size*batch_size) DeviceSpan next_token_scores_; // shape (beam_size*batch_size, vocab_size) diff --git a/src/generators.h b/src/generators.h index defe1eb4ce..0e678cbbf9 100644 --- a/src/generators.h +++ b/src/generators.h @@ -42,6 +42,11 @@ struct State; struct Search; struct Tokenizer; +template +bool contains(const T& t, V&& v) { + return std::find(t.begin(), t.end(), v) != t.end(); +} + template DeviceSpan WrapTensor(DeviceInterface& device, OrtValue& value) { auto info = value.GetTensorTypeAndShapeInfo(); diff --git a/src/models/logits.cpp b/src/models/logits.cpp index a12f69c763..739140d8a8 100644 --- a/src/models/logits.cpp +++ b/src/models/logits.cpp @@ -13,13 +13,6 @@ Logits::Logits(State& state) type_{model_.session_info_.GetOutputDataType(model_.config_->model.decoder.outputs.logits)} { output_raw_ = std::make_unique(model_.p_device_inputs_, type_); - if (model_.p_device_inputs_->GetType() == DeviceType::CUDA && !model_.config_->model.eos_token_ids.empty()) { - auto& cpu_ids = model_.config_->model.eos_token_ids; - cuda_eos_token_ids_ = model_.p_device_->Allocate(cpu_ids.size()); - copy(std::span{cpu_ids}, cuda_eos_token_ids_.CpuSpan()); - cuda_eos_token_ids_.CopyCpuToDevice(); - } - input_sequence_lengths.resize(state_.params_->search.batch_size); if (IsOpenVINOStatefulModel(state.model_)) { @@ -82,23 +75,6 @@ DeviceSpan Logits::Get() { if (logits_.empty() || logits_of_last_token->GetTensorMutableRawData() != logits_.Span().data()) logits_ = WrapTensor(*model_.p_device_inputs_, *logits_of_last_token); - // TODO: This functionality may have to be moved to DeviceInterface to make the code platform agnostic - if (model_.p_device_inputs_->GetType() == DeviceType::CUDA) { - if (!cuda_eos_token_ids_.empty()) - model_.p_device_inputs_->LaunchHandleEOSArray( - logits_.Span().data(), - static_cast(shape_[0]) /* batch_beam_size*/, - static_cast(shape_[2]) /* vocab_size */, - cuda_eos_token_ids_.Span().data(), - static_cast(cuda_eos_token_ids_.size())); - return logits_; - } else if (model_.p_device_inputs_->GetType() == DeviceType::DML) { - HandleEOSArray(logits_.CopyDeviceToCpu()); - logits_.CopyCpuToDevice(); - return logits_; - } - - HandleEOSArray(logits_.Span()); return logits_; } @@ -132,26 +108,6 @@ void Logits::Update(const DeviceSpan& next_tokens, size_t new_kv_length state_.outputs_[output_index_] = output_raw_->GetOrtTensor(); } -void Logits::HandleEOSArray(std::span batched_logits) { - if (model_.config_->model.eos_token_ids.empty()) - return; - - const size_t vocab_size = shape_[2]; - size_t vocab_index = 0; // Simpler math to have this index go up by vocab_size for every logit chunk we process - - for (int index = 0; index < shape_[0]; index++) { - auto logits = batched_logits.subspan(vocab_index, vocab_size); - float max = std::numeric_limits::lowest(); - for (auto id : model_.config_->model.eos_token_ids) { - max = std::max(max, logits[id]); - logits[id] = std::numeric_limits::lowest(); // Set all EOS token options to never happen (the first will get the max of all) - } - - logits[model_.config_->model.eos_token_id] = max; // Set the score of the primary EOS token to the highest of any of the EOS tokens - vocab_index += vocab_size; - } -} - void Logits::Add() { output_index_ = state_.outputs_.size(); diff --git a/src/models/logits.h b/src/models/logits.h index 9ccdce3c21..82c433180a 100644 --- a/src/models/logits.h +++ b/src/models/logits.h @@ -16,8 +16,6 @@ struct Logits { void Update(const DeviceSpan& next_tokens, size_t new_kv_length); private: - void HandleEOSArray(std::span logits); - State& state_; const Model& model_{state_.model_}; size_t output_index_{~0U}; @@ -37,8 +35,6 @@ struct Logits { // OrtValue wrapped in a DeviceMemory object to make it universal DeviceSpan logits_; - DeviceSpan cuda_eos_token_ids_; // eos_token_ids from params, but in cuda accessible memory - // Set to true when prefill will generate the already 'trimmed' logits required for sampling. bool trimmed_prefill_logits_ = false; }; diff --git a/src/search.cpp b/src/search.cpp index 32cf4b266a..93019f26cd 100644 --- a/src/search.cpp +++ b/src/search.cpp @@ -53,6 +53,7 @@ DeviceSpan Search_Cpu::GetLogits() const { void Search_Cpu::SetLogits(DeviceSpan logits) { next_token_scores_ = logits; + next_token_scores_.CopyDeviceToCpu(); // To the device->cpu copy once here as all later calls use CpuSpan() } DeviceSpan GreedySearch_Cpu::GetNextTokens() { @@ -246,7 +247,7 @@ bool GreedySearch_Cpu::PadIfAlreadyEOS(size_t batch_id) { void GreedySearch_Cpu::SetNextToken(size_t batch_id, int32_t token) { next_tokens_[batch_id] = token; - if (token == params_->config.model.eos_token_id) { + if (contains(params_->config.model.eos_token_id, token)) { eos_seen_[batch_id] = true; if (g_log.enabled && g_log.hit_eos) Log("hit_eos", "EOS seen on batch " + std::to_string(batch_id)); @@ -396,7 +397,8 @@ void Search_Cpu::ApplyMinLength(int min_length) { const int batch_beam_size = params_->BatchBeamSize(); for (int i = 0; i < batch_beam_size; i++) { std::span const beam_token_scores = GetScores(i); - beam_token_scores[params_->config.model.eos_token_id] = std::numeric_limits::lowest(); + for (auto token_id : params_->config.model.eos_token_id) + beam_token_scores[token_id] = std::numeric_limits::lowest(); } } diff --git a/src/smartptrs.h b/src/smartptrs.h index 12612997a4..b67fe507f1 100644 --- a/src/smartptrs.h +++ b/src/smartptrs.h @@ -120,7 +120,6 @@ struct DeviceInterface { virtual bool UpdatePositionIds(void* /*position_ids*/, int /*batch_beam_size*/, int /*total_length*/, int /*new_kv_length*/, ONNXTensorElementDataType /*type*/) { return false; } virtual bool UpdateAttentionMask(void* /*next_mask_data*/, void* /*mask_data*/, int /*batch_beam_size*/, int /*new_kv_length*/, int /*total_length*/, int /*max_length*/, bool /*update_only*/, ONNXTensorElementDataType /*type*/) { return false; } - virtual void LaunchHandleEOSArray(float* /*batch_logits*/, int /*batch_beam_size*/, int /*vocab_size*/, const int32_t* /*eos_token_ids*/, int /*eos_token_ids_count*/) { assert(false); } virtual void UpdateCacheIndirectionKernelLauncher(int32_t* /*tgt_indir_cache*/, const int32_t* /*src_indir_cache*/, const int32_t* /*beam_ids*/, int /*batch_size*/, int /*beam_width*/, int /*input_seq_length*/, int /*max_seq_length*/, int /*current_length*/) { assert(false); } virtual void ReorderPastStatesKernelLauncher(void* /*out_buffer*/, const void* /*in_buffer*/, int /*batch_size*/, int /*num_heads*/, int /*max_length*/, int /*head_size*/, int /*chunk_size*/) { assert(false); } virtual void LaunchCopyCrossQKSingleDecodeStep(float* /*cross_qk_buffer_data*/, float** /*qk_layer_pointers*/, int /*token_index*/, int /*batch_beam_size*/, int /*num_layers*/, int /*num_heads*/, int /*num_alignment_heads*/, const int* /*alignment_heads*/, int /*frames*/, int /*max_length*/) { assert(false); }