From f6ee1bc7c1c56dbc6afc8b379f568c896387b48b Mon Sep 17 00:00:00 2001 From: Harlan Wilton Date: Wed, 12 Aug 2026 15:28:15 +1000 Subject: [PATCH] fix(transformers): forward embedding device --- src/embeddings/transformers-js.ts | 5 +++++ test/transformers-js.test.ts | 34 +++++++++++++++++++++++++++++++ 2 files changed, 39 insertions(+) create mode 100644 test/transformers-js.test.ts diff --git a/src/embeddings/transformers-js.ts b/src/embeddings/transformers-js.ts index 629d33e..96b7918 100644 --- a/src/embeddings/transformers-js.ts +++ b/src/embeddings/transformers-js.ts @@ -1,3 +1,4 @@ +import type { DeviceType } from '@huggingface/transformers' import type { EmbeddingConfig, EmbeddingProvider, ResolvedEmbedding } from '../types' import { rm } from 'node:fs/promises' import { env, pipeline } from '@huggingface/transformers' @@ -17,6 +18,8 @@ export interface TransformersEmbeddingOptions { model?: string /** Embedding dimensions (auto-detected for known models) */ dimensions?: number + /** Device used to run the model. Omit it to use transformers.js defaults. */ + device?: DeviceType /** Called with model download progress (initiate → download → progress → done → ready) */ onProgress?: (info: TransformersProgressInfo) => void } @@ -65,6 +68,8 @@ export function transformersJs(options: TransformersEmbeddingOptions = {}): Embe const pipelineOpts: Record = { dtype: 'fp32' } if (options.onProgress) pipelineOpts.progress_callback = options.onProgress + if (options.device) + pipelineOpts.device = options.device const extractor = await pipeline('feature-extraction', model, pipelineOpts) .catch(async (err) => { diff --git a/test/transformers-js.test.ts b/test/transformers-js.test.ts new file mode 100644 index 0000000..1d4c286 --- /dev/null +++ b/test/transformers-js.test.ts @@ -0,0 +1,34 @@ +import { pipeline } from '@huggingface/transformers' +import { beforeEach, describe, expect, it, vi } from 'vitest' +import { transformersJs } from '../src/embeddings/transformers-js' + +vi.mock('@huggingface/transformers', () => ({ + env: {}, + pipeline: vi.fn(async () => vi.fn()), +})) + +describe('transformersJs', () => { + beforeEach(() => { + vi.mocked(pipeline).mockClear() + }) + + it('runs the model on the selected device', async () => { + await transformersJs({ model: 'bge-small-en-v1.5', device: 'webgpu' }).resolve() + + expect(pipeline).toHaveBeenCalledWith( + 'feature-extraction', + 'Xenova/bge-small-en-v1.5', + expect.objectContaining({ device: 'webgpu' }), + ) + }) + + it('leaves device selection to transformers.js when omitted', async () => { + await transformersJs({ model: 'bge-small-en-v1.5' }).resolve() + + expect(pipeline).toHaveBeenCalledWith( + 'feature-extraction', + 'Xenova/bge-small-en-v1.5', + expect.not.objectContaining({ device: expect.anything() }), + ) + }) +})