Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion package.json
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
2 changes: 1 addition & 1 deletion src/embeddings/model-info.ts
Original file line number Diff line number Diff line change
Expand Up @@ -122,7 +122,7 @@ export function getModelMaxTokens(model: string): number | undefined {
const MODEL_MAPPINGS: Record<string, Record<string, string>> = {
'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',
Expand Down
58 changes: 55 additions & 3 deletions src/embeddings/transformers-js.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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<string, TransformersDevice>
/**
* Data type or per-file data type map. Defaults to `'fp32'`.
*/
dtype?: TransformersDtype | Record<string, TransformersDtype>
/** Called with model download progress (initiate β†’ download β†’ progress β†’ done β†’ ready) */
onProgress?: (info: TransformersProgressInfo) => void
}
Expand Down Expand Up @@ -50,6 +90,12 @@ async function clearCorruptedCache(error: unknown, model: string): Promise<boole
* path: 'vectors.db',
* embeddings: transformersJs({ model: 'bge-base-en-v1.5' }),
* })
*
* // Run inference with WebGPU
* const fast = await sqliteVec({
* path: 'vectors.db',
* embeddings: transformersJs({ model: 'bge-base-en-v1.5', device: 'webgpu' }),
* })
* ```
*/
export function transformersJs(options: TransformersEmbeddingOptions = {}): EmbeddingConfig {
Expand All @@ -62,7 +108,9 @@ export function transformersJs(options: TransformersEmbeddingOptions = {}): Embe
if (cached)
return cached

const pipelineOpts: Record<string, unknown> = { dtype: 'fp32' }
const pipelineOpts: Record<string, unknown> = { dtype: options.dtype ?? 'fp32' }
if (options.device !== undefined)
pipelineOpts.device = options.device
if (options.onProgress)
pipelineOpts.progress_callback = options.onProgress

Expand All @@ -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 })
Expand Down
78 changes: 78 additions & 0 deletions test/embeddings-transformers-js.test.ts
Original file line number Diff line number Diff line change
@@ -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()
})
})
32 changes: 32 additions & 0 deletions test/model-info.test.ts
Original file line number Diff line number Diff line change
@@ -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)
})
})
Loading