Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Empty file.
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
from sparknlp.annotator import *


class AlbertZeroShotClassifier:
@staticmethod
def get_default_model():
return AlbertForZeroShotClassification.pretrained() \
.setInputCols(["token", "sentence"]) \
.setOutputCol("category") \
.setCaseSensitive(True)

@staticmethod
def get_pretrained_model(name, language, bucket=None):
return AlbertForZeroShotClassification.pretrained(name, language, bucket) \
.setInputCols(["token", "sentence"]) \
.setOutputCol("category") \
.setCaseSensitive(True)
19 changes: 19 additions & 0 deletions nlu/components/embeddings/mxbai/MxbaiEmbeddings.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,19 @@
from sparknlp.annotator import *


class MxbaiEmbeddings:
@staticmethod
def get_default_model():
from sparknlp.annotator import MxbaiEmbeddings
return MxbaiEmbeddings.pretrained() \
.setInputCols(["document"]) \
.setOutputCol("mxbai_embeddings")

# @staticmethod
# def get_pretrained_model(name, language, bucket=None):
# return MxbaiEmbeddings.pretrained(name,language,bucket) \
# .setInputCols(["document"]) \
# .setOutputCol("sentence_embeddings")



Empty file.
10 changes: 9 additions & 1 deletion nlu/spellbook.py
Original file line number Diff line number Diff line change
Expand Up @@ -2438,6 +2438,9 @@ class Spellbook:
'el.stopwords.iso': 'stopwords_iso'},
'eml': {'eml.embed.w2v_cc_300d': 'w2v_cc_300d'},
'en': {

'en.classify_zero_shot.albert.onnx' : 'albert_zero_shot_classifier_onnx',
'en.classify_zero_shot.albert.tf' : 'albert_zero_shot_classifier_tf',
'en.distilbert.zero_shot_classifier': 'distilbert_base_zero_shot_classifier_uncased_mnli',
'en.deberta.zero_shot_classifier': 'deberta_base_zero_shot_classifier_mnli_anli_v3',
'en.classify_image.convnext.tiny': 'image_classifier_convnext_tiny_224_local',
Expand Down Expand Up @@ -4556,6 +4559,7 @@ class Spellbook:
'en.dep.untyped': 'dependency_conllu',
'en.dep.untyped.conllu': 'dependency_conllu',
'en.e2e': 'multiclassifierdl_use_e2e',
'en.embed.mxbai' : 'mxbai_large_v1',
'en.embed': 'glove_100d',
'en.embed.Bible_roberta_base': 'roberta_embeddings_Bible_roberta_base',
'en.embed.COVID_SciBERT': 'bert_embeddings_COVID_SciBERT',
Expand Down Expand Up @@ -11525,7 +11529,10 @@ class Spellbook:
# Map every nlp_ref to an Annotator class. Language Agnostic and includes HC+OS
# For models with no pretrained weight, i.e. most OCR annotators, it maps AnnoId to Class

nlp_ref_to_anno_class = {'579_stmodel_product_rem_v3a': 'MPNetEmbeddings',
nlp_ref_to_anno_class = {
'albert_zero_shot_classifier_onnx': 'AlbertForZeroShotClassification',
'albert_zero_shot_classifier_tf': 'AlbertForZeroShotClassification',
'579_stmodel_product_rem_v3a': 'MPNetEmbeddings',
'abbreviation_category_mapper': 'ChunkMapperModel',
'abbreviation_mapper': 'ChunkMapperModel',
'abbreviation_mapper_augmented': 'ChunkMapperModel',
Expand Down Expand Up @@ -16051,6 +16058,7 @@ class Spellbook:
'genericclassifier_sdoh_tobacco_usage_sbiobert_cased_mli': 'GenericClassifierModel',
'github_issues_mpnet_southern_sotho_e10': 'MPNetEmbeddings',
'github_issues_preprocessed_mpnet_southern_sotho_e10': 'MPNetEmbeddings',
'mxbai_large_v1': 'MxbaiEmbeddings',
'glove_100d': 'WordEmbeddingsModel',
'glove_6B_100': 'WordEmbeddingsModel',
'glove_6B_300': 'WordEmbeddingsModel',
Expand Down
13 changes: 8 additions & 5 deletions nlu/universe/annotator_class_universe.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
from nlu.universe.atoms import JslAnnoId, JslAnnoPyClass
from nlu.universe.feature_node_ids import OCR_NODE_IDS, NLP_NODE_IDS, NLP_HC_NODE_IDS

import sparknlp_jsl.llm


class AnnoClassRef:
# Reference of every Annotator class name in OS/HC/OCR
Expand All @@ -13,7 +15,7 @@ class AnnoClassRef:
HC_A_N = NLP_HC_NODE_IDS
# Map AnnoID to PyCLass
JSL_anno2_py_class: Dict[JslAnnoId, JslAnnoPyClass] = {

A_N.ALBERT_FOR_ZERO_SHOT_CLASSIFICATION: 'AlbertForZeroShotClassification',
A_N.E5_SENTENCE_EMBEDDINGS: 'E5Embeddings',
A_N.BGE_SENTENCE_EMBEDDINGS: 'BGEEmbeddings',
A_N.INSTRUCTOR_SENTENCE_EMBEDDINGS: 'InstructorEmbeddings',
Expand Down Expand Up @@ -97,14 +99,15 @@ class AnnoClassRef:
A_N.ALBERT_EMBEDDINGS: 'AlbertEmbeddings',
A_N.ALBERT_FOR_TOKEN_CLASSIFICATION: 'AlbertForTokenClassification',
A_N.BERT_EMBEDDINGS: 'BertEmbeddings',
A_N.MXBAI_EMBEDDINGS: 'MxbaiEmbeddings',
A_N.BERT_FOR_TOKEN_CLASSIFICATION: 'BertForTokenClassification',
A_N.BERT_SENTENCE_EMBEDDINGS: 'BertSentenceEmbeddings',
A_N.DISTIL_BERT_EMBEDDINGS: 'DistilBertEmbeddings',
A_N.DISTIL_BERT_FOR_SEQUENCE_CLASSIFICATION: 'DistilBertForSequenceClassification',
A_N.DISTIL_BERT_FOR_ZERO_SHOT_CLASSIFICATION: 'DistilBertForZeroShotClassification',

A_N.DEBERTA_FOR_ZERO_SHOT_CLASSIFICATION: 'DeBertaForZeroShotClassification',

A_N.BERT_FOR_SEQUENCE_CLASSIFICATION: 'BertForSequenceClassification',
A_N.XLM_ROBERTA_FOR_ZERO_SHOT_CLASSIFICATION: 'XlmRoBertaForZeroShotClassification',
A_N.BERT_FOR_ZERO_SHOT_CLASSIFICATION: 'BertForZeroShotClassification',
Expand Down Expand Up @@ -132,7 +135,7 @@ class AnnoClassRef:
A_N.ALBERT_FOR_SEQUENCE_CLASSIFICATION: 'AlbertForSequenceClassification',
A_N.XLNET_FOR_SEQUENCE_CLASSIFICATION: 'XlnetForSequenceClassification',
A_N.GPT2: 'GPT2Transformer',
A_N.OPENAI_COMPLETION : 'OpenAICompletion',
A_N.OPENAI_COMPLETION: 'OpenAICompletion',
A_N.OPENAI_EMBEDDINGS: 'OpenAIEmbeddings',
A_N.DEBERTA_WORD_EMBEDDINGS: 'DeBertaEmbeddings',
A_N.DEBERTA_FOR_TOKEN_CLASSIFICATION: 'DeBertaForTokenClassification',
Expand Down Expand Up @@ -253,7 +256,7 @@ class AnnoClassRef:
JSL_anno_HC_ref_2_py_class: Dict[JslAnnoId, JslAnnoPyClass] = {
HC_A_N.MEDICAL_QUESTION_ANSWERING: 'MedicalQuestionAnswering',
HC_A_N.MEDICAL_TEXT_GENERATOR: 'MedicalTextGenerator',
HC_A_N.MEDICAL_SUMMARIZER:'MedicalSummarizer',
HC_A_N.MEDICAL_SUMMARIZER: 'MedicalSummarizer',
HC_A_N.ZERO_SHOT_NER: 'ZeroShotNerModel',
HC_A_N.CHUNK_MAPPER_MODEL: 'ChunkMapperModel',
HC_A_N.ASSERTION_DL: 'AssertionDLModel',
Expand Down Expand Up @@ -321,7 +324,7 @@ class AnnoClassRef:
OCR_NODE_IDS.IMAGE_SPLIT_REGIONS: 'ImageSplitRegions',
OCR_NODE_IDS.VISUAL_DOCUMENT_NER: 'VisualDocumentNer',
OCR_NODE_IDS.HOCR_TOKENIZER: 'HocrTokenizer',
OCR_NODE_IDS.FORM_RELATION_EXTRACTOR: 'FormRelationExtractor',
OCR_NODE_IDS.FORM_RELATION_EXTRACTOR: 'FormRelationExtractor',
OCR_NODE_IDS.IMAGE_DRAW_REGIONS: 'ImageDrawRegions',
OCR_NODE_IDS.POSITION_FINDER: 'PositionFinder',
OCR_NODE_IDS.IMAGE2PDF: 'ImageToPdf',
Expand Down
46 changes: 46 additions & 0 deletions nlu/universe/component_universes.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
from nlu.components.chunkers.contextual_parser.contextual_parser import ContextualParser
from nlu.components.chunkers.default_chunker.default_chunker import DefaultChunker
from nlu.components.chunkers.ngram.ngram import NGram
from nlu.components.classifiers.albert_zero_shot_classification.albert_zero_shot import AlbertZeroShotClassifier
from nlu.components.classifiers.asr.wav2Vec import Wav2Vec
from nlu.components.classifiers.asr_hubert.hubert import Hubert
from nlu.components.classifiers.asr_whisper.whisper import Whisper
Expand Down Expand Up @@ -96,6 +97,7 @@
from nlu.components.embeddings.xlm.xlm import XLM
from nlu.components.embeddings.xlnet.spark_nlp_xlnet import SparkNLPXlnet
from nlu.components.embeddings_chunks.chunk_embedder.chunk_embedder import ChunkEmbedder
from nlu.components.embeddings.mxbai.MxbaiEmbeddings import MxbaiEmbeddings
from nlu.components.lemmatizers.lemmatizer.spark_nlp_lemmatizer import SparkNLPLemmatizer
from nlu.components.matchers.regex_matcher.regex_matcher import RegexMatcher
from nlu.components.normalizers.document_normalizer.spark_nlp_document_normalizer import SparkNLPDocumentNormalizer
Expand Down Expand Up @@ -1988,6 +1990,27 @@ class ComponentUniverse:
is_storage_ref_producer=True,
has_storage_ref=True
),

A.MXBAI_EMBEDDINGS: partial(NluComponent,
name=A.MXBAI_EMBEDDINGS,
type=T.DOCUMENT_EMBEDDING,
get_default_model=MxbaiEmbeddings.get_default_model,
pdf_extractor_methods={'default': default_sentence_embedding_config,
'default_full': default_full_config, },
pdf_col_name_substitutor=substitute_sent_embed_cols,
output_level=L.INPUT_DEPENDENT_DOCUMENT_EMBEDDING,
node=NLP_FEATURE_NODES.nodes[A.MXBAI_EMBEDDINGS],
description='Converts Word Embeddings to Sentence/Document Embeddings',
provider=ComponentBackends.open_source,
license=Licenses.open_source,
computation_context=ComputeContexts.spark,
output_context=ComputeContexts.spark,
jsl_anno_class_id=A.MXBAI_EMBEDDINGS,
jsl_anno_py_class=ACR.JSL_anno2_py_class[
A.MXBAI_EMBEDDINGS],
is_storage_ref_producer=True,
has_storage_ref=True
),
A.STEMMER: partial(NluComponent,
name=A.STEMMER,
type=T.TOKEN_NORMALIZER,
Expand Down Expand Up @@ -3152,6 +3175,29 @@ class ComponentUniverse:
),


A.ALBERT_FOR_ZERO_SHOT_CLASSIFICATION: partial(NluComponent,
name=A.ALBERT_FOR_ZERO_SHOT_CLASSIFICATION,
type=T.TRANSFORMER_SEQUENCE_CLASSIFIER,
get_default_model=AlbertZeroShotClassifier.get_default_model,
get_pretrained_model=AlbertZeroShotClassifier.get_pretrained_model,
pdf_extractor_methods={
'default': default_seq_classifier_config,
'default_full': default_full_config, },
pdf_col_name_substitutor=substitute_seq_bert_classifier_cols,
output_level=L.INPUT_DEPENDENT_DOCUMENT_CLASSIFIER,
node=NLP_FEATURE_NODES.nodes[
A.ALBERT_FOR_ZERO_SHOT_CLASSIFICATION],
description='ALBERT Zero Shot Classifier.',
provider=ComponentBackends.open_source,
license=Licenses.open_source,
computation_context=ComputeContexts.spark,
output_context=ComputeContexts.spark,
jsl_anno_class_id=A.ALBERT_FOR_ZERO_SHOT_CLASSIFICATION,
jsl_anno_py_class=ACR.JSL_anno2_py_class[
A.ALBERT_FOR_ZERO_SHOT_CLASSIFICATION],
),


A.DEBERTA_FOR_ZERO_SHOT_CLASSIFICATION: partial(NluComponent,
name=A.DEBERTA_FOR_ZERO_SHOT_CLASSIFICATION,
type=T.TRANSFORMER_SEQUENCE_CLASSIFIER,
Expand Down
2 changes: 2 additions & 0 deletions nlu/universe/feature_node_ids.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,7 @@ class NLP_NODE_IDS:
SENTENCE_DETECTOR = JslAnnoId('sentence_detector')
SENTENCE_DETECTOR_DL = JslAnnoId('sentence_detector_dl')
SENTENCE_EMBEDDINGS_CONVERTER = JslAnnoId('sentence_embeddings_converter')
MXBAI_EMBEDDINGS = JslAnnoId('mxbai_embeddings')
STEMMER = JslAnnoId('stemmer')
STOP_WORDS_CLEANER = JslAnnoId('stop_words_cleaner')
SYMMETRIC_DELETE_SPELLCHECKER = JslAnnoId('symmetric_delete_spellchecker')
Expand Down Expand Up @@ -117,6 +118,7 @@ class NLP_NODE_IDS:
MPNET_SENTENCE_EMBEDDINGS = JslAnnoId('mpnet_sentence_embeddings')
MPNET_FOR_SEQUENCE_CLASSIFICATION = JslAnnoId('mpnet_for_sequence_classification')
DISTIL_BERT_FOR_ZERO_SHOT_CLASSIFICATION = JslAnnoId('distil_bert_zero_shot')
ALBERT_FOR_ZERO_SHOT_CLASSIFICATION = JslAnnoId('albert_bert_zero_shot')
XLM_ROBERTA_FOR_ZERO_SHOT_CLASSIFICATION = JslAnnoId('xlm_roberta_zero_shot')

DEBERTA_FOR_ZERO_SHOT_CLASSIFICATION = JslAnnoId('deberta_zero_shot')
Expand Down
3 changes: 3 additions & 0 deletions nlu/universe/feature_node_universes.py
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,7 @@ class NLP_FEATURE_NODES: # or Mode Node?
A.INSTRUCTOR_SENTENCE_EMBEDDINGS: NlpFeatureNode(A.INSTRUCTOR_SENTENCE_EMBEDDINGS, [F.DOCUMENT], [F.SENTENCE_EMBEDDINGS]),

A.E5_SENTENCE_EMBEDDINGS: NlpFeatureNode(A.E5_SENTENCE_EMBEDDINGS, [F.DOCUMENT],[F.SENTENCE_EMBEDDINGS]),
A.MXBAI_EMBEDDINGS: NlpFeatureNode(A.MXBAI_EMBEDDINGS, [F.DOCUMENT],[F.SENTENCE_EMBEDDINGS]),
A.BGE_SENTENCE_EMBEDDINGS: NlpFeatureNode(A.BGE_SENTENCE_EMBEDDINGS, [F.DOCUMENT], [F.SENTENCE_EMBEDDINGS]),
A.MPNET_SENTENCE_EMBEDDINGS: NlpFeatureNode(A.MPNET_SENTENCE_EMBEDDINGS, [F.DOCUMENT], [F.SENTENCE_EMBEDDINGS]),

Expand Down Expand Up @@ -246,6 +247,8 @@ class NLP_FEATURE_NODES: # or Mode Node?
[F.SEQUENCE_CLASSIFICATION]),
A.BERT_FOR_ZERO_SHOT_CLASSIFICATION: NlpFeatureNode(A.BERT_FOR_ZERO_SHOT_CLASSIFICATION, [F.DOCUMENT, F.TOKEN],
[F.SEQUENCE_CLASSIFICATION]),
A.ALBERT_FOR_ZERO_SHOT_CLASSIFICATION: NlpFeatureNode(A.BERT_FOR_ZERO_SHOT_CLASSIFICATION, [F.DOCUMENT, F.TOKEN],
[F.SEQUENCE_CLASSIFICATION]),
A.BART_FOR_ZERO_SHOT_CLASSIFICATION: NlpFeatureNode(A.BART_FOR_ZERO_SHOT_CLASSIFICATION, [F.DOCUMENT, F.TOKEN],
[F.SEQUENCE_CLASSIFICATION]),

Expand Down
24 changes: 22 additions & 2 deletions tests/base_model_test.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import pytest

from tests.utils import all_tests, one_per_lib, NluTest, model_and_output_levels_test
from tests.utils.model_test import quick_test


def model_id(model_to_test: NluTest) -> str:
Expand All @@ -15,7 +16,10 @@ def one_test_per_lib():
return one_per_lib


@pytest.mark.skip(reason="Use run_tests.py instead until pytest-xdist issue is fixed")
def quick_test_data():
return quick_test

# @pytest.mark.skip(reason="Use run_tests.py instead until pytest-xdist issue is fixed")
@pytest.mark.parametrize("model_to_test", all_annotator_tests(), ids=model_id)
def test_model_all_annotators(model_to_test: NluTest):
model_and_output_levels_test(
Expand All @@ -29,7 +33,7 @@ def test_model_all_annotators(model_to_test: NluTest):
)


@pytest.mark.skip(reason="Local testing")
# @pytest.mark.skip(reason="Local testing")
@pytest.mark.parametrize("model_to_test", one_test_per_lib(), ids=model_id)
def test_one_per_lib(model_to_test: NluTest):
model_and_output_levels_test(
Expand All @@ -41,3 +45,19 @@ def test_one_per_lib(model_to_test: NluTest):
library=model_to_test.library,
pipe_params=model_to_test.pipe_params
)




@pytest.mark.skip(reason="Local testing")
@pytest.mark.parametrize("model_to_test", quick_test_data(), ids=model_id)
def test_quick(model_to_test: NluTest):
model_and_output_levels_test(
nlu_ref=model_to_test.nlu_ref,
lang=model_to_test.lang,
test_group=model_to_test.test_group,
output_levels=model_to_test.output_levels,
input_data_type=model_to_test.input_data_type,
library=model_to_test.library,
pipe_params=model_to_test.pipe_params
)
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
# import tests.secrets as sct

import os
import sys

# sys.path.append(os.getcwd())
import unittest
import nlu

os.environ["PYTHONPATH"] = "F:/Work/repos/nlu_new/nlu"
os.environ['PYSPARK_PYTHON'] = sys.executable
os.environ['PYSPARK_DRIVER_PYTHON'] = sys.executable
from johnsnowlabs import nlp, visual

# nlp.install(json_license_path="license.json")

nlp.start()

class EmbeddingTests(unittest.TestCase):
def test_mxbai_embeddings_model(self):

res = nlu.load("en.embed.mxbai").predict('This is an example sentence', output_level='document')
print(res)


if __name__ == "__main__":
EmbeddingTests().test_mxbai_embeddings_model()
Original file line number Diff line number Diff line change
@@ -1,7 +1,17 @@
import sys
import os
# sys.path.append(os.getcwd())
import unittest
import nlu

from nlu import *
os.environ["PYTHONPATH"] = "F:/Work/repos/nlu_new/nlu"
os.environ['PYSPARK_PYTHON'] = sys.executable
os.environ['PYSPARK_DRIVER_PYTHON'] = sys.executable
from johnsnowlabs import nlp, visual

# nlp.install(json_license_path="license.json")

nlp.start()

class TestE5SentenceEmbeddings(unittest.TestCase):
def test_e5_embeds(self):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import unittest
import nlu

# os.environ["PYTHONPATH"] = "F:/Work/repos/nlu_new/nlu"
os.environ["PYTHONPATH"] = "F:/Work/repos/nlu_new/nlu"
os.environ['PYSPARK_PYTHON'] = sys.executable
os.environ['PYSPARK_DRIVER_PYTHON'] = sys.executable
from johnsnowlabs import nlp, visual
Expand Down
5 changes: 5 additions & 0 deletions tests/utils/model_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -264,6 +264,7 @@ class NluTest(BaseModel):
param_val='translate English to French')]),
NluTest(nlu_ref="match.chunks", lang='en', test_group='matcher', input_data_type='generic',
library='open_source'),
NluTest(nlu_ref="en.classify_zero_shot.albert.onnx", lang='en', test_group='chunker', input_data_type='generic', library='open_source'),

]

Expand All @@ -277,4 +278,8 @@ class NluTest(BaseModel):
]


quick_test = [
NluTest(nlu_ref="en.classify_zero_shot.abert.onnx", lang='en', test_group='chunker', input_data_type='generic', library='open_source'),
]

all_tests = ocr_tests + medical_tests + nlp_tests
4 changes: 2 additions & 2 deletions tests/utils/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,9 @@
import pandas as pd
import sparknlp

import _secrets as secrets
import nlu
from test_data import get_test_data
from .test_data import get_test_data
from . import _secrets as secrets

os.environ['PYSPARK_PYTHON'] = sys.executable
os.environ['PYSPARK_DRIVER_PYTHON'] = sys.executable
Expand Down