summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-03 12:42:48 +0800
committerGitHub <noreply@github.com>2024-09-03 12:42:48 +0800
commitb2fef25a5269f23ef09018a2c3d55156532fc481 (patch)
treee622d45563b13c1d4ce1f271763dbf780b07aa80
parent476d29c40aaf2b7a49e2ba5068a12f789defcb75 (diff)
downloadaichat-b2fef25a5269f23ef09018a2c3d55156532fc481.tar.gz
refactor: gemini use batch embedding api (#827)
-rw-r--r--models.yaml1
-rw-r--r--src/client/gemini.rs37
2 files changed, 28 insertions, 10 deletions
diff --git a/models.yaml b/models.yaml
index 6d703c9..5a79454 100644
--- a/models.yaml
+++ b/models.yaml
@@ -112,6 +112,7 @@
output_price: 0
max_tokens_per_chunk: 2048
default_chunk_size: 1500
+ max_batch_size: 100
# Links:
# - https://docs.anthropic.com/claude/docs/models-overview
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 9578e23..1bb31cc 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -74,20 +74,33 @@ fn prepare_embeddings(self_: &GeminiClient, data: EmbeddingsData) -> Result<Requ
.unwrap_or_else(|_| API_BASE.to_string());
let url = format!(
- "{}/models/{}:embedContent?key={}",
+ "{}/models/{}:batchEmbedContents?key={}",
api_base.trim_end_matches('/'),
self_.model.name(),
api_key
);
+ let model_id = format!("models/{}", self_.model.name());
+
+ let requests: Vec<_> = data
+ .texts
+ .iter()
+ .map(|text| {
+ json!({
+ "model": model_id,
+ "content": {
+ "parts": [
+ {
+ "text": text
+ }
+ ]
+ },
+ })
+ })
+ .collect();
+
let body = json!({
- "content": {
- "parts": [
- {
- "text": data.texts[0],
- }
- ]
- }
+ "requests": requests,
});
let request_data = RequestData::new(url, body);
@@ -104,13 +117,17 @@ async fn embeddings(builder: RequestBuilder, _model: &Model) -> Result<Embedding
}
let res_body: EmbeddingsResBody =
serde_json::from_value(data).context("Invalid embeddings data")?;
- let output = vec![res_body.embedding.values];
+ let output = res_body
+ .embeddings
+ .into_iter()
+ .map(|embedding| embedding.values)
+ .collect();
Ok(output)
}
#[derive(Deserialize)]
struct EmbeddingsResBody {
- embedding: EmbeddingsResBodyEmbedding,
+ embeddings: Vec<EmbeddingsResBodyEmbedding>,
}
#[derive(Deserialize)]