Last active
November 14, 2018 16:37
-
-
Save tteofili/ce26241bd086e61a8c26a7544c175298 to your computer and use it in GitHub Desktop.
Anserini Reranker based on mean averaged word embeddings nearest neighbour
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| package io.anserini.rerank.lib; | |
| import io.anserini.rerank.Reranker; | |
| import io.anserini.rerank.RerankerContext; | |
| import io.anserini.rerank.ScoredDocuments; | |
| import org.apache.lucene.analysis.Analyzer; | |
| import org.apache.lucene.analysis.TokenStream; | |
| import org.apache.lucene.analysis.tokenattributes.CharTermAttribute; | |
| import org.apache.lucene.document.Document; | |
| import org.apache.lucene.index.IndexReader; | |
| import org.apache.lucene.search.ScoreDoc; | |
| import org.apache.lucene.search.TopDocs; | |
| import org.apache.lucene.util.BytesRef; | |
| import org.deeplearning4j.models.embeddings.loader.WordVectorSerializer; | |
| import org.deeplearning4j.models.word2vec.Word2Vec; | |
| import org.jetbrains.annotations.Nullable; | |
| import org.nd4j.linalg.api.ndarray.INDArray; | |
| import org.nd4j.linalg.factory.Nd4j; | |
| import org.nd4j.linalg.ops.transforms.Transforms; | |
| import org.slf4j.Logger; | |
| import org.slf4j.LoggerFactory; | |
| import java.io.IOException; | |
| import java.util.Arrays; | |
| import java.util.Collection; | |
| import java.util.LinkedList; | |
| import java.util.List; | |
| /** | |
| * Averaged Word Embeddings reranking model. | |
| */ | |
| public class AWEReranker implements Reranker { | |
| private final Logger log = LoggerFactory.getLogger(getClass()); | |
| private final Analyzer analyzer; | |
| private final String contentField; | |
| private final Word2Vec vec; | |
| public AWEReranker(Analyzer analyzer, String contentField, String model) { | |
| this.analyzer = analyzer; | |
| this.contentField = contentField; | |
| try { | |
| vec = WordVectorSerializer.readWord2VecModel(model); | |
| log.info("using word embeddings model from {}", model); | |
| } catch (Exception e) { | |
| throw new RuntimeException(e); | |
| } | |
| } | |
| public AWEReranker(Analyzer analyzer, String contentField, Word2Vec vec) { | |
| this.analyzer = analyzer; | |
| this.contentField = contentField; | |
| this.vec = vec; | |
| } | |
| @Override | |
| public ScoredDocuments rerank(ScoredDocuments docs, RerankerContext context) { | |
| if (docs.ids.length > 0) { | |
| if (vec != null) { | |
| int size = docs.ids.length; | |
| ScoreDoc[] scoreDocs = new ScoreDoc[size]; | |
| for (int i = 0; i < size; i++) { | |
| scoreDocs[i] = new ScoreDoc(docs.ids[i], docs.scores[i]); | |
| } | |
| TopDocs topDocs = new TopDocs(size, scoreDocs, docs.scores[0]); | |
| try { | |
| long start = System.currentTimeMillis(); | |
| log.debug("reranking {} docs", size); | |
| Collection<String> queryTokens = context.getQueryTokens(); | |
| kNNRerank(size, topDocs, queryTokens, "vector", | |
| false, context.getIndexSearcher().getIndexReader(), vec); | |
| log.debug("reranked {} docs in {}s", topDocs.totalHits, (double) (System.currentTimeMillis() - start) / 1000d); | |
| return ScoredDocuments.fromTopDocs(topDocs, context.getIndexSearcher()); | |
| } catch (IOException e) { | |
| throw new RuntimeException(e); | |
| } | |
| } | |
| } | |
| return docs; | |
| } | |
| private void kNNRerank(int k, TopDocs docs, Collection<String> queryTokens, String vectorField, | |
| boolean exact, IndexReader reader, Word2Vec vec) throws IOException { | |
| INDArray queryVector = averageWordVectors(queryTokens, vec); | |
| if (queryVector.maxNumber().doubleValue() == 0d) { | |
| log.warn("reranking skipped for query {}", queryTokens); | |
| return; | |
| } | |
| List<Integer> toDiscard = new LinkedList<>(); | |
| for (int j = 0; j < docs.scoreDocs.length; j++) { | |
| INDArray documentVector = averageWordVectors(reader, docs.scoreDocs[j].doc, vectorField, | |
| contentField, analyzer, vec); | |
| double similarity = Transforms.cosineSim(queryVector, documentVector); | |
| if (Double.isNaN(similarity) || similarity < 0.0) { | |
| log.warn("wrong similarity {} for {} and {}", similarity, queryVector, documentVector); | |
| toDiscard.add(docs.scoreDocs[j].doc); | |
| } else { | |
| if (exact) { | |
| docs.scoreDocs[j].score = (float) similarity; | |
| } else { | |
| docs.scoreDocs[j].score += (float) similarity; | |
| } | |
| } | |
| } | |
| if (!toDiscard.isEmpty()) { | |
| docs.scoreDocs = Arrays.stream(docs.scoreDocs).filter(e -> !toDiscard.contains(e.doc)).toArray(ScoreDoc[]::new); | |
| } | |
| Arrays.parallelSort(docs.scoreDocs, 0, docs.scoreDocs.length, (o1, o2) -> { // rerank scoreDocs | |
| return -1 * Double.compare(o1.score, o2.score); | |
| }); | |
| if (docs.scoreDocs.length > k) { | |
| docs.scoreDocs = Arrays.copyOfRange(docs.scoreDocs, 0, k); // retain only the top k nearest neighbours | |
| } | |
| if (docs.scoreDocs.length > 0) { | |
| docs.setMaxScore(docs.scoreDocs[0].score); | |
| } | |
| docs.totalHits = docs.scoreDocs.length; | |
| } | |
| private static Collection<String> getTokens(Analyzer analyzer, String field, String text) throws IOException { | |
| Collection<String> tokens = new LinkedList<>(); | |
| TokenStream ts = analyzer.tokenStream(field, text); | |
| ts.reset(); | |
| ts.addAttribute(CharTermAttribute.class); | |
| while (ts.incrementToken()) { | |
| CharTermAttribute charTermAttribute = ts.getAttribute(CharTermAttribute.class); | |
| String token = new String(charTermAttribute.buffer(), 0, charTermAttribute.length()); | |
| tokens.add(token); | |
| } | |
| ts.end(); | |
| ts.close(); | |
| return tokens; | |
| } | |
| private INDArray averageWordVectors(IndexReader reader, int doc, String vectorField, String contentField, Analyzer analyzer, | |
| Word2Vec word2Vec) throws IOException { | |
| INDArray vector; | |
| Document document = reader.document(doc); | |
| if (document != null) { | |
| BytesRef binaryValue = document.getBinaryValue(vectorField); | |
| if (binaryValue != null) { | |
| vector = Nd4j.fromByteArray(binaryValue.bytes); | |
| } else { | |
| vector = averageWordVectors(getTokens(analyzer, contentField, document.get(contentField)), word2Vec); | |
| } | |
| } else { | |
| vector = Nd4j.zeros(word2Vec.getLayerSize()); | |
| } | |
| return vector; | |
| } | |
| private INDArray averageWordVectors(Collection<String> words, Word2Vec word2Vec) { | |
| INDArray denseDocumentVector; | |
| try { | |
| denseDocumentVector = word2Vec.getWordVectorsMean(words); | |
| } catch (Exception e) { | |
| denseDocumentVector = Nd4j.zeros(word2Vec.getLayerSize()); | |
| double i = 0d; | |
| for (String token : words) { | |
| INDArray wordVector = fetchWordVector(word2Vec, token); | |
| if (wordVector != null) { | |
| denseDocumentVector.addi(wordVector); | |
| i++; | |
| } else { | |
| INDArray unkVector = word2Vec.getLookupTable().vector(word2Vec.getUNK()); | |
| if (unkVector != null) { | |
| denseDocumentVector.addi(unkVector); | |
| i++; | |
| } | |
| } | |
| } | |
| denseDocumentVector.divi(i); | |
| } | |
| return denseDocumentVector; | |
| } | |
| @Nullable | |
| private INDArray fetchWordVector(Word2Vec word2Vec, String termString) { | |
| INDArray wordVector = word2Vec.getLookupTable().vector(termString); | |
| double accuracy = 0.9; | |
| if (wordVector == null) { | |
| List<String> strings = word2Vec.similarWordsInVocabTo(termString, accuracy); | |
| if (!strings.isEmpty()) { | |
| wordVector = word2Vec.getLookupTable().vector(strings.get(0)); | |
| } | |
| } | |
| if (wordVector == null) { | |
| String[] strings = word2Vec.wordsNearest(termString, 1).toArray(new String[0]); | |
| if (strings.length > 0) { | |
| String nearestWord = strings[0]; | |
| if (word2Vec.similarity(termString, nearestWord) > accuracy) { | |
| wordVector = word2Vec.getLookupTable().vector(nearestWord); | |
| } | |
| } | |
| } | |
| return wordVector; | |
| } | |
| } |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment