summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-05-08 13:46:26 +0800
committerGitHub <noreply@github.com>2024-05-08 13:46:26 +0800
commit7762cd6bed7f95fa855144e01aa94e932869fe12 (patch)
tree66e3d9371f4db081a2ce87e763c5c717c355adfc /src/client
parent1c6c740381d0da4bcb70c6531bf93b406b94d9a6 (diff)
downloadaichat-7762cd6bed7f95fa855144e01aa94e932869fe12.tar.gz
refactor: model pass_max_tokens (#493)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/bedrock.rs6
-rw-r--r--src/client/claude.rs2
-rw-r--r--src/client/cloudflare.rs2
-rw-r--r--src/client/cohere.rs2
-rw-r--r--src/client/ernie.rs2
-rw-r--r--src/client/model.rs40
-rw-r--r--src/client/ollama.rs2
-rw-r--r--src/client/openai.rs2
-rw-r--r--src/client/qianwen.rs2
-rw-r--r--src/client/replicate.rs2
-rw-r--r--src/client/vertexai.rs2
11 files changed, 32 insertions, 32 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 {