Vector Scoring Plugin in Solr 8.6

Viewed 421

I am trying to upgrade the VectorScoringPlugin to Solr 8.6. Upon checking the migration logs of Lucene 8.x, I understood that the classes used in the plugin viz.CustomScoreQuery and CustomScoreProvider are deprecated and instead, we have to use FunctionScoreQuery along with DoubleValuesSource. I did a lot of searching but could not find any example of implementing a custom scorer using the above-mentioned classes. I stumbled upon these 2 threads [thread1, thread2] on the java-lucene forum that essentially talks about the same problem, and the solution mentioned is to implement a custom DoubleValuesSource class that will have the custom logic. Below is the implementation;

import java.io.IOException;
import java.util.ArrayList;
import java.util.List;

import org.apache.lucene.analysis.payloads.PayloadHelper;
import org.apache.lucene.index.LeafReaderContext;
import org.apache.lucene.index.PostingsEnum;
import org.apache.lucene.index.Terms;
import org.apache.lucene.index.TermsEnum;
import org.apache.lucene.search.DocIdSetIterator;
import org.apache.lucene.search.DoubleValues;
import org.apache.lucene.search.DoubleValuesSource;
import org.apache.lucene.search.IndexSearcher;
import org.apache.lucene.util.BytesRef;
import org.apache.solr.common.SolrException;

public class CustomDoubleValueSource extends DoubleValuesSource {

    List<Double> vector;
    private String field;
    private boolean cosine;
    double queryVectorNorm = 0;

    public CustomDoubleValueSource(String field, String Vector, boolean cosine) {
        // TODO Auto-generated constructor stub
        super();
        this.field = field;
        this.cosine = cosine;
        this.vector = new ArrayList<Double>();
        String[] vectorArray = Vector.split(",");
        for (int i = 0; i < vectorArray.length; i++) {
            double v = Double.parseDouble(vectorArray[i]);
            vector.add(v);
            if (cosine) {
                queryVectorNorm += Math.pow(v, 2.0);
            }
        }
        
        System.out.println("Vector size:"+this.vector.size());
    }

    @Override
    public boolean isCacheable(LeafReaderContext ctx) {
        // TODO Auto-generated method stub
        return false;
    }

    @Override
    public DoubleValues getValues(LeafReaderContext ctx, DoubleValues scores) throws IOException {

        Terms terms = ctx.reader().terms(field);
        TermsEnum te = terms == null ? null : terms.iterator();
        
        System.out.println("Term size:"+terms.size());

        if (vector.size() != terms.size()) {
            throw new SolrException(SolrException.ErrorCode.BAD_REQUEST,
                    "indexed and input vector array must have same length");
        }

        final PostingsEnum pe = te.postings(null);
        // TODO Auto-generated method stub
        return new DoubleValues() {

            @Override
            public double doubleValue() throws IOException {
                // TODO Auto-generated method stub
                float score = 0;
                double docVectorNorm = 0;
                BytesRef text;
                while ((text = te.next()) != null) {
                    String term = text.utf8ToString();
                    float payloadValue = 0f;
                    PostingsEnum postings = te.postings(null, PostingsEnum.ALL);
                    while (postings.nextDoc() != DocIdSetIterator.NO_MORE_DOCS) {
                        int freq = postings.freq();
                        while (freq-- > 0)
                            postings.nextPosition();

                        BytesRef payload = postings.getPayload();
                        payloadValue = PayloadHelper.decodeFloat(payload.bytes, payload.offset);

                        if (cosine)
                            docVectorNorm += Math.pow(payloadValue, 2.0);
                    }

                    score = (float) (score + payloadValue * (vector.get(Integer.parseInt(term))));
                }

                if (cosine) {
                    if ((docVectorNorm == 0) || (queryVectorNorm == 0))
                        return 0f;
                    return (float) (score / (Math.sqrt(docVectorNorm) * Math.sqrt(queryVectorNorm)));
                }

                return score;
            }

            @Override
            public boolean advanceExact(int doc) throws IOException {
                // TODO Auto-generated method stub
                if (pe.docID() > doc)
                    return false;
                
                return pe.docID() == doc || pe.advance(doc) == doc;
            }
        };
    }

    @Override
    public boolean needsScores() {
        // TODO Auto-generated method stub
        return true;
    }

    @Override
    public DoubleValuesSource rewrite(IndexSearcher reader) throws IOException {
        // TODO Auto-generated method stub
        return null;
    }

    @Override
    public int hashCode() {
        // TODO Auto-generated method stub
        return 0;
    }

    @Override
    public boolean equals(Object obj) {
        // TODO Auto-generated method stub
        return false;
    }

    @Override
    public String toString() {
        // TODO Auto-generated method stub
        return null;
    }

}

and I am using the above custom class as follows;

