whisper : add support for --carry-initial-prompt (#3395)

* Add support for --carry-initial-prompt

* PR fixes for ruby and go

* Refactoring for readability

* WIP 1

* WIP 2

* PR fixes

* More PR fixes

* PR fix

* Further simplification

* d'oh

* One more logic fix

* Update src/whisper.cpp

Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>

* Truncate prompt_past0 upon initialization

* Slight simplification

---------

Co-authored-by: Georgi Gerganov <ggerganov@gmail.com>
This commit is contained in:
Andreas Lubbe
2025-10-10 19:51:15 +03:00
committed by GitHub
co-authored by Georgi Gerganov
parent a0ca50f3b9
commit 85871a9469
8 changed files with 257 additions and 162 deletions
+67 -26
View File
@@ -140,6 +140,10 @@ static void whisper_log_callback_default(ggml_log_level level, const char * text
} while (0)
#define WHISPER_MAX_DECODERS 8
// temperature below which we condition on past text history
static constexpr float WHISPER_HISTORY_CONDITIONING_TEMP_CUTOFF = 0.5f;
#define WHISPER_MAX_NODES 4096
static std::string format(const char * fmt, ...) {
@@ -882,7 +886,10 @@ struct whisper_state {
std::vector<float> logits;
std::vector<whisper_segment> result_all;
std::vector<whisper_token> prompt_past;
// prompt history split into static prefix (prompt_past0) and dynamic rolling context (prompt_past1)
std::vector<whisper_token> prompt_past0; // static carried initial prompt (if enabled)
std::vector<whisper_token> prompt_past1; // dynamic context from decoded output
int lang_id = 0; // english by default
@@ -5922,9 +5929,10 @@ struct whisper_full_params whisper_full_default_params(enum whisper_sampling_str
/* suppress_regex =*/ nullptr,
/*.initial_prompt =*/ nullptr,
/*.prompt_tokens =*/ nullptr,
/*.prompt_n_tokens =*/ 0,
/*.initial_prompt =*/ nullptr,
/*.carry_initial_prompt =*/ false,
/*.prompt_tokens =*/ nullptr,
/*.prompt_n_tokens =*/ 0,
/*.language =*/ "en",
/*.detect_language =*/ false,
@@ -6880,17 +6888,22 @@ int whisper_full_with_state(
decoder.rng = std::mt19937(j);
}
// the accumulated text context so far
auto & prompt_past = state->prompt_past;
// the accumulated text context split into static (prompt_past0) and dynamic (prompt_past1)
auto & prompt_past0 = state->prompt_past0;
auto & prompt_past1 = state->prompt_past1;
if (params.no_context) {
prompt_past.clear();
prompt_past0.clear();
prompt_past1.clear();
}
// calculate the maximum context budget for prompt history
const int max_prompt_ctx = std::min(params.n_max_text_ctx, whisper_n_text_ctx(ctx)/2);
// prepare prompt
{
std::vector<whisper_token> prompt_tokens;
// initial prompt
// tokenize the initial prompt
if (!params.prompt_tokens && params.initial_prompt) {
prompt_tokens.resize(1024);
int n_needed = whisper_tokenize(ctx, params.initial_prompt, prompt_tokens.data(), prompt_tokens.size());
@@ -6902,14 +6915,25 @@ int whisper_full_with_state(
params.prompt_tokens = prompt_tokens.data();
params.prompt_n_tokens = prompt_tokens.size();
}
// prepend the prompt tokens to the prompt_past
if (params.prompt_tokens && params.prompt_n_tokens > 0) {
// parse tokens from the pointer
for (int i = 0; i < params.prompt_n_tokens; i++) {
prompt_past.push_back(params.prompt_tokens[i]);
if (params.carry_initial_prompt) {
if (prompt_past0.empty()) {
const int max_tokens = std::max(1, max_prompt_ctx - 1);
if (params.prompt_n_tokens > max_tokens) {
WHISPER_LOG_WARN("%s: initial prompt is too long (%d tokens), will use only the last %d tokens\n",
__func__, params.prompt_n_tokens, max_tokens);
}
const int n_tokens = std::min(params.prompt_n_tokens, max_tokens);
prompt_past0.assign(params.prompt_tokens + (params.prompt_n_tokens - n_tokens), params.prompt_tokens + params.prompt_n_tokens);
}
} else {
for (int i = 0; i < params.prompt_n_tokens; ++i) {
prompt_past1.push_back(params.prompt_tokens[i]);
}
std::rotate(prompt_past1.begin(), prompt_past1.end() - params.prompt_n_tokens, prompt_past1.end());
}
std::rotate(prompt_past.begin(), prompt_past.end() - params.prompt_n_tokens, prompt_past.end());
}
}
@@ -6995,7 +7019,8 @@ int whisper_full_with_state(
// if there is a very short audio segment left to process, we remove any past prompt since it tends
// to confuse the decoder and often make it repeat or hallucinate stuff
if (seek > seek_start && seek + 500 >= seek_end) {
prompt_past.clear();
prompt_past0.clear();
prompt_past1.clear();
}
int best_decoder_id = 0;
@@ -7056,12 +7081,25 @@ int whisper_full_with_state(
{
prompt.clear();
// if we have already generated some text, use it as a prompt to condition the next generation
if (!prompt_past.empty() && t_cur < 0.5f && params.n_max_text_ctx > 0) {
int n_take = std::min(std::min(params.n_max_text_ctx, whisper_n_text_ctx(ctx)/2), int(prompt_past.size()));
if (params.n_max_text_ctx > 0 && t_cur < WHISPER_HISTORY_CONDITIONING_TEMP_CUTOFF) {
const bool can_take0 = params.carry_initial_prompt && !prompt_past0.empty();
const bool can_take1 = !prompt_past1.empty();
prompt = { whisper_token_prev(ctx) };
prompt.insert(prompt.begin() + 1, prompt_past.end() - n_take, prompt_past.end());
if (max_prompt_ctx > 0 && (can_take0 || can_take1)) {
// Always start with previous token marker to connect continuity
prompt.push_back(whisper_token_prev(ctx));
// Take static tokens (initial prompt) first
int n_take0 = 0;
if (can_take0) {
n_take0 = prompt_past0.size();
prompt.insert(prompt.end(), prompt_past0.end() - n_take0, prompt_past0.end());
}
// Fill remaining budget with dynamic tokens (rolling context)
const int n_take1 = std::min<int>(max_prompt_ctx - n_take0 - 1, prompt_past1.size());
prompt.insert(prompt.end(), prompt_past1.end() - n_take1, prompt_past1.end());
}
}
// init new transcription with sot, language (opt) and task tokens
@@ -7543,14 +7581,17 @@ int whisper_full_with_state(
//WHISPER_LOG_DEBUG("prompt_init.size() = %d, prompt.size() = %d, result_len = %d, seek_delta = %d\n", prompt_init.size(), prompt.size(), result_len, seek_delta);
// update prompt_past
prompt_past.clear();
if (prompt.front() == whisper_token_prev(ctx)) {
prompt_past.insert(prompt_past.end(), prompt.begin() + 1, prompt.end() - prompt_init.size());
// update prompt_past1
prompt_past1.clear();
if (!params.carry_initial_prompt && !prompt.empty() && prompt.front() == whisper_token_prev(ctx)) {
prompt_past1.insert(prompt_past1.end(), prompt.begin() + 1, prompt.end() - prompt_init.size());
}
for (int i = 0; i < result_len && !is_no_speech; ++i) {
prompt_past.push_back(tokens_cur[i].id);
// Add newly decoded tokens to the rolling context
if (!is_no_speech) {
for (int i = 0; i < result_len; ++i) {
prompt_past1.push_back(tokens_cur[i].id);
}
}
if (!tokens_cur.empty() && ctx->model.n_loaded > 0 && !is_no_speech) {