diff --git a/cspell.config.yaml b/cspell.config.yaml
index 7c37dcb36..bca7b2a20 100644
--- a/cspell.config.yaml
+++ b/cspell.config.yaml
@@ -40,6 +40,7 @@ words:
- demi
- dotenvx
- dtype
+ - dtypes
- duckdb
- DuckDBWASM
- DuckDBWASMQ
diff --git a/packages/provider-transformers/package.json b/packages/provider-transformers/package.json
index ad2c9077b..fff7f69f2 100644
--- a/packages/provider-transformers/package.json
+++ b/packages/provider-transformers/package.json
@@ -22,10 +22,6 @@
"./worker": {
"types": "./dist/worker/index.d.ts",
"import": "./dist/worker/index.mjs"
- },
- "./types": {
- "types": "./dist/types/index.d.ts",
- "import": "./dist/types/index.mjs"
}
},
"module": "./dist/index.mjs",
diff --git a/packages/provider-transformers/playground/src/App.vue b/packages/provider-transformers/playground/src/App.vue
index b306b348e..d3823d065 100644
--- a/packages/provider-transformers/playground/src/App.vue
+++ b/packages/provider-transformers/playground/src/App.vue
@@ -23,7 +23,6 @@ onMounted(async () => {
async function execute() {
const result = await embed({
...transformersProvider.embed(modelId.value),
- model: modelId.value,
input: input.value,
})
diff --git a/packages/provider-transformers/src/index.ts b/packages/provider-transformers/src/index.ts
index 7451aa1a7..f0eb0d43d 100644
--- a/packages/provider-transformers/src/index.ts
+++ b/packages/provider-transformers/src/index.ts
@@ -9,7 +9,7 @@ export type Loadable
= P & {
loadEmbed: (model: (string & {}) | T, options?: T2) => Promise
}
-export function createEmbedProvider(createOptions: CreateProviderOptions): Loadable, T, T2> {
+export function createEmbedProvider & LoadOptions>(createOptions: CreateProviderOptions): Loadable, T, T2> {
let worker: Worker
let isReady = false
@@ -40,6 +40,12 @@ export function createEmbedProvider & Partial as any,
loadEmbed: loadModel,
}
}
export function createTransformers(options: { embedWorkerURL: string }) {
return merge(
- createEmbedProvider<'Xenova/all-MiniLM-L6-v2', CreateProviderOptions & LoadOptions & { model: string }>({ baseURL: `xsai-provider-ext:///?worker-url=${options.embedWorkerURL}&other=` }),
+ createEmbedProvider<'Xenova/all-MiniLM-L6-v2', Omit & LoadOptions>({ baseURL: `xsai-provider-ext:///?worker-url=${options.embedWorkerURL}&other=` }),
)
}
diff --git a/packages/provider-transformers/src/types/index.ts b/packages/provider-transformers/src/types/index.ts
index b5f166427..685a80bfe 100644
--- a/packages/provider-transformers/src/types/index.ts
+++ b/packages/provider-transformers/src/types/index.ts
@@ -1,12 +1,14 @@
-import type { FeatureExtractionPipelineOptions, pipeline, ProgressInfo } from '@huggingface/transformers'
-import type { PipelineOptionsFrom } from '@proj-airi/utils-transformers/types'
+import type { FeatureExtractionPipelineOptions } from '@huggingface/transformers'
+import type { ModelSpecificPretrainedOptions, PretrainedOptions, ProgressInfo } from '@proj-airi/utils-transformers/types'
export enum MessageStatus {
Loading = 'loading',
Ready = 'ready',
}
-export type LoadOptions = Omit>, 'progress_callback'>
+export type LoadOptions = Omit & { onProgress?: LoadOptionProgressCallback }
+export type LoadOptionProgressCallback = (progress: ProgressInfo) => void | Promise
+export type { ProgressInfo }
export interface WorkerMessageBaseEvent {
type: T
diff --git a/packages/utils-transformers/package.json b/packages/utils-transformers/package.json
index 57046c381..9fe09d509 100644
--- a/packages/utils-transformers/package.json
+++ b/packages/utils-transformers/package.json
@@ -32,6 +32,7 @@
"typecheck": "tsc --noEmit"
},
"dependencies": {
- "@huggingface/transformers": "^3.3.3"
+ "@huggingface/transformers": "^3.3.3",
+ "onnxruntime-common": "^1.20.1"
}
}
diff --git a/packages/utils-transformers/src/types/core.ts b/packages/utils-transformers/src/types/core.ts
new file mode 100644
index 000000000..74cbab46e
--- /dev/null
+++ b/packages/utils-transformers/src/types/core.ts
@@ -0,0 +1,36 @@
+export interface InitiateProgressInfo {
+ status: 'initiate'
+ name: string
+ file: string
+}
+
+export interface DownloadProgressInfo {
+ status: 'download'
+ name: string
+ file: string
+}
+
+export interface ProgressStatusInfo {
+ status: 'progress'
+ name: string
+ file: string
+ progress: number
+ loaded: number
+ total: number
+}
+
+export interface DoneProgressInfo {
+ status: 'done'
+ name: string
+ file: string
+}
+
+export interface ReadyProgressInfo {
+ status: 'ready'
+ task: string
+ model: string
+}
+
+export type ProgressInfo = InitiateProgressInfo | DownloadProgressInfo | ProgressStatusInfo | DoneProgressInfo | ReadyProgressInfo
+
+export type ProgressCallback = (progress: ProgressInfo) => void
diff --git a/packages/utils-transformers/src/types/devices.ts b/packages/utils-transformers/src/types/devices.ts
new file mode 100644
index 000000000..d33993963
--- /dev/null
+++ b/packages/utils-transformers/src/types/devices.ts
@@ -0,0 +1,3 @@
+import type { pipeline } from '@huggingface/transformers'
+
+export type Device = Extract>[2]['device']>, Record>, 'webgpu' | 'wasm'>
diff --git a/packages/utils-transformers/src/types/dtypes.ts b/packages/utils-transformers/src/types/dtypes.ts
new file mode 100644
index 000000000..490fd19cb
--- /dev/null
+++ b/packages/utils-transformers/src/types/dtypes.ts
@@ -0,0 +1,3 @@
+import type { pipeline } from '@huggingface/transformers'
+
+export type DType = Record>[2]['dtype']>, string>[string]>
diff --git a/packages/utils-transformers/src/types/hub.ts b/packages/utils-transformers/src/types/hub.ts
new file mode 100644
index 000000000..0315c8206
--- /dev/null
+++ b/packages/utils-transformers/src/types/hub.ts
@@ -0,0 +1,76 @@
+import type { PretrainedConfig } from '@huggingface/transformers'
+import type { InferenceSession } from 'onnxruntime-common'
+import type { ProgressCallback } from './core'
+import type { Device } from './devices'
+import type { DType } from './dtypes'
+
+/**
+ * Options for loading a pretrained model.
+ */
+export interface ModelSpecificPretrainedOptions {
+ /**
+ * In case the relevant files are located inside a subfolder of the model repo on huggingface.co,
+ * you can specify the folder name here.
+ *
+ * @default 'onnx'
+ */
+ subfolder?: string
+ /**
+ * If specified, load the model with this name (excluding the .onnx suffix). Currently only valid for encoder- or decoder-only models.
+ *
+ * @default null
+ */
+ model_file_name?: string
+ /**
+ * The device to run the model on. If not specified, the device will be chosen from the environment settings.
+ */
+ device?: Device
+ /**
+ * The data type to use for the model. If not specified, the data type will be chosen from the environment settings.
+ */
+ dtype?: DType
+ /**
+ * Whether to load the model using the external data format (used for models >= 2GB in size).
+ *
+ * @default false
+ */
+ use_external_data_format?: boolean
+ /**
+ * User-specified session options passed to the runtime. If not provided, suitable defaults will be chosen.
+ */
+ session_options?: InferenceSession.SessionOptions
+}
+
+/**
+ * Options for loading a pretrained model.
+ */
+export interface PretrainedOptions {
+ /**
+ * If specified, this function will be called during model construction, to provide the user with progress updates.
+ */
+ progress_callback?: ProgressCallback
+ /**
+ * Configuration for the model to use instead of an automatically loaded configuration. Configuration can be automatically loaded when:
+ * - The model is a model provided by the library (loaded with the *model id* string of a pretrained model).
+ * - The model is loaded by supplying a local directory as `pretrained_model_name_or_path` and a configuration JSON file named *config.json* is found in the directory.
+ */
+ config?: PretrainedConfig
+ /**
+ * Path to a directory in which a downloaded pretrained model configuration should be cached if the standard cache should not be used.
+ */
+ cache_dir?: string
+ /**
+ * Whether or not to only look at local files (e.g., not try downloading the model).
+ *
+ * @default false
+ */
+ local_files_only?: boolean
+ /**
+ * The specific model version to use. It can be a branch name, a tag name, or a commit id,
+ * since we use a git-based system for storing models and other artifacts on huggingface.co, so `revision` can be any identifier allowed by git.
+ * NOTE: This setting is ignored for local requests.
+ *
+ * @default 'main'
+ */
+ revision?: string
+}
diff --git a/packages/utils-transformers/src/types/index.ts b/packages/utils-transformers/src/types/index.ts
index 6ff6b2190..625306612 100644
--- a/packages/utils-transformers/src/types/index.ts
+++ b/packages/utils-transformers/src/types/index.ts
@@ -1,7 +1,10 @@
-import type { AutoModel, pipeline } from '@huggingface/transformers'
+import type { AutoModel } from '@huggingface/transformers'
+
+export * from './core'
+export * from './devices'
+export * from './dtypes'
+export * from './hub'
-export type DType = Record>[2]['dtype']>, string>[string]>
-export type Device = Extract>[2]['device']>, Record>, 'webgpu' | 'wasm'>
export type PretrainedConfig = NonNullable[1]>['config']
export type PretrainedConfigFrom = T extends { from_pretrained: (...args: any) => any } ? NonNullable[1]>['config'] : never
export type PipelineOptionsFrom = T extends (...args: any) => any ? NonNullable[2]> : never
diff --git a/pnpm-lock.yaml b/pnpm-lock.yaml
index fd348b06b..1696194a4 100644
--- a/pnpm-lock.yaml
+++ b/pnpm-lock.yaml
@@ -1041,6 +1041,9 @@ importers:
'@huggingface/transformers':
specifier: ^3.3.3
version: 3.3.3
+ onnxruntime-common:
+ specifier: ^1.20.1
+ version: 1.20.1
services/discord-bot:
dependencies:
@@ -14566,7 +14569,7 @@ snapshots:
dependencies:
'@typeschema/core': 0.14.0(@types/json-schema@7.0.15)
optionalDependencies:
- '@typeschema/valibot': 0.14.0(@gcornut/valibot-json-schema@0.42.0(esbuild@0.25.0)(typescript@5.7.3))(@types/json-schema@7.0.15)(valibot@1.0.0-beta.9(typescript@5.7.3))
+ '@typeschema/valibot': 0.14.0(@gcornut/valibot-json-schema@0.42.0(esbuild@0.24.2)(typescript@5.7.3))(@types/json-schema@7.0.15)(valibot@1.0.0-beta.9(typescript@5.7.3))
'@typeschema/zod': 0.14.0(@types/json-schema@7.0.15)(zod-to-json-schema@3.24.3(zod@3.24.2))(zod@3.24.2)
transitivePeerDependencies:
- '@types/json-schema'