From db9422740ceca32956d9628b6326b8be206344e2 Mon Sep 17 00:00:00 2001 From: Slavi Pantaleev Date: Thu, 3 Oct 2024 10:36:01 +0300 Subject: [PATCH] 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 --- docs/providers.md | 5 +++- docs/sample-provider-configs/openai-o1.yml | 24 +++++++++++++++++++ src/agent/provider/anthropic/controller.rs | 2 +- src/agent/provider/groq/mod.rs | 2 +- src/agent/provider/localai/mod.rs | 2 +- src/agent/provider/ollama/mod.rs | 2 +- src/agent/provider/openai/config.rs | 4 ++-- src/agent/provider/openai/controller.rs | 14 +++++++---- src/agent/provider/openai_compat/config.rs | 4 ++-- .../provider/openai_compat/controller.rs | 9 ++++--- src/agent/provider/openai_compat/mod.rs | 2 +- src/agent/provider/openrouter/mod.rs | 2 +- src/agent/provider/togetherai/mod.rs | 2 +- src/conversation/llm/tokenization.rs | 13 +++++----- 14 files changed, 62 insertions(+), 25 deletions(-) create mode 100644 docs/sample-provider-configs/openai-o1.yml diff --git a/docs/providers.md b/docs/providers.md index 3c87d4f..6aa9274 100644 --- a/docs/providers.md +++ b/docs/providers.md @@ -125,7 +125,10 @@ For services which are not fully compatible with the OpenAI API, consider using - create a room-local agent: `!bai agent create-room-local openai my-openai-agent` - create a global agent: `!bai agent create-global openai my-openai-agent` -💡 When creating an agent, the bot will show you an up-to-date sample configuration for this provider which looks [like this](./sample-provider-configs/openai.yml). +💡 When creating an agent, the bot will show you an up-to-date sample configuration for this provider which: + +- in the general case looks [like this](./sample-provider-configs/openai.yml) +- for the [o1](https://platform.openai.com/docs/models/o1) models needs to look [like this](./sample-provider-configs/openai-o1.yml) ### OpenAI Compatible diff --git a/docs/sample-provider-configs/openai-o1.yml b/docs/sample-provider-configs/openai-o1.yml new file mode 100644 index 0000000..db62598 --- /dev/null +++ b/docs/sample-provider-configs/openai-o1.yml @@ -0,0 +1,24 @@ +base_url: https://api.openai.com/v1 +api_key: YOUR_API_KEY_HERE +text_generation: + model_id: o1-mini + # o1 models do not support a system prompt + prompt: null + temperature: 1.0 + # o1 models do not support max_response_tokens. + # They use `max_completion_tokens` as an alternative, + # but we don't support it yet (see https://github.com/64bit/async-openai/issues/272). + max_response_tokens: null + max_context_tokens: 128000 +speech_to_text: + model_id: whisper-1 +text_to_speech: + model_id: tts-1-hd + voice: onyx + speed: 1.0 + response_format: opus +image_generation: + model_id: dall-e-3 + style: vivid + size: 1024x1024 + quality: standard diff --git a/src/agent/provider/anthropic/controller.rs b/src/agent/provider/anthropic/controller.rs index 0e69f6d..4d43980 100644 --- a/src/agent/provider/anthropic/controller.rs +++ b/src/agent/provider/anthropic/controller.rs @@ -129,7 +129,7 @@ impl ControllerTrait for Controller { &text_generation_config.model_id, &prompt_message, conversation_messages, - text_generation_config.max_response_tokens, + Some(text_generation_config.max_response_tokens), text_generation_config.max_context_tokens, ); diff --git a/src/agent/provider/groq/mod.rs b/src/agent/provider/groq/mod.rs index 1cf93ab..ee25e40 100644 --- a/src/agent/provider/groq/mod.rs +++ b/src/agent/provider/groq/mod.rs @@ -15,7 +15,7 @@ pub fn default_config() -> Config { if let Some(ref mut config) = config.text_generation.as_mut() { config.model_id = "llama3-70b-8192".to_owned(); config.max_context_tokens = 131_072; - config.max_response_tokens = 4096; + config.max_response_tokens = Some(4096); } if let Some(ref mut config) = config.speech_to_text.as_mut() { diff --git a/src/agent/provider/localai/mod.rs b/src/agent/provider/localai/mod.rs index 5f0e49a..9be909b 100644 --- a/src/agent/provider/localai/mod.rs +++ b/src/agent/provider/localai/mod.rs @@ -13,7 +13,7 @@ pub fn default_config() -> Config { if let Some(ref mut config) = config.text_generation.as_mut() { config.model_id = "gpt-4".to_owned(); config.max_context_tokens = 128_000; - config.max_response_tokens = 4096; + config.max_response_tokens = Some(4096); } if let Some(ref mut config) = config.text_to_speech.as_mut() { diff --git a/src/agent/provider/ollama/mod.rs b/src/agent/provider/ollama/mod.rs index 918a3b9..1f80d42 100644 --- a/src/agent/provider/ollama/mod.rs +++ b/src/agent/provider/ollama/mod.rs @@ -17,7 +17,7 @@ pub fn default_config() -> Config { if let Some(ref mut config) = config.text_generation.as_mut() { config.model_id = "gemma2:2b".to_owned(); config.max_context_tokens = 128_000; - config.max_response_tokens = 4096; + config.max_response_tokens = Some(4096); } config diff --git a/src/agent/provider/openai/config.rs b/src/agent/provider/openai/config.rs index bcfd623..fd2ca2a 100644 --- a/src/agent/provider/openai/config.rs +++ b/src/agent/provider/openai/config.rs @@ -56,7 +56,7 @@ pub struct TextGenerationConfig { pub temperature: f32, #[serde(default)] - pub max_response_tokens: u32, + pub max_response_tokens: Option, #[serde(default)] pub max_context_tokens: u32, @@ -68,7 +68,7 @@ impl Default for TextGenerationConfig { model_id: default_text_model_id(), prompt: Some(default_prompt().to_owned()), temperature: super::super::default_temperature(), - max_response_tokens: 16_384, + max_response_tokens: Some(16_384), max_context_tokens: 128_000, } } diff --git a/src/agent/provider/openai/controller.rs b/src/agent/provider/openai/controller.rs index 9597599..b1911bc 100644 --- a/src/agent/provider/openai/controller.rs +++ b/src/agent/provider/openai/controller.rs @@ -131,12 +131,18 @@ impl ControllerTrait for Controller { .temperature_override .unwrap_or(text_generation_config.temperature); - let request = CreateChatCompletionRequestArgs::default() - .max_tokens(text_generation_config.max_response_tokens) + let mut request_builder = CreateChatCompletionRequestArgs::default(); + + request_builder .model(&text_generation_config.model_id) .temperature(temperature) - .messages(openai_conversation_messages) - .build()?; + .messages(openai_conversation_messages); + + if let Some(max_response_tokens) = text_generation_config.max_response_tokens { + request_builder.max_tokens(max_response_tokens); + } + + let request = request_builder.build()?; if let Ok(request_as_json) = serde_json::to_string(&request) { tracing::trace!( diff --git a/src/agent/provider/openai_compat/config.rs b/src/agent/provider/openai_compat/config.rs index 27698b9..4ae4e52 100644 --- a/src/agent/provider/openai_compat/config.rs +++ b/src/agent/provider/openai_compat/config.rs @@ -66,7 +66,7 @@ pub struct TextGenerationConfig { pub temperature: f32, #[serde(default)] - pub max_response_tokens: u32, + pub max_response_tokens: Option, #[serde(default)] pub max_context_tokens: u32, @@ -78,7 +78,7 @@ impl Default for TextGenerationConfig { model_id: default_text_model_id(), prompt: Some(default_prompt().to_owned()), temperature: super::super::default_temperature(), - max_response_tokens: 4096, + max_response_tokens: Some(4096), max_context_tokens: 128_000, } } diff --git a/src/agent/provider/openai_compat/controller.rs b/src/agent/provider/openai_compat/controller.rs index d1e88b4..d56c9ab 100644 --- a/src/agent/provider/openai_compat/controller.rs +++ b/src/agent/provider/openai_compat/controller.rs @@ -131,12 +131,15 @@ impl ControllerTrait for Controller { let max_tokens = text_generation_config .max_response_tokens - .try_into() - .expect("Failed converting max_response_tokens from u32 to i32"); + .map(|max_response_tokens| { + max_response_tokens + .try_into() + .expect("Failed converting max_response_tokens from u32 to i32") + }); let request = ChatBody { model: text_generation_config.model_id.clone(), - max_tokens: Some(max_tokens), + max_tokens, temperature: Some(temperature), top_p: None, n: Some(1), diff --git a/src/agent/provider/openai_compat/mod.rs b/src/agent/provider/openai_compat/mod.rs index 9fc6c12..85b1061 100644 --- a/src/agent/provider/openai_compat/mod.rs +++ b/src/agent/provider/openai_compat/mod.rs @@ -56,7 +56,7 @@ pub fn default_config() -> Config { if let Some(text_generation) = &mut config.text_generation { text_generation.model_id = "some-model".to_string(); - text_generation.max_response_tokens = 4096; + text_generation.max_response_tokens = Some(4096); text_generation.max_context_tokens = 128_000; } diff --git a/src/agent/provider/openrouter/mod.rs b/src/agent/provider/openrouter/mod.rs index 56b0c43..0dad010 100644 --- a/src/agent/provider/openrouter/mod.rs +++ b/src/agent/provider/openrouter/mod.rs @@ -14,7 +14,7 @@ pub fn default_config() -> Config { if let Some(ref mut config) = config.text_generation.as_mut() { config.model_id = "mattshumer/reflection-70b:free".to_owned(); config.max_context_tokens = 8192; - config.max_response_tokens = 2048; + config.max_response_tokens = Some(2048); } config diff --git a/src/agent/provider/togetherai/mod.rs b/src/agent/provider/togetherai/mod.rs index ccf0a2c..e22117a 100644 --- a/src/agent/provider/togetherai/mod.rs +++ b/src/agent/provider/togetherai/mod.rs @@ -14,7 +14,7 @@ pub fn default_config() -> Config { if let Some(ref mut config) = config.text_generation.as_mut() { config.model_id = "meta-llama/Meta-Llama-3.1-405B-Instruct-Turbo".to_owned(); config.max_context_tokens = 8192; - config.max_response_tokens = 2048; + config.max_response_tokens = Some(2048); } config diff --git a/src/conversation/llm/tokenization.rs b/src/conversation/llm/tokenization.rs index 28330f1..93cbea0 100644 --- a/src/conversation/llm/tokenization.rs +++ b/src/conversation/llm/tokenization.rs @@ -16,7 +16,7 @@ pub fn shorten_messages_list_to_context_size( model: &str, prompt_message: &Option, mut messages: Vec, - max_response_tokens: u32, + max_response_tokens: Option, max_context_tokens: u32, ) -> Vec { // 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 = 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 = 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());