mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-09-30 11:36:38 +02:00
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:
co-authored by
Georgi Gerganov
parent
a0ca50f3b9
commit
85871a9469
+67
-26
@@ -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) {
|
||||
|
||||
Reference in New Issue
Block a user