documents / extractor.worker.js
mxkuzn's picture
solvi documents: cited answers from documents, model and decisions in your browser
1499ec4 verified
Raw History Blame Contribute Delete
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 });