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

@@ -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 room-local agent: `!bai agent create-room-local openai my-openai-agent`
- create a global agent: `!bai agent create-global 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 ### OpenAI Compatible

View File

@@ -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

View File

@@ -129,7 +129,7 @@ impl ControllerTrait for Controller {
&text_generation_config.model_id, &text_generation_config.model_id,
&prompt_message, &prompt_message,
conversation_messages, conversation_messages,
text_generation_config.max_response_tokens, Some(text_generation_config.max_response_tokens),
text_generation_config.max_context_tokens, text_generation_config.max_context_tokens,
); );

View File

@@ -15,7 +15,7 @@ pub fn default_config() -> Config {
if let Some(ref mut config) = config.text_generation.as_mut() { if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "llama3-70b-8192".to_owned(); config.model_id = "llama3-70b-8192".to_owned();
config.max_context_tokens = 131_072; 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() { if let Some(ref mut config) = config.speech_to_text.as_mut() {

View File

@@ -13,7 +13,7 @@ pub fn default_config() -> Config {
if let Some(ref mut config) = config.text_generation.as_mut() { if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "gpt-4".to_owned(); config.model_id = "gpt-4".to_owned();
config.max_context_tokens = 128_000; 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() { if let Some(ref mut config) = config.text_to_speech.as_mut() {

View File

@@ -17,7 +17,7 @@ pub fn default_config() -> Config {
if let Some(ref mut config) = config.text_generation.as_mut() { if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "gemma2:2b".to_owned(); config.model_id = "gemma2:2b".to_owned();
config.max_context_tokens = 128_000; config.max_context_tokens = 128_000;
config.max_response_tokens = 4096; config.max_response_tokens = Some(4096);
} }
config config

View File

@@ -56,7 +56,7 @@ pub struct TextGenerationConfig {
pub temperature: f32, pub temperature: f32,
#[serde(default)] #[serde(default)]
pub max_response_tokens: u32, pub max_response_tokens: Option<u32>,
#[serde(default)] #[serde(default)]
pub max_context_tokens: u32, pub max_context_tokens: u32,
@@ -68,7 +68,7 @@ impl Default for TextGenerationConfig {
model_id: default_text_model_id(), model_id: default_text_model_id(),
prompt: Some(default_prompt().to_owned()), prompt: Some(default_prompt().to_owned()),
temperature: super::super::default_temperature(), temperature: super::super::default_temperature(),
max_response_tokens: 16_384, max_response_tokens: Some(16_384),
max_context_tokens: 128_000, max_context_tokens: 128_000,
} }
} }

View File

@@ -131,12 +131,18 @@ impl ControllerTrait for Controller {
.temperature_override .temperature_override
.unwrap_or(text_generation_config.temperature); .unwrap_or(text_generation_config.temperature);
let request = CreateChatCompletionRequestArgs::default() let mut request_builder = CreateChatCompletionRequestArgs::default();
.max_tokens(text_generation_config.max_response_tokens)
request_builder
.model(&text_generation_config.model_id) .model(&text_generation_config.model_id)
.temperature(temperature) .temperature(temperature)
.messages(openai_conversation_messages) .messages(openai_conversation_messages);
.build()?;
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) { if let Ok(request_as_json) = serde_json::to_string(&request) {
tracing::trace!( tracing::trace!(

View File

@@ -66,7 +66,7 @@ pub struct TextGenerationConfig {
pub temperature: f32, pub temperature: f32,
#[serde(default)] #[serde(default)]
pub max_response_tokens: u32, pub max_response_tokens: Option<u32>,
#[serde(default)] #[serde(default)]
pub max_context_tokens: u32, pub max_context_tokens: u32,
@@ -78,7 +78,7 @@ impl Default for TextGenerationConfig {
model_id: default_text_model_id(), model_id: default_text_model_id(),
prompt: Some(default_prompt().to_owned()), prompt: Some(default_prompt().to_owned()),
temperature: super::super::default_temperature(), temperature: super::super::default_temperature(),
max_response_tokens: 4096, max_response_tokens: Some(4096),
max_context_tokens: 128_000, max_context_tokens: 128_000,
} }
} }

View File

@@ -131,12 +131,15 @@ impl ControllerTrait for Controller {
let max_tokens = text_generation_config let max_tokens = text_generation_config
.max_response_tokens .max_response_tokens
.try_into() .map(|max_response_tokens| {
.expect("Failed converting max_response_tokens from u32 to i32"); max_response_tokens
.try_into()
.expect("Failed converting max_response_tokens from u32 to i32")
});
let request = ChatBody { let request = ChatBody {
model: text_generation_config.model_id.clone(), model: text_generation_config.model_id.clone(),
max_tokens: Some(max_tokens), max_tokens,
temperature: Some(temperature), temperature: Some(temperature),
top_p: None, top_p: None,
n: Some(1), n: Some(1),

View File

@@ -56,7 +56,7 @@ pub fn default_config() -> Config {
if let Some(text_generation) = &mut config.text_generation { if let Some(text_generation) = &mut config.text_generation {
text_generation.model_id = "some-model".to_string(); 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; text_generation.max_context_tokens = 128_000;
} }

View File

@@ -14,7 +14,7 @@ pub fn default_config() -> Config {
if let Some(ref mut config) = config.text_generation.as_mut() { if let Some(ref mut config) = config.text_generation.as_mut() {
config.model_id = "mattshumer/reflection-70b:free".to_owned(); config.model_id = "mattshumer/reflection-70b:free".to_owned();
config.max_context_tokens = 8192; config.max_context_tokens = 8192;
config.max_response_tokens = 2048; config.max_response_tokens = Some(2048);
} }
config config

View File

@@ -14,7 +14,7 @@ pub fn default_config() -> Config {
if let Some(ref mut config) = config.text_generation.as_mut() { 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.model_id = "meta-llama/Meta-Llama-3.1-405B-Instruct-Turbo".to_owned();
config.max_context_tokens = 8192; config.max_context_tokens = 8192;
config.max_response_tokens = 2048; config.max_response_tokens = Some(2048);
} }
config config

View File

@@ -16,7 +16,7 @@ pub fn shorten_messages_list_to_context_size(
model: &str, model: &str,
prompt_message: &Option<Message>, prompt_message: &Option<Message>,
mut messages: Vec<Message>, mut messages: Vec<Message>,
max_response_tokens: u32, max_response_tokens: Option<u32>,
max_context_tokens: u32, max_context_tokens: u32,
) -> Vec<Message> { ) -> Vec<Message> {
// Loading the tokenization data is an expensive process, so // 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 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. // 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 { 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 { } else {
0 0
}; };
@@ -98,7 +99,7 @@ pub mod test {
let bpe = super::get_bpe_for_model(model); 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 { let prompt = super::Message {
author: super::Author::Prompt, author: super::Author::Prompt,
@@ -173,7 +174,7 @@ pub mod test {
&Some(prompt), &Some(prompt),
conversation_messages, conversation_messages,
max_response_tokens, 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()); assert_eq!(2, new_conversation_messages.len());
@@ -195,7 +196,7 @@ pub mod test {
let bpe = super::get_bpe_for_model(model); 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 { let prompt = super::Message {
author: super::Author::User, author: super::Author::User,
@@ -269,7 +270,7 @@ pub mod test {
&Some(prompt), &Some(prompt),
conversation_messages, conversation_messages,
max_response_tokens, 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()); assert_eq!(2, new_conversation_messages.len());