summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/bedrock.rs12
-rw-r--r--src/client/common.rs3
-rw-r--r--src/client/prompt_format.rs4
3 files changed, 8 insertions, 11 deletions
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<Value> {
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<Self, Self::Err> {
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<Option<(St
.find(|(name, _)| client == *name)
{
None => 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] <<SYS>>",
system_post_message: "<</SYS>> [/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") {