diff --git a/chrome/content/zotero/xpcom/embeddings.js b/chrome/content/zotero/xpcom/embeddings.js index 61154a43fb..9227842d31 100644 --- a/chrome/content/zotero/xpcom/embeddings.js +++ b/chrome/content/zotero/xpcom/embeddings.js @@ -763,6 +763,54 @@ Zotero.Embeddings = new function () { (score - minScore) / (maxDisplayScore - minScore))); }; + /** + * The item fields whose text is embedded, and so the fields a query can + * match literally. + * + * @return {Number[]} + */ + this.getIndexedFieldIDs = function () { + return [...new Set([ + Zotero.ItemFields.getID('title'), + Zotero.ItemFields.getID('abstractNote'), + ...Zotero.ItemFields.getTypeFieldsFromBase('title') + ])]; + }; + + // Items among the given ones whose indexed text contains every word of the + // query, matched as substrings the way quick search matches them. + async function _findLiteralMatches(queryText, itemIDs) { + let terms = Zotero.Embeddings.normalizeQuery(queryText).split(/\s+/).filter(Boolean); + if (!terms.length || !itemIDs.length) { + return new Set(); + } + let fieldIDs = Zotero.Embeddings.getIndexedFieldIDs(); + let matched = null; + for (let term of terms) { + let ids = new Set(); + let chunkSize = 500; + for (let i = 0; i < itemIDs.length; i += chunkSize) { + let chunk = itemIDs.slice(i, i + chunkSize); + let rows = await Zotero.DB.columnQueryAsync( + "SELECT DISTINCT itemID FROM itemData " + + "JOIN itemDataValues USING (valueID) " + + "WHERE fieldID IN (" + fieldIDs.join(',') + ") " + + "AND itemID IN (" + chunk.map(() => '?').join(',') + ") " + + "AND value LIKE ? ESCAPE '\\'", + [...chunk, '%' + term.replace(/[\\%_]/g, '\\$&') + '%'] + ); + for (let id of rows || []) { + ids.add(id); + } + } + matched = matched ? new Set([...matched].filter(id => ids.has(id))) : ids; + if (!matched.size) { + break; + } + } + return matched; + } + // mozStorage returns a BLOB as an array of byte values; reinterpret those // bytes as the stored Float32 embedding vector. function _blobToVector(blob) { @@ -773,9 +821,12 @@ Zotero.Embeddings = new function () { /** * Score a given set of items by similarity to a query. Items without a * stored embedding aren't scored, and neither are items scoring below the - * model's minimum, which aren't matches (see minScore in MODELS). Used to - * apply semantic ranking within an existing result scope (e.g. the current - * collection) rather than the whole library. + * model's minimum, which aren't matches (see minScore in MODELS) -- unless + * nothing clears it, in which case items whose text contains the query's + * words are scored, since the model missing what an item says literally + * shouldn't leave the search with nothing to show. Used to apply semantic + * ranking within an existing result scope (e.g. the current collection) + * rather than the whole library. * * @param {String} queryText * @param {Number[]} itemIDs - Candidate item IDs to score @@ -787,6 +838,7 @@ Zotero.Embeddings = new function () { */ this.scoreItemIDs = async function (queryText, itemIDs, { shouldCancel } = {}) { let scores = new Map(); + let belowFloor = new Map(); if (!itemIDs.length || !this.isEnabled()) { return scores; } @@ -839,6 +891,18 @@ Zotero.Embeddings = new function () { if (dot >= minScore) { scores.set(row.itemID, dot); } + else { + belowFloor.set(row.itemID, dot); + } + } + } + // Nothing was similar enough to the query to be a match, so fall back to + // the items that use its words, ranked by their own scores, which leave + // the relevance bar empty + if (!scores.size && belowFloor.size) { + let literal = await _findLiteralMatches(queryText, [...belowFloor.keys()]); + for (let itemID of literal) { + scores.set(itemID, belowFloor.get(itemID)); } } if (generation !== _modelGeneration) { @@ -1104,11 +1168,7 @@ Zotero.Embeddings.Indexing = new function () { // // @return {Promise} - libraryID -> [itemID, ...] async function _getEligibleItemIDs() { - let fieldIDs = [...new Set([ - Zotero.ItemFields.getID('title'), - Zotero.ItemFields.getID('abstractNote'), - ...Zotero.ItemFields.getTypeFieldsFromBase('title') - ])]; + let fieldIDs = Zotero.Embeddings.getIndexedFieldIDs(); let rows = await Zotero.DB.queryAsync( "SELECT libraryID, itemID, value FROM itemData " + "JOIN itemDataValues USING (valueID) " diff --git a/test/tests/embeddingsTest.js b/test/tests/embeddingsTest.js index 3612334a0c..850e27a54a 100644 --- a/test/tests/embeddingsTest.js +++ b/test/tests/embeddingsTest.js @@ -89,6 +89,67 @@ describe("Zotero.Embeddings", function () { Zotero.Prefs.clear('embeddings.model'); } }); + + it("should fall back to text matches when nothing clears the minimum", async function () { + Zotero.Prefs.set('embeddings.model', 'bge-small-en-v1.5'); + await Zotero.Embeddings.initDB(); + let mean = Zotero.Embeddings.getMeanVector(); + + let axis = (index, scale = 1) => { + let vector = Float32Array.from(mean); + vector[index] += scale; + return vector; + }; + let store = async (item, vector) => { + let blob = new Uint8Array(vector.buffer, vector.byteOffset, vector.byteLength); + await Zotero.DB.queryAsync( + "REPLACE INTO embeddings.itemEmbeddings (itemID, embedding, sourceHash) " + + "VALUES (?, ?, 'hash')", + [item.id, blob], { debugParams: false } + ); + }; + // Neither item is close to the query, but one says the word + let literal = await createDataObject('item', + { title: 'Migratory timing in Arctic-breeding shorebirds' }); + await store(literal, axis(1)); + let unrelated = await createDataObject('item', + { title: 'Guild regulation in early modern Nuremberg' }); + await store(unrelated, axis(2)); + + let stubs = [ + sinon.stub(Zotero.Embeddings, 'isEnabled').returns(true), + sinon.stub(Zotero.Embeddings, 'getModelVersion').returns('test-model/1'), + sinon.stub(Zotero.Embeddings, 'embedQuery').resolves(axis(0)) + ]; + await Zotero.DB.queryAsync( + "REPLACE INTO embeddings.itemEmbeddingsMeta (key, value) " + + "VALUES ('modelVersion', 'test-model/1')" + ); + try { + let scores = await Zotero.Embeddings.scoreItemIDs('birds', + [literal.id, unrelated.id]); + assert.isTrue(scores.has(literal.id)); + assert.isBelow(scores.get(literal.id), 0.2); + assert.isFalse(scores.has(unrelated.id)); + + // Every word has to appear + scores = await Zotero.Embeddings.scoreItemIDs('breeding penguins', + [literal.id, unrelated.id]); + assert.isFalse(scores.has(literal.id)); + + // With a real match to show, text matches stay out of it + let match = await createDataObject('item', { title: 'Birds' }); + await store(match, axis(0)); + scores = await Zotero.Embeddings.scoreItemIDs('birds', + [literal.id, unrelated.id, match.id]); + assert.isTrue(scores.has(match.id)); + assert.isFalse(scores.has(literal.id)); + } + finally { + stubs.forEach(stub => stub.restore()); + Zotero.Prefs.clear('embeddings.model'); + } + }); }); describe("#getScoreFraction()", function () {