Add vectorization via Vertex AI (#4311) * Add vectorization via Vertex AI * Enable batching for Google vectors * Add batch embeddings for Vertex * Split embed methods for Vertex/AI Studio
Signed| @@ -36,6 +36,7 @@ import { slashCommandReturnHelper } from '../../slash-commands/SlashCommandRetur | |||
| 36 | import { generateWebLlmChatPrompt, isWebLlmSupported } from '../shared.js'; | 36 | import { generateWebLlmChatPrompt, isWebLlmSupported } from '../shared.js'; |
| 37 | import { WebLlmVectorProvider } from './webllm.js'; | 37 | import { WebLlmVectorProvider } from './webllm.js'; |
| 38 | import { removeReasoningFromString } from '../../reasoning.js'; | 38 | import { removeReasoningFromString } from '../../reasoning.js'; |
| 39 | import { oai_settings } from '../../openai.js'; | ||
| 39 | 40 | ||
| 40 | /** | 41 | /** |
| 41 | * @typedef {object} HashedMessage | 42 | * @typedef {object} HashedMessage |
| @@ -50,7 +51,7 @@ export const EXTENSION_PROMPT_TAG = '3_vectors'; | |||
| 50 | export const EXTENSION_PROMPT_TAG_DB = '4_vectors_data_bank'; | 51 | export const EXTENSION_PROMPT_TAG_DB = '4_vectors_data_bank'; |
| 51 | 52 | ||
| 52 | // Force solo chunks for sources that don't support batching. | 53 | // Force solo chunks for sources that don't support batching. |
| 53 | const getBatchSize = () => ['transformers', 'palm', 'ollama'].includes(settings.source) ? 1 : 5; | 54 | const getBatchSize = () => ['transformers', 'ollama'].includes(settings.source) ? 1 : 5; |
| 54 | 55 | ||
| 55 | const settings = { | 56 | const settings = { |
| 56 | // For both | 57 | // For both |
| @@ -796,6 +797,14 @@ function getVectorsRequestBody(args = {}) { | |||
| 796 | break; | 797 | break; |
| 797 | case 'palm': | 798 | case 'palm': |
| 798 | body.model = extension_settings.vectors.google_model; | 799 | body.model = extension_settings.vectors.google_model; |
| 800 | body.api = 'makersuite'; | ||
| 801 | break; | ||
| 802 | case 'vertexai': | ||
| 803 | body.model = extension_settings.vectors.google_model; | ||
| 804 | body.api = 'vertexai'; | ||
| 805 | body.vertexai_auth_mode = oai_settings.vertexai_auth_mode; | ||
| 806 | body.vertexai_region = oai_settings.vertexai_region; | ||
| 807 | body.vertexai_express_project_id = oai_settings.vertexai_express_project_id; | ||
| 799 | break; | 808 | break; |
| 800 | default: | 809 | default: |
| 801 | break; | 810 | break; |
| @@ -881,6 +890,7 @@ async function insertVectorItems(collectionId, items) { | |||
| 881 | function throwIfSourceInvalid() { | 890 | function throwIfSourceInvalid() { |
| 882 | if (settings.source === 'openai' && !secret_state[SECRET_KEYS.OPENAI] || | 891 | if (settings.source === 'openai' && !secret_state[SECRET_KEYS.OPENAI] || |
| 883 | settings.source === 'palm' && !secret_state[SECRET_KEYS.MAKERSUITE] || | 892 | settings.source === 'palm' && !secret_state[SECRET_KEYS.MAKERSUITE] || |
| 893 | settings.source === 'vertexai' && !secret_state[SECRET_KEYS.VERTEXAI] && !secret_state[SECRET_KEYS.VERTEXAI_SERVICE_ACCOUNT] || | ||
| 884 | settings.source === 'mistral' && !secret_state[SECRET_KEYS.MISTRALAI] || | 894 | settings.source === 'mistral' && !secret_state[SECRET_KEYS.MISTRALAI] || |
| 885 | settings.source === 'togetherai' && !secret_state[SECRET_KEYS.TOGETHERAI] || | 895 | settings.source === 'togetherai' && !secret_state[SECRET_KEYS.TOGETHERAI] || |
| 886 | settings.source === 'nomicai' && !secret_state[SECRET_KEYS.NOMICAI] || | 896 | settings.source === 'nomicai' && !secret_state[SECRET_KEYS.NOMICAI] || |
| @@ -1098,7 +1108,7 @@ function toggleSettings() { | |||
| 1098 | $('#nomicai_apiKey').toggle(settings.source === 'nomicai'); | 1108 | $('#nomicai_apiKey').toggle(settings.source === 'nomicai'); |
| 1099 | $('#webllm_vectorsModel').toggle(settings.source === 'webllm'); | 1109 | $('#webllm_vectorsModel').toggle(settings.source === 'webllm'); |
| 1100 | $('#koboldcpp_vectorsModel').toggle(settings.source === 'koboldcpp'); | 1110 | $('#koboldcpp_vectorsModel').toggle(settings.source === 'koboldcpp'); |
| 1101 | $('#google_vectorsModel').toggle(settings.source === 'palm'); | 1111 | $('#google_vectorsModel').toggle(settings.source === 'palm' || settings.source === 'vertexai'); |
| 1102 | $('#vector_altEndpointUrl').toggle(vectorApiRequiresUrl.includes(settings.source)); | 1112 | $('#vector_altEndpointUrl').toggle(vectorApiRequiresUrl.includes(settings.source)); |
| 1103 | if (settings.source === 'webllm') { | 1113 | if (settings.source === 'webllm') { |
| 1104 | loadWebLlmModels(); | 1114 | loadWebLlmModels(); |
| @@ -13,6 +13,7 @@ | |||
| 13 | <option value="cohere">Cohere</option> | 13 | <option value="cohere">Cohere</option> |
| 14 | <option value="extras">Extras (deprecated)</option> | 14 | <option value="extras">Extras (deprecated)</option> |
| 15 | <option value="palm">Google AI Studio</option> | 15 | <option value="palm">Google AI Studio</option> |
| 16 | <option value="vertexai">Google Vertex AI</option> | ||
| 16 | <option value="koboldcpp">KoboldCpp</option> | 17 | <option value="koboldcpp">KoboldCpp</option> |
| 17 | <option value="llamacpp">llama.cpp</option> | 18 | <option value="llamacpp">llama.cpp</option> |
| 18 | <option value="transformers" data-i18n="Local (Transformers)">Local (Transformers)</option> | 19 | <option value="transformers" data-i18n="Local (Transformers)">Local (Transformers)</option> |
| @@ -136,6 +137,7 @@ | |||
| 136 | Vectorization Model | 137 | Vectorization Model |
| 137 | </label> | 138 | </label> |
| 138 | <select id="vectors_google_model" class="text_pole"> | 139 | <select id="vectors_google_model" class="text_pole"> |
| 140 | <option value="gemini-embedding-001">gemini-embedding-001</option> | ||
| 139 | <option value="gemini-embedding-exp-03-07">gemini-embedding-exp-03-07</option> | 141 | <option value="gemini-embedding-exp-03-07">gemini-embedding-exp-03-07</option> |
| 140 | <option value="text-embedding-004">text-embedding-004</option> | 142 | <option value="text-embedding-004">text-embedding-004</option> |
| 141 | <option value="embedding-001">embedding-001</option> | 143 | <option value="embedding-001">embedding-001</option> |
| @@ -11,7 +11,8 @@ import { getNomicAIBatchVector, getNomicAIVector } from '../vectors/nomicai-vect | |||
| 11 | import { getOpenAIVector, getOpenAIBatchVector } from '../vectors/openai-vectors.js'; | 11 | import { getOpenAIVector, getOpenAIBatchVector } from '../vectors/openai-vectors.js'; |
| 12 | import { getTransformersVector, getTransformersBatchVector } from '../vectors/embedding.js'; | 12 | import { getTransformersVector, getTransformersBatchVector } from '../vectors/embedding.js'; |
| 13 | import { getExtrasVector, getExtrasBatchVector } from '../vectors/extras-vectors.js'; | 13 | import { getExtrasVector, getExtrasBatchVector } from '../vectors/extras-vectors.js'; |
| 14 | import { getMakerSuiteVector, getMakerSuiteBatchVector } from '../vectors/makersuite-vectors.js'; | 14 | import { getMakerSuiteVector, getMakerSuiteBatchVector } from '../vectors/google-vectors.js'; |
| 15 | import { getVertexVector, getVertexBatchVector } from '../vectors/google-vectors.js'; | ||
| 15 | import { getCohereVector, getCohereBatchVector } from '../vectors/cohere-vectors.js'; | 16 | import { getCohereVector, getCohereBatchVector } from '../vectors/cohere-vectors.js'; |
| 16 | import { getLlamaCppVector, getLlamaCppBatchVector } from '../vectors/llamacpp-vectors.js'; | 17 | import { getLlamaCppVector, getLlamaCppBatchVector } from '../vectors/llamacpp-vectors.js'; |
| 17 | import { getVllmVector, getVllmBatchVector } from '../vectors/vllm-vectors.js'; | 18 | import { getVllmVector, getVllmBatchVector } from '../vectors/vllm-vectors.js'; |
| @@ -32,6 +33,7 @@ const SOURCES = [ | |||
| 32 | 'vllm', | 33 | 'vllm', |
| 33 | 'webllm', | 34 | 'webllm', |
| 34 | 'koboldcpp', | 35 | 'koboldcpp', |
| 36 | 'vertexai', | ||
| 35 | ]; | 37 | ]; |
| 36 | 38 | ||
| 37 | /** | 39 | /** |
| @@ -56,7 +58,9 @@ async function getVector(source, sourceSettings, text, isQuery, directories) { | |||
| 56 | case 'extras': | 58 | case 'extras': |
| 57 | return getExtrasVector(text, sourceSettings.extrasUrl, sourceSettings.extrasKey); | 59 | return getExtrasVector(text, sourceSettings.extrasUrl, sourceSettings.extrasKey); |
| 58 | case 'palm': | 60 | case 'palm': |
| 59 | return getMakerSuiteVector(text, directories, sourceSettings.model); | 61 | return getMakerSuiteVector(text, sourceSettings.model, sourceSettings.request); |
| 62 | case 'vertexai': | ||
| 63 | return getVertexVector(text, sourceSettings.model, sourceSettings.request); | ||
| 60 | case 'cohere': | 64 | case 'cohere': |
| 61 | return getCohereVector(text, isQuery, directories, sourceSettings.model); | 65 | return getCohereVector(text, isQuery, directories, sourceSettings.model); |
| 62 | case 'llamacpp': | 66 | case 'llamacpp': |
| @@ -105,7 +109,10 @@ async function getBatchVector(source, sourceSettings, texts, isQuery, directorie | |||
| 105 | results.push(...await getExtrasBatchVector(batch, sourceSettings.extrasUrl, sourceSettings.extrasKey)); | 109 | results.push(...await getExtrasBatchVector(batch, sourceSettings.extrasUrl, sourceSettings.extrasKey)); |
| 106 | break; | 110 | break; |
| 107 | case 'palm': | 111 | case 'palm': |
| 108 | results.push(...await getMakerSuiteBatchVector(batch, directories, sourceSettings.model)); | 112 | results.push(...await getMakerSuiteBatchVector(batch, sourceSettings.model, sourceSettings.request)); |
| 113 | break; | ||
| 114 | case 'vertexai': | ||
| 115 | results.push(...await getVertexBatchVector(batch, sourceSettings.model, sourceSettings.request)); | ||
| 109 | break; | 116 | break; |
| 110 | case 'cohere': | 117 | case 'cohere': |
| 111 | results.push(...await getCohereBatchVector(batch, isQuery, directories, sourceSettings.model)); | 118 | results.push(...await getCohereBatchVector(batch, isQuery, directories, sourceSettings.model)); |
| @@ -178,8 +185,10 @@ function getSourceSettings(source, request) { | |||
| 178 | model: getConfigValue('extensions.models.embedding', ''), | 185 | model: getConfigValue('extensions.models.embedding', ''), |
| 179 | }; | 186 | }; |
| 180 | case 'palm': | 187 | case 'palm': |
| 188 | case 'vertexai': | ||
| 181 | return { | 189 | return { |
| 182 | model: String(request.body.model || 'text-embedding-004'), | 190 | model: String(request.body.model || 'text-embedding-004'), |
| 191 | request: request, // Pass the request object to get API key and URL | ||
| 183 | }; | 192 | }; |
| 184 | case 'mistral': | 193 | case 'mistral': |
| 185 | return { | 194 | return { |
| @@ -0,0 +1,101 @@ | |||
| 1 | import fetch from 'node-fetch'; | ||
| 2 | import { getGoogleApiConfig } from '../endpoints/google.js'; | ||
| 3 | |||
| 4 | /** | ||
| 5 | * Gets the vector for the given text from Google AI Studio | ||
| 6 | * @param {string[]} texts - The array of texts to get the vector for | ||
| 7 | * @param {string} model - The model to use for embedding | ||
| 8 | * @param {import('express').Request} request - The request object to get API key and URL | ||
| 9 | * @returns {Promise<number[][]>} - The array of vectors for the texts | ||
| 10 | */ | ||
| 11 | export async function getMakerSuiteBatchVector(texts, model, request) { | ||
| 12 | const { url, headers, apiName } = await getGoogleApiConfig(request, model, 'batchEmbedContents'); | ||
| 13 | |||
| 14 | const body = { | ||
| 15 | requests: texts.map(text => ({ | ||
| 16 | model: `models/${model}`, | ||
| 17 | content: { parts: [{ text }] }, | ||
| 18 | })), | ||
| 19 | }; | ||
| 20 | |||
| 21 | const response = await fetch(url, { | ||
| 22 | body: JSON.stringify(body), | ||
| 23 | method: 'POST', | ||
| 24 | headers: headers, | ||
| 25 | }); | ||
| 26 | |||
| 27 | if (!response.ok) { | ||
| 28 | const text = await response.text(); | ||
| 29 | console.warn(`${apiName} batch request failed`, response.statusText, text); | ||
| 30 | throw new Error(`${apiName} batch request failed`); | ||
| 31 | } | ||
| 32 | |||
| 33 | /** @type {any} */ | ||
| 34 | const data = await response.json(); | ||
| 35 | if (!Array.isArray(data?.embeddings)) { | ||
| 36 | throw new Error(`${apiName} did not return an array`); | ||
| 37 | } | ||
| 38 | |||
| 39 | const embeddings = data.embeddings.map(embedding => embedding.values); | ||
| 40 | return embeddings; | ||
| 41 | } | ||
| 42 | |||
| 43 | /** | ||
| 44 | * Gets the vector for the given text from Google Vertex AI | ||
| 45 | * @param {string[]} texts - The array of texts to get the vector for | ||
| 46 | * @param {string} model - The model to use for embedding | ||
| 47 | * @param {import('express').Request} request - The request object to get API key and URL | ||
| 48 | * @returns {Promise<number[][]>} - The array of vectors for the texts | ||
| 49 | */ | ||
| 50 | export async function getVertexBatchVector(texts, model, request) { | ||
| 51 | const { url, headers, apiName } = await getGoogleApiConfig(request, model, 'predict'); | ||
| 52 | |||
| 53 | const body = { | ||
| 54 | instances: texts.map(text => ({ content: text })), | ||
| 55 | }; | ||
| 56 | |||
| 57 | const response = await fetch(url, { | ||
| 58 | body: JSON.stringify(body), | ||
| 59 | method: 'POST', | ||
| 60 | headers: headers, | ||
| 61 | }); | ||
| 62 | |||
| 63 | if (!response.ok) { | ||
| 64 | const text = await response.text(); | ||
| 65 | console.warn(`${apiName} batch request failed`, response.statusText, text); | ||
| 66 | throw new Error(`${apiName} batch request failed`); | ||
| 67 | } | ||
| 68 | |||
| 69 | /** @type {any} */ | ||
| 70 | const data = await response.json(); | ||
| 71 | if (!Array.isArray(data?.predictions)) { | ||
| 72 | throw new Error(`${apiName} did not return an array`); | ||
| 73 | } | ||
| 74 | |||
| 75 | const embeddings = data.predictions.map(p => p.embeddings.values); | ||
| 76 | return embeddings; | ||
| 77 | } | ||
| 78 | |||
| 79 | /** | ||
| 80 | * Gets the vector for the given text from Google AI Studio | ||
| 81 | * @param {string} text - The text to get the vector for | ||
| 82 | * @param {string} model - The model to use for embedding | ||
| 83 | * @param {import('express').Request} request - The request object to get API key and URL | ||
| 84 | * @returns {Promise<number[]>} - The vector for the text | ||
| 85 | */ | ||
| 86 | export async function getMakerSuiteVector(text, model, request) { | ||
| 87 | const [embedding] = await getMakerSuiteBatchVector([text], model, request); | ||
| 88 | return embedding; | ||
| 89 | } | ||
| 90 | |||
| 91 | /** | ||
| 92 | * Gets the vector for the given text from Google Vertex AI | ||
| 93 | * @param {string} text - The text to get the vector for | ||
| 94 | * @param {string} model - The model to use for embedding | ||
| 95 | * @param {import('express').Request} request - The request object to get API key and URL | ||
| 96 | * @returns {Promise<number[]>} - The vector for the text | ||
| 97 | */ | ||
| 98 | export async function getVertexVector(text, model, request) { | ||
| 99 | const [embedding] = await getVertexBatchVector([text], model, request); | ||
| 100 | return embedding; | ||
| 101 | } | ||
| @@ -1,61 +0,0 @@ | |||
| 1 | import fetch from 'node-fetch'; | ||
| 2 | import { SECRET_KEYS, readSecret } from '../endpoints/secrets.js'; | ||
| 3 | import { trimTrailingSlash } from '../util.js'; | ||
| 4 | const API_MAKERSUITE = 'https://generativelanguage.googleapis.com'; | ||
| 5 | |||
| 6 | /** | ||
| 7 | * Gets the vector for the given text from gecko model | ||
| 8 | * @param {string[]} texts - The array of texts to get the vector for | ||
| 9 | * @param {import('../users.js').UserDirectoryList} directories - The directories object for the user | ||
| 10 | * @param {string} model - The model to use for embedding | ||
| 11 | * @returns {Promise<number[][]>} - The array of vectors for the texts | ||
| 12 | */ | ||
| 13 | export async function getMakerSuiteBatchVector(texts, directories, model) { | ||
| 14 | const promises = texts.map(text => getMakerSuiteVector(text, directories, model)); | ||
| 15 | return await Promise.all(promises); | ||
| 16 | } | ||
| 17 | |||
| 18 | /** | ||
| 19 | * Gets the vector for the given text from Gemini API text-embedding-004 model | ||
| 20 | * @param {string} text - The text to get the vector for | ||
| 21 | * @param {import('../users.js').UserDirectoryList} directories - The directories object for the user | ||
| 22 | * @param {string} model - The model to use for embedding (default is 'text-embedding-004') | ||
| 23 | * @returns {Promise<number[]>} - The vector for the text | ||
| 24 | */ | ||
| 25 | export async function getMakerSuiteVector(text, directories, model) { | ||
| 26 | const key = readSecret(directories, SECRET_KEYS.MAKERSUITE); | ||
| 27 | |||
| 28 | if (!key) { | ||
| 29 | console.warn('No Google AI Studio key found'); | ||
| 30 | throw new Error('No Google AI Studio key found'); | ||
| 31 | } | ||
| 32 | |||
| 33 | const apiUrl = trimTrailingSlash(API_MAKERSUITE); | ||
| 34 | const url = `${apiUrl}/v1beta/models/${model}:embedContent?key=${key}`; | ||
| 35 | const body = { | ||
| 36 | content: { | ||
| 37 | parts: [ | ||
| 38 | { text: text }, | ||
| 39 | ], | ||
| 40 | }, | ||
| 41 | }; | ||
| 42 | |||
| 43 | const response = await fetch(url, { | ||
| 44 | body: JSON.stringify(body), | ||
| 45 | method: 'POST', | ||
| 46 | headers: { | ||
| 47 | 'Content-Type': 'application/json', | ||
| 48 | }, | ||
| 49 | }); | ||
| 50 | |||
| 51 | if (!response.ok) { | ||
| 52 | const text = await response.text(); | ||
| 53 | console.warn('Google AI Studio request failed', response.statusText, text); | ||
| 54 | throw new Error('Google AI Studio request failed'); | ||
| 55 | } | ||
| 56 | |||
| 57 | /** @type {any} */ | ||
| 58 | const data = await response.json(); | ||
| 59 | // noinspection JSValidateTypes | ||
| 60 | return data['embedding']['values']; | ||
| 61 | } | ||