sparknlp.annotator.classifier_dl.cross_encoder#

Contains classes for CrossEncoder.

Module Contents#

Classes#

CrossEncoder

CrossEncoder brings cross-encoder relevance scoring (as in

class CrossEncoder(classname='com.johnsnowlabs.nlp.annotators.classifier.dl.CrossEncoder', java_model=None)[source]#

CrossEncoder brings cross-encoder relevance scoring (as in sentence-transformers CrossEncoder) into Spark NLP as a first-class annotator.

It takes two row-aligned document columns, jointly encodes each row’s pair as a single sequence [CLS] text_a [SEP] text_b [SEP], runs one forward pass through a BERT-family transformer with a single-logit regression head, and writes one score per row to a single output column. The logit is squashed with a sigmoid, so every score lands in [0, 1]

Pretrained models can be loaded with pretrained() of the companion object:

>>> crossEncoder = CrossEncoder.pretrained() \
...     .setInputCols(["document1", "document2"]) \
...     .setOutputCol("score")

The default model is "cross_encoder_ms_marco_minilm_l6_v2", if no name is provided.

For available pretrained models please see the Models Hub.

To see which models are compatible and how to import them see Import Transformers into Spark NLP 🚀.

Input Annotation types

Output Annotation type

DOCUMENT, DOCUMENT

CATEGORY

Parameters:
batchSize

Batch size. Large values allows faster processing but requires more memory, by default 8

caseSensitive

Whether to ignore case in tokens for embeddings matching, by default False

Examples

>>> import sparknlp
>>> from sparknlp.base import *
>>> from sparknlp.annotator import *
>>> from pyspark.ml import Pipeline
>>> document = MultiDocumentAssembler() \
...     .setInputCols(["query", "passage"]) \
...     .setOutputCols(["document1", "document2"])
>>> crossEncoder = CrossEncoder.pretrained() \
...     .setInputCols(["document1", "document2"]) \
...     .setOutputCol("score")
>>> pipeline = Pipeline().setStages([document, crossEncoder])
>>> data = spark.createDataFrame([
...     ["How many people live in Berlin?", "Berlin is well known for its museums."]
... ]).toDF("query", "passage")
>>> result = pipeline.fit(data).transform(data)
>>> result.select("score.result").show(truncate=False)
name = 'CrossEncoder'[source]#
inputAnnotatorTypes[source]#
outputAnnotatorType = 'category'[source]#
static loadSavedModel(folder, spark_session)[source]#

Loads a locally saved model.

Parameters:
folderstr

Folder of the saved model

spark_sessionpyspark.sql.SparkSession

The current SparkSession

Returns:
CrossEncoder

The restored model

static pretrained(name='cross_encoder_ms_marco_minilm_l6_v2', lang='en', remote_loc=None)[source]#

Downloads and loads a pretrained model.

Parameters:
namestr, optional

Name of the pretrained model, by default “cross_encoder_ms_marco_minilm_l6_v2”

langstr, optional

Language of the pretrained model, by default “en”

remote_locstr, optional

Optional remote address of the resource, by default None. Will use Spark NLPs repositories otherwise.

Returns:
CrossEncoder

The restored model