Skip to content

Commit fa2e7a4

Browse files
authored
Fix spancat tests on GPU (#8872)
* Fix spancat tests on GPU * Fix more spancat tests
1 parent 77d698d commit fa2e7a4

1 file changed

Lines changed: 14 additions & 12 deletions

File tree

spacy/tests/pipeline/test_spancat.py

Lines changed: 14 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,11 @@
11
import pytest
2-
from numpy.testing import assert_equal
2+
from numpy.testing import assert_equal, assert_array_equal
3+
from thinc.api import get_current_ops
34
from spacy.language import Language
45
from spacy.training import Example
56
from spacy.util import fix_random_seed, registry
67

8+
OPS = get_current_ops()
79

810
SPAN_KEY = "labeled_spans"
911

@@ -116,22 +118,22 @@ def test_ngram_suggester(en_tokenizer):
116118
for span in spans:
117119
assert 0 <= span[0] < len(doc)
118120
assert 0 < span[1] <= len(doc)
119-
spans_set.add((span[0], span[1]))
121+
spans_set.add((int(span[0]), int(span[1])))
120122
# spans are unique
121123
assert spans.shape[0] == len(spans_set)
122124
offset += ngrams.lengths[i]
123125
# the number of spans is correct
124-
assert_equal(ngrams.lengths, [max(0, len(doc) - (size - 1)) for doc in docs])
126+
assert_array_equal(OPS.to_numpy(ngrams.lengths), [max(0, len(doc) - (size - 1)) for doc in docs])
125127

126128
# test 1-3-gram suggestions
127129
ngram_suggester = registry.misc.get("spacy.ngram_suggester.v1")(sizes=[1, 2, 3])
128130
docs = [
129131
en_tokenizer(text) for text in ["a", "a b", "a b c", "a b c d", "a b c d e"]
130132
]
131133
ngrams = ngram_suggester(docs)
132-
assert_equal(ngrams.lengths, [1, 3, 6, 9, 12])
133-
assert_equal(
134-
ngrams.data,
134+
assert_array_equal(OPS.to_numpy(ngrams.lengths), [1, 3, 6, 9, 12])
135+
assert_array_equal(
136+
OPS.to_numpy(ngrams.data),
135137
[
136138
# doc 0
137139
[0, 1],
@@ -176,13 +178,13 @@ def test_ngram_suggester(en_tokenizer):
176178
ngram_suggester = registry.misc.get("spacy.ngram_suggester.v1")(sizes=[1])
177179
docs = [en_tokenizer(text) for text in ["", "a", ""]]
178180
ngrams = ngram_suggester(docs)
179-
assert_equal(ngrams.lengths, [len(doc) for doc in docs])
181+
assert_array_equal(OPS.to_numpy(ngrams.lengths), [len(doc) for doc in docs])
180182

181183
# test all empty docs
182184
ngram_suggester = registry.misc.get("spacy.ngram_suggester.v1")(sizes=[1])
183185
docs = [en_tokenizer(text) for text in ["", "", ""]]
184186
ngrams = ngram_suggester(docs)
185-
assert_equal(ngrams.lengths, [len(doc) for doc in docs])
187+
assert_array_equal(OPS.to_numpy(ngrams.lengths), [len(doc) for doc in docs])
186188

187189

188190
def test_ngram_sizes(en_tokenizer):
@@ -195,12 +197,12 @@ def test_ngram_sizes(en_tokenizer):
195197
]
196198
ngrams_1 = size_suggester(docs)
197199
ngrams_2 = range_suggester(docs)
198-
assert_equal(ngrams_1.lengths, [1, 3, 6, 9, 12])
199-
assert_equal(ngrams_1.lengths, ngrams_2.lengths)
200-
assert_equal(ngrams_1.data, ngrams_2.data)
200+
assert_array_equal(OPS.to_numpy(ngrams_1.lengths), [1, 3, 6, 9, 12])
201+
assert_array_equal(OPS.to_numpy(ngrams_1.lengths), OPS.to_numpy(ngrams_2.lengths))
202+
assert_array_equal(OPS.to_numpy(ngrams_1.data), OPS.to_numpy(ngrams_2.data))
201203

202204
# one more variation
203205
suggester_factory = registry.misc.get("spacy.ngram_range_suggester.v1")
204206
range_suggester = suggester_factory(min_size=2, max_size=4)
205207
ngrams_3 = range_suggester(docs)
206-
assert_equal(ngrams_3.lengths, [0, 1, 3, 6, 9])
208+
assert_array_equal(OPS.to_numpy(ngrams_3.lengths), [0, 1, 3, 6, 9])

0 commit comments

Comments
 (0)