diff options
| author | sigoden <sigoden@gmail.com> | 2024-04-23 16:46:48 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-04-23 16:46:48 +0800 |
| commit | d1aafa11153ab689c21c2c57c47da52337d8e8d1 (patch) | |
| tree | dc20dc033e9d376aab09941835a842b22fe32c02 /src/client/vertexai.rs | |
| parent | 1cc89eff514669d6459f62cc4eed8e0b34d8ef0c (diff) | |
| download | aichat-d1aafa11153ab689c21c2c57c47da52337d8e8d1.tar.gz | |
feat: customize model's max_output_tokens (#428)
Diffstat (limited to 'src/client/vertexai.rs')
| -rw-r--r-- | src/client/vertexai.rs | 16 |
1 files changed, 9 insertions, 7 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index 88035ec..eceeb2c 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -81,9 +81,9 @@ impl VertexAIClient { let block_threshold = self.config.block_threshold.clone(); - let body = build_body(data, self.model.name.clone(), block_threshold)?; + let body = build_body(data, &self.model, block_threshold)?; - let model = self.model.name.clone(); + let model = &self.model.name; let url = format!("{api_base}/{}:{}", model, func); @@ -176,7 +176,7 @@ fn check_error(data: &Value) -> Result<()> { pub(crate) fn build_body( data: SendData, - _model: String, + model: &Model, block_threshold: Option<String>, ) -> Result<Value> { let SendData { @@ -228,7 +228,7 @@ pub(crate) fn build_body( ); } - let mut body = json!({ "contents": contents }); + let mut body = json!({ "contents": contents, "generationConfig": {} }); if let Some(block_threshold) = block_threshold { body["safetySettings"] = json!([ @@ -239,10 +239,12 @@ pub(crate) fn build_body( ]); } + if let Some(max_output_tokens) = model.max_output_tokens { + body["generationConfig"]["maxOutputTokens"] = max_output_tokens.into(); + } + if let Some(temperature) = temperature { - body["generationConfig"] = json!({ - "temperature": temperature, - }); + body["generationConfig"]["temperature"] = temperature.into(); } Ok(body) |