import org.apache.lucene.queries.function.FunctionScoreQuery;
import org.apache.lucene.search.Query;
import org.apache.solr.common.SolrException;
import org.apache.solr.common.params.SolrParams;
import org.apache.solr.request.SolrQueryRequest;
import org.apache.solr.schema.FieldType;
import org.apache.solr.search.QParser;
import org.apache.solr.search.QParserPlugin;
import org.apache.solr.search.QueryParsing;
import org.apache.solr.search.SyntaxError;

public class VectorQParserPlugin extends QParserPlugin {
    @Override
    public QParser createParser(String qstr, SolrParams localParams, SolrParams params, SolrQueryRequest req) {
        return new QParser(qstr, localParams, params, req) {
            @Override
            public Query parse() throws SyntaxError {
                String field = localParams.get(QueryParsing.F);
                String vector = localParams.get("vector");
                boolean cosine = localParams.getBool("cosine", true);

                if (field == null) {
                    throw new SolrException(SolrException.ErrorCode.BAD_REQUEST, "'f' not specified");
                }

                if (vector == null) {
                    throw new SolrException(SolrException.ErrorCode.BAD_REQUEST, "vector missing");
                }
                
                System.out.println("FIELD:"+field);

                Query subQuery = subQuery(localParams.get(QueryParsing.V), null).getQuery();

                FieldType ft = req.getCore().getLatestSchema().getFieldType(field);
                
                if(ft != null) {
                    System.out.println("in here");
                    VectorQuery q = new VectorQuery(subQuery);
                    q.setQueryString(localParams.toLocalParamsString()); 
                    query = q;
                }
            
                System.out.println("QUERY:"+query);
                if (query == null) {
                    throw new SolrException(SolrException.ErrorCode.BAD_REQUEST, "Query is null");
                }

                return FunctionScoreQuery.boostByValue(query, new CustomDoubleValueSource(field, vector, cosine));

            }
        };
    }
}

Also, I upgraded the custom Query implementation to 8.6 specific to avoid Query does not implement createWieght error.

Below is the implementation for VectorQuery;

import java.io.IOException;
import org.apache.lucene.index.LeafReaderContext;
import org.apache.lucene.search.ConstantScoreScorer;
import org.apache.lucene.search.ConstantScoreWeight;
import org.apache.lucene.search.DocIdSetIterator;
import org.apache.lucene.search.IndexSearcher;
import org.apache.lucene.search.Query;
import org.apache.lucene.search.ScoreMode;
import org.apache.lucene.search.Scorer;
import org.apache.lucene.search.Weight;
public class VectorQuery extends Query {
    String queryStr = "";
    Query q;
    public VectorQuery(Query subQuery) {
        this.q = subQuery;
    }
    
    public void setQueryString(String queryString){
        this.queryStr = queryString;
    }

    public Weight createWeight(IndexSearcher searcher, ScoreMode needsScores, float boost) throws IOException {
        Weight w;
        if(q == null){
            w =  new ConstantScoreWeight(this, boost) {
                @Override
                public Scorer scorer(LeafReaderContext context) throws IOException {
                    return new ConstantScoreScorer(this, score(), needsScores, DocIdSetIterator.all(context.reader().maxDoc()));
                }

                @Override
                public boolean isCacheable(LeafReaderContext ctx) {
                    // TODO Auto-generated method stub
                    return false;
                }
            };
        }else{
            w = searcher.createWeight(q, needsScores, boost);
        }
        return w;
    }

    @Override
    public String toString(String field) {
        return queryStr;
    }

    @Override
    public boolean equals(Object other) {
        return sameClassAs(other) &&
                queryStr.equals(other.toString());
    }

    @Override
    public int hashCode() {
        return classHash() ^ queryStr.hashCode();
    }

}

I added print statements to check the execution flow and the CustomDoubleValueSource class is invoked. Below is the screenshot of the log.

enter image description here

But the execution flow does not reach to getValues method. I am getting following error;

