|
| 1 | +from typing import Callable |
| 2 | +from functools import partial |
| 3 | +import copy |
| 4 | + |
| 5 | +import torch |
| 6 | +from transformers import AutoModelForSequenceClassification |
| 7 | +from transformers.models.distilbert.modeling_distilbert import ( |
| 8 | + DistilBertForSequenceClassification, |
| 9 | +) |
| 10 | +from transformers.modeling_outputs import SequenceClassifierOutput |
| 11 | + |
| 12 | +import spacy |
| 13 | +from thinc.api import Model |
| 14 | + |
| 15 | +from spacy_transformers.data_classes import HFObjects, WordpieceBatch |
| 16 | +from spacy_transformers.layers.hf_wrapper import HFWrapper |
| 17 | +from spacy_transformers.layers.transformer_model import _convert_transformer_inputs |
| 18 | +from spacy_transformers.layers.transformer_model import _convert_transformer_outputs |
| 19 | +from spacy_transformers.layers.transformer_model import forward |
| 20 | +from spacy_transformers.layers.transformer_model import huggingface_from_pretrained |
| 21 | +from spacy_transformers.layers.transformer_model import huggingface_tokenize |
| 22 | +from spacy_transformers.layers.transformer_model import set_pytorch_transformer |
| 23 | +from spacy_transformers.span_getters import get_strided_spans |
| 24 | + |
| 25 | + |
| 26 | +def test_model_for_sequence_classification(): |
| 27 | + # adapted from https://github.com/KennethEnevoldsen/spacy-wrap/ |
| 28 | + class ClassificationTransformerModel(Model): |
| 29 | + def __init__( |
| 30 | + self, |
| 31 | + name: str, |
| 32 | + get_spans: Callable, |
| 33 | + tokenizer_config: dict = {}, |
| 34 | + transformer_config: dict = {}, |
| 35 | + mixed_precision: bool = False, |
| 36 | + grad_scaler_config: dict = {}, |
| 37 | + ): |
| 38 | + hf_model = HFObjects(None, None, None, tokenizer_config, transformer_config) |
| 39 | + wrapper = HFWrapper( |
| 40 | + hf_model, |
| 41 | + convert_inputs=_convert_transformer_inputs, |
| 42 | + convert_outputs=_convert_transformer_outputs, |
| 43 | + mixed_precision=mixed_precision, |
| 44 | + grad_scaler_config=grad_scaler_config, |
| 45 | + model_cls=AutoModelForSequenceClassification, |
| 46 | + ) |
| 47 | + super().__init__( |
| 48 | + "clf_transformer", |
| 49 | + forward, |
| 50 | + init=init, |
| 51 | + layers=[wrapper], |
| 52 | + dims={"nO": None}, |
| 53 | + attrs={ |
| 54 | + "get_spans": get_spans, |
| 55 | + "name": name, |
| 56 | + "set_transformer": set_pytorch_transformer, |
| 57 | + "has_transformer": False, |
| 58 | + "flush_cache_chance": 0.0, |
| 59 | + }, |
| 60 | + ) |
| 61 | + |
| 62 | + @property |
| 63 | + def tokenizer(self): |
| 64 | + return self.layers[0].shims[0]._hfmodel.tokenizer |
| 65 | + |
| 66 | + @property |
| 67 | + def transformer(self): |
| 68 | + return self.layers[0].shims[0]._hfmodel.transformer |
| 69 | + |
| 70 | + @property |
| 71 | + def _init_tokenizer_config(self): |
| 72 | + return self.layers[0].shims[0]._hfmodel._init_tokenizer_config |
| 73 | + |
| 74 | + @property |
| 75 | + def _init_transformer_config(self): |
| 76 | + return self.layers[0].shims[0]._hfmodel._init_transformer_config |
| 77 | + |
| 78 | + def copy(self): |
| 79 | + """ |
| 80 | + Create a copy of the model, its attributes, and its parameters. Any child |
| 81 | + layers will also be deep-copied. The copy will receive a distinct `model.id` |
| 82 | + value. |
| 83 | + """ |
| 84 | + copied = ClassificationTransformerModel(self.name, self.attrs["get_spans"]) |
| 85 | + params = {} |
| 86 | + for name in self.param_names: |
| 87 | + params[name] = self.get_param(name) if self.has_param(name) else None |
| 88 | + copied.params = copy.deepcopy(params) |
| 89 | + copied.dims = copy.deepcopy(self._dims) |
| 90 | + copied.layers[0] = copy.deepcopy(self.layers[0]) |
| 91 | + for name in self.grad_names: |
| 92 | + copied.set_grad(name, self.get_grad(name).copy()) |
| 93 | + return copied |
| 94 | + |
| 95 | + def init(model: Model, X=None, Y=None): |
| 96 | + if model.attrs["has_transformer"]: |
| 97 | + return |
| 98 | + name = model.attrs["name"] |
| 99 | + tok_cfg = model._init_tokenizer_config |
| 100 | + trf_cfg = model._init_transformer_config |
| 101 | + hf_model = huggingface_from_pretrained( |
| 102 | + name, tok_cfg, trf_cfg, model_cls=AutoModelForSequenceClassification |
| 103 | + ) |
| 104 | + model.attrs["set_transformer"](model, hf_model) |
| 105 | + tokenizer = model.tokenizer |
| 106 | + texts = ["hello world", "foo bar"] |
| 107 | + token_data = huggingface_tokenize(tokenizer, texts) |
| 108 | + wordpieces = WordpieceBatch.from_batch_encoding(token_data) |
| 109 | + model.layers[0].initialize(X=wordpieces) |
| 110 | + |
| 111 | + model = ClassificationTransformerModel( |
| 112 | + "sgugger/tiny-distilbert-classification", |
| 113 | + get_spans=partial(get_strided_spans, window=128, stride=96), |
| 114 | + ) |
| 115 | + model.initialize() |
| 116 | + |
| 117 | + assert isinstance(model.transformer, DistilBertForSequenceClassification) |
| 118 | + nlp = spacy.blank("en") |
| 119 | + doc = nlp.make_doc("some text") |
| 120 | + assert isinstance(model.predict([doc]).model_output, SequenceClassifierOutput) |
| 121 | + |
| 122 | + b = model.to_bytes() |
| 123 | + model_re = ClassificationTransformerModel( |
| 124 | + "sgugger/tiny-distilbert-classification", |
| 125 | + get_spans=partial(get_strided_spans, window=128, stride=96), |
| 126 | + ).from_bytes(b) |
| 127 | + assert isinstance(model_re.transformer, DistilBertForSequenceClassification) |
| 128 | + assert isinstance(model_re.predict([doc]).model_output, SequenceClassifierOutput) |
| 129 | + assert torch.equal( |
| 130 | + model.predict([doc]).model_output.logits, |
| 131 | + model_re.predict([doc]).model_output.logits, |
| 132 | + ) |
| 133 | + # Note that model.to_bytes() != model_re.to_bytes(), but this is also not |
| 134 | + # true for the default models. |
0 commit comments