diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-24 16:12:38 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-24 16:12:38 +0800 |
| commit | a17f349daa0609402c51205929b9f2b32f6fb1bb (patch) | |
| tree | 380dd9c158f8d9278355634a8d474df3a90f8d25 /src/client | |
| parent | 040c48b9b392e3329d3f1de8eeb1b8773129773c (diff) | |
| download | aichat-a17f349daa0609402c51205929b9f2b32f6fb1bb.tar.gz | |
feat: support customizing `top_p` parameter (#434)
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/claude.rs | 4 | ||||
| -rw-r--r-- | src/client/cohere.rs | 4 | ||||
| -rw-r--r-- | src/client/common.rs | 1 | ||||
| -rw-r--r-- | src/client/ernie.rs | 4 | ||||
| -rw-r--r-- | src/client/ollama.rs | 4 | ||||
| -rw-r--r-- | src/client/openai.rs | 12 | ||||
| -rw-r--r-- | src/client/qianwen.rs | 53 | ||||
| -rw-r--r-- | src/client/vertexai.rs | 7 |
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) } |
