QuoteSearch / worker.js
ruidiao's picture
fix: bump IndexedDB version to invalidate old 256-dim cache
1886487
Raw History Blame Contribute Delete
16.7 kB
const EMBEDDING_DIM = 128;
const INDEX_PATH = 'data/quotes_index.bin';
const WEIGHTS_PATH = 'data/potion_weights.bin';
const TOKENIZER_PATH = 'data/potion_tokenizer.json';
const APPROX_MODEL_SIZE_MB = 8;
const DB_NAME = 'QuoteSearchDB';
const DB_VERSION = 2;
const STORE_NAME = 'quoteIndex';
let weights;
let wholeWordTrie;
let subwordTrie;
let unkId;
let indexData;
function openDB() {
return new Promise((resolve, reject) => {
const request = indexedDB.open(DB_NAME, DB_VERSION);
request.onupgradeneeded = (event) => {
const db = event.target.result;
if (db.objectStoreNames.contains(STORE_NAME)) {
db.deleteObjectStore(STORE_NAME);
}
db.createObjectStore(STORE_NAME, { keyPath: 'id' });
};
request.onsuccess = (event) => resolve(event.target.result);
request.onerror = (event) => reject('IndexedDB error: ' + event.target.errorCode);
});
}
async function getFromDB(key) {
const db = await openDB();
return new Promise((resolve, reject) => {
const transaction = db.transaction([STORE_NAME], 'readonly');
const request = transaction.objectStore(STORE_NAME).get(key);
request.onsuccess = () => resolve(request.result ? request.result.value : null);
request.onerror = () => reject('Error getting data from DB');
});
}
async function putInDB(key, value) {
const db = await openDB();
return new Promise((resolve, reject) => {
const transaction = db.transaction([STORE_NAME], 'readwrite');
const request = transaction.objectStore(STORE_NAME).put({ id: key, value: value });
request.onsuccess = () => resolve();
request.onerror = () => reject('Error putting data in DB');
});
}
async function deleteFromDB(key) {
const db = await openDB();
return new Promise((resolve, reject) => {
const transaction = db.transaction([STORE_NAME], 'readwrite');
const request = transaction.objectStore(STORE_NAME).delete(key);
request.onsuccess = () => resolve();
request.onerror = () => reject('Error deleting data from DB');
});
}
function buildTrie(entries) {
const root = {};
for (const { token, id } of entries) {
let node = root;
for (const ch of token) {
if (!node[ch]) node[ch] = {};
node = node[ch];
}
node._id = id;
}
return root;
}
function longestPrefix(trie, str, start) {
let node = trie;
let matchId = null;
let matchLen = 0;
for (let i = start; i < str.length; i++) {
const ch = str[i];
if (!node[ch]) break;
node = node[ch];
const len = i - start + 1;
if (node._id !== undefined) {
matchId = node._id;
matchLen = len;
}
}
return { id: matchId, length: matchLen };
}
function normalize(text) {
let t = text;
t = t.replace(/[\x00-\x08\x0b\x0c\x0e-\x1f\x7f]/g, '');
t = t.replace(/([\u3400-\u4dbf\uf900-\ufaff\u4e00-\u9fff])/g, ' $1 ');
t = t.normalize('NFKD').replace(/[\u0300-\u036f]/g, '');
t = t.toLowerCase();
return t;
}
function tokenizeWord(word) {
if (!word) return [];
const tokens = [];
let pos = 0;
let match = longestPrefix(wholeWordTrie, word, pos);
if (!match.id) return [unkId];
tokens.push(match.id);
pos += match.length;
while (pos < word.length) {
match = longestPrefix(subwordTrie, word, pos);
if (!match.id) { tokens.push(unkId); break; }
tokens.push(match.id);
pos += match.length;
}
return tokens;
}
function tokenize(text) {
text = normalize(text);
const words = text.match(/[^\s]+/g) || [];
const ids = [];
for (const w of words) ids.push(...tokenizeWord(w));
return ids;
}
function halfToFloat(h) {
const s = (h >> 15) & 1 ? -1 : 1;
const e = (h >> 10) & 31;
const m = h & 1023;
if (e === 0) return s * Math.pow(2, -14) * (m / 1024);
if (e === 31) return m ? NaN : s * Infinity;
return s * Math.pow(2, e - 15) * (1 + m / 1024);
}
function computeEmbedding(text) {
const ids = tokenize(text);
const dim = EMBEDDING_DIM;
const emb = new Float32Array(dim);
for (const id of ids) {
const offset = id * dim;
for (let i = 0; i < dim; i++) emb[i] += weights[offset + i];
}
if (ids.length > 0) for (let i = 0; i < dim; i++) emb[i] /= ids.length;
let norm = 0;
for (let i = 0; i < dim; i++) norm += emb[i] * emb[i];
norm = Math.sqrt(norm);
if (norm > 0) for (let i = 0; i < dim; i++) emb[i] /= norm;
return emb;
}
async function loadModel() {
self.postMessage({ type: 'loading', payload: 'Downloading model weights...' });
const resp = await fetch(WEIGHTS_PATH);
const contentLength = resp.headers.get('Content-Length');
const total = parseInt(contentLength, 10);
let loaded = 0;
const reader = resp.body.getReader();
const chunks = [];
while (true) {
const { done, value } = await reader.read();
if (done) break;
chunks.push(value);
loaded += value.length;
self.postMessage({
type: 'progress',
payload: {
status: 'Downloading weights',
progress: (loaded / total) * 100,
file: 'potion_weights.bin',
detail: `Downloading model: ${Math.floor((loaded / total) * 100)}% (${(loaded / (1024 * 1024)).toFixed(2)}MB / ${(total / (1024 * 1024)).toFixed(2)}MB)`
}
});
}
const buffer = await new Response(new Blob(chunks)).arrayBuffer();
const uint16view = new Uint16Array(buffer);
weights = new Float32Array(uint16view.length);
for (let i = 0; i < uint16view.length; i++) {
weights[i] = halfToFloat(uint16view[i]);
}
self.postMessage({ type: 'loading', payload: 'Loading tokenizer...' });
const tokResp = await fetch(TOKENIZER_PATH);
const tokData = await tokResp.json();
unkId = tokData.unk_id;
const wholeWordEntries = [];
const subwordEntries = [];
for (const token in tokData.vocab) {
const id = tokData.vocab[token];
if (token.startsWith('##')) {
subwordEntries.push({ token: token.slice(2), id });
} else {
wholeWordEntries.push({ token, id });
}
}
wholeWordTrie = buildTrie(wholeWordEntries);
subwordTrie = buildTrie(subwordEntries);
}
async function loadIndex() {
try {
self.postMessage({ type: 'loading', payload: 'Checking for cached index file...' });
const cachedIndex = await getFromDB('quoteIndexData');
if (cachedIndex) {
indexData = cachedIndex;
self.postMessage({ type: 'loading', payload: 'Index loaded from cache.' });
} else {
self.postMessage({ type: 'loading', payload: 'Downloading index file (this may take a while)...' });
const response = await fetch(INDEX_PATH);
const contentLength = response.headers.get('Content-Length');
const total = parseInt(contentLength, 10);
let loaded = 0;
const reader = response.body.getReader();
const chunks = [];
while (true) {
const { done, value } = await reader.read();
if (done) break;
chunks.push(value);
loaded += value.length;
self.postMessage({
type: 'progress',
payload: {
status: 'Downloading Index',
progress: (loaded / total) * 100,
file: INDEX_PATH,
detail: `Downloading index file: ${Math.floor((loaded / total) * 100)}% (${(loaded / (1024 * 1024)).toFixed(2)}MB / ${(total / (1024 * 1024)).toFixed(2)}MB)`
}
});
}
const buffer = await new Response(new Blob(chunks)).arrayBuffer();
let offset = 0;
const numQuotes = new Uint32Array(buffer.slice(offset, offset + 4))[0];
offset += 4;
const embeddingDim = new Uint16Array(buffer.slice(offset, offset + 2))[0];
offset += 2;
const scale = new Float32Array(buffer.slice(offset, offset + 4))[0];
offset += 4;
const metadataSize = new Uint32Array(buffer.slice(offset, offset + 4))[0];
offset += 4;
let metadataFormat = 0;
if (offset + 1 <= buffer.byteLength) {
metadataFormat = new Uint8Array(buffer.slice(offset, offset + 1))[0];
offset += 1;
}
let metadataBytes = buffer.slice(offset, offset + metadataSize);
offset += metadataSize;
async function decodeMetadata(bytes, format) {
const decoder = new TextDecoder('utf-8');
if (format === 0) {
return JSON.parse(decoder.decode(bytes));
} else if (format === 1) {
if (typeof DecompressionStream !== 'undefined') {
const ds = new DecompressionStream('gzip');
const decompressed = await new Response(new Blob([bytes]).stream().pipeThrough(ds)).arrayBuffer();
return JSON.parse(decoder.decode(decompressed));
} else {
throw new Error('Gzip decompression not available.');
}
} else {
throw new Error('Unknown metadata format: ' + format);
}
}
const metadata = await decodeMetadata(metadataBytes, metadataFormat);
const quantizedEmbeddings = new Int8Array(buffer.slice(offset));
const embeddings = new Float32Array(quantizedEmbeddings.length);
const totalValues = quantizedEmbeddings.length;
const updateInterval = Math.floor(totalValues / 100);
for (let i = 0; i < totalValues; i++) {
embeddings[i] = quantizedEmbeddings[i] / scale;
if (updateInterval > 0 && i % updateInterval === 0) {
self.postMessage({
type: 'progress',
payload: {
status: 'Processing index (de-quantizing)',
progress: (i / totalValues) * 100,
file: INDEX_PATH,
detail: `Processing index: ${Math.floor((i / totalValues) * 100)}%`
}
});
}
}
indexData = {
metadata,
embeddings: reshape(embeddings, [numQuotes, embeddingDim]),
embeddingsByteLength: quantizedEmbeddings.byteLength
};
await putInDB('quoteIndexData', indexData);
}
} catch (error) {
console.error('Error loading index:', error);
self.postMessage({ type: 'error', payload: error.message });
throw error;
}
}
self.onmessage = async (event) => {
const { type, payload } = event.data;
if (type === 'search') {
if (!indexData) {
self.postMessage({ type: 'loading', payload: 'Downloading index before running your search...' });
await loadIndex();
}
if (!weights) {
self.postMessage({ type: 'loading', payload: 'Downloading model before running your search...' });
await loadModel();
}
self.postMessage({ type: 'loading', payload: 'Searching...' });
const results = await search(payload);
self.postMessage({ type: 'results', payload: results });
} else if (type === 'deleteData') {
self.postMessage({ type: 'loading', payload: 'Deleting cached data...' });
const cleanup = async () => {
try { await deleteFromDB('quoteIndexData'); } catch (e) { }
await new Promise((resolve, reject) => {
const deleteRequest = indexedDB.deleteDatabase(DB_NAME);
deleteRequest.onsuccess = () => resolve();
deleteRequest.onerror = () => reject(deleteRequest.error || new Error('Failed to delete IndexedDB'));
deleteRequest.onblocked = () => resolve();
});
try {
if (typeof caches !== 'undefined' && caches.keys) {
const cacheNames = await caches.keys();
for (const cacheName of cacheNames) {
if (cacheName.startsWith('transformers-cache') || cacheName.includes('nomic')) {
await caches.delete(cacheName);
}
}
}
} catch (e) { console.warn('Cache cleanup error', e); }
try {
if (typeof localStorage !== 'undefined') {
const keysToClear = [];
for (let i = 0; i < localStorage.length; i++) {
const key = localStorage.key(i);
if (!key) continue;
if (key.startsWith('transformers') || key.includes('nomic') || key.includes('hf_')) {
keysToClear.push(key);
}
}
for (const k of keysToClear) localStorage.removeItem(k);
}
} catch (e) { }
weights = null;
wholeWordTrie = null;
subwordTrie = null;
unkId = null;
indexData = null;
};
const TIMEOUT_MS = 8000;
try {
await Promise.race([
cleanup(),
new Promise((_, reject) => setTimeout(() => reject(new Error('cleanup-timeout')), TIMEOUT_MS))
]);
self.postMessage({ type: 'dataDeleted', payload: 'Cached data deleted successfully.' });
} catch (error) {
if (error && error.message === 'cleanup-timeout') {
console.warn('deleteData: cleanup timed out');
self.postMessage({ type: 'dataDeleted', payload: 'Cached data deletion attempted (timed out).' });
} else {
console.error('deleteData error:', error);
self.postMessage({ type: 'error', payload: 'Failed to delete cached data: ' + (error && error.message ? error.message : String(error)) });
}
}
} else if (type === 'getIndexSize') {
try {
let totalSize = 0;
let indexCached = false;
let modelCached = false;
const cachedIndex = await getFromDB('quoteIndexData');
if (cachedIndex) {
totalSize += (cachedIndex.metadata ? JSON.stringify(cachedIndex.metadata).length : 0) + (cachedIndex.embeddingsByteLength || 0);
indexCached = true;
} else {
const response = await fetch(INDEX_PATH, { method: 'HEAD' });
const cl = response.headers.get('Content-Length');
totalSize += parseInt(cl, 10);
}
totalSize += APPROX_MODEL_SIZE_MB * 1024 * 1024;
if (weights) modelCached = true;
self.postMessage({ type: 'indexSize', payload: { size: totalSize, indexCached, modelCached } });
} catch (error) {
self.postMessage({ type: 'error', payload: 'Failed to get index size: ' + error.message });
}
}
};
async function search(query) {
if (!weights || !indexData) return [];
const queryEmbedding = computeEmbedding(query);
const similarities = [];
for (let i = 0; i < indexData.embeddings.length; i++) {
similarities.push({ index: i, similarity: cosineSimilarity(queryEmbedding, indexData.embeddings[i]) });
}
similarities.sort((a, b) => b.similarity - a.similarity);
return similarities.slice(0, 30).map(item => indexData.metadata[item.index]);
}
function cosineSimilarity(vecA, vecB) {
let dotProduct = 0;
let normA = 0;
let normB = 0;
for (let i = 0; i < vecA.length; i++) {
dotProduct += vecA[i] * vecB[i];
normA += vecA[i] * vecA[i];
normB += vecB[i] * vecB[i];
}
return dotProduct / (Math.sqrt(normA) * Math.sqrt(normB));
}
function reshape(array, shape) {
const reshaped = [];
let offset = 0;
for (let i = 0; i < shape[0]; i++) {
reshaped.push(array.slice(offset, offset + shape[1]));
offset += shape[1];
}
return reshaped;
}