From 9c6c9f10a27d0993636b453f39d8934c95c5c2b2 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 23 Apr 2024 18:14:47 +0800 Subject: feat: builtin models can be overwrited by models config (#429) --- src/client/ernie.rs | 26 ++++++++++++++++---------- 1 file changed, 16 insertions(+), 10 deletions(-) (limited to 'src/client/ernie.rs') diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 0b10022..b0a0087 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,6 +1,6 @@ use super::{ - patch_system_message, Client, ErnieClient, ExtraConfig, Model, PromptType, ReplyHandler, - SendData, + patch_system_message, Client, ErnieClient, ExtraConfig, Model, ModelConfig, PromptType, + ReplyHandler, SendData, }; use crate::utils::PromptKind; @@ -73,6 +73,8 @@ pub struct ErnieConfig { pub name: Option, pub api_key: Option, pub secret_key: Option, + #[serde(default)] + pub models: Vec, pub extra: Option, } @@ -106,14 +108,18 @@ impl ErnieClient { pub fn list_models(local_config: &ErnieConfig) -> Vec { let client_name = Self::name(local_config); - MODELS - .into_iter() - .map(|(name, _, max_input_tokens, max_output_tokens)| { - Model::new(client_name, name) - .set_max_input_tokens(Some(max_input_tokens)) - .set_max_output_tokens(Some(max_output_tokens)) - }) // ERNIE tokenizer is different from cl100k_base - .collect() + if local_config.models.is_empty() { + MODELS + .into_iter() + .map(|(name, _, max_input_tokens, max_output_tokens)| { + Model::new(client_name, name) + .set_max_input_tokens(Some(max_input_tokens)) + .set_max_output_tokens(Some(max_output_tokens)) + }) // ERNIE tokenizer is different from cl100k_base + .collect() + } else { + Model::from_config(client_name, &local_config.models) + } } fn request_builder(&self, client: &ReqwestClient, data: SendData) -> Result { -- cgit v1.2.3