From b08f388423b73bec19e3adc1945d5ee851f59928 Mon Sep 17 00:00:00 2001 From: GADDE SAI SHAILESH Date: Wed, 6 Nov 2024 12:33:20 -0500 Subject: [PATCH 1/4] Added MxBai Embeddings --- .../embeddings/mxbai/MxbaiEmbeddings.py | 19 ++++++++ nlu/components/embeddings/mxbai/__init__.py | 0 nlu/spellbook.py | 2 + nlu/universe/annotator_class_universe.py | 1 + nlu/universe/component_universes.py | 44 +++++++++++++++++++ nlu/universe/feature_node_ids.py | 1 + nlu/universe/feature_node_universes.py | 1 + .../sentence_embeddings/mxbai_embeddings.py | 27 ++++++++++++ .../sentence_embeddings/sentence_e5_tests.py | 12 ++++- .../assertion_tests.py | 2 +- 10 files changed, 107 insertions(+), 2 deletions(-) create mode 100644 nlu/components/embeddings/mxbai/MxbaiEmbeddings.py create mode 100644 nlu/components/embeddings/mxbai/__init__.py create mode 100644 tests/nlu_core_tests/component_tests/embed_tests/sentence_embeddings/mxbai_embeddings.py diff --git a/nlu/components/embeddings/mxbai/MxbaiEmbeddings.py b/nlu/components/embeddings/mxbai/MxbaiEmbeddings.py new file mode 100644 index 00000000..09385b67 --- /dev/null +++ b/nlu/components/embeddings/mxbai/MxbaiEmbeddings.py @@ -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") + + + diff --git a/nlu/components/embeddings/mxbai/__init__.py b/nlu/components/embeddings/mxbai/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/nlu/spellbook.py b/nlu/spellbook.py index dd2225bb..e3396d9b 100644 --- a/nlu/spellbook.py +++ b/nlu/spellbook.py @@ -4556,6 +4556,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', @@ -16051,6 +16052,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', diff --git a/nlu/universe/annotator_class_universe.py b/nlu/universe/annotator_class_universe.py index 0eeecc13..5cf1848a 100644 --- a/nlu/universe/annotator_class_universe.py +++ b/nlu/universe/annotator_class_universe.py @@ -97,6 +97,7 @@ 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', diff --git a/nlu/universe/component_universes.py b/nlu/universe/component_universes.py index 0ace75e6..33acca29 100644 --- a/nlu/universe/component_universes.py +++ b/nlu/universe/component_universes.py @@ -96,6 +96,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 @@ -1988,6 +1989,49 @@ 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.E5_SENTENCE_EMBEDDINGS: partial(NluComponent, + name=A.E5_SENTENCE_EMBEDDINGS, + type=T.DOCUMENT_EMBEDDING, + get_default_model=E5.get_default_model, + get_pretrained_model=E5.get_pretrained_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.E5_SENTENCE_EMBEDDINGS], + description='Sentence-level embeddings using E5. E5, a weakly supervised text embedding model that can generate text embeddings tailored to any task (e.g., classification, retrieval, clustering, text evaluation, etc.).', + provider=ComponentBackends.open_source, + license=Licenses.open_source, + computation_context=ComputeContexts.spark, + output_context=ComputeContexts.spark, + jsl_anno_class_id=A.E5_SENTENCE_EMBEDDINGS, + jsl_anno_py_class=ACR.JSL_anno2_py_class[A.E5_SENTENCE_EMBEDDINGS], + has_storage_ref=True, + is_storage_ref_producer=True, + ), + A.STEMMER: partial(NluComponent, name=A.STEMMER, type=T.TOKEN_NORMALIZER, diff --git a/nlu/universe/feature_node_ids.py b/nlu/universe/feature_node_ids.py index 49617140..257095ba 100644 --- a/nlu/universe/feature_node_ids.py +++ b/nlu/universe/feature_node_ids.py @@ -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') diff --git a/nlu/universe/feature_node_universes.py b/nlu/universe/feature_node_universes.py index 21a0d3ae..2c34d929 100644 --- a/nlu/universe/feature_node_universes.py +++ b/nlu/universe/feature_node_universes.py @@ -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]), diff --git a/tests/nlu_core_tests/component_tests/embed_tests/sentence_embeddings/mxbai_embeddings.py b/tests/nlu_core_tests/component_tests/embed_tests/sentence_embeddings/mxbai_embeddings.py new file mode 100644 index 00000000..8faa70a5 --- /dev/null +++ b/tests/nlu_core_tests/component_tests/embed_tests/sentence_embeddings/mxbai_embeddings.py @@ -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() diff --git a/tests/nlu_core_tests/component_tests/embed_tests/sentence_embeddings/sentence_e5_tests.py b/tests/nlu_core_tests/component_tests/embed_tests/sentence_embeddings/sentence_e5_tests.py index 5c8dda98..edec64e2 100644 --- a/tests/nlu_core_tests/component_tests/embed_tests/sentence_embeddings/sentence_e5_tests.py +++ b/tests/nlu_core_tests/component_tests/embed_tests/sentence_embeddings/sentence_e5_tests.py @@ -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): diff --git a/tests/nlu_hc_tests/component_tests/few_shot_assertion_classifier/assertion_tests.py b/tests/nlu_hc_tests/component_tests/few_shot_assertion_classifier/assertion_tests.py index a2e61e77..c9395770 100644 --- a/tests/nlu_hc_tests/component_tests/few_shot_assertion_classifier/assertion_tests.py +++ b/tests/nlu_hc_tests/component_tests/few_shot_assertion_classifier/assertion_tests.py @@ -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 From b2d0b8a1a81489571aea887c5b78f6a4d992bb0e Mon Sep 17 00:00:00 2001 From: GADDE SAI SHAILESH Date: Wed, 6 Nov 2024 12:36:52 -0500 Subject: [PATCH 2/4] Addedd MXBai Embeddings --- nlu/universe/component_universes.py | 24 +----------------------- 1 file changed, 1 insertion(+), 23 deletions(-) diff --git a/nlu/universe/component_universes.py b/nlu/universe/component_universes.py index 33acca29..1ac9bd74 100644 --- a/nlu/universe/component_universes.py +++ b/nlu/universe/component_universes.py @@ -2009,29 +2009,7 @@ class ComponentUniverse: A.MXBAI_EMBEDDINGS], is_storage_ref_producer=True, has_storage_ref=True - ), - - A.E5_SENTENCE_EMBEDDINGS: partial(NluComponent, - name=A.E5_SENTENCE_EMBEDDINGS, - type=T.DOCUMENT_EMBEDDING, - get_default_model=E5.get_default_model, - get_pretrained_model=E5.get_pretrained_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.E5_SENTENCE_EMBEDDINGS], - description='Sentence-level embeddings using E5. E5, a weakly supervised text embedding model that can generate text embeddings tailored to any task (e.g., classification, retrieval, clustering, text evaluation, etc.).', - provider=ComponentBackends.open_source, - license=Licenses.open_source, - computation_context=ComputeContexts.spark, - output_context=ComputeContexts.spark, - jsl_anno_class_id=A.E5_SENTENCE_EMBEDDINGS, - jsl_anno_py_class=ACR.JSL_anno2_py_class[A.E5_SENTENCE_EMBEDDINGS], - has_storage_ref=True, - is_storage_ref_producer=True, - ), - + ),git A.STEMMER: partial(NluComponent, name=A.STEMMER, type=T.TOKEN_NORMALIZER, From 8bafc5389bfb529b6c0b5bc902ec2b3b4c1b8bd7 Mon Sep 17 00:00:00 2001 From: Faisal Adnan <68223517+faisaladnanpeltops@users.noreply.github.com> Date: Wed, 27 Nov 2024 17:32:20 +0100 Subject: [PATCH 3/4] fix typo --- nlu/universe/component_universes.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/nlu/universe/component_universes.py b/nlu/universe/component_universes.py index 1ac9bd74..d385bd69 100644 --- a/nlu/universe/component_universes.py +++ b/nlu/universe/component_universes.py @@ -2009,7 +2009,7 @@ class ComponentUniverse: A.MXBAI_EMBEDDINGS], is_storage_ref_producer=True, has_storage_ref=True - ),git + ), A.STEMMER: partial(NluComponent, name=A.STEMMER, type=T.TOKEN_NORMALIZER, From 58bea7ba117f9f62e1e37cda41545c93e73bea26 Mon Sep 17 00:00:00 2001 From: C-K-Loan Date: Wed, 29 Jan 2025 00:42:52 +0100 Subject: [PATCH 4/4] AlbertForZeroShotclassification --- .../__init__.py | 0 .../albert_zero_shot.py | 17 +++++++++++++ nlu/spellbook.py | 8 ++++++- nlu/universe/annotator_class_universe.py | 12 ++++++---- nlu/universe/component_universes.py | 24 +++++++++++++++++++ nlu/universe/feature_node_ids.py | 1 + nlu/universe/feature_node_universes.py | 2 ++ tests/base_model_test.py | 24 +++++++++++++++++-- tests/utils/model_test.py | 5 ++++ tests/utils/test_utils.py | 4 ++-- 10 files changed, 87 insertions(+), 10 deletions(-) create mode 100644 nlu/components/classifiers/albert_zero_shot_classification/__init__.py create mode 100644 nlu/components/classifiers/albert_zero_shot_classification/albert_zero_shot.py diff --git a/nlu/components/classifiers/albert_zero_shot_classification/__init__.py b/nlu/components/classifiers/albert_zero_shot_classification/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/nlu/components/classifiers/albert_zero_shot_classification/albert_zero_shot.py b/nlu/components/classifiers/albert_zero_shot_classification/albert_zero_shot.py new file mode 100644 index 00000000..7aea3464 --- /dev/null +++ b/nlu/components/classifiers/albert_zero_shot_classification/albert_zero_shot.py @@ -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) diff --git a/nlu/spellbook.py b/nlu/spellbook.py index e3396d9b..80f660e0 100644 --- a/nlu/spellbook.py +++ b/nlu/spellbook.py @@ -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', @@ -11526,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', diff --git a/nlu/universe/annotator_class_universe.py b/nlu/universe/annotator_class_universe.py index 5cf1848a..beeff56c 100644 --- a/nlu/universe/annotator_class_universe.py +++ b/nlu/universe/annotator_class_universe.py @@ -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 @@ -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', @@ -105,7 +107,7 @@ class AnnoClassRef: 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', @@ -133,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', @@ -254,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', @@ -322,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', diff --git a/nlu/universe/component_universes.py b/nlu/universe/component_universes.py index d385bd69..aa527da8 100644 --- a/nlu/universe/component_universes.py +++ b/nlu/universe/component_universes.py @@ -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 @@ -3174,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, diff --git a/nlu/universe/feature_node_ids.py b/nlu/universe/feature_node_ids.py index 257095ba..87c84e21 100644 --- a/nlu/universe/feature_node_ids.py +++ b/nlu/universe/feature_node_ids.py @@ -118,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') diff --git a/nlu/universe/feature_node_universes.py b/nlu/universe/feature_node_universes.py index 2c34d929..f90bf2f1 100644 --- a/nlu/universe/feature_node_universes.py +++ b/nlu/universe/feature_node_universes.py @@ -247,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]), diff --git a/tests/base_model_test.py b/tests/base_model_test.py index eee4fa45..e53c2f3a 100644 --- a/tests/base_model_test.py +++ b/tests/base_model_test.py @@ -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: @@ -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( @@ -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( @@ -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 + ) \ No newline at end of file diff --git a/tests/utils/model_test.py b/tests/utils/model_test.py index b20a3539..054704fe 100644 --- a/tests/utils/model_test.py +++ b/tests/utils/model_test.py @@ -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'), ] @@ -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 diff --git a/tests/utils/test_utils.py b/tests/utils/test_utils.py index 8c303b2a..e1497c2e 100644 --- a/tests/utils/test_utils.py +++ b/tests/utils/test_utils.py @@ -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