2020-10-21 16:55:09.578 ERROR (qtp1962826816-19) [   x:example_vector] o.a.s.s.HttpSolrCall null:java.lang.NullPointerException
        at org.apache.lucene.queries.function.FunctionScoreQuery$MultiplicativeBoostValuesSource.getValues(FunctionScoreQuery.java:261)
        at org.apache.lucene.queries.function.FunctionScoreQuery$FunctionScoreWeight.scorer(FunctionScoreQuery.java:224)
        at org.apache.lucene.search.Weight.bulkScorer(Weight.java:181)
        at org.apache.lucene.search.IndexSearcher.search(IndexSearcher.java:658)
        at org.apache.lucene.search.IndexSearcher.search(IndexSearcher.java:445)
        at org.apache.solr.search.SolrIndexSearcher.buildAndRunCollectorChain(SolrIndexSearcher.java:208)
        at org.apache.solr.search.SolrIndexSearcher.getDocListNC(SolrIndexSearcher.java:1593)
        at org.apache.solr.search.SolrIndexSearcher.getDocListC(SolrIndexSearcher.java:1410)
        at org.apache.solr.search.SolrIndexSearcher.search(SolrIndexSearcher.java:593)
        at org.apache.solr.handler.component.QueryComponent.doProcessUngroupedSearch(QueryComponent.java:1513)
        at org.apache.solr.handler.component.QueryComponent.process(QueryComponent.java:403)
        at org.apache.solr.handler.component.SearchHandler.handleRequestBody(SearchHandler.java:331)
        at org.apache.solr.handler.RequestHandlerBase.handleRequest(RequestHandlerBase.java:214)
        at org.apache.solr.core.SolrCore.execute(SolrCore.java:2606)
        at org.apache.solr.servlet.HttpSolrCall.execute(HttpSolrCall.java:812)
        at org.apache.solr.servlet.HttpSolrCall.call(HttpSolrCall.java:588)
        at org.apache.solr.servlet.SolrDispatchFilter.doFilter(SolrDispatchFilter.java:415)
        at org.apache.solr.servlet.SolrDispatchFilter.doFilter(SolrDispatchFilter.java:345)
        at org.eclipse.jetty.servlet.ServletHandler$CachedChain.doFilter(ServletHandler.java:1596)
        at org.eclipse.jetty.servlet.ServletHandler.doHandle(ServletHandler.java:545)
        at org.eclipse.jetty.server.handler.ScopedHandler.handle(ScopedHandler.java:143)
        at org.eclipse.jetty.security.SecurityHandler.handle(SecurityHandler.java:590)
        at org.eclipse.jetty.server.handler.HandlerWrapper.handle(HandlerWrapper.java:127)
        at org.eclipse.jetty.server.handler.ScopedHandler.nextHandle(ScopedHandler.java:235)
        at org.eclipse.jetty.server.session.SessionHandler.doHandle(SessionHandler.java:1610)
        at org.eclipse.jetty.server.handler.ScopedHandler.nextHandle(ScopedHandler.java:233)
        at org.eclipse.jetty.server.handler.ContextHandler.doHandle(ContextHandler.java:1300)
        at org.eclipse.jetty.server.handler.ScopedHandler.nextScope(ScopedHandler.java:188)
        at org.eclipse.jetty.servlet.ServletHandler.doScope(ServletHandler.java:485)
        at org.eclipse.jetty.server.session.SessionHandler.doScope(SessionHandler.java:1580)
        at org.eclipse.jetty.server.handler.ScopedHandler.nextScope(ScopedHandler.java:186)
        at org.eclipse.jetty.server.handler.ContextHandler.doScope(ContextHandler.java:1215)
        at org.eclipse.jetty.server.handler.ScopedHandler.handle(ScopedHandler.java:141)
        at org.eclipse.jetty.server.handler.ContextHandlerCollection.handle(ContextHandlerCollection.java:221)
        at org.eclipse.jetty.server.handler.InetAccessHandler.handle(InetAccessHandler.java:177)
        at org.eclipse.jetty.server.handler.HandlerCollection.handle(HandlerCollection.java:146)
        at org.eclipse.jetty.server.handler.HandlerWrapper.handle(HandlerWrapper.java:127)
        at org.eclipse.jetty.rewrite.handler.RewriteHandler.handle(RewriteHandler.java:322)
        at org.eclipse.jetty.server.handler.HandlerWrapper.handle(HandlerWrapper.java:127)
        at org.eclipse.jetty.server.Server.handle(Server.java:500)
        at org.eclipse.jetty.server.HttpChannel.lambda$handle$1(HttpChannel.java:383)
        at org.eclipse.jetty.server.HttpChannel.dispatch(HttpChannel.java:547)
        at org.eclipse.jetty.server.HttpChannel.handle(HttpChannel.java:375)
        at org.eclipse.jetty.server.HttpConnection.onFillable(HttpConnection.java:273)
        at org.eclipse.jetty.io.AbstractConnection$ReadCallback.succeeded(AbstractConnection.java:311)
        at org.eclipse.jetty.io.FillInterest.fillable(FillInterest.java:103)
        at org.eclipse.jetty.io.ChannelEndPoint$2.run(ChannelEndPoint.java:117)
        at org.eclipse.jetty.util.thread.QueuedThreadPool.runJob(QueuedThreadPool.java:806)
        at org.eclipse.jetty.util.thread.QueuedThreadPool$Runner.run(QueuedThreadPool.java:938)
        at java.lang.Thread.run(Thread.java:748)

I strongly suspect this has to do something with the unimplemented methods in CustomDoubleValuesSource. I read the java docs of DoubleValuesSource but they are not descriptive or do not contain any example.

Appreciate any help that will enable me to move forward :)

0 Answers
Related