private-voice-agent / src /services /pocket-runtime.ts
aut4rk's picture
Deploy static voice agent
009fd18 verified
Raw History Blame
17.7 kB
const SAMPLE_RATE = 24_000;
const MAX_REFERENCE_SAMPLES = SAMPLE_RATE * 10;
const DEFAULT_MODEL_BASE_URL = "/pocket-tts";
const WORKER_MODULE_URL = new URL("../workers/pocket-bootstrap.ts", import.meta.url);
type Deferred<T> = {
promise: Promise<T>;
resolve: (value: T | PromiseLike<T>) => void;
reject: (reason?: unknown) => void;
};
type WorkerMessage =
| { type: "loaded" }
| { type: "voices_loaded"; defaultVoice?: string | null }
| { type: "voice_encoded"; voiceName?: string | null }
| { type: "voice_set"; voiceName?: string | null }
| { type: "audio_chunk"; data?: ArrayLike<number> }
| { type: "stream_ended" }
| { type: "error"; error?: string }
| { type: "status"; status?: string; state?: string };
type PendingGeneration = {
mode: "blob" | "stream";
chunks: Float32Array[];
chunkTask: Promise<void>;
onAudioChunk?: (chunk: Float32Array) => void | Promise<void>;
resolve: (value?: Blob) => void;
reject: (error: Error) => void;
abortCleanup?: () => void;
};
const createDeferred = <T>(): Deferred<T> => {
let resolve!: Deferred<T>["resolve"];
let reject!: Deferred<T>["reject"];
const promise = new Promise<T>((innerResolve, innerReject) => {
resolve = innerResolve;
reject = innerReject;
});
return { promise, resolve, reject };
};
const resolveModelBaseUrl = (modelId?: string): string => {
if (!modelId) {
return DEFAULT_MODEL_BASE_URL;
}
if (modelId.startsWith("https://") || modelId.startsWith("http://")) {
return modelId.replace(/\/+$/, "");
}
if (
modelId.startsWith("/") ||
modelId.startsWith("./") ||
modelId.startsWith("../")
) {
return modelId.replace(/\/+$/, "");
}
return `https://huggingface.co/spaces/${modelId.replace(/^spaces\//, "")}/resolve/main`;
};
const isWindowsPlatform = (): boolean => {
if (typeof navigator === "undefined") {
return false;
}
const platform =
navigator.userAgentData?.platform ??
navigator.platform ??
navigator.userAgent ??
"";
return /win/i.test(platform);
};
const toMono = (audioBuffer: AudioBuffer): Float32Array => {
if (audioBuffer.numberOfChannels === 1) {
return audioBuffer.getChannelData(0).slice();
}
const length = audioBuffer.length;
const mono = new Float32Array(length);
for (let channel = 0; channel < audioBuffer.numberOfChannels; channel += 1) {
const input = audioBuffer.getChannelData(channel);
for (let index = 0; index < length; index += 1) {
mono[index] += input[index] / audioBuffer.numberOfChannels;
}
}
return mono;
};
const resample = (
input: Float32Array,
fromRate: number,
toRate: number,
): Float32Array => {
if (fromRate === toRate) {
return input;
}
const ratio = fromRate / toRate;
const outputLength = Math.max(1, Math.round(input.length / ratio));
const output = new Float32Array(outputLength);
for (let index = 0; index < outputLength; index += 1) {
const sourceIndex = index * ratio;
const lower = Math.floor(sourceIndex);
const upper = Math.min(lower + 1, input.length - 1);
const weight = sourceIndex - lower;
output[index] = input[lower] * (1 - weight) + input[upper] * weight;
}
return output;
};
const encodeWav = (chunks: Float32Array[]): Blob => {
const totalSamples = chunks.reduce((sum, chunk) => sum + chunk.length, 0);
const pcm = new Float32Array(totalSamples);
let offset = 0;
for (const chunk of chunks) {
pcm.set(chunk, offset);
offset += chunk.length;
}
const buffer = new ArrayBuffer(44 + pcm.length * 2);
const view = new DataView(buffer);
const writeString = (start: number, value: string) => {
for (let index = 0; index < value.length; index += 1) {
view.setUint8(start + index, value.charCodeAt(index));
}
};
writeString(0, "RIFF");
view.setUint32(4, 36 + pcm.length * 2, true);
writeString(8, "WAVE");
writeString(12, "fmt ");
view.setUint32(16, 16, true);
view.setUint16(20, 1, true);
view.setUint16(22, 1, true);
view.setUint32(24, SAMPLE_RATE, true);
view.setUint32(28, SAMPLE_RATE * 2, true);
view.setUint16(32, 2, true);
view.setUint16(34, 16, true);
writeString(36, "data");
view.setUint32(40, pcm.length * 2, true);
for (let index = 0; index < pcm.length; index += 1) {
const sample = Math.max(-1, Math.min(1, pcm[index] ?? 0));
view.setInt16(
44 + index * 2,
sample < 0 ? sample * 0x8000 : sample * 0x7fff,
true,
);
}
return new Blob([buffer], { type: "audio/wav" });
};
export class BrowserPocketTTSRuntime {
#worker: Worker | null = null;
#audioContext: AudioContext | null = null;
#initializePromise: Promise<void> | null = null;
#readyDeferred: Deferred<void> | null = null;
#voiceDeferred: Deferred<void> | null = null;
#generation: PendingGeneration | null = null;
#modelBaseUrl = resolveModelBaseUrl();
#warmed = false;
#customVoiceReady = false;
#defaultVoiceName: string | null = null;
#threadCount = 1;
#retriedSingleThread = false;
#statusCallback: ((message: string) => void) | null = null;
#lastWorkerStatus = "Worker not started.";
async initialize(options: {
modelId: string;
onProgress?: (message: string) => void;
}): Promise<void> {
this.#modelBaseUrl = resolveModelBaseUrl(options.modelId);
this.#threadCount = this.#resolveInitialThreadCount();
this.#statusCallback = options.onProgress ?? null;
if (this.#initializePromise) {
return this.#initializePromise;
}
this.#initializePromise = this.#startWorkerWithFallback().catch((error) => {
this.#initializePromise = null;
throw error;
});
return this.#initializePromise;
}
async warmup(options: { onProgress?: (message: string) => void } = {}): Promise<void> {
await this.initialize({
modelId: this.#modelBaseUrl,
onProgress: options.onProgress,
});
if (this.#warmed) {
return;
}
options.onProgress?.("Compiling voice runtime...");
await this.#generateBlob("Benchmark.");
this.#warmed = true;
options.onProgress?.("Voice runtime warmed.");
}
async bootstrapFromUtterance(
audio: Blob,
): Promise<{ embeddingId?: string }> {
await this.initialize({ modelId: this.#modelBaseUrl });
if (!this.#worker) {
throw new Error("Pocket TTS worker failed to initialize.");
}
const audioData = await this.#decodeReferenceAudio(audio);
const deferred = createDeferred<void>();
this.#voiceDeferred = deferred;
this.#worker.postMessage(
{
type: "encode_voice",
data: {
audio: audioData,
},
},
[audioData.buffer],
);
await deferred.promise;
this.#customVoiceReady = true;
return {
embeddingId: `custom-${Date.now()}`,
};
}
async synthesize(options: {
text: string;
signal?: AbortSignal;
referenceAudio?: Blob;
}): Promise<Blob> {
await this.initialize({ modelId: this.#modelBaseUrl });
if (options.signal?.aborted) {
throw new Error("Speech synthesis was cancelled.");
}
if (!this.#customVoiceReady && options.referenceAudio) {
await this.bootstrapFromUtterance(options.referenceAudio);
}
if (!this.#customVoiceReady && !this.#defaultVoiceName) {
throw new Error("No Pocket TTS voice is ready.");
}
return this.#generateBlob(options.text, options.signal);
}
async stream(options: {
text: string;
signal?: AbortSignal;
referenceAudio?: Blob;
onAudioChunk: (chunk: Float32Array) => void | Promise<void>;
}): Promise<void> {
await this.initialize({ modelId: this.#modelBaseUrl });
if (options.signal?.aborted) {
throw new Error("Speech synthesis was cancelled.");
}
if (!this.#customVoiceReady && options.referenceAudio) {
await this.bootstrapFromUtterance(options.referenceAudio);
}
if (!this.#customVoiceReady && !this.#defaultVoiceName) {
throw new Error("No Pocket TTS voice is ready.");
}
await this.#generateStream(options.text, options.onAudioChunk, options.signal);
}
async #startWorkerWithFallback(): Promise<void> {
try {
await this.#startWorker();
} catch (error) {
if (this.#threadCount > 1 && !this.#retriedSingleThread) {
console.warn(
"Pocket TTS worker failed with multithreaded WASM. Retrying single-threaded.",
error,
);
this.#retriedSingleThread = true;
this.#threadCount = 1;
this.#resetWorkerState();
await this.#startWorker();
return;
}
throw error;
}
}
async #startWorker(): Promise<void> {
this.#readyDeferred = createDeferred<void>();
this.#lastWorkerStatus = `Starting worker (${this.#threadCount} thread${this.#threadCount === 1 ? "" : "s"})`;
console.info("[PocketRuntime] starting worker", {
modelBaseUrl: this.#modelBaseUrl,
threadCount: this.#threadCount,
crossOriginIsolated: globalThis.crossOriginIsolated === true,
});
this.#worker = new Worker(WORKER_MODULE_URL, { type: "module" });
this.#worker.onmessage = (event: MessageEvent<WorkerMessage>) => {
this.#handleWorkerMessage(event.data);
};
this.#worker.onmessageerror = (event) => {
console.error("[PocketRuntime] worker message error", event);
this.#rejectPending(
new Error(
`Pocket TTS worker message error. Last status: ${this.#lastWorkerStatus}.`,
),
);
};
this.#worker.onerror = (event) => {
const location = [event.filename, event.lineno, event.colno]
.filter(Boolean)
.join(":");
const error = new Error(
[
event.message || "Pocket TTS worker crashed.",
location ? `at ${location}` : null,
`last status: ${this.#lastWorkerStatus}`,
`thread count: ${this.#threadCount}`,
]
.filter(Boolean)
.join(" "),
);
console.error("[PocketRuntime] worker crashed", {
message: event.message,
filename: event.filename,
lineno: event.lineno,
colno: event.colno,
lastStatus: this.#lastWorkerStatus,
threadCount: this.#threadCount,
});
this.#rejectPending(error);
};
console.info("[PocketRuntime] posting load request", {
modelBaseUrl: this.#modelBaseUrl,
threadCount: this.#threadCount,
});
this.#worker.postMessage({
type: "load",
data: {
modelBaseUrl: this.#modelBaseUrl,
threadCount: this.#threadCount,
},
});
return this.#readyDeferred.promise;
}
#handleWorkerMessage(message: WorkerMessage): void {
if (message.type === "status" && message.status) {
this.#lastWorkerStatus = message.status;
console.info("[PocketRuntime] worker status", message.status, message.state);
}
switch (message.type) {
case "loaded":
this.#lastWorkerStatus = "Worker loaded";
console.info("[PocketRuntime] worker loaded");
this.#readyDeferred?.resolve();
this.#readyDeferred = null;
break;
case "voices_loaded":
this.#lastWorkerStatus = `Voices loaded (${message.defaultVoice ?? "none"})`;
this.#defaultVoiceName = message.defaultVoice ?? null;
if (this.#defaultVoiceName) {
this.#statusCallback?.(`Voice ready (${this.#defaultVoiceName}).`);
}
break;
case "voice_encoded":
this.#lastWorkerStatus = `Voice encoded (${message.voiceName ?? "custom"})`;
this.#voiceDeferred?.resolve();
this.#voiceDeferred = null;
this.#customVoiceReady = true;
this.#statusCallback?.("Voice profile encoded.");
break;
case "voice_set":
this.#lastWorkerStatus = `Voice selected (${message.voiceName ?? "unknown"})`;
if (message.voiceName === "custom") {
this.#customVoiceReady = true;
}
if (message.voiceName) {
this.#statusCallback?.(`Voice selected (${message.voiceName}).`);
}
break;
case "audio_chunk":
if (this.#generation?.chunks && message.data) {
const generation = this.#generation;
const audioChunk = Float32Array.from(message.data);
if (generation.mode === "blob") {
generation.chunks.push(audioChunk);
}
if (generation.onAudioChunk) {
generation.chunkTask = generation.chunkTask
.then(() => generation.onAudioChunk?.(audioChunk))
.catch((error) => {
this.#rejectPending(
error instanceof Error
? error
: new Error("Pocket TTS audio streaming failed."),
);
});
}
}
break;
case "stream_ended":
this.#lastWorkerStatus = "Audio stream ended";
if (this.#generation) {
const generation = this.#generation;
this.#generation = null;
generation.abortCleanup?.();
void generation.chunkTask.finally(() => {
if (generation.mode === "blob") {
generation.resolve(encodeWav(generation.chunks));
return;
}
generation.resolve();
});
}
break;
case "error":
this.#lastWorkerStatus = message.error || "Pocket TTS failed";
console.error("[PocketRuntime] worker error", message.error);
this.#rejectPending(new Error(message.error || "Pocket TTS failed."));
break;
case "status":
if (message.status) {
this.#statusCallback?.(message.status);
}
break;
default:
break;
}
}
#rejectPending(error: Error): void {
console.error("[PocketRuntime] rejecting pending work", error);
this.#readyDeferred?.reject(error);
this.#readyDeferred = null;
this.#voiceDeferred?.reject(error);
this.#voiceDeferred = null;
if (this.#generation) {
const generation = this.#generation;
this.#generation = null;
generation.abortCleanup?.();
generation.reject(error);
}
this.#worker?.terminate();
this.#worker = null;
}
async #generateBlob(text: string, signal?: AbortSignal): Promise<Blob> {
return new Promise<Blob>((resolve, reject) => {
this.#startGeneration({
text,
signal,
mode: "blob",
resolve,
reject,
});
});
}
async #generateStream(
text: string,
onAudioChunk: (chunk: Float32Array) => void | Promise<void>,
signal?: AbortSignal,
): Promise<void> {
return new Promise<void>((resolve, reject) => {
this.#startGeneration({
text,
signal,
mode: "stream",
resolve,
reject,
onAudioChunk,
});
});
}
#startGeneration(options: {
text: string;
signal?: AbortSignal;
mode: "blob" | "stream";
resolve: (value?: Blob) => void;
reject: (error: Error) => void;
onAudioChunk?: (chunk: Float32Array) => void | Promise<void>;
}): void {
if (!this.#worker) {
throw new Error("Pocket TTS worker failed to initialize.");
}
if (this.#generation) {
throw new Error("Pocket TTS only supports one generation at a time.");
}
const pending: PendingGeneration = {
mode: options.mode,
chunks: [],
chunkTask: Promise.resolve(),
onAudioChunk: options.onAudioChunk,
resolve: options.resolve,
reject: options.reject,
};
if (options.signal) {
const handleAbort = () => {
if (this.#generation !== pending) {
return;
}
this.#worker?.postMessage({ type: "stop" });
this.#generation = null;
pending.abortCleanup?.();
options.reject(new Error("Speech synthesis was cancelled."));
};
options.signal.addEventListener("abort", handleAbort, { once: true });
pending.abortCleanup = () => {
options.signal?.removeEventListener("abort", handleAbort);
};
}
this.#generation = pending;
this.#worker.postMessage({
type: "generate",
data: {
text: options.text,
},
});
}
async #decodeReferenceAudio(audio: Blob): Promise<Float32Array> {
const audioContext = await this.#getAudioContext();
const input = await audio.arrayBuffer();
const decoded = await audioContext.decodeAudioData(input.slice(0));
const mono = toMono(decoded);
const resampled = resample(mono, decoded.sampleRate, SAMPLE_RATE);
return resampled.slice(0, MAX_REFERENCE_SAMPLES);
}
async #getAudioContext(): Promise<AudioContext> {
if (this.#audioContext) {
return this.#audioContext;
}
this.#audioContext = new AudioContext({ sampleRate: SAMPLE_RATE });
return this.#audioContext;
}
#resetWorkerState(): void {
this.#worker?.terminate();
this.#worker = null;
this.#readyDeferred = null;
this.#voiceDeferred = null;
this.#generation = null;
this.#lastWorkerStatus = "Worker reset";
}
#resolveInitialThreadCount(): number {
if (typeof navigator === "undefined") {
return 1;
}
if (globalThis.crossOriginIsolated !== true) {
return 1;
}
// ORT's threaded WASM worker path is currently unstable here on Windows.
// Start single-threaded instead of crashing during benchmark and retrying.
if (isWindowsPlatform()) {
this.#statusCallback?.(
"Voice runtime is using single-threaded mode on Windows for compatibility.",
);
return 1;
}
return Math.min(navigator.hardwareConcurrency || 4, 8);
}
}