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

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