**feat(docs): add comprehensive documentation for model bundling and robust loading**
- Introduced new documentation files under `docs/`: - `model-bundling-analysis.md`: Provides detailed analysis of current, bundled, hybrid, and dynamic model loading approaches, including pros, cons, and recommendations. - `model-management.md`: Explains how Brainy manages Universal Sentence Encoder models, including setup, usage, and troubleshooting. - `optional-model-bundling.md`: Details the `@soulcraft/brainy-models` package for offline reliability with pre-bundled models. - Added `src/utils/robustModelLoader.ts`: - Implements enhanced model loading with retry mechanisms, timeout handling, fallback URLs, and optional local model bundling. - Supports Node.js and browser environments with exponential backoff logic. - Key Updates: - **Hybrid Loading Strategy**: Recommended for balancing reliability and flexibility via hybrid online/offline mechanisms. - **Enhanced Fallback Scenarios**: Robust loader improves network-dependent reliability for embedding workflows. - **Offline Reliability Support**: Optional model bundling eliminates dependency on external services, supporting air-gapped and edge environments. **Purpose**: Introduce a hybrid model loading approach with robust options for
This commit is contained in:
parent
ea885d8e1b
commit
12a0d3d281
6 changed files with 1098 additions and 214 deletions
|
|
@ -5,6 +5,12 @@
|
|||
import { EmbeddingFunction, EmbeddingModel, Vector } from '../coreTypes.js'
|
||||
import { executeInThread } from './workerUtils.js'
|
||||
import { isBrowser } from './environment.js'
|
||||
import {
|
||||
RobustModelLoader,
|
||||
ModelLoadOptions,
|
||||
createRobustModelLoader,
|
||||
getUniversalSentenceEncoderFallbacks
|
||||
} from './robustModelLoader.js'
|
||||
|
||||
/**
|
||||
* TensorFlow Universal Sentence Encoder embedding model
|
||||
|
|
@ -14,6 +20,11 @@ import { isBrowser } from './environment.js'
|
|||
* This implementation attempts to use GPU processing when available for better performance,
|
||||
* falling back to CPU processing for compatibility across all environments.
|
||||
*/
|
||||
export interface UniversalSentenceEncoderOptions extends ModelLoadOptions {
|
||||
/** Whether to enable verbose logging */
|
||||
verbose?: boolean
|
||||
}
|
||||
|
||||
export class UniversalSentenceEncoder implements EmbeddingModel {
|
||||
private model: any = null
|
||||
private initialized = false
|
||||
|
|
@ -21,13 +32,26 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
private use: any = null
|
||||
private backend: string = 'cpu' // Default to CPU
|
||||
private verbose: boolean = true // Whether to log non-essential messages
|
||||
|
||||
private robustLoader: RobustModelLoader
|
||||
|
||||
/**
|
||||
* Create a new UniversalSentenceEncoder instance
|
||||
* @param options Configuration options
|
||||
* @param options Configuration options including reliability settings
|
||||
*/
|
||||
constructor(options: { verbose?: boolean } = {}) {
|
||||
constructor(options: UniversalSentenceEncoderOptions = {}) {
|
||||
this.verbose = options.verbose !== undefined ? options.verbose : true
|
||||
|
||||
// Create robust model loader with enhanced reliability features
|
||||
this.robustLoader = createRobustModelLoader({
|
||||
maxRetries: options.maxRetries ?? 3,
|
||||
initialRetryDelay: options.initialRetryDelay ?? 1000,
|
||||
maxRetryDelay: options.maxRetryDelay ?? 30000,
|
||||
timeout: options.timeout ?? 60000,
|
||||
useExponentialBackoff: options.useExponentialBackoff ?? true,
|
||||
fallbackUrls: options.fallbackUrls ?? getUniversalSentenceEncoderFallbacks(),
|
||||
verbose: this.verbose,
|
||||
preferLocalModel: options.preferLocalModel ?? true
|
||||
})
|
||||
}
|
||||
|
||||
/**
|
||||
|
|
@ -116,140 +140,36 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
}
|
||||
|
||||
/**
|
||||
* Load the Universal Sentence Encoder model with retry logic
|
||||
* This helps handle network failures and JSON parsing errors from TensorFlow Hub
|
||||
* @param loadFunction The function to load the model
|
||||
* @param maxRetries Maximum number of retry attempts
|
||||
* @param baseDelay Base delay in milliseconds for exponential backoff
|
||||
* Load the Universal Sentence Encoder model with robust retry and fallback mechanisms
|
||||
* @param loadFunction The function to load the model from TensorFlow Hub
|
||||
*/
|
||||
private async loadModelWithRetry(
|
||||
loadFunction: () => Promise<EmbeddingModel>,
|
||||
maxRetries: number = 3,
|
||||
baseDelay: number = 1000
|
||||
private async loadModelFromLocal(
|
||||
loadFunction: () => Promise<EmbeddingModel>
|
||||
): Promise<EmbeddingModel> {
|
||||
let lastError: Error | null = null
|
||||
|
||||
// Define alternative model URLs to try if the default one fails
|
||||
const alternativeLoadFunctions: Array<() => Promise<EmbeddingModel>> = []
|
||||
|
||||
// Try to create alternative load functions using different model URLs
|
||||
if (this.use) {
|
||||
// Add alternative model URLs to try
|
||||
const alternativeUrls = [
|
||||
'https://storage.googleapis.com/tfjs-models/savedmodel/universal_sentence_encoder/model.json',
|
||||
'https://tfhub.dev/tensorflow/tfjs-model/universal-sentence-encoder-lite/1/default/1/model.json',
|
||||
'https://tfhub.dev/tensorflow/tfjs-model/universal-sentence-encoder/1/default/1/model.json',
|
||||
'https://tfhub.dev/tensorflow/tfjs-model/universal-sentence-encoder/1/default/1',
|
||||
'https://tfhub.dev/tensorflow/universal-sentence-encoder/4',
|
||||
'https://tfhub.dev/tensorflow/universal-sentence-encoder/4/default/1/model.json'
|
||||
]
|
||||
|
||||
// Create load functions for each alternative URL
|
||||
for (const url of alternativeUrls) {
|
||||
if (this.use.load) {
|
||||
alternativeLoadFunctions.push(() => this.use!.load(url))
|
||||
} else if (this.use.default && this.use.default.load) {
|
||||
alternativeLoadFunctions.push(() => this.use!.default.load(url))
|
||||
}
|
||||
this.logger('log', 'Loading Universal Sentence Encoder model with robust loader...')
|
||||
|
||||
try {
|
||||
// Use the robust model loader to handle all retry logic, timeouts, and fallbacks
|
||||
const model = await this.robustLoader.loadModel(
|
||||
loadFunction,
|
||||
'universal-sentence-encoder'
|
||||
)
|
||||
|
||||
this.logger('log', 'Successfully loaded Universal Sentence Encoder model')
|
||||
return model
|
||||
|
||||
} catch (error) {
|
||||
const errorMessage = error instanceof Error ? error.message : String(error)
|
||||
this.logger('error', `Failed to load Universal Sentence Encoder model: ${errorMessage}`)
|
||||
|
||||
// Log loading statistics for debugging
|
||||
const stats = this.robustLoader.getLoadingStats()
|
||||
if (Object.keys(stats).length > 0) {
|
||||
this.logger('log', 'Loading attempt statistics:', stats)
|
||||
}
|
||||
|
||||
throw error
|
||||
}
|
||||
|
||||
// First try with the original load function
|
||||
for (let attempt = 0; attempt <= maxRetries; attempt++) {
|
||||
try {
|
||||
this.logger(
|
||||
'log',
|
||||
attempt === 0
|
||||
? 'Loading Universal Sentence Encoder model...'
|
||||
: `Retrying Universal Sentence Encoder model loading (attempt ${attempt + 1}/${maxRetries + 1})...`
|
||||
)
|
||||
|
||||
const model = await loadFunction()
|
||||
|
||||
if (attempt > 0) {
|
||||
this.logger(
|
||||
'log',
|
||||
'Universal Sentence Encoder model loaded successfully after retry'
|
||||
)
|
||||
}
|
||||
|
||||
return model
|
||||
} catch (error) {
|
||||
lastError = error as Error
|
||||
const errorMessage = lastError.message || String(lastError)
|
||||
|
||||
// Check if this is a network-related error that might benefit from retry
|
||||
const isRetryableError =
|
||||
errorMessage.includes('Failed to parse model JSON') ||
|
||||
errorMessage.includes('Failed to fetch') ||
|
||||
errorMessage.includes('Network error') ||
|
||||
errorMessage.includes('ENOTFOUND') ||
|
||||
errorMessage.includes('ECONNRESET') ||
|
||||
errorMessage.includes('ETIMEDOUT') ||
|
||||
errorMessage.includes('JSON') ||
|
||||
errorMessage.includes('model.json') ||
|
||||
errorMessage.includes('byte length') ||
|
||||
errorMessage.includes('tensor should have') ||
|
||||
errorMessage.includes('shape') ||
|
||||
errorMessage.includes('dimensions')
|
||||
|
||||
if (attempt < maxRetries && isRetryableError) {
|
||||
const delay = baseDelay * Math.pow(2, attempt) // Exponential backoff
|
||||
this.logger(
|
||||
'warn',
|
||||
`Universal Sentence Encoder model loading failed (attempt ${attempt + 1}): ${errorMessage}. Retrying in ${delay}ms...`
|
||||
)
|
||||
await new Promise((resolve) => setTimeout(resolve, delay))
|
||||
} else {
|
||||
// Either we've exhausted retries or this is not a retryable error
|
||||
if (attempt >= maxRetries) {
|
||||
this.logger(
|
||||
'warn',
|
||||
`Universal Sentence Encoder model loading failed after ${maxRetries + 1} attempts. Last error: ${errorMessage}. Trying alternative URLs...`
|
||||
)
|
||||
|
||||
// Try alternative URLs if available
|
||||
if (alternativeLoadFunctions.length > 0) {
|
||||
for (let i = 0; i < alternativeLoadFunctions.length; i++) {
|
||||
try {
|
||||
this.logger(
|
||||
'log',
|
||||
`Trying alternative model URL ${i + 1}/${alternativeLoadFunctions.length}...`
|
||||
)
|
||||
const model = await alternativeLoadFunctions[i]()
|
||||
this.logger(
|
||||
'log',
|
||||
`Successfully loaded Universal Sentence Encoder from alternative URL ${i + 1}`
|
||||
)
|
||||
return model
|
||||
} catch (altError) {
|
||||
this.logger(
|
||||
'warn',
|
||||
`Failed to load from alternative URL ${i + 1}: ${altError}`
|
||||
)
|
||||
// Continue to the next alternative
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If we get here, all alternatives failed
|
||||
this.logger(
|
||||
'error',
|
||||
`Universal Sentence Encoder model loading failed after trying all alternatives. Last error: ${errorMessage}`
|
||||
)
|
||||
} else {
|
||||
this.logger(
|
||||
'error',
|
||||
`Universal Sentence Encoder model loading failed with non-retryable error: ${errorMessage}`
|
||||
)
|
||||
}
|
||||
throw lastError
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// This should never be reached, but just in case
|
||||
throw lastError || new Error('Unknown error during model loading')
|
||||
}
|
||||
|
||||
/**
|
||||
|
|
@ -277,8 +197,6 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
// Add polyfills for TensorFlow.js compatibility
|
||||
this.addServerCompatibilityPolyfills()
|
||||
|
||||
// TensorFlow.js will use its default EPSILON value
|
||||
|
||||
// CRITICAL: Ensure TextEncoder/TextDecoder are available before TensorFlow.js loads
|
||||
try {
|
||||
// Get the appropriate global object for the current environment
|
||||
|
|
@ -303,7 +221,8 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
globalObj.TextEncoder = util.TextEncoder
|
||||
}
|
||||
if (!globalObj.TextDecoder) {
|
||||
globalObj.TextDecoder = util.TextDecoder as unknown as typeof TextDecoder
|
||||
globalObj.TextDecoder =
|
||||
util.TextDecoder as unknown as typeof TextDecoder
|
||||
}
|
||||
}
|
||||
} catch (utilError) {
|
||||
|
|
@ -363,7 +282,9 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
} catch (error) {
|
||||
this.logger('error', 'Failed to initialize TensorFlow.js:', error)
|
||||
// No fallback allowed - throw error
|
||||
throw new Error(`Universal Sentence Encoder initialization failed: ${error}`)
|
||||
throw new Error(
|
||||
`Universal Sentence Encoder initialization failed: ${error}`
|
||||
)
|
||||
}
|
||||
|
||||
// Set the backend
|
||||
|
|
@ -375,13 +296,18 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
const loadFunction = findUSELoadFunction(this.use)
|
||||
|
||||
if (!loadFunction) {
|
||||
this.logger('error', 'Could not find Universal Sentence Encoder load function')
|
||||
throw new Error('Universal Sentence Encoder load function not found. Fallback mechanisms are not allowed.')
|
||||
this.logger(
|
||||
'error',
|
||||
'Could not find Universal Sentence Encoder load function'
|
||||
)
|
||||
throw new Error(
|
||||
'Universal Sentence Encoder load function not found. Fallback mechanisms are not allowed.'
|
||||
)
|
||||
}
|
||||
|
||||
try {
|
||||
// Load the model with retry logic for network failures
|
||||
this.model = await this.loadModelWithRetry(loadFunction)
|
||||
// Load the model from local files first, falling back to default loading if necessary
|
||||
this.model = await this.loadModelFromLocal(loadFunction)
|
||||
this.initialized = true
|
||||
} catch (modelError) {
|
||||
this.logger(
|
||||
|
|
@ -390,7 +316,9 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
modelError
|
||||
)
|
||||
// No fallback allowed - throw error
|
||||
throw new Error(`Universal Sentence Encoder model loading failed: ${modelError}`)
|
||||
throw new Error(
|
||||
`Universal Sentence Encoder model loading failed: ${modelError}`
|
||||
)
|
||||
}
|
||||
|
||||
// Restore original console.warn
|
||||
|
|
@ -402,7 +330,9 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
error
|
||||
)
|
||||
// No fallback allowed - throw error
|
||||
throw new Error(`Universal Sentence Encoder initialization failed: ${error}`)
|
||||
throw new Error(
|
||||
`Universal Sentence Encoder initialization failed: ${error}`
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -410,15 +340,6 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
* Embed text into a vector using Universal Sentence Encoder
|
||||
* @param data Text to embed
|
||||
*/
|
||||
/**
|
||||
* This method has been removed as we should always use Universal Sentence Encoder
|
||||
* and never fall back to alternative vector generation methods
|
||||
* @deprecated
|
||||
*/
|
||||
private generateFallbackVector(text: string): Vector {
|
||||
throw new Error('Fallback vector generation is not allowed. Universal Sentence Encoder must be used for all embeddings.')
|
||||
}
|
||||
|
||||
public async embed(data: string | string[]): Promise<Vector> {
|
||||
if (!this.initialized) {
|
||||
await this.init()
|
||||
|
|
@ -430,7 +351,7 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
if (typeof data === 'string') {
|
||||
// Handle empty string case
|
||||
if (data.trim() === '') {
|
||||
// Return a zero vector of appropriate dimension (512 is the default for USE)
|
||||
// Return a zero vector of 512 dimensions (standard for Universal Sentence Encoder)
|
||||
return new Array(512).fill(0)
|
||||
}
|
||||
textToEmbed = [data]
|
||||
|
|
@ -453,11 +374,9 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
)
|
||||
}
|
||||
|
||||
// Ensure the model is available - no fallbacks allowed
|
||||
// Ensure the model is available
|
||||
if (!this.model) {
|
||||
throw new Error(
|
||||
'Universal Sentence Encoder model is not available. Fallback mechanisms are not allowed.'
|
||||
)
|
||||
throw new Error('Universal Sentence Encoder model is not available')
|
||||
}
|
||||
|
||||
// Get embeddings
|
||||
|
|
@ -471,14 +390,14 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
|
||||
// Get the first embedding
|
||||
let embedding = embeddingArray[0]
|
||||
|
||||
// Ensure the embedding is exactly 512 dimensions
|
||||
|
||||
// Always ensure the embedding is exactly 512 dimensions
|
||||
if (embedding.length !== 512) {
|
||||
this.logger(
|
||||
'warn',
|
||||
`Embedding dimension mismatch: expected 512, got ${embedding.length}. Standardizing...`
|
||||
`Embedding dimension mismatch: expected 512, got ${embedding.length}. Standardizing to 512 dimensions.`
|
||||
)
|
||||
|
||||
|
||||
// If the embedding is too short, pad with zeros
|
||||
if (embedding.length < 512) {
|
||||
const paddedEmbedding = new Array(512).fill(0)
|
||||
|
|
@ -486,27 +405,15 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
paddedEmbedding[i] = embedding[i]
|
||||
}
|
||||
embedding = paddedEmbedding
|
||||
}
|
||||
}
|
||||
// If the embedding is too long, truncate
|
||||
else if (embedding.length > 512) {
|
||||
// Special handling for 1536-dimensional vectors (common with newer models)
|
||||
if (embedding.length === 1536) {
|
||||
// Take every third value to reduce from 1536 to 512
|
||||
const reducedEmbedding = new Array(512).fill(0)
|
||||
for (let i = 0; i < 512; i++) {
|
||||
reducedEmbedding[i] = embedding[i * 3]
|
||||
}
|
||||
embedding = reducedEmbedding
|
||||
} else {
|
||||
// For other dimensions, just truncate
|
||||
embedding = embedding.slice(0, 512)
|
||||
}
|
||||
embedding = embedding.slice(0, 512)
|
||||
}
|
||||
}
|
||||
|
||||
return embedding
|
||||
} catch (error) {
|
||||
// No fallback - throw the error
|
||||
this.logger(
|
||||
'error',
|
||||
'Failed to embed text with Universal Sentence Encoder:',
|
||||
|
|
@ -543,11 +450,9 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
return dataArray.map(() => new Array(512).fill(0))
|
||||
}
|
||||
|
||||
// Ensure the model is available - no fallbacks allowed
|
||||
// Ensure the model is available
|
||||
if (!this.model) {
|
||||
throw new Error(
|
||||
'Universal Sentence Encoder model is not available. Fallback mechanisms are not allowed.'
|
||||
)
|
||||
throw new Error('Universal Sentence Encoder model is not available')
|
||||
}
|
||||
|
||||
// Get embeddings for all texts in a single batch operation
|
||||
|
|
@ -564,9 +469,9 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
if (embedding.length !== 512) {
|
||||
this.logger(
|
||||
'warn',
|
||||
`Batch embedding dimension mismatch: expected 512, got ${embedding.length}. Standardizing...`
|
||||
`Batch embedding dimension mismatch: expected 512, got ${embedding.length}. Standardizing to 512 dimensions.`
|
||||
)
|
||||
|
||||
|
||||
// If the embedding is too short, pad with zeros
|
||||
if (embedding.length < 512) {
|
||||
const paddedEmbedding = new Array(512).fill(0)
|
||||
|
|
@ -574,21 +479,10 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
paddedEmbedding[i] = embedding[i]
|
||||
}
|
||||
return paddedEmbedding
|
||||
}
|
||||
}
|
||||
// If the embedding is too long, truncate
|
||||
else if (embedding.length > 512) {
|
||||
// Special handling for 1536-dimensional vectors (common with newer models)
|
||||
if (embedding.length === 1536) {
|
||||
// Take every third value to reduce from 1536 to 512
|
||||
const reducedEmbedding = new Array(512).fill(0)
|
||||
for (let i = 0; i < 512; i++) {
|
||||
reducedEmbedding[i] = embedding[i * 3]
|
||||
}
|
||||
return reducedEmbedding
|
||||
} else {
|
||||
// For other dimensions, just truncate
|
||||
return embedding.slice(0, 512)
|
||||
}
|
||||
return embedding.slice(0, 512)
|
||||
}
|
||||
}
|
||||
return embedding
|
||||
|
|
@ -612,13 +506,14 @@ export class UniversalSentenceEncoder implements EmbeddingModel {
|
|||
|
||||
return results
|
||||
} catch (error) {
|
||||
// No fallback - throw the error
|
||||
this.logger(
|
||||
'error',
|
||||
'Failed to batch embed text with Universal Sentence Encoder:',
|
||||
error
|
||||
)
|
||||
throw new Error(`Universal Sentence Encoder batch embedding failed: ${error}`)
|
||||
throw new Error(
|
||||
`Universal Sentence Encoder batch embedding failed: ${error}`
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -693,7 +588,8 @@ function findUSELoadFunction(
|
|||
} else if (
|
||||
sentenceEncoderModule.default &&
|
||||
sentenceEncoderModule.default.UniversalSentenceEncoder &&
|
||||
typeof sentenceEncoderModule.default.UniversalSentenceEncoder.load === 'function'
|
||||
typeof sentenceEncoderModule.default.UniversalSentenceEncoder.load ===
|
||||
'function'
|
||||
) {
|
||||
loadFunction = sentenceEncoderModule.default.UniversalSentenceEncoder.load
|
||||
}
|
||||
|
|
@ -734,7 +630,7 @@ function findUSELoadFunction(
|
|||
if (loadFunction) {
|
||||
return async () => await loadFunction()
|
||||
}
|
||||
|
||||
|
||||
return null
|
||||
}
|
||||
|
||||
|
|
@ -805,17 +701,19 @@ let sharedModel: UniversalSentenceEncoder | null = null
|
|||
let sharedModelInitialized = false
|
||||
let sharedModelVerbose = true
|
||||
|
||||
export function createTensorFlowEmbeddingFunction(options: { verbose?: boolean } = {}): EmbeddingFunction {
|
||||
export function createTensorFlowEmbeddingFunction(
|
||||
options: { verbose?: boolean } = {}
|
||||
): EmbeddingFunction {
|
||||
// Update verbose setting if provided
|
||||
if (options.verbose !== undefined) {
|
||||
sharedModelVerbose = options.verbose
|
||||
}
|
||||
|
||||
|
||||
// Create the shared model if it doesn't exist yet
|
||||
if (!sharedModel) {
|
||||
sharedModel = new UniversalSentenceEncoder({ verbose: sharedModelVerbose })
|
||||
}
|
||||
|
||||
|
||||
return async (data: any): Promise<Vector> => {
|
||||
try {
|
||||
// Initialize the model if it hasn't been initialized yet
|
||||
|
|
@ -832,7 +730,12 @@ export function createTensorFlowEmbeddingFunction(options: { verbose?: boolean }
|
|||
|
||||
return await sharedModel!.embed(data)
|
||||
} catch (error) {
|
||||
logIfNotTest('error', 'Failed to use Universal Sentence Encoder:', [error], sharedModelVerbose)
|
||||
logIfNotTest(
|
||||
'error',
|
||||
'Failed to use Universal Sentence Encoder:',
|
||||
[error],
|
||||
sharedModelVerbose
|
||||
)
|
||||
// No fallback - Universal Sentence Encoder is required
|
||||
throw new Error(
|
||||
`Universal Sentence Encoder is required and no fallbacks are allowed: ${error}`
|
||||
|
|
@ -849,7 +752,9 @@ export function createTensorFlowEmbeddingFunction(options: { verbose?: boolean }
|
|||
* @param options Configuration options
|
||||
* @param options.verbose Whether to log non-essential messages (default: true)
|
||||
*/
|
||||
export function getDefaultEmbeddingFunction(options: { verbose?: boolean } = {}): EmbeddingFunction {
|
||||
export function getDefaultEmbeddingFunction(
|
||||
options: { verbose?: boolean } = {}
|
||||
): EmbeddingFunction {
|
||||
return createTensorFlowEmbeddingFunction(options)
|
||||
}
|
||||
|
||||
|
|
@ -859,7 +764,8 @@ export function getDefaultEmbeddingFunction(options: { verbose?: boolean } = {})
|
|||
* TensorFlow.js is required for this to work
|
||||
* Uses CPU for compatibility
|
||||
*/
|
||||
export const defaultEmbeddingFunction: EmbeddingFunction = getDefaultEmbeddingFunction()
|
||||
export const defaultEmbeddingFunction: EmbeddingFunction =
|
||||
getDefaultEmbeddingFunction()
|
||||
|
||||
/**
|
||||
* Creates a batch embedding function that uses UniversalSentenceEncoder
|
||||
|
|
@ -874,19 +780,21 @@ let sharedBatchModel: UniversalSentenceEncoder | null = null
|
|||
let sharedBatchModelInitialized = false
|
||||
let sharedBatchModelVerbose = true
|
||||
|
||||
export function createBatchEmbeddingFunction(options: { verbose?: boolean } = {}): (
|
||||
dataArray: string[]
|
||||
) => Promise<Vector[]> {
|
||||
export function createBatchEmbeddingFunction(
|
||||
options: { verbose?: boolean } = {}
|
||||
): (dataArray: string[]) => Promise<Vector[]> {
|
||||
// Update verbose setting if provided
|
||||
if (options.verbose !== undefined) {
|
||||
sharedBatchModelVerbose = options.verbose
|
||||
}
|
||||
|
||||
|
||||
// Create the shared model if it doesn't exist yet
|
||||
if (!sharedBatchModel) {
|
||||
sharedBatchModel = new UniversalSentenceEncoder({ verbose: sharedBatchModelVerbose })
|
||||
sharedBatchModel = new UniversalSentenceEncoder({
|
||||
verbose: sharedBatchModelVerbose
|
||||
})
|
||||
}
|
||||
|
||||
|
||||
return async (dataArray: string[]): Promise<Vector[]> => {
|
||||
try {
|
||||
// Initialize the model if it hasn't been initialized yet
|
||||
|
|
@ -903,7 +811,12 @@ export function createBatchEmbeddingFunction(options: { verbose?: boolean } = {}
|
|||
|
||||
return await sharedBatchModel!.embedBatch(dataArray)
|
||||
} catch (error) {
|
||||
logIfNotTest('error', 'Failed to use Universal Sentence Encoder batch embedding:', [error], sharedBatchModelVerbose)
|
||||
logIfNotTest(
|
||||
'error',
|
||||
'Failed to use Universal Sentence Encoder batch embedding:',
|
||||
[error],
|
||||
sharedBatchModelVerbose
|
||||
)
|
||||
// No fallback - Universal Sentence Encoder is required
|
||||
throw new Error(
|
||||
`Universal Sentence Encoder is required for batch embedding and no fallbacks are allowed: ${error}`
|
||||
|
|
@ -920,9 +833,9 @@ export function createBatchEmbeddingFunction(options: { verbose?: boolean } = {}
|
|||
* @param options Configuration options
|
||||
* @param options.verbose Whether to log non-essential messages (default: true)
|
||||
*/
|
||||
export function getDefaultBatchEmbeddingFunction(options: { verbose?: boolean } = {}): (
|
||||
dataArray: string[]
|
||||
) => Promise<Vector[]> {
|
||||
export function getDefaultBatchEmbeddingFunction(
|
||||
options: { verbose?: boolean } = {}
|
||||
): (dataArray: string[]) => Promise<Vector[]> {
|
||||
return createBatchEmbeddingFunction(options)
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue