Skip to content

Commit 326286f

Browse files
authored
Support alternate model loaders in HFShim and HFWrapper (#332)
* Support alternate model loaders in HFShim and HFWrapper In order to support other model loaders such as `AutoModelForSequenceClassification`, make the loading config / model / tokenizer classes configurable in `huggingface_from_pretrained`, `HFWrapper` and `HFShim`. * Simplify and reformat test
1 parent d682d5e commit 326286f

4 files changed

Lines changed: 160 additions & 8 deletions

File tree

spacy_transformers/layers/hf_shim.py

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,14 @@ def __init__(
2424
optimizer: Any = None,
2525
mixed_precision: bool = False,
2626
grad_scaler_config: dict = {},
27+
config_cls = AutoConfig,
28+
model_cls = AutoModel,
29+
tokenizer_cls = AutoTokenizer,
2730
):
2831
self._hfmodel = model
32+
self.config_cls = config_cls
33+
self.model_cls = model_cls
34+
self.tokenizer_cls = tokenizer_cls
2935

3036
# Enable gradient scaling when mixed precision is enabled and gradient
3137
# scaling is not explicitly disabled in the configuration.
@@ -86,18 +92,18 @@ def from_bytes(self, bytes_data):
8692
with make_tempdir() as temp_dir:
8793
config_file = temp_dir / "config.json"
8894
srsly.write_json(config_file, config_dict)
89-
config = AutoConfig.from_pretrained(config_file)
95+
config = self.config_cls.from_pretrained(config_file)
9096
for x, x_bytes in tok_dict.items():
9197
Path(temp_dir / x).write_bytes(x_bytes)
92-
tokenizer = AutoTokenizer.from_pretrained(str(temp_dir.absolute()))
98+
tokenizer = self.tokenizer_cls.from_pretrained(str(temp_dir.absolute()))
9399
vocab_file_contents = None
94100
if hasattr(tokenizer, "vocab_file"):
95101
vocab_file_name = tokenizer.vocab_files_names["vocab_file"]
96102
vocab_file_path = str((temp_dir / vocab_file_name).absolute())
97103
with open(vocab_file_path, "rb") as fileh:
98104
vocab_file_contents = fileh.read()
99105

100-
transformer = AutoModel.from_config(config)
106+
transformer = self.model_cls.from_config(config)
101107
self._hfmodel = HFObjects(
102108
tokenizer,
103109
transformer,

spacy_transformers/layers/hf_wrapper.py

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,10 @@
11
from typing import Callable, Optional, Any
22
from thinc.layers.pytorchwrapper import forward as pt_forward
33
from thinc.layers.pytorchwrapper import convert_pytorch_default_inputs, convert_pytorch_default_outputs
4-
54
from thinc.api import registry, Model
65

6+
from transformers import AutoConfig, AutoModel, AutoTokenizer
7+
78
from .hf_shim import HFShim
89

910

@@ -14,6 +15,9 @@ def HFWrapper(
1415
convert_outputs: Optional[Callable] = None,
1516
mixed_precision: bool = False,
1617
grad_scaler_config: dict = {},
18+
config_cls = AutoConfig,
19+
model_cls = AutoModel,
20+
tokenizer_cls = AutoTokenizer,
1721
) -> Model[Any, Any]:
1822
"""Wrap a PyTorch HF model, so that it has the same API as Thinc models.
1923
To optimize the model, you'll need to create a PyTorch optimizer and call
@@ -50,6 +54,9 @@ def HFWrapper(
5054
hf_model,
5155
mixed_precision=mixed_precision,
5256
grad_scaler_config=grad_scaler_config,
57+
config_cls=config_cls,
58+
model_cls=model_cls,
59+
tokenizer_cls=tokenizer_cls,
5360
)
5461
],
5562
dims={"nI": None, "nO": None},

spacy_transformers/layers/transformer_model.py

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -234,7 +234,12 @@ def backprop(d_model_output: ModelOutput) -> ArgsKwargs:
234234

235235

236236
def huggingface_from_pretrained(
237-
source: Union[Path, str], tok_config: Dict, trf_config: Dict
237+
source: Union[Path, str],
238+
tok_config: Dict,
239+
trf_config: Dict,
240+
config_cls = AutoConfig,
241+
model_cls = AutoModel,
242+
tokenizer_cls = AutoTokenizer,
238243
) -> HFObjects:
239244
"""Create a Huggingface transformer model from pretrained weights. Will
240245
download the model if it is not already downloaded.
@@ -248,14 +253,14 @@ def huggingface_from_pretrained(
248253
str_path = str(source.absolute())
249254
else:
250255
str_path = source
251-
tokenizer = AutoTokenizer.from_pretrained(str_path, **tok_config)
256+
tokenizer = tokenizer_cls.from_pretrained(str_path, **tok_config)
252257
vocab_file_contents = None
253258
if hasattr(tokenizer, "vocab_file"):
254259
with open(tokenizer.vocab_file, "rb") as fileh:
255260
vocab_file_contents = fileh.read()
256261
trf_config["return_dict"] = True
257-
config = AutoConfig.from_pretrained(str_path, **trf_config)
258-
transformer = AutoModel.from_pretrained(str_path, config=config)
262+
config = config_cls.from_pretrained(str_path, **trf_config)
263+
transformer = model_cls.from_pretrained(str_path, config=config)
259264
ops = get_current_ops()
260265
if isinstance(ops, CupyOps):
261266
transformer.cuda()
Lines changed: 134 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,134 @@
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

Comments
 (0)