Spaces:
Running
Running
| import type { | |
| ASRAdapter, | |
| ASRResult, | |
| RuntimeWarmupOptions, | |
| } from "../app/types"; | |
| interface TransformerPipeline { | |
| (input: string): Promise<{ text?: string } | string>; | |
| } | |
| type TransformersProgressInfo = { | |
| status?: "initiate" | "download" | "progress" | "done" | "ready"; | |
| file?: string; | |
| progress?: number; | |
| task?: string; | |
| model?: string; | |
| }; | |
| const formatProgress = (progress: TransformersProgressInfo): string => { | |
| switch (progress.status) { | |
| case "initiate": | |
| return `Preparing ${progress.file ?? "speech model"}...`; | |
| case "download": | |
| return `Downloading ${progress.file ?? "speech model"}...`; | |
| case "progress": | |
| return `Downloading ${progress.file ?? "speech model"} (${Math.round(progress.progress ?? 0)}%)`; | |
| case "done": | |
| return `Loaded ${progress.file ?? "speech model"}.`; | |
| case "ready": | |
| return `Speech model ready (${progress.model ?? progress.task ?? "loaded"}).`; | |
| default: | |
| return "Loading speech model..."; | |
| } | |
| }; | |
| export class TransformersAsrAdapter implements ASRAdapter { | |
| #pipeline: TransformerPipeline | null = null; | |
| #device: "webgpu" | "wasm" = "wasm"; | |
| constructor(private readonly modelId: string) {} | |
| async initialize(options: RuntimeWarmupOptions = {}): Promise<void> { | |
| if (this.#pipeline) { | |
| return; | |
| } | |
| const { env, pipeline } = await import("@huggingface/transformers"); | |
| env.allowLocalModels = false; | |
| env.useFS = false; | |
| env.useFSCache = false; | |
| env.useBrowserCache = true; | |
| const preferredDevice = | |
| typeof navigator !== "undefined" && "gpu" in navigator ? "webgpu" : "wasm"; | |
| try { | |
| options.onProgress?.(`Loading ${this.modelId} on ${preferredDevice.toUpperCase()}...`); | |
| this.#pipeline = (await pipeline( | |
| "automatic-speech-recognition", | |
| this.modelId, | |
| { | |
| device: preferredDevice, | |
| dtype: preferredDevice === "webgpu" ? "fp32" : "q8", | |
| progress_callback: (progress: TransformersProgressInfo) => { | |
| options.onProgress?.(formatProgress(progress)); | |
| }, | |
| }, | |
| )) as TransformerPipeline; | |
| this.#device = preferredDevice; | |
| } catch (error) { | |
| if (preferredDevice !== "webgpu") { | |
| throw error; | |
| } | |
| console.warn( | |
| "ASR WebGPU initialization failed, falling back to WASM.", | |
| error, | |
| ); | |
| options.onProgress?.("Speech WebGPU setup failed. Falling back to WASM..."); | |
| this.#pipeline = (await pipeline( | |
| "automatic-speech-recognition", | |
| this.modelId, | |
| { | |
| device: "wasm", | |
| dtype: "q8", | |
| progress_callback: (progress: TransformersProgressInfo) => { | |
| options.onProgress?.(formatProgress(progress)); | |
| }, | |
| }, | |
| )) as TransformerPipeline; | |
| this.#device = "wasm"; | |
| } | |
| options.onProgress?.(`Speech ready on ${this.#device.toUpperCase()}.`); | |
| } | |
| async warmup(options: RuntimeWarmupOptions = {}): Promise<void> { | |
| await this.initialize(options); | |
| options.onProgress?.("Speech runtime warmed."); | |
| } | |
| async transcribe(audio: Blob): Promise<ASRResult> { | |
| await this.initialize(); | |
| if (!this.#pipeline) { | |
| throw new Error("ASR pipeline failed to initialize."); | |
| } | |
| const url = URL.createObjectURL(audio); | |
| try { | |
| const result = await this.#pipeline(url); | |
| const text = typeof result === "string" ? result : result.text ?? ""; | |
| if (!text.trim()) { | |
| throw new Error("Speech recognition returned an empty transcript."); | |
| } | |
| return { text: text.trim() }; | |
| } finally { | |
| URL.revokeObjectURL(url); | |
| } | |
| } | |
| getDevice(): "webgpu" | "wasm" { | |
| return this.#device; | |
| } | |
| } | |