diff --git a/README.md b/README.md index 58a18d7..f1c4e16 100644 --- a/README.md +++ b/README.md @@ -324,6 +324,20 @@ ollama({ model: 'nomic-embed-text' }) transformersJs({ model: 'Xenova/all-MiniLM-L6-v2' }) ``` +### Transformers.js runtime options + +Retriv passes `device` and `dtype` to Transformers.js. Retriv uses `fp32` when +you omit `dtype`. Transformers.js selects the device when you omit `device`. + +```ts +transformersJs({ model: 'bge-base-en-v1.5', device: 'webgpu' }) + +transformersJs({ model: 'bge-base-en-v1.5', dtype: 'q8' }) +``` + +Support depends on your Transformers.js version and runtime. Test the selected +combination on the target system. + ## API ### SearchProvider diff --git a/package.json b/package.json index 8f8be99..0f56842 100644 --- a/package.json +++ b/package.json @@ -140,7 +140,7 @@ "@ai-sdk/google": "^3.0.0", "@ai-sdk/mistral": "^3.0.0", "@ai-sdk/openai": "^3.0.0", - "@huggingface/transformers": "^3.0.0", + "@huggingface/transformers": "^3.0.0 || ^4.0.0", "@libsql/client": "^0.14.0 || ^0.15.0 || ^0.16.0 || ^0.17.0", "@upstash/vector": "^1.0.0", "ai": "^4.0.0 || ^5.0.0 || ^6.0.0", diff --git a/src/embeddings/model-info.ts b/src/embeddings/model-info.ts index 0667f2e..662a03e 100644 --- a/src/embeddings/model-info.ts +++ b/src/embeddings/model-info.ts @@ -122,7 +122,7 @@ export function getModelMaxTokens(model: string): number | undefined { const MODEL_MAPPINGS: Record> = { 'transformers.js': { 'bge-base-en-v1.5': 'Xenova/bge-base-en-v1.5', - 'bge-large-en-v1.5': 'onnx-community/bge-large-en-v1.5', + 'bge-large-en-v1.5': 'Xenova/bge-large-en-v1.5', 'bge-small-en-v1.5': 'Xenova/bge-small-en-v1.5', 'bge-m3': 'Xenova/bge-m3', 'all-MiniLM-L6-v2': 'Xenova/all-MiniLM-L6-v2', diff --git a/src/embeddings/transformers-js.ts b/src/embeddings/transformers-js.ts index 629d33e..9a3e6ca 100644 --- a/src/embeddings/transformers-js.ts +++ b/src/embeddings/transformers-js.ts @@ -12,11 +12,51 @@ export interface TransformersProgressInfo { total?: number } +/** Execution device supported by Transformers.js */ +export type TransformersDevice + = | 'auto' + | 'cpu' + | 'gpu' + | 'wasm' + | 'webgpu' + | 'cuda' + | 'dml' + | 'coreml' + | 'webnn' + | 'webnn-npu' + | 'webnn-gpu' + | 'webnn-cpu' + +/** Quantization level supported by Transformers.js */ +export type TransformersDtype + = | 'auto' + | 'fp32' + | 'fp16' + | 'q8' + | 'int8' + | 'uint8' + | 'q4' + | 'bnb4' + | 'q4f16' + | 'q2' + | 'q2f16' + | 'q1' + | 'q1f16' + export interface TransformersEmbeddingOptions { /** Model name (e.g., 'bge-base-en-v1.5' or 'Xenova/bge-base-en-v1.5') */ model?: string /** Embedding dimensions (auto-detected for known models) */ dimensions?: number + /** + * Execution device or per-file device map. + * Transformers.js selects the device when this option is omitted. + */ + device?: TransformersDevice | Record + /** + * Data type or per-file data type map. Defaults to `'fp32'`. + */ + dtype?: TransformersDtype | Record /** Called with model download progress (initiate → download → progress → done → ready) */ onProgress?: (info: TransformersProgressInfo) => void } @@ -50,6 +90,12 @@ async function clearCorruptedCache(error: unknown, model: string): Promise = { dtype: 'fp32' } + const pipelineOpts: Record = { dtype: options.dtype ?? 'fp32' } + if (options.device !== undefined) + pipelineOpts.device = options.device if (options.onProgress) pipelineOpts.progress_callback = options.onProgress @@ -73,9 +121,13 @@ export function transformersJs(options: TransformersEmbeddingOptions = {}): Embe throw err }) - const dimensions = options.dimensions ?? getModelDimensions(model) + let dimensions = options.dimensions ?? getModelDimensions(model) + if (!dimensions) { + const probe = await extractor(['dimension probe'], { pooling: 'mean', normalize: true }) + dimensions = (probe.data as Float32Array).length + } if (!dimensions) - throw new Error(`Unknown dimensions for model ${model}. Please specify dimensions option.`) + throw new Error(`Could not determine dimensions for model ${model}. Please specify the dimensions option.`) const embedder: EmbeddingProvider = async (texts) => { const output = await extractor(texts, { pooling: 'mean', normalize: true }) diff --git a/test/embeddings-transformers-js.test.ts b/test/embeddings-transformers-js.test.ts new file mode 100644 index 0000000..e39c300 --- /dev/null +++ b/test/embeddings-transformers-js.test.ts @@ -0,0 +1,78 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest' + +const { pipelineMock } = vi.hoisted(() => ({ pipelineMock: vi.fn() })) + +vi.mock('@huggingface/transformers', () => ({ + env: { cacheDir: undefined }, + pipeline: pipelineMock, +})) + +const { transformersJs } = await import('../src/embeddings/transformers-js') + +describe('transformersJs pipeline options', () => { + beforeEach(() => { + pipelineMock.mockReset() + pipelineMock.mockResolvedValue(async () => ({ data: new Float32Array(0) })) + }) + + it('defaults to fp32 and leaves device unset', async () => { + await transformersJs({ model: 'bge-small-en-v1.5' }).resolve() + + const [task, model, opts] = pipelineMock.mock.calls[0]! + expect(task).toBe('feature-extraction') + expect(model).toBe('Xenova/bge-small-en-v1.5') + expect(opts).toEqual({ dtype: 'fp32' }) + }) + + it('forwards device and dtype to the pipeline', async () => { + await transformersJs({ + model: 'bge-base-en-v1.5', + device: 'coreml', + dtype: 'q8', + }).resolve() + + const [, model, opts] = pipelineMock.mock.calls[0]! + expect(model).toBe('Xenova/bge-base-en-v1.5') + expect(opts).toMatchObject({ device: 'coreml', dtype: 'q8' }) + }) + + it('forwards per-file device and dtype maps', async () => { + const device = { 'model.onnx': 'webgpu' } as const + const dtype = { 'model.onnx': 'q8' } as const + + await transformersJs({ model: 'bge-base-en-v1.5', device, dtype }).resolve() + + const [, , opts] = pipelineMock.mock.calls[0]! + expect(opts).toMatchObject({ device, dtype }) + }) + + it('forwards device without overriding the default dtype', async () => { + await transformersJs({ model: 'bge-small-en-v1.5', device: 'webgpu' }).resolve() + + const [, , opts] = pipelineMock.mock.calls[0]! + expect(opts).toEqual({ dtype: 'fp32', device: 'webgpu' }) + }) + + it('resolves dimensions for the selected model', async () => { + const resolved = await transformersJs({ model: 'bge-large-en-v1.5' }).resolve() + expect(resolved.dimensions).toBe(1024) + }) + + it('probes dimensions for a model missing from the registry', async () => { + pipelineMock.mockResolvedValue(async () => ({ data: new Float32Array(384) })) + + const resolved = await transformersJs({ model: 'some-org/unlisted-model' }).resolve() + + expect(resolved.dimensions).toBe(384) + }) + + it('prefers an explicit dimensions option over probing', async () => { + const extractor = vi.fn(async () => ({ data: new Float32Array(384) })) + pipelineMock.mockResolvedValue(extractor) + + const resolved = await transformersJs({ model: 'some-org/unlisted-model', dimensions: 512 }).resolve() + + expect(resolved.dimensions).toBe(512) + expect(extractor).not.toHaveBeenCalled() + }) +}) diff --git a/test/model-info.test.ts b/test/model-info.test.ts new file mode 100644 index 0000000..eb9ecfc --- /dev/null +++ b/test/model-info.test.ts @@ -0,0 +1,32 @@ +import { describe, expect, it } from 'vitest' +import { DEFAULT_MODELS, getModelDimensions, getModelMaxTokens, resolveModelForPreset } from '../src/embeddings/model-info' + +describe('transformers.js preset mapping', () => { + const presets = [ + ['bge-small-en-v1.5', 'Xenova/bge-small-en-v1.5', 384], + ['bge-base-en-v1.5', 'Xenova/bge-base-en-v1.5', 768], + ['bge-large-en-v1.5', 'Xenova/bge-large-en-v1.5', 1024], + ['bge-m3', 'Xenova/bge-m3', 1024], + ['all-MiniLM-L6-v2', 'Xenova/all-MiniLM-L6-v2', 384], + ] as const + + it.each(presets)('maps %s to a repo with published weights', (preset, repo, dims) => { + expect(resolveModelForPreset(preset, 'transformers.js')).toBe(repo) + expect(getModelDimensions(preset)).toBe(dims) + }) + + it('passes through fully-qualified repo ids untouched', () => { + expect(resolveModelForPreset('Xenova/bge-base-en-v1.5', 'transformers.js')) + .toBe('Xenova/bge-base-en-v1.5') + }) + + it('resolves dimensions and max tokens through repo prefixes', () => { + expect(getModelDimensions('Xenova/bge-large-en-v1.5')).toBe(1024) + expect(getModelMaxTokens('Xenova/bge-large-en-v1.5')).toBe(512) + }) + + it('keeps the transformers.js default consistent with its mapping', () => { + const fallback = DEFAULT_MODELS['transformers.js'] + expect(getModelDimensions(fallback.model)).toBe(fallback.dimensions) + }) +})