Add support for OpenAI's o1 models by making max_response_tokens optional

The other prerequisite seems to be not using a `prompt` (`prompt: null`),
but we already supported this.

It'd be nice to add an optional `max_completion_tokens` parameter as
well, for the benefit of the o1 models, but this is not yet supported by
async-openai.
Possibly tracked here: https://github.com/64bit/async-openai/issues/272
This commit is contained in:
Slavi Pantaleev
2024-10-03 10:36:01 +03:00
parent 90fbad5b64
commit db9422740c
14 changed files with 62 additions and 25 deletions

View File

@@ -16,7 +16,7 @@ pub fn shorten_messages_list_to_context_size(
model: &str,
prompt_message: &Option<Message>,
mut messages: Vec<Message>,
max_response_tokens: u32,
max_response_tokens: Option<u32>,
max_context_tokens: u32,
) -> Vec<Message> {
// Loading the tokenization data is an expensive process, so
@@ -26,7 +26,8 @@ pub fn shorten_messages_list_to_context_size(
// We want to retain the prompt in all cases, so we always count it first.
// We also always reserve enough tokens for the maximum response we expect.
let mut current_context_length: u32 = if let Some(prompt_message) = prompt_message {
calculate_token_size_for_message(&bpe, model, prompt_message) + max_response_tokens
calculate_token_size_for_message(&bpe, model, prompt_message)
+ max_response_tokens.unwrap_or(0)
} else {
0
};
@@ -98,7 +99,7 @@ pub mod test {
let bpe = super::get_bpe_for_model(model);
let max_response_tokens: u32 = 5;
let max_response_tokens: Option<u32> = Some(5);
let prompt = super::Message {
author: super::Author::Prompt,
@@ -173,7 +174,7 @@ pub mod test {
&Some(prompt),
conversation_messages,
max_response_tokens,
prompt_length + max_response_tokens + forth_length + third_length,
prompt_length + max_response_tokens.unwrap_or(0) + forth_length + third_length,
);
assert_eq!(2, new_conversation_messages.len());
@@ -195,7 +196,7 @@ pub mod test {
let bpe = super::get_bpe_for_model(model);
let max_response_tokens: u32 = 5;
let max_response_tokens: Option<u32> = Some(5);
let prompt = super::Message {
author: super::Author::User,
@@ -269,7 +270,7 @@ pub mod test {
&Some(prompt),
conversation_messages,
max_response_tokens,
prompt_length + max_response_tokens + forth_length + third_length,
prompt_length + max_response_tokens.unwrap_or(0) + forth_length + third_length,
);
assert_eq!(2, new_conversation_messages.len());