sparknlp.annotator.classifier_dl.cross_encoder#
Contains classes for CrossEncoder.
Module Contents#
Classes#
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-transformersCrossEncoder) 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, DOCUMENTCATEGORY- 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)
- 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