From 5eae392dbd4fc789829ac2f47a24614dbc5db5e9 Mon Sep 17 00:00:00 2001 From: sigoden Date: Wed, 1 May 2024 06:01:10 +0800 Subject: refactore: add models for openai-compatible platforms (#471) --- src/client/bedrock.rs | 12 ++++-------- src/client/common.rs | 3 ++- src/client/prompt_format.rs | 4 ++-- 3 files changed, 8 insertions(+), 11 deletions(-) (limited to 'src/client') diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index fb4dad8..d6ab696 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -2,7 +2,7 @@ use super::claude::{claude_build_body, claude_extract_completion}; use super::{ catch_error, generate_prompt, BedrockClient, Client, CompletionDetails, ExtraConfig, Model, ModelConfig, PromptAction, PromptFormat, PromptKind, SendData, SseHandler, - LLAMA2_PROMPT_FORMAT, LLAMA3_PROMPT_FORMAT, + LLAMA3_PROMPT_FORMAT, MISTRAL_PROMPT_FORMAT, }; use crate::utils::{base64_decode, encode_uri, hex_encode, hmac_sha256, sha256}; @@ -140,7 +140,7 @@ async fn send_message( match model_category { ModelCategory::Anthropic => claude_extract_completion(&data), - ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => llama_extract_completion(&data), + ModelCategory::MetaLlama3 => llama_extract_completion(&data), ModelCategory::Mistral => mistral_extrat_completion(&data), } } @@ -183,7 +183,7 @@ async fn send_message_streaming( } } } - ModelCategory::MetaLlama2 | ModelCategory::MetaLlama3 => { + ModelCategory::MetaLlama3 => { if let Some(text) = data["generation"].as_str() { handler.text(text)?; } @@ -220,7 +220,6 @@ fn build_body(data: SendData, model: &Model, model_category: &ModelCategory) -> body["anthropic_version"] = "bedrock-2023-05-31".into(); Ok(body) } - ModelCategory::MetaLlama2 => meta_llama_build_body(data, model, LLAMA2_PROMPT_FORMAT), ModelCategory::MetaLlama3 => meta_llama_build_body(data, model, LLAMA3_PROMPT_FORMAT), ModelCategory::Mistral => mistral_build_body(data, model), } @@ -256,7 +255,7 @@ fn mistral_build_body(data: SendData, model: &Model) -> Result { top_p, stream: _, } = data; - let prompt = generate_prompt(&messages, LLAMA2_PROMPT_FORMAT)?; + let prompt = generate_prompt(&messages, MISTRAL_PROMPT_FORMAT)?; let mut body = json!({ "prompt": prompt }); if let Some(v) = model.max_output_tokens { @@ -294,7 +293,6 @@ fn mistral_extrat_completion(data: &Value) -> Result<(String, CompletionDetails) #[derive(Debug, Clone, Copy, PartialEq, Eq)] enum ModelCategory { Anthropic, - MetaLlama2, MetaLlama3, Mistral, } @@ -305,8 +303,6 @@ impl FromStr for ModelCategory { fn from_str(s: &str) -> std::result::Result { if s.starts_with("anthropic.") { Ok(ModelCategory::Anthropic) - } else if s.starts_with("meta.llama2") { - Ok(ModelCategory::MetaLlama2) } else if s.starts_with("meta.llama3") { Ok(ModelCategory::MetaLlama3) } else if s.starts_with("mistral") { diff --git a/src/client/common.rs b/src/client/common.rs index 91d6bca..e889bc8 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -390,10 +390,11 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result Ok(None), - Some((name, _)) => { + Some((name, api_base)) => { let mut config = json!({ "type": "openai-compatible", "name": name, + "api_base": api_base, }); let prompts = if ALL_CLIENT_MODELS.iter().any(|v| &v.platform == name) { vec![("api_key", "API Key:", false, PromptKind::String)] diff --git a/src/client/prompt_format.rs b/src/client/prompt_format.rs index fca87ac..61647a8 100644 --- a/src/client/prompt_format.rs +++ b/src/client/prompt_format.rs @@ -22,7 +22,7 @@ pub const GENERIC_PROMPT_FORMAT: PromptFormat<'static> = PromptFormat { end: "### Assistant\n", }; -pub const LLAMA2_PROMPT_FORMAT: PromptFormat<'static> = PromptFormat { +pub const MISTRAL_PROMPT_FORMAT: PromptFormat<'static> = PromptFormat { begin: "", system_pre_message: "[INST] <>", system_post_message: "<> [/INST]", @@ -136,7 +136,7 @@ pub fn smart_prompt_format(model_name: &str) -> PromptFormat<'static> { || model_name.contains("mistral") || model_name.contains("mixtral") { - LLAMA2_PROMPT_FORMAT + MISTRAL_PROMPT_FORMAT } else if model_name.contains("phi3") || model_name.contains("phi-3") { PHI3_PROMPT_FORMAT } else if model_name.contains("command-r") { -- cgit v1.2.3