summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-03-06 08:35:40 +0800
committerGitHub <noreply@github.com>2024-03-06 08:35:40 +0800
commit8e5d4e55b1a5158a1f35adaf044dec545159f426 (patch)
treeb360cc4e9ad8acccea3c68b37057798cd1d657c0 /src/client
parentbe4e5e569a61c54d8ac8fb77144e9b1d01e3b81f (diff)
downloadaichat-8e5d4e55b1a5158a1f35adaf044dec545159f426.tar.gz
refactor: rename model's `max_tokens` to `max_input_tokens` (#339)
BREAKING CHANGE: rename model's `max_tokens` to `max_input_tokens`
Diffstat (limited to 'src/client')
-rw-r--r--src/client/azure_openai.rs6
-rw-r--r--src/client/claude.rs15
-rw-r--r--src/client/gemini.rs9
-rw-r--r--src/client/localai.rs6
-rw-r--r--src/client/mistral.rs5
-rw-r--r--src/client/model.rs22
-rw-r--r--src/client/ollama.rs6
-rw-r--r--src/client/openai.rs10
-rw-r--r--src/client/qianwen.rs14
-rw-r--r--src/client/vertexai.rs9
10 files changed, 54 insertions, 48 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index 5c9f3ef..e553f3f 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -28,8 +28,8 @@ impl AzureOpenAIClient {
("api_key", "API Key:", true, PromptKind::String),
("models[].name", "Model Name:", true, PromptKind::String),
(
- "models[].max_tokens",
- "Max Tokens:",
+ "models[].max_input_tokens",
+ "Max Input Tokens:",
true,
PromptKind::Integer,
),
@@ -43,7 +43,7 @@ impl AzureOpenAIClient {
.iter()
.map(|v| {
Model::new(client_name, &v.name)
- .set_max_tokens(v.max_tokens)
+ .set_max_input_tokens(v.max_input_tokens)
.set_capabilities(v.capabilities)
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 325bce6..78d3640 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -20,11 +20,12 @@ use serde_json::{json, Value};
const API_BASE: &str = "https://api.anthropic.com/v1/messages";
const MODELS: [(&str, usize, &str); 5] = [
- ("claude-3-opus-20240229", 204096, "text,vision"),
- ("claude-3-sonnet-20240229", 204096, "text,vision"),
- ("claude-2.1", 204096, "text"),
- ("claude-2.0", 104096, "text"),
- ("claude-instant-1.2", 104096, "text"),
+ // https://docs.anthropic.com/claude/docs/models-overview
+ ("claude-3-opus-20240229", 200000, "text,vision"),
+ ("claude-3-sonnet-20240229", 200000, "text,vision"),
+ ("claude-2.1", 200000, "text"),
+ ("claude-2.0", 100000, "text"),
+ ("claude-instant-1.2", 100000, "text"),
];
const TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2);
@@ -66,10 +67,10 @@ impl ClaudeClient {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens, capabilities)| {
+ .map(|(name, max_input_tokens, capabilities)| {
Model::new(client_name, name)
.set_capabilities(capabilities.into())
- .set_max_tokens(Some(max_tokens))
+ .set_max_input_tokens(Some(max_input_tokens))
.set_tokens_count_factors(TOKENS_COUNT_FACTORS)
})
.collect()
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 1c76fd8..68db15a 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -11,8 +11,9 @@ use serde::Deserialize;
const API_BASE: &str = "https://generativelanguage.googleapis.com/v1beta/models/";
const MODELS: [(&str, usize, &str); 2] = [
- ("gemini-pro", 32768, "text"),
- ("gemini-pro-vision", 16384, "vision"),
+ // https://ai.google.dev/models/gemini
+ ("gemini-pro", 30720, "text"),
+ ("gemini-pro-vision", 12288, "vision"),
];
const TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2);
@@ -54,10 +55,10 @@ impl GeminiClient {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens, capabilities)| {
+ .map(|(name, max_input_tokens, capabilities)| {
Model::new(client_name, name)
.set_capabilities(capabilities.into())
- .set_max_tokens(Some(max_tokens))
+ .set_max_input_tokens(Some(max_input_tokens))
.set_tokens_count_factors(TOKENS_COUNT_FACTORS)
})
.collect()
diff --git a/src/client/localai.rs b/src/client/localai.rs
index 7795039..0e9db0e 100644
--- a/src/client/localai.rs
+++ b/src/client/localai.rs
@@ -28,8 +28,8 @@ impl LocalAIClient {
("api_key", "API Key:", false, PromptKind::String),
("models[].name", "Model Name:", true, PromptKind::String),
(
- "models[].max_tokens",
- "Max Tokens:",
+ "models[].max_input_tokens",
+ "Max Input Tokens:",
false,
PromptKind::Integer,
),
@@ -44,7 +44,7 @@ impl LocalAIClient {
.map(|v| {
Model::new(client_name, &v.name)
.set_capabilities(v.capabilities)
- .set_max_tokens(v.max_tokens)
+ .set_max_input_tokens(v.max_input_tokens)
.set_extra_fields(v.extra_fields.clone())
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
diff --git a/src/client/mistral.rs b/src/client/mistral.rs
index 1bb4889..ba96116 100644
--- a/src/client/mistral.rs
+++ b/src/client/mistral.rs
@@ -11,6 +11,7 @@ use serde::Deserialize;
const API_URL: &str = "https://api.mistral.ai/v1/chat/completions";
const MODELS: [(&str, usize, &str); 5] = [
+ // https://docs.mistral.ai/platform/endpoints/
("mistral-small-latest", 32000, "text"),
("mistral-medium-latest", 32000, "text"),
("mistral-larget-latest", 32000, "text"),
@@ -39,10 +40,10 @@ impl MistralClient {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens, capabilities)| {
+ .map(|(name, max_input_tokens, capabilities)| {
Model::new(client_name, name)
.set_capabilities(capabilities.into())
- .set_max_tokens(Some(max_tokens))
+ .set_max_input_tokens(Some(max_input_tokens))
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
.collect()
diff --git a/src/client/model.rs b/src/client/model.rs
index b29166e..ce181a5 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -11,7 +11,7 @@ pub type TokensCountFactors = (usize, usize); // (per-messages, bias)
pub struct Model {
pub client_name: String,
pub name: String,
- pub max_tokens: Option<usize>,
+ pub max_input_tokens: Option<usize>,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
pub tokens_count_factors: TokensCountFactors,
pub capabilities: ModelCapabilities,
@@ -29,7 +29,7 @@ impl Model {
client_name: client_name.into(),
name: name.into(),
extra_fields: None,
- max_tokens: None,
+ max_input_tokens: None,
tokens_count_factors: Default::default(),
capabilities: ModelCapabilities::Text,
}
@@ -83,10 +83,10 @@ impl Model {
self
}
- pub fn set_max_tokens(mut self, max_tokens: Option<usize>) -> Self {
- match max_tokens {
- None | Some(0) => self.max_tokens = None,
- _ => self.max_tokens = max_tokens,
+ pub fn set_max_input_tokens(mut self, max_input_tokens: Option<usize>) -> Self {
+ match max_input_tokens {
+ None | Some(0) => self.max_input_tokens = None,
+ _ => self.max_input_tokens = max_input_tokens,
}
self
}
@@ -122,12 +122,12 @@ impl Model {
}
}
- pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> {
+ pub fn max_input_tokens_limit(&self, messages: &[Message]) -> Result<()> {
let (_, bias) = self.tokens_count_factors;
let total_tokens = self.total_tokens(messages) + bias;
- if let Some(max_tokens) = self.max_tokens {
- if total_tokens >= max_tokens {
- bail!("Exceed max tokens limit")
+ if let Some(max_input_tokens) = self.max_input_tokens {
+ if total_tokens >= max_input_tokens {
+ bail!("Exceed max input tokens limit")
}
}
Ok(())
@@ -147,7 +147,7 @@ impl Model {
#[derive(Debug, Clone, Deserialize)]
pub struct ModelConfig {
pub name: String,
- pub max_tokens: Option<usize>,
+ pub max_input_tokens: Option<usize>,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
#[serde(deserialize_with = "deserialize_capabilities")]
#[serde(default = "default_capabilities")]
diff --git a/src/client/ollama.rs b/src/client/ollama.rs
index e24cffd..5ecf6c8 100644
--- a/src/client/ollama.rs
+++ b/src/client/ollama.rs
@@ -52,8 +52,8 @@ impl OllamaClient {
("api_key", "API Key:", false, PromptKind::String),
("models[].name", "Model Name:", true, PromptKind::String),
(
- "models[].max_tokens",
- "Max Tokens:",
+ "models[].max_input_tokens",
+ "Max Input Tokens:",
false,
PromptKind::Integer,
),
@@ -68,7 +68,7 @@ impl OllamaClient {
.map(|v| {
Model::new(client_name, &v.name)
.set_capabilities(v.capabilities)
- .set_max_tokens(v.max_tokens)
+ .set_max_input_tokens(v.max_input_tokens)
.set_extra_fields(v.extra_fields.clone())
.set_tokens_count_factors(TOKENS_COUNT_FACTORS)
})
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 8644ae6..2c8aabf 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -12,14 +12,14 @@ use serde_json::{json, Value};
const API_BASE: &str = "https://api.openai.com/v1";
-const MODELS: [(&str, usize, &str); 7] = [
+const MODELS: [(&str, usize, &str); 5] = [
+ // https://platform.openai.com/docs/models/gpt-3-5-turbo
("gpt-3.5-turbo", 16385, "text"),
("gpt-3.5-turbo-1106", 16385, "text"),
+ // https://platform.openai.com/docs/models/gpt-4-and-gpt-4-turbo
("gpt-4-turbo-preview", 128000, "text"),
("gpt-4-vision-preview", 128000, "text,vision"),
("gpt-4-1106-preview", 128000, "text"),
- ("gpt-4", 8192, "text"),
- ("gpt-4-32k", 32768, "text"),
];
pub const OPENAI_TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2);
@@ -46,10 +46,10 @@ impl OpenAIClient {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens, capabilities)| {
+ .map(|(name, max_input_tokens, capabilities)| {
Model::new(client_name, name)
.set_capabilities(capabilities.into())
- .set_max_tokens(Some(max_tokens))
+ .set_max_input_tokens(Some(max_input_tokens))
.set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS)
})
.collect()
diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs
index 70c251a..1858030 100644
--- a/src/client/qianwen.rs
+++ b/src/client/qianwen.rs
@@ -25,10 +25,12 @@ const API_URL_VL: &str =
"https://dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation";
const MODELS: [(&str, usize, &str); 6] = [
- ("qwen-turbo", 8192, "text"),
- ("qwen-plus", 32768, "text"),
- ("qwen-max", 8192, "text"),
- ("qwen-max-longcontext", 30720, "text"),
+ // https://help.aliyun.com/zh/dashscope/developer-reference/api-details
+ ("qwen-turbo", 6000, "text"),
+ ("qwen-plus", 30000, "text"),
+ ("qwen-max", 6000, "text"),
+ ("qwen-max-longcontext", 28000, "text"),
+ // https://help.aliyun.com/zh/dashscope/developer-reference/tongyi-qianwen-vl-plus-api
("qwen-vl-plus", 0, "text,vision"),
("qwen-vl-max", 0, "text,vision"),
];
@@ -78,10 +80,10 @@ impl QianwenClient {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens, capabilities)| {
+ .map(|(name, max_input_tokens, capabilities)| {
Model::new(client_name, name)
.set_capabilities(capabilities.into())
- .set_max_tokens(Some(max_tokens))
+ .set_max_input_tokens(Some(max_input_tokens))
})
.collect()
}
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 94becdb..0f92b4b 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -14,9 +14,10 @@ use serde::Deserialize;
use serde_json::{json, Value};
use std::path::PathBuf;
+// https://cloud.google.com/vertex-ai/generative-ai/docs/learn/models
const MODELS: [(&str, usize, &str); 5] = [
- ("gemini-1.0-pro", 32760, "text"),
- ("gemini.1.0-pro-vision", 16384, "text,vision"),
+ ("gemini-1.0-pro", 24568, "text"),
+ ("gemini.1.0-pro-vision", 14336, "text,vision"),
("gemini-1.0-ultra", 8192, "text"),
("gemini.1.0-ultra-vision", 8192, "text,vision"),
("gemini-1.5-pro", 1000000, "text"),
@@ -66,10 +67,10 @@ impl VertexAIClient {
let client_name = Self::name(local_config);
MODELS
.into_iter()
- .map(|(name, max_tokens, capabilities)| {
+ .map(|(name, max_input_tokens, capabilities)| {
Model::new(client_name, name)
.set_capabilities(capabilities.into())
- .set_max_tokens(Some(max_tokens))
+ .set_max_input_tokens(Some(max_input_tokens))
.set_tokens_count_factors(TOKENS_COUNT_FACTORS)
})
.collect()