Download extractor.worker.js from solvi-ai/documents: direct link, hf CLI and curl.
- Browser
- Download file 13.1 kB
-
https://huggingface.co/spaces/solvi-ai/documents/resolve/main/extractor.worker.js
- Command line
-
hf download hf://spaces/solvi-ai/documents/extractor.worker.js
-
curl -L -o extractor.worker.js https://huggingface.co/spaces/solvi-ai/documents/resolve/main/extractor.worker.js
13.1 kB
| // Field extractor in a Web Worker: onnxruntime-web (WebGPU, WASM fallback) + a byte-level BPE tokenizer with offsets. | |
| // Reproduces solvi.extract_long.LongSpanExtractor.predict exactly: | |
| // windows "[CLS] description(<=64 tokens) [SEP] chunk [SEP]", room = max_len - len(desc) - 3, step = room - stride; | |
| // per window softmax of the start and end logits over the window; best span start <= end < start + max_span inside the | |
| // chunk by ps[s] * pe[e]; best over windows; leading whitespace stripped; present if score >= threshold. | |
| // Messages in: {type: "init", key, repo, onnx, backend?} {type: "extract", id, key, text, fields: [{name, desc}]} | |
| // {type: "cancel", id} {type: "probe", repo, onnx} | |
| // Messages out: progress / ready / error / field / window / done. | |
| import * as ort from "https://cdn.jsdelivr.net/npm/onnxruntime-web@1.30.0/dist/ort.webgpu.min.mjs"; | |
| import { ByteLevelBPE } from "./tokenizer.js"; | |
| const ORT_VERSION = "1.30.0"; | |
| const CACHE = "solvi-documents-v1"; | |
| ort.env.wasm.wasmPaths = `https://cdn.jsdelivr.net/npm/onnxruntime-web@${ORT_VERSION}/dist/`; | |
| ort.env.wasm.numThreads = self.crossOriginIsolated ? Math.min(8, Math.max(1, (navigator.hardwareConcurrency || 4) - 1)) : 1; | |
| ort.env.logLevel = "error"; | |
| const models = new Map(); // key -> {session, tok, cfg, backend, file} | |
| const cancelled = new Set(); | |
| let queue = Promise.resolve(); | |
| const post = (m) => self.postMessage(m); | |
| // ------------------------------------------------------------------------------------------------ downloads + cache | |
| async function openCache() { | |
| try { return await caches.open(CACHE); } catch { return null; } // no Cache API (e.g. some private modes): download each visit | |
| } | |
| async function readBody(resp, total, onProgress) { | |
| const reader = resp.body.getReader(); | |
| let buf = total > 0 ? new Uint8Array(total) : null; | |
| const parts = []; | |
| let loaded = 0, last = 0; | |
| for (;;) { | |
| const { done, value } = await reader.read(); | |
| if (done) break; | |
| if (buf && loaded + value.length <= buf.length) buf.set(value, loaded); | |
| else { if (buf) { parts.push(buf.subarray(0, loaded)); buf = null; } parts.push(value); } | |
| loaded += value.length; | |
| const now = performance.now(); | |
| if (onProgress && now - last > 120) { last = now; onProgress(loaded, total); } | |
| } | |
| if (onProgress) onProgress(loaded, total); | |
| if (buf) return loaded === buf.length ? buf : buf.subarray(0, loaded); | |
| const out = new Uint8Array(loaded); | |
| let o = 0; | |
| for (const p of parts) { out.set(p, o); o += p.length; } | |
| return out; | |
| } | |
| // -> {bytes, fromCache} | |
| async function fetchCached(url, onProgress) { | |
| const cache = await openCache(); | |
| if (cache) { | |
| const hit = await cache.match(url); | |
| if (hit) { | |
| const total = +hit.headers.get("x-solvi-size") || +hit.headers.get("content-length") || 0; | |
| const bytes = await readBody(hit, total, onProgress && ((l, t) => onProgress(l, t, true))); | |
| if (!total || bytes.length === total) return { bytes, fromCache: true }; | |
| await cache.delete(url); // truncated entry: download again | |
| } | |
| } | |
| const resp = await fetch(url, { mode: "cors" }); | |
| if (!resp.ok) throw new Error(`HTTP ${resp.status} for ${url}`); | |
| const total = +resp.headers.get("content-length") || 0; | |
| let body = resp; | |
| let put = null; | |
| if (cache && resp.body) { | |
| const [a, b] = resp.body.tee(); | |
| body = new Response(a); | |
| const headers = { "content-type": "application/octet-stream", "x-solvi-size": String(total) }; | |
| put = cache.put(url, new Response(b, { headers })).catch((e) => post({ type: "warn", message: "Could not cache " + url + ": " + e })); | |
| } | |
| const bytes = await readBody(body, total, onProgress && ((l, t) => onProgress(l, t, false))); | |
| if (put) await put; | |
| return { bytes, fromCache: false }; | |
| } | |
| async function fetchJson(url) { | |
| const { bytes } = await fetchCached(url); | |
| return JSON.parse(new TextDecoder().decode(bytes)); | |
| } | |
| // ------------------------------------------------------------------------------------------------------ init | |
| // -> "none" | "no-f16" | "ok". An fp16 graph needs the shader-f16 feature (missing e.g. on some Linux/Vulkan setups). | |
| async function webGPUState() { | |
| try { | |
| if (!self.navigator?.gpu) return "none"; | |
| const a = await navigator.gpu.requestAdapter({ powerPreference: "high-performance" }); | |
| if (!a) return "none"; | |
| return a.features.has("shader-f16") ? "ok" : "no-f16"; | |
| } catch { return "none"; } | |
| } | |
| async function init(msg) { | |
| const { key, repo, onnx } = msg; | |
| if (models.has(key) && models.get(key).file === onnx) { post({ type: "ready", key, ...models.get(key).info }); return; } | |
| const t0 = performance.now(); | |
| post({ type: "progress", key, phase: "files", loaded: 0, total: 0 }); | |
| const [cfg, tj] = await Promise.all([fetchJson(repo + "solvi_extract.json"), fetchJson(repo + "tokenizer.json")]); | |
| const tok = new ByteLevelBPE(tj); | |
| const gpuState = msg.backend === "wasm" ? "off" : await webGPUState(); | |
| const needsF16 = /fp16/.test(onnx); | |
| let backend = gpuState === "ok" || (gpuState === "no-f16" && !needsF16) ? "webgpu" : "wasm"; | |
| const t1 = performance.now(); | |
| const { bytes, fromCache } = await fetchCached(repo + onnx, (loaded, total, cached) => | |
| post({ type: "progress", key, phase: cached ? "cache" : "download", loaded, total })); | |
| const t2 = performance.now(); | |
| post({ type: "progress", key, phase: "session", backend, loaded: bytes.length, total: bytes.length }); | |
| // one model at a time: two ModernBERT-large sessions do not fit in one WebAssembly heap (4 GB) | |
| for (const [k, m] of models) { | |
| try { await m.session.release(); } catch { /* ignore */ } | |
| models.delete(k); | |
| if (k !== key) post({ type: "unloaded", key: k }); | |
| } | |
| const create = async (be) => { | |
| try { | |
| return await ort.InferenceSession.create(bytes, { executionProviders: [be], graphOptimizationLevel: "all" }); | |
| } catch (e) { | |
| // right after a reload the previous page's heap may not be freed yet: wait and try once more | |
| if (!/bad_alloc|out of memory|memory access out of bounds/i.test(String(e?.message || e))) throw e; | |
| await new Promise((r) => setTimeout(r, 4000)); | |
| return await ort.InferenceSession.create(bytes, { executionProviders: [be], graphOptimizationLevel: "all" }); | |
| } | |
| }; | |
| // create + warm-up on a short input (compiles kernels/shaders; fails early if an operator is missing on this backend) | |
| const start = async (be) => { | |
| const s = await create(be); | |
| const tw = performance.now(); | |
| await runWindow(s, [tok.cls, tok.sep, 100, tok.sep]); | |
| return { s, warm: performance.now() - tw }; | |
| }; | |
| let session, warmupMs, gpuFailed = null; | |
| try { | |
| ({ s: session, warm: warmupMs } = await start(backend)); | |
| } catch (e) { | |
| if (backend !== "webgpu") { | |
| const text = String(e?.message || e); | |
| const fp16 = /float16|fp16|MLFloat16|not implemented|Could not find an implementation/i.test(text); | |
| post({ type: "error", key, code: fp16 ? "fp16-wasm" : "session", backend, message: text.slice(0, 600) }); | |
| return; | |
| } | |
| gpuFailed = String(e?.message || e).slice(0, 300); // WebGPU could not run this graph: fall back to the CPU | |
| backend = "wasm"; | |
| post({ type: "progress", key, phase: "session", backend, loaded: bytes.length, total: bytes.length }); | |
| try { | |
| ({ s: session, warm: warmupMs } = await start(backend)); | |
| } catch (e2) { | |
| post({ type: "error", key, code: "session", backend, message: String(e2?.message || e2).slice(0, 600) }); | |
| return; | |
| } | |
| } | |
| const t4 = performance.now(); | |
| const t3 = t4 - warmupMs; | |
| const info = { backend, gpuState, gpuFailed, fromCache, threads: ort.env.wasm.numThreads, file: onnx, size: bytes.length, cfg, | |
| filesMs: t1 - t0, downloadMs: t2 - t1, sessionMs: t3 - t2, warmupMs: t4 - t3, totalMs: t4 - t0 }; | |
| models.set(key, { session, tok, cfg, backend, file: onnx, info, encCache: new Map() }); | |
| post({ type: "ready", key, ...info }); | |
| } | |
| // ------------------------------------------------------------------------------------------------------ predict | |
| async function runWindow(session, seq) { | |
| const n = seq.length; | |
| const ids = new BigInt64Array(n), att = new BigInt64Array(n); | |
| for (let i = 0; i < n; i++) { ids[i] = BigInt(seq[i]); att[i] = 1n; } | |
| const out = await session.run({ | |
| input_ids: new ort.Tensor("int64", ids, [1, n]), | |
| attention_mask: new ort.Tensor("int64", att, [1, n]), | |
| }); | |
| const t = out.logits ?? out[session.outputNames[0]]; | |
| const data = t.data; // Float32Array [1, n, 2] | |
| const res = data.slice ? data.slice(0, n * 2) : Float32Array.from(data); | |
| t.dispose?.(); | |
| return res; | |
| } | |
| function softmaxCol(lg, n, col) { | |
| let mx = -Infinity; | |
| for (let i = 0; i < n; i++) mx = Math.max(mx, lg[2 * i + col]); | |
| const p = new Float64Array(n); | |
| let s = 0; | |
| for (let i = 0; i < n; i++) { p[i] = Math.exp(lg[2 * i + col] - mx); s += p[i]; } | |
| for (let i = 0; i < n; i++) p[i] /= s; | |
| return p; | |
| } | |
| function encodeText(m, text) { | |
| let e = m.encCache.get(text); | |
| if (!e) { | |
| e = m.tok.encode(text); | |
| if (m.encCache.size > 16) m.encCache.clear(); | |
| m.encCache.set(text, e); | |
| } | |
| return e; | |
| } | |
| async function predict(m, text, desc, onWindow, isCancelled) { | |
| const { tok, cfg, session } = m; | |
| const enc = encodeText(m, text); | |
| const ids = enc.ids; | |
| const d = tok.encode(desc.normalize("NFC")).ids.slice(0, 64); | |
| const room = cfg.max_len - d.length - 3; | |
| const step = Math.max(1, room - cfg.stride); | |
| const nWin = ids.length <= room ? 1 : 1 + Math.ceil((ids.length - room) / step); | |
| let best = { s: 0, e: 0, score: -1 }, nul = 0, w = 0; | |
| for (let a = 0; a < Math.max(1, ids.length); a += step) { | |
| if (isCancelled()) return null; | |
| const chunk = ids.slice(a, a + room); | |
| if (chunk.length) { | |
| const seq = [tok.cls, ...d, tok.sep, ...chunk, tok.sep]; | |
| const n = seq.length; | |
| const lg = await runWindow(session, seq); | |
| const ps = softmaxCol(lg, n, 0), pe = softmaxCol(lg, n, 1); | |
| nul = Math.max(nul, ps[0] * pe[0]); | |
| const c0 = d.length + 2, c1 = c0 + chunk.length; | |
| let bs = -1, be = -1, bsc = -1; | |
| for (let s = c0; s < c1; s++) { | |
| const lim = Math.min(c1, s + cfg.max_span); | |
| const pss = ps[s]; | |
| for (let e = s; e < lim; e++) { | |
| const sc = pss * pe[e]; | |
| if (sc > bsc) { bsc = sc; bs = s; be = e; } | |
| } | |
| } | |
| if (bsc > best.score) best = { s: a + bs - c0, e: a + be - c0, score: bsc }; | |
| } | |
| w++; | |
| onWindow?.(w, nWin); | |
| if (a + room >= ids.length) break; | |
| } | |
| let st = 0, en = 0, st16 = 0, en16 = 0; | |
| if (best.score >= 0 && ids.length) { | |
| st = enc.offsets[best.s][0]; en = enc.offsets[best.e][1]; | |
| st16 = enc.offsets16[best.s][0]; en16 = enc.offsets16[best.e][1]; | |
| while (st < en && /\s/.test(text[st16])) { st++; st16++; } // BPE offsets include the leading space | |
| } | |
| return { start: st, end: en, start16: st16, end16: en16, score: best.score, nullScore: nul, windows: w, tokens: ids.length, | |
| value: text.slice(st16, en16) }; | |
| } | |
| async function extract(msg) { | |
| const { id, key, fields } = msg; | |
| const text = msg.text; | |
| const m = models.get(key); | |
| if (!m) { post({ type: "error", id, key, code: "not-ready", message: "model not loaded" }); return; } | |
| const t0 = performance.now(); | |
| for (const f of fields) { | |
| if (cancelled.has(id)) break; | |
| const t = performance.now(); | |
| const r = await predict(m, text, f.desc, (w, n) => post({ type: "window", id, name: f.name, w, n }), () => cancelled.has(id)); | |
| if (!r) break; | |
| const thr = m.cfg.thr?.[f.name] ?? m.cfg.thr_default; | |
| post({ type: "field", id, key, name: f.name, desc: f.desc, result: { ...r, threshold: thr, present: r.score >= thr }, | |
| ms: performance.now() - t }); | |
| } | |
| post({ type: "done", id, key, cancelled: cancelled.has(id), ms: performance.now() - t0 }); | |
| cancelled.delete(id); | |
| } | |
| async function probe(msg) { // does a file exist? (HEAD, follows redirects) | |
| try { | |
| const r = await fetch(msg.repo + msg.onnx, { method: "HEAD" }); | |
| post({ type: "probe", id: msg.id, ok: r.ok, size: +r.headers.get("content-length") || +r.headers.get("x-linked-size") || 0 }); | |
| } catch { post({ type: "probe", id: msg.id, ok: false }); } | |
| } | |
| self.onmessage = (ev) => { | |
| const msg = ev.data; | |
| if (msg.type === "cancel") { cancelled.add(msg.id); return; } | |
| if (msg.type === "probe") { probe(msg); return; } | |
| if (msg.type === "cached") { | |
| (async () => { | |
| const c = await openCache(); | |
| const hit = c ? !!(await c.match(msg.url)) : false; | |
| post({ type: "cached", id: msg.id, hit }); | |
| })(); | |
| return; | |
| } | |
| queue = queue.then(async () => { | |
| try { | |
| if (msg.type === "init") await init(msg); | |
| else if (msg.type === "extract") await extract(msg); | |
| } catch (e) { | |
| post({ type: "error", id: msg.id, key: msg.key, code: "exception", message: String(e?.stack || e).slice(0, 800) }); | |
| } | |
| }); | |
| }; | |
| post({ type: "hello", ort: ORT_VERSION, crossOriginIsolated: self.crossOriginIsolated, threads: ort.env.wasm.numThreads }); | |