summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-24 16:12:38 +0800
committerGitHub <noreply@github.com>2024-04-24 16:12:38 +0800
commita17f349daa0609402c51205929b9f2b32f6fb1bb (patch)
tree380dd9c158f8d9278355634a8d474df3a90f8d25 /src/client
parent040c48b9b392e3329d3f1de8eeb1b8773129773c (diff)
downloadaichat-a17f349daa0609402c51205929b9f2b32f6fb1bb.tar.gz
feat: support customizing `top_p` parameter (#434)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/claude.rs4
-rw-r--r--src/client/cohere.rs4
-rw-r--r--src/client/common.rs1
-rw-r--r--src/client/ernie.rs4
-rw-r--r--src/client/ollama.rs4
-rw-r--r--src/client/openai.rs12
-rw-r--r--src/client/qianwen.rs53
-rw-r--r--src/client/vertexai.rs7
8 files changed, 55 insertions, 34 deletions
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 283ae99..054731c 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -140,6 +140,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
let SendData {
mut messages,
temperature,
+ top_p,
stream,
} = data;
@@ -205,6 +206,9 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
if let Some(v) = temperature {
body["temperature"] = v.into();
}
+ if let Some(v) = top_p {
+ body["top_p"] = v.into();
+ }
if stream {
body["stream"] = true.into();
}
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index 4d713b6..cfab0fa 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -110,6 +110,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
let SendData {
mut messages,
temperature,
+ top_p,
stream,
} = data;
@@ -173,6 +174,9 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
if let Some(temperature) = temperature {
body["temperature"] = temperature.into();
}
+ if let Some(top_p) = top_p {
+ body["p"] = top_p.into();
+ }
if stream {
body["stream"] = true.into();
}
diff --git a/src/client/common.rs b/src/client/common.rs
index 62c5f92..593140d 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -323,6 +323,7 @@ pub struct ExtraConfig {
pub struct SendData {
pub messages: Vec<Message>,
pub temperature: Option<f64>,
+ pub top_p: Option<f64>,
pub stream: bool,
}
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 68d1ace..ffe10f0 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -230,6 +230,7 @@ fn build_body(data: SendData, model: &Model) -> Value {
let SendData {
mut messages,
temperature,
+ top_p,
stream,
} = data;
@@ -242,6 +243,9 @@ fn build_body(data: SendData, model: &Model) -> Value {
if let Some(temperature) = temperature {
body["temperature"] = temperature.into();
}
+ if let Some(top_p) = top_p {
+ body["top_p"] = top_p.into();
+ }
if let Some(max_output_tokens) = model.max_output_tokens {
body["max_output_tokens"] = max_output_tokens.into();
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index 055730a..434e487 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -122,6 +122,7 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
let SendData {
messages,
temperature,
+ top_p,
stream,
} = data;
@@ -185,6 +186,9 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
if let Some(temperature) = temperature {
body["options"]["temperature"] = temperature.into();
}
+ if let Some(top_p) = top_p {
+ body["options"]["top_p"] = top_p.into();
+ }
Ok(body)
}
diff --git a/src/client/openai.rs b/src/client/openai.rs
index f35c6ab..2c5d99b 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -130,6 +130,7 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value {
let SendData {
messages,
temperature,
+ top_p,
stream,
} = data;
@@ -139,13 +140,16 @@ pub fn openai_build_body(data: SendData, model: &Model) -> Value {
});
if let Some(max_tokens) = model.max_output_tokens {
- body["max_tokens"] = json!(max_tokens);
+ body["max_tokens"] = max_tokens.into();
} else if model.name == "gpt-4-vision-preview" {
// The default max_tokens of gpt-4-vision-preview is only 16, we need to make it larger
- body["max_tokens"] = json!(4096);
+ body["max_tokens"] = 4096.into();
}
- if let Some(v) = temperature {
- body["temperature"] = v.into();
+ if let Some(temperature) = temperature {
+ body["temperature"] = temperature.into();
+ }
+ if let Some(top_p) = top_p {
+ body["top_p"] = top_p.into();
}
if stream {
body["stream"] = true.into();
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index fb97964..c5a72e0 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -130,7 +130,6 @@ async fn send_message_streaming(
is_vl: bool,
) -> Result<()> {
let mut es = builder.eventsource()?;
- let mut offset = 0;
while let Some(event) = es.next().await {
match event {
@@ -139,12 +138,10 @@ async fn send_message_streaming(
let data: Value = serde_json::from_str(&message.data)?;
catch_error(&data)?;
if is_vl {
- let text =
- data["output"]["choices"][0]["message"]["content"][0]["text"].as_str();
- if let Some(text) = text {
- let text = &text[offset..];
+ if let Some(text) =
+ data["output"]["choices"][0]["message"]["content"][0]["text"].as_str()
+ {
handler.text(text)?;
- offset += text.len();
}
} else if let Some(text) = data["output"]["text"].as_str() {
handler.text(text)?;
@@ -169,11 +166,12 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
let SendData {
messages,
temperature,
+ top_p,
stream,
} = data;
let mut has_upload = false;
- let (input, parameters) = if is_vl {
+ let input = if is_vl {
let messages: Vec<Value> = messages
.into_iter()
.map(|message| {
@@ -199,40 +197,37 @@ fn build_body(data: SendData, model: &Model, is_vl: bool) -> Result<(Value, bool
})
.collect();
- let input = json!({
+ json!({
"messages": messages,
- });
-
- let mut parameters = json!({});
- if let Some(v) = temperature {
- parameters["temperature"] = v.into();
- }
- (input, parameters)
+ })
} else {
- let input = json!({
+ json!({
"messages": messages,
- });
+ })
+ };
- let mut parameters = json!({});
- if stream {
- parameters["incremental_output"] = true.into();
- }
+ let mut parameters = json!({});
+ if stream {
+ parameters["incremental_output"] = true.into();
+ }
- if let Some(max_tokens) = model.max_output_tokens {
- parameters["max_tokens"] = max_tokens.into();
- }
+ if let Some(max_tokens) = model.max_output_tokens {
+ parameters["max_tokens"] = max_tokens.into();
+ }
- if let Some(v) = temperature {
- parameters["temperature"] = v.into();
- }
- (input, parameters)
- };
+ if let Some(temperature) = temperature {
+ parameters["temperature"] = temperature.into();
+ }
+ if let Some(top_p) = top_p {
+ parameters["top_p"] = top_p.into();
+ }
let body = json!({
"model": &model.name,
"input": input,
"parameters": parameters
});
+
Ok((body, has_upload))
}
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index d9c1019..e0ae567 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -158,7 +158,8 @@ pub(crate) fn build_body(
let SendData {
mut messages,
temperature,
- ..
+ top_p,
+ stream: _,
} = data;
patch_system_message(&mut messages);
@@ -223,6 +224,10 @@ pub(crate) fn build_body(
body["generationConfig"]["temperature"] = temperature.into();
}
+ if let Some(top_p) = top_p {
+ body["generationConfig"]["topP"] = top_p.into();
+ }
+
Ok(body)
}