Spaces:
Running
Running
Download pnc-worker.js from RyeAI/ekko-tiny-browser: direct link, hf CLI and curl.
- Browser
- Download file 12.1 kB
-
https://huggingface.co/spaces/RyeAI/ekko-tiny-browser/resolve/main/pnc-worker.js
- Command line
-
hf download hf://spaces/RyeAI/ekko-tiny-browser/pnc-worker.js
-
curl -L -o pnc-worker.js https://huggingface.co/spaces/RyeAI/ekko-tiny-browser/resolve/main/pnc-worker.js
12.1 kB
| /* global ort */ | |
| function signedAsset(path) { | |
| const base = new URL(self.location.href); | |
| const url = new URL(path, base); | |
| const signature = base.searchParams.get("__sign"); | |
| if (signature && url.origin === base.origin) url.searchParams.set("__sign", signature); | |
| return url.href; | |
| } | |
| importScripts(signedAsset("ort.wasm.min.js")); | |
| ort.env.wasm.numThreads = 1; | |
| ort.env.wasm.proxy = false; | |
| ort.env.wasm.wasmPaths = { | |
| mjs: signedAsset("ort-wasm-simd-threaded.mjs"), | |
| wasm: signedAsset("ort-wasm-simd-threaded.wasm"), | |
| }; | |
| let initialization = null; | |
| let session = null; | |
| let vocabulary = null; | |
| let labels = null; | |
| let maximumLength = 128; | |
| const INFERENCE_WINDOW_TOKENS = 128; | |
| let pendingRequest = null; | |
| let processing = false; | |
| async function sha256(bytes) { | |
| const digest = await crypto.subtle.digest("SHA-256", bytes); | |
| return Array.from(new Uint8Array(digest), (value) => value.toString(16).padStart(2, "0")).join(""); | |
| } | |
| async function fetchVerified(file, completedBytes, totalBytes) { | |
| const url = signedAsset(file.url); | |
| let cache = null; | |
| try { cache = await self.caches?.open("ekko-browser-models-v1"); } catch { /* Storage is optional. */ } | |
| for (let attempt = 0; attempt < 2; attempt += 1) { | |
| let cached = null; | |
| if (!attempt && cache) { | |
| try { cached = await cache.match(url); } catch { /* Use the network. */ } | |
| } | |
| const response = cached || await fetch(url, {cache: attempt ? "reload" : "force-cache"}); | |
| if (!response.ok || !response.body) { | |
| throw new Error(`Kunne ikke hente ${file.path} (${response.status})`); | |
| } | |
| const reader = response.body.getReader(); | |
| const bytes = new Uint8Array(file.bytes); | |
| let loaded = 0; | |
| while (true) { | |
| const {done, value} = await reader.read(); | |
| if (done) break; | |
| if (loaded + value.byteLength > bytes.length) { | |
| await reader.cancel(); | |
| loaded = bytes.length + 1; | |
| break; | |
| } | |
| bytes.set(value, loaded); | |
| loaded += value.byteLength; | |
| self.postMessage({ | |
| type: "pnc-progress", | |
| cached: Boolean(cached), | |
| file: file.path, | |
| loaded, | |
| fileBytes: file.bytes, | |
| overall: Math.min(1, (completedBytes + loaded) / totalBytes), | |
| }); | |
| } | |
| if (loaded === file.bytes && await sha256(bytes) === file.sha256) { | |
| if (!cached && cache) { | |
| try { await cache.put(url, new Response(bytes)); } catch { /* Keep using the verified bytes. */ } | |
| } | |
| return bytes; | |
| } | |
| if (cached) { | |
| try { await cache.delete(url); } catch { /* Retry without this entry. */ } | |
| } | |
| if (attempt) throw new Error(`${file.path} fejlede størrelses- eller SHA-256-kontrol`); | |
| } | |
| } | |
| async function initialize(configUrl) { | |
| if (session) return; | |
| const response = await fetch(signedAsset(configUrl), {cache: "no-cache"}); | |
| if (!response.ok) throw new Error("Kunne ikke hente PnC-konfigurationen"); | |
| const rootConfig = await response.json(); | |
| const config = rootConfig.pnc; | |
| if (!config?.files?.length) throw new Error("PnC mangler i modelkonfigurationen"); | |
| const totalBytes = config.files.reduce((sum, file) => sum + file.bytes, 0); | |
| let completedBytes = 0; | |
| const downloaded = new Map(); | |
| for (const file of config.files) { | |
| const bytes = await fetchVerified(file, completedBytes, totalBytes); | |
| downloaded.set(file.path, bytes); | |
| completedBytes += bytes.byteLength; | |
| } | |
| const tokenizer = JSON.parse(new TextDecoder().decode(downloaded.get("tokenizer.json"))); | |
| const modelConfig = JSON.parse(new TextDecoder().decode(downloaded.get("config.json"))); | |
| vocabulary = tokenizer.model.vocab; | |
| labels = modelConfig.id2label; | |
| // Shorter context restores sentence stops in long, continuous dictation. | |
| maximumLength = Math.min(config.maxLength || 128, INFERENCE_WINDOW_TOKENS); | |
| session = await ort.InferenceSession.create(downloaded.get("pnc.int8.onnx"), { | |
| executionProviders: ["wasm"], | |
| graphOptimizationLevel: "all", | |
| }); | |
| self.postMessage({type: "pnc-ready", bytes: totalBytes, model: config.name}); | |
| } | |
| function splitBasicToken(word) { | |
| const parts = []; | |
| let current = ""; | |
| const flush = () => { | |
| if (current) parts.push(current); | |
| current = ""; | |
| }; | |
| for (const character of Array.from(word)) { | |
| if (/\s/u.test(character)) { | |
| flush(); | |
| } else if (/[\p{P}\p{S}]/u.test(character)) { | |
| flush(); | |
| parts.push(character); | |
| } else if (!/[\p{Cc}\p{Cf}]/u.test(character)) { | |
| current += character; | |
| } | |
| } | |
| flush(); | |
| return parts; | |
| } | |
| function wordPiece(token) { | |
| const characters = Array.from(token); | |
| if (!characters.length) return []; | |
| if (characters.length > 100) return [vocabulary["[UNK]"]]; | |
| const pieces = []; | |
| let start = 0; | |
| while (start < characters.length) { | |
| let end = characters.length; | |
| let identifier; | |
| while (start < end) { | |
| const candidate = `${start ? "##" : ""}${characters.slice(start, end).join("")}`; | |
| if (Object.hasOwn(vocabulary, candidate)) { | |
| identifier = vocabulary[candidate]; | |
| break; | |
| } | |
| end -= 1; | |
| } | |
| if (identifier === undefined) return [vocabulary["[UNK]"]]; | |
| pieces.push(identifier); | |
| start = end; | |
| } | |
| return pieces; | |
| } | |
| function encodeWords(words) { | |
| const identifiers = [vocabulary["[CLS]"]]; | |
| const wordIds = [null]; | |
| words.forEach((word, wordIndex) => { | |
| const basicTokens = splitBasicToken(word); | |
| const pieces = basicTokens.length | |
| ? basicTokens.flatMap((token) => wordPiece(token)) | |
| : [vocabulary["[UNK]"]]; | |
| for (const identifier of pieces) { | |
| identifiers.push(identifier); | |
| wordIds.push(wordIndex); | |
| } | |
| }); | |
| identifiers.push(vocabulary["[SEP]"]); | |
| wordIds.push(null); | |
| return {identifiers, wordIds}; | |
| } | |
| function windowTokenIndex(words) { | |
| const {wordIds} = encodeWords(words); | |
| const counts = new Array(words.length).fill(0); | |
| let specialTokens = 0; | |
| for (const wordId of wordIds) { | |
| if (wordId === null) specialTokens += 1; | |
| else counts[wordId] += 1; | |
| } | |
| const prefix = [0]; | |
| for (const count of counts) prefix.push(prefix.at(-1) + count); | |
| return {prefix, specialTokens}; | |
| } | |
| function maximumWindowEnd(words, start, prefix, specialTokens) { | |
| const tokenBudget = prefix[start] + maximumLength - specialTokens; | |
| let low = start + 1; | |
| let high = Math.min(words.length, start + maximumLength); | |
| while (low < high) { | |
| const middle = Math.floor((low + high + 1) / 2); | |
| if (prefix[middle] <= tokenBudget) { | |
| low = middle; | |
| } else { | |
| high = middle - 1; | |
| } | |
| } | |
| return low; | |
| } | |
| function capitalize(word) { | |
| const characters = Array.from(word); | |
| if (!characters.length) return word; | |
| return characters[0].toLocaleUpperCase("da-DK") | |
| + characters.slice(1).join("").toLocaleLowerCase("da-DK"); | |
| } | |
| function capitalizeInitial(word) { | |
| const characters = Array.from(word); | |
| const index = characters.findIndex((character) => /\p{L}/u.test(character)); | |
| if (index < 0) return word; | |
| characters[index] = characters[index].toLocaleUpperCase("da-DK"); | |
| return characters.join(""); | |
| } | |
| function stabilizeRenderedWords(words) { | |
| if (!words.length) return []; | |
| const output = [...words]; | |
| output[0] = capitalizeInitial(output[0]); | |
| for (let index = 0; index < output.length - 1; index += 1) { | |
| if (!/[.!?]+$/u.test(output[index])) continue; | |
| if (/\p{L}\.\p{L}/u.test(output[index])) continue; // f.eks. | |
| if (/\p{Lu}/u.test(output[index + 1])) continue; // CPR, iPhone | |
| output[index + 1] = capitalizeInitial(output[index + 1]); | |
| } | |
| return output; | |
| } | |
| async function predictWindow(words) { | |
| const {identifiers, wordIds} = encodeWords(words); | |
| const shape = [1, identifiers.length]; | |
| const inputIds = BigInt64Array.from(identifiers, (value) => BigInt(value)); | |
| const attentionMask = new BigInt64Array(identifiers.length).fill(1n); | |
| const output = await session.run({ | |
| input_ids: new ort.Tensor("int64", inputIds, shape), | |
| attention_mask: new ort.Tensor("int64", attentionMask, shape), | |
| }); | |
| const logits = output.logits; | |
| const labelCount = logits.dims.at(-1); | |
| const rendered = [...words]; | |
| const seen = new Set(); | |
| wordIds.forEach((wordIndex, tokenIndex) => { | |
| if (wordIndex === null || seen.has(wordIndex)) return; | |
| seen.add(wordIndex); | |
| let prediction = 0; | |
| let best = -Infinity; | |
| for (let labelIndex = 0; labelIndex < labelCount; labelIndex += 1) { | |
| const value = logits.data[tokenIndex * labelCount + labelIndex]; | |
| if (value > best) { | |
| best = value; | |
| prediction = labelIndex; | |
| } | |
| } | |
| const label = labels[String(prediction)]; | |
| const upper = label.endsWith("|U"); | |
| const mark = upper ? label.slice(0, -2) : label; | |
| rendered[wordIndex] = `${upper ? capitalize(words[wordIndex]) : words[wordIndex]}${mark === "O" ? "" : mark}`; | |
| }); | |
| return rendered; | |
| } | |
| async function punctuateWords(words) { | |
| if (!words.length) return []; | |
| const {prefix, specialTokens} = windowTokenIndex(words); | |
| const output = []; | |
| let nextWord = 0; | |
| const overlapWords = 12; | |
| while (nextWord < words.length) { | |
| let contextStart = Math.max(0, nextWord - overlapWords); | |
| let windowEnd = maximumWindowEnd(words, contextStart, prefix, specialTokens); | |
| if (windowEnd <= nextWord) { | |
| contextStart = nextWord; | |
| windowEnd = maximumWindowEnd(words, contextStart, prefix, specialTokens); | |
| } | |
| const rendered = await predictWindow(words.slice(contextStart, windowEnd)); | |
| const emitEnd = windowEnd === words.length | |
| ? windowEnd | |
| : Math.max(nextWord + 1, windowEnd - overlapWords); | |
| output.push(...rendered.slice(nextWord - contextStart, emitEnd - contextStart)); | |
| nextWord = emitEnd; | |
| } | |
| if (output.length !== words.length) throw new Error("PnC ændrede antallet af ord"); | |
| return output; | |
| } | |
| async function punctuateSegments(segments, contextWords = [], capitalizeFirst = false) { | |
| const counts = []; | |
| const sourceWords = [...contextWords]; | |
| segments.forEach((segment) => { | |
| const words = String(segment.text || "").trim().split(/\s+/).filter(Boolean); | |
| counts.push(words.length); | |
| sourceWords.push(...words); | |
| }); | |
| const renderedWords = stabilizeRenderedWords(await punctuateWords(sourceWords)); | |
| if (capitalizeFirst && renderedWords.length > contextWords.length) { | |
| renderedWords[contextWords.length] = capitalizeInitial(renderedWords[contextWords.length]); | |
| } | |
| let offset = contextWords.length; | |
| return segments.map((segment, segmentIndex) => { | |
| const start = offset; | |
| const rendered = renderedWords.slice(start, start + counts[segmentIndex]); | |
| offset += counts[segmentIndex]; | |
| const timedWords = Array.isArray(segment.words) && segment.words.length === rendered.length | |
| && segment.words.every((word, index) => word.text === sourceWords[start + index]); | |
| const words = timedWords | |
| ? segment.words.map((word, wordIndex) => ({...word, text: rendered[wordIndex]})) | |
| : []; | |
| return { | |
| ...segment, | |
| rawText: segment.rawText || segment.text, | |
| text: rendered.join(" "), | |
| words, | |
| }; | |
| }); | |
| } | |
| async function drainRequests() { | |
| if (processing) return; | |
| processing = true; | |
| try { | |
| while (pendingRequest) { | |
| const request = pendingRequest; | |
| pendingRequest = null; | |
| await initialization; | |
| const segments = await punctuateSegments( | |
| request.segments, request.contextWords, request.capitalizeFirst, | |
| ); | |
| self.postMessage({type: "pnc-result", requestId: request.requestId, segments}); | |
| } | |
| } catch (error) { | |
| self.postMessage({ | |
| type: "pnc-error", | |
| message: error instanceof Error ? error.message : String(error), | |
| }); | |
| } finally { | |
| processing = false; | |
| } | |
| } | |
| self.onmessage = (event) => { | |
| const message = event.data || {}; | |
| if (message.type === "init") { | |
| initialization ||= initialize(message.configUrl || "model-config.json"); | |
| initialization.catch((error) => self.postMessage({ | |
| type: "pnc-error", | |
| message: error instanceof Error ? error.message : String(error), | |
| })); | |
| } else if (message.type === "punctuate") { | |
| initialization ||= initialize(message.configUrl || "model-config.json"); | |
| pendingRequest = message; | |
| drainRequests(); | |
| } | |
| }; | |