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

@@ -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,
);

View File

@@ -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() {

View File

@@ -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() {

View File

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

View File

@@ -56,7 +56,7 @@ pub struct TextGenerationConfig {
pub temperature: f32,
#[serde(default)]
pub max_response_tokens: u32,
pub max_response_tokens: Option<u32>,
#[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,
}
}

View File

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

View File

@@ -66,7 +66,7 @@ pub struct TextGenerationConfig {
pub temperature: f32,
#[serde(default)]
pub max_response_tokens: u32,
pub max_response_tokens: Option<u32>,
#[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,
}
}

View File

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

View File

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

View File

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

View File

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