diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-08 13:46:26 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-08 13:46:26 +0800 |
| commit | 7762cd6bed7f95fa855144e01aa94e932869fe12 (patch) | |
| tree | 66e3d9371f4db081a2ce87e763c5c717c355adfc /src | |
| parent | 1c6c740381d0da4bcb70c6531bf93b406b94d9a6 (diff) | |
| download | aichat-7762cd6bed7f95fa855144e01aa94e932869fe12.tar.gz | |
refactor: model pass_max_tokens (#493)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/bedrock.rs | 6 | ||||
| -rw-r--r-- | src/client/claude.rs | 2 | ||||
| -rw-r--r-- | src/client/cloudflare.rs | 2 | ||||
| -rw-r--r-- | src/client/cohere.rs | 2 | ||||
| -rw-r--r-- | src/client/ernie.rs | 2 | ||||
| -rw-r--r-- | src/client/model.rs | 40 | ||||
| -rw-r--r-- | src/client/ollama.rs | 2 | ||||
| -rw-r--r-- | src/client/openai.rs | 2 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 2 | ||||
| -rw-r--r-- | src/client/replicate.rs | 2 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 2 | ||||
| -rw-r--r-- | src/config/mod.rs | 6 | ||||
| -rw-r--r-- | src/serve.rs | 4 |
13 files changed, 37 insertions, 37 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index 6abd939..b07152b 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -172,7 +172,7 @@ async fn send_message_streaming( let data: Value = decode_chunk(message.payload()).ok_or_else(|| { anyhow!("Invalid chunk data: {}", hex_encode(message.payload())) })?; - debug!("bedrock chunk: {data}"); + // debug!("bedrock chunk: {data}"); match model_category { ModelCategory::Anthropic => { if let Some(typ) = data["type"].as_str() { @@ -235,7 +235,7 @@ fn meta_llama_build_body(data: SendData, model: &Model, pt: PromptFormat) -> Res let prompt = generate_prompt(&messages, pt)?; let mut body = json!({ "prompt": prompt }); - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["max_gen_len"] = v.into(); } if let Some(v) = temperature { @@ -258,7 +258,7 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result<Value> { let prompt = generate_prompt(&messages, MISTRAL_PROMPT_FORMAT)?; let mut body = json!({ "prompt": prompt }); - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["max_tokens"] = v.into(); } if let Some(v) = temperature { diff --git a/src/client/claude.rs b/src/client/claude.rs index 0a230e9..89742f3 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -142,7 +142,7 @@ pub fn claude_build_body(data: SendData, model: &Model) -> Result<Value> { if let Some(v) = system_message { body["system"] = v.into(); } - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["max_tokens"] = v.into(); } if let Some(v) = temperature { diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs index 9758032..5a4bf8c 100644 --- a/src/client/cloudflare.rs +++ b/src/client/cloudflare.rs @@ -88,7 +88,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> { "messages": messages, }); - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["max_tokens"] = v.into(); } if let Some(v) = temperature { diff --git a/src/client/cohere.rs b/src/client/cohere.rs index e0ef6f0..b5d6647 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -135,7 +135,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> { body["chat_history"] = messages.into(); } - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["max_tokens"] = v.into(); } if let Some(v) = temperature { diff --git a/src/client/ernie.rs b/src/client/ernie.rs index d3002ff..28cb857 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -128,7 +128,7 @@ fn build_body(data: SendData, model: &Model) -> Value { "messages": messages, }); - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["max_output_tokens"] = v.into(); } if let Some(v) = temperature { diff --git a/src/client/model.rs b/src/client/model.rs index b24a546..d0d86e4 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -14,11 +14,11 @@ pub struct Model { pub name: String, pub max_input_tokens: Option<usize>, pub max_output_tokens: Option<isize>, - pub ref_max_output_tokens: Option<isize>, + pub pass_max_tokens: bool, pub input_price: Option<f64>, pub output_price: Option<f64>, - pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>, pub capabilities: ModelCapabilities, + pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>, } impl Default for Model { @@ -32,13 +32,13 @@ impl Model { Self { client_name: client_name.into(), name: name.into(), - extra_fields: None, max_input_tokens: None, max_output_tokens: None, - ref_max_output_tokens: None, + pass_max_tokens: false, input_price: None, output_price: None, capabilities: ModelCapabilities::Text, + extra_fields: None, } } @@ -49,8 +49,7 @@ impl Model { let mut model = Model::new(client_name, &v.name); model .set_max_input_tokens(v.max_input_tokens) - .set_max_output_tokens(v.max_output_tokens) - .set_ref_max_output_tokens(v.ref_max_output_tokens) + .set_max_tokens(v.max_output_tokens, v.pass_max_tokens) .set_input_price(v.input_price) .set_output_price(v.output_price) .set_supports_vision(v.supports_vision) @@ -97,7 +96,7 @@ impl Model { pub fn description(&self) -> String { let max_input_tokens = format_option_value(&self.max_input_tokens); - let max_output_tokens = format_option_value(&self.show_max_output_tokens()); + let max_output_tokens = format_option_value(&self.max_output_tokens); let input_price = format_option_value(&self.input_price); let output_price = format_option_value(&self.output_price); let vision = if self.capabilities.contains(ModelCapabilities::Vision) { @@ -115,8 +114,12 @@ impl Model { self.capabilities.contains(ModelCapabilities::Vision) } - pub fn show_max_output_tokens(&self) -> Option<isize> { - self.max_output_tokens.or(self.ref_max_output_tokens) + pub fn max_tokens_param(&self) -> Option<isize> { + if self.pass_max_tokens { + self.max_output_tokens + } else { + None + } } pub fn set_max_input_tokens(&mut self, max_input_tokens: Option<usize>) -> &mut Self { @@ -127,19 +130,16 @@ impl Model { self } - pub fn set_max_output_tokens(&mut self, max_output_tokens: Option<isize>) -> &mut Self { + pub fn set_max_tokens( + &mut self, + max_output_tokens: Option<isize>, + pass_max_tokens: bool, + ) -> &mut Self { match max_output_tokens { None | Some(0) => self.max_output_tokens = None, _ => self.max_output_tokens = max_output_tokens, } - self - } - - pub fn set_ref_max_output_tokens(&mut self, ref_max_output_tokens: Option<isize>) -> &mut Self { - match ref_max_output_tokens { - None | Some(0) => self.ref_max_output_tokens = None, - _ => self.ref_max_output_tokens = ref_max_output_tokens, - } + self.pass_max_tokens = pass_max_tokens; self } @@ -237,12 +237,12 @@ pub struct ModelConfig { pub name: String, pub max_input_tokens: Option<usize>, pub max_output_tokens: Option<isize>, - #[serde(rename = "max_output_tokens?")] - pub ref_max_output_tokens: Option<isize>, pub input_price: Option<f64>, pub output_price: Option<f64>, #[serde(default)] pub supports_vision: bool, + #[serde(default)] + pub pass_max_tokens: bool, pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>, } diff --git a/src/client/ollama.rs b/src/client/ollama.rs index b61417a..6408d2e 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -159,7 +159,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> { "options": {}, }); - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["options"]["num_predict"] = v.into(); } if let Some(v) = temperature { diff --git a/src/client/openai.rs b/src/client/openai.rs index 08bb94d..0b111db 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -90,7 +90,7 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value { "messages": messages, }); - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["max_tokens"] = v.into(); } if let Some(v) = temperature { diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index 3fa17e6..7391a38 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -173,7 +173,7 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool parameters["incremental_output"] = true.into(); } - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { parameters["max_tokens"] = v.into(); } if let Some(v) = temperature { diff --git a/src/client/replicate.rs b/src/client/replicate.rs index a20ce71..34cfd94 100644 --- a/src/client/replicate.rs +++ b/src/client/replicate.rs @@ -148,7 +148,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> { "prompt_template": "{prompt}" }); - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { input["max_tokens"] = v.into(); input["max_new_tokens"] = v.into(); } diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 9907a7f..4d06934 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -201,7 +201,7 @@ pub(crate) fn gemini_build_body( body["safetySettings"] = safety_settings; } - if let Some(v) = model.max_output_tokens { + if let Some(v) = model.max_tokens_param() { body["generationConfig"]["maxOutputTokens"] = v.into(); } if let Some(v) = temperature { diff --git a/src/config/mod.rs b/src/config/mod.rs index f33ee73..5ac44a8 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -422,7 +422,7 @@ impl Config { ( "max_output_tokens", self.model - .max_output_tokens + .max_tokens_param() .map(|v| format!("{v} (current model)")) .unwrap_or_else(|| "-".into()), ), @@ -523,7 +523,7 @@ impl Config { (values, args[0]) } else if args.len() == 2 { let values = match args[0] { - "max_output_tokens" => match self.model.show_max_output_tokens() { + "max_output_tokens" => match self.model.max_output_tokens { Some(v) => vec![v.to_string()], None => vec![], }, @@ -564,7 +564,7 @@ impl Config { match key { "max_output_tokens" => { let value = parse_value(value)?; - self.model.set_max_output_tokens(value); + self.model.set_max_tokens(value, true); } "temperature" => { let value = parse_value(value)?; diff --git a/src/serve.rs b/src/serve.rs index 5f748e8..4413c89 100644 --- a/src/serve.rs +++ b/src/serve.rs @@ -93,7 +93,7 @@ impl Server { "id": id, "max_input_tokens": model.max_input_tokens, "max_output_tokens": model.max_output_tokens, - "max_output_tokens?": model.ref_max_output_tokens, + "pass_max_tokens": model.pass_max_tokens, "input_price": model.input_price, "output_price": model.output_price, "supports_vision": model.supports_vision(), @@ -244,7 +244,7 @@ impl Server { let mut client = init_client(&config)?; if max_tokens.is_some() { - client.model_mut().set_max_output_tokens(max_tokens); + client.model_mut().set_max_tokens(max_tokens, true); } let abort = create_abort_signal(); let http_client = client.build_client()?; |
