Skip to content

Commit 7aacc21

Browse files
authored
Merge pull request #231 from svlandeg/fix/listener_ref
Explicit reference to TransformerListener
2 parents 769becc + eb2d546 commit 7aacc21

4 files changed

Lines changed: 70 additions & 5 deletions

File tree

spacy_transformers/architectures.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -36,10 +36,10 @@ def transformer_listener_tok2vec_v1(
3636
never have multiple upstream Transformer components, so the wildcard
3737
string will almost always be fine.
3838
"""
39-
return chain(
40-
TransformerListener(upstream_name=upstream),
41-
trfs2arrays(pooling, grad_factor)
42-
)
39+
listener = TransformerListener(upstream_name=upstream)
40+
model = chain(listener, trfs2arrays(pooling, grad_factor))
41+
model.set_ref("listener", listener)
42+
return model
4343

4444

4545
@registry.architectures.register("spacy-transformers.Tok2VecTransformer.v1")

spacy_transformers/layers/listener.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@ class TransformerListener(Model):
1515
_outputs: Optional[List[TransformerData]]
1616
_backprop: Optional[Callable[[List[TransformerData]], List[Doc]]]
1717

18-
def __init__(self, upstream_name):
18+
def __init__(self, upstream_name: str):
1919
Model.__init__(self, name=self.name, forward=forward, dims={"nO": None})
2020
self.upstream_name = upstream_name
2121
self._batch_id = None

spacy_transformers/tests/regression/__init__.py

Whitespace-only changes.
Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,65 @@
1+
from spacy.training.example import Example
2+
from spacy.util import make_tempdir
3+
from spacy import util
4+
from thinc.api import Config
5+
6+
7+
TRAIN_DATA = [
8+
("I'm so happy.", {"cats": {"POSITIVE": 1.0, "NEGATIVE": 0.0}}),
9+
("I'm so angry", {"cats": {"POSITIVE": 0.0, "NEGATIVE": 1.0}}),
10+
]
11+
12+
13+
cfg_string = """
14+
[nlp]
15+
lang = "en"
16+
pipeline = ["transformer","textcat"]
17+
18+
[components]
19+
20+
[components.textcat]
21+
factory = "textcat"
22+
23+
[components.textcat.model]
24+
@architectures = "spacy.TextCatEnsemble.v2"
25+
26+
[components.textcat.model.tok2vec]
27+
@architectures = "spacy-transformers.TransformerListener.v1"
28+
grad_factor = 1.0
29+
30+
[components.textcat.model.tok2vec.pooling]
31+
@layers = "reduce_mean.v1"
32+
33+
[components.transformer]
34+
factory = "transformer"
35+
"""
36+
37+
38+
def test_transformer_pipeline_textcat():
39+
"""Test that a pipeline with just a transformer+textcat runs and trains properly.
40+
This used to throw an error because of shape inference issues -
41+
cf https://github.com/explosion/spaCy/issues/6401"""
42+
orig_config = Config().from_str(cfg_string)
43+
nlp = util.load_model_from_config(orig_config, auto_fill=True, validate=True)
44+
assert nlp.pipe_names == ["transformer", "textcat"]
45+
train_examples = []
46+
47+
for text, annotations in TRAIN_DATA:
48+
train_examples.append(Example.from_dict(nlp.make_doc(text), annotations))
49+
optimizer = nlp.initialize(get_examples=lambda: train_examples)
50+
51+
for i in range(2):
52+
losses = {}
53+
nlp.update(train_examples, sgd=optimizer, losses=losses)
54+
55+
doc = nlp("We're interested at underwater basket weaving.")
56+
cats1 = doc.cats
57+
58+
# ensure IO goes OK
59+
with make_tempdir() as d:
60+
file_path = d / "trained_nlp"
61+
nlp.to_disk(file_path)
62+
nlp2 = util.load_model_from_path(file_path)
63+
doc2 = nlp2("We're interested at underwater basket weaving.")
64+
cats2 = doc2.cats
65+
assert cats1 == cats2

0 commit comments

Comments
 (0)