Skip to content

Instantly share code, notes, and snippets.

@tteofili
Last active November 14, 2018 16:37
Show Gist options
  • Select an option

  • Save tteofili/ce26241bd086e61a8c26a7544c175298 to your computer and use it in GitHub Desktop.

Select an option

Save tteofili/ce26241bd086e61a8c26a7544c175298 to your computer and use it in GitHub Desktop.
Anserini Reranker based on mean averaged word embeddings nearest neighbour
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