Skip to content

Commit 941a591

Browse files
authored
Pass excludes when serializing vocab (#8824)
* Pass excludes when serializing vocab Additional minor bug fix: * Deserialize vocab in `EntityLinker.from_disk` * Add test for excluding strings on load * Fix formatting
1 parent 175847f commit 941a591

8 files changed

Lines changed: 45 additions & 26 deletions

File tree

spacy/language.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1909,7 +1909,7 @@ def to_disk(
19091909
if not hasattr(proc, "to_disk"):
19101910
continue
19111911
serializers[name] = lambda p, proc=proc: proc.to_disk(p, exclude=["vocab"])
1912-
serializers["vocab"] = lambda p: self.vocab.to_disk(p)
1912+
serializers["vocab"] = lambda p: self.vocab.to_disk(p, exclude=exclude)
19131913
util.to_disk(path, serializers, exclude)
19141914

19151915
def from_disk(
@@ -1940,7 +1940,7 @@ def deserialize_meta(path: Path) -> None:
19401940

19411941
def deserialize_vocab(path: Path) -> None:
19421942
if path.exists():
1943-
self.vocab.from_disk(path)
1943+
self.vocab.from_disk(path, exclude=exclude)
19441944

19451945
path = util.ensure_path(path)
19461946
deserializers = {}
@@ -1978,7 +1978,7 @@ def to_bytes(self, *, exclude: Iterable[str] = SimpleFrozenList()) -> bytes:
19781978
DOCS: https://spacy.io/api/language#to_bytes
19791979
"""
19801980
serializers = {}
1981-
serializers["vocab"] = lambda: self.vocab.to_bytes()
1981+
serializers["vocab"] = lambda: self.vocab.to_bytes(exclude=exclude)
19821982
serializers["tokenizer"] = lambda: self.tokenizer.to_bytes(exclude=["vocab"])
19831983
serializers["meta.json"] = lambda: srsly.json_dumps(self.meta)
19841984
serializers["config.cfg"] = lambda: self.config.to_bytes()
@@ -2014,7 +2014,7 @@ def deserialize_meta(b):
20142014
b, interpolate=False
20152015
)
20162016
deserializers["meta.json"] = deserialize_meta
2017-
deserializers["vocab"] = self.vocab.from_bytes
2017+
deserializers["vocab"] = lambda b: self.vocab.from_bytes(b, exclude=exclude)
20182018
deserializers["tokenizer"] = lambda b: self.tokenizer.from_bytes(
20192019
b, exclude=["vocab"]
20202020
)

spacy/pipeline/attributeruler.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -276,7 +276,7 @@ def to_bytes(self, exclude: Iterable[str] = SimpleFrozenList()) -> bytes:
276276
DOCS: https://spacy.io/api/attributeruler#to_bytes
277277
"""
278278
serialize = {}
279-
serialize["vocab"] = self.vocab.to_bytes
279+
serialize["vocab"] = lambda: self.vocab.to_bytes(exclude=exclude)
280280
serialize["patterns"] = lambda: srsly.msgpack_dumps(self.patterns)
281281
return util.to_bytes(serialize, exclude)
282282

@@ -296,7 +296,7 @@ def load_patterns(b):
296296
self.add_patterns(srsly.msgpack_loads(b))
297297

298298
deserialize = {
299-
"vocab": lambda b: self.vocab.from_bytes(b),
299+
"vocab": lambda b: self.vocab.from_bytes(b, exclude=exclude),
300300
"patterns": load_patterns,
301301
}
302302
util.from_bytes(bytes_data, deserialize, exclude)
@@ -313,7 +313,7 @@ def to_disk(
313313
DOCS: https://spacy.io/api/attributeruler#to_disk
314314
"""
315315
serialize = {
316-
"vocab": lambda p: self.vocab.to_disk(p),
316+
"vocab": lambda p: self.vocab.to_disk(p, exclude=exclude),
317317
"patterns": lambda p: srsly.write_msgpack(p, self.patterns),
318318
}
319319
util.to_disk(path, serialize, exclude)
@@ -334,7 +334,7 @@ def load_patterns(p):
334334
self.add_patterns(srsly.read_msgpack(p))
335335

336336
deserialize = {
337-
"vocab": lambda p: self.vocab.from_disk(p),
337+
"vocab": lambda p: self.vocab.from_disk(p, exclude=exclude),
338338
"patterns": load_patterns,
339339
}
340340
util.from_disk(path, deserialize, exclude)

spacy/pipeline/entity_linker.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -412,7 +412,7 @@ def to_bytes(self, *, exclude=tuple()):
412412
serialize = {}
413413
if hasattr(self, "cfg") and self.cfg is not None:
414414
serialize["cfg"] = lambda: srsly.json_dumps(self.cfg)
415-
serialize["vocab"] = self.vocab.to_bytes
415+
serialize["vocab"] = lambda: self.vocab.to_bytes(exclude=exclude)
416416
serialize["kb"] = self.kb.to_bytes
417417
serialize["model"] = self.model.to_bytes
418418
return util.to_bytes(serialize, exclude)
@@ -436,7 +436,7 @@ def load_model(b):
436436
deserialize = {}
437437
if hasattr(self, "cfg") and self.cfg is not None:
438438
deserialize["cfg"] = lambda b: self.cfg.update(srsly.json_loads(b))
439-
deserialize["vocab"] = lambda b: self.vocab.from_bytes(b)
439+
deserialize["vocab"] = lambda b: self.vocab.from_bytes(b, exclude=exclude)
440440
deserialize["kb"] = lambda b: self.kb.from_bytes(b)
441441
deserialize["model"] = load_model
442442
util.from_bytes(bytes_data, deserialize, exclude)
@@ -453,7 +453,7 @@ def to_disk(
453453
DOCS: https://spacy.io/api/entitylinker#to_disk
454454
"""
455455
serialize = {}
456-
serialize["vocab"] = lambda p: self.vocab.to_disk(p)
456+
serialize["vocab"] = lambda p: self.vocab.to_disk(p, exclude=exclude)
457457
serialize["cfg"] = lambda p: srsly.write_json(p, self.cfg)
458458
serialize["kb"] = lambda p: self.kb.to_disk(p)
459459
serialize["model"] = lambda p: self.model.to_disk(p)
@@ -480,6 +480,7 @@ def load_model(p):
480480

481481
deserialize = {}
482482
deserialize["cfg"] = lambda p: self.cfg.update(deserialize_config(p))
483+
deserialize["vocab"] = lambda p: self.vocab.from_disk(p, exclude=exclude)
483484
deserialize["kb"] = lambda p: self.kb.from_disk(p)
484485
deserialize["model"] = load_model
485486
util.from_disk(path, deserialize, exclude)

spacy/pipeline/lemmatizer.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -269,7 +269,7 @@ def to_disk(
269269
DOCS: https://spacy.io/api/lemmatizer#to_disk
270270
"""
271271
serialize = {}
272-
serialize["vocab"] = lambda p: self.vocab.to_disk(p)
272+
serialize["vocab"] = lambda p: self.vocab.to_disk(p, exclude=exclude)
273273
serialize["lookups"] = lambda p: self.lookups.to_disk(p)
274274
util.to_disk(path, serialize, exclude)
275275

@@ -285,7 +285,7 @@ def from_disk(
285285
DOCS: https://spacy.io/api/lemmatizer#from_disk
286286
"""
287287
deserialize = {}
288-
deserialize["vocab"] = lambda p: self.vocab.from_disk(p)
288+
deserialize["vocab"] = lambda p: self.vocab.from_disk(p, exclude=exclude)
289289
deserialize["lookups"] = lambda p: self.lookups.from_disk(p)
290290
util.from_disk(path, deserialize, exclude)
291291
self._validate_tables()
@@ -300,7 +300,7 @@ def to_bytes(self, *, exclude: Iterable[str] = SimpleFrozenList()) -> bytes:
300300
DOCS: https://spacy.io/api/lemmatizer#to_bytes
301301
"""
302302
serialize = {}
303-
serialize["vocab"] = self.vocab.to_bytes
303+
serialize["vocab"] = lambda: self.vocab.to_bytes(exclude=exclude)
304304
serialize["lookups"] = self.lookups.to_bytes
305305
return util.to_bytes(serialize, exclude)
306306

@@ -316,7 +316,7 @@ def from_bytes(
316316
DOCS: https://spacy.io/api/lemmatizer#from_bytes
317317
"""
318318
deserialize = {}
319-
deserialize["vocab"] = lambda b: self.vocab.from_bytes(b)
319+
deserialize["vocab"] = lambda b: self.vocab.from_bytes(b, exclude=exclude)
320320
deserialize["lookups"] = lambda b: self.lookups.from_bytes(b)
321321
util.from_bytes(bytes_data, deserialize, exclude)
322322
self._validate_tables()

spacy/pipeline/trainable_pipe.pyx

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -273,7 +273,7 @@ cdef class TrainablePipe(Pipe):
273273
serialize = {}
274274
if hasattr(self, "cfg") and self.cfg is not None:
275275
serialize["cfg"] = lambda: srsly.json_dumps(self.cfg)
276-
serialize["vocab"] = self.vocab.to_bytes
276+
serialize["vocab"] = lambda: self.vocab.to_bytes(exclude=exclude)
277277
serialize["model"] = self.model.to_bytes
278278
return util.to_bytes(serialize, exclude)
279279
@@ -296,7 +296,7 @@ cdef class TrainablePipe(Pipe):
296296
deserialize = {}
297297
if hasattr(self, "cfg") and self.cfg is not None:
298298
deserialize["cfg"] = lambda b: self.cfg.update(srsly.json_loads(b))
299-
deserialize["vocab"] = lambda b: self.vocab.from_bytes(b)
299+
deserialize["vocab"] = lambda b: self.vocab.from_bytes(b, exclude=exclude)
300300
deserialize["model"] = load_model
301301
util.from_bytes(bytes_data, deserialize, exclude)
302302
return self
@@ -313,7 +313,7 @@ cdef class TrainablePipe(Pipe):
313313
serialize = {}
314314
if hasattr(self, "cfg") and self.cfg is not None:
315315
serialize["cfg"] = lambda p: srsly.write_json(p, self.cfg)
316-
serialize["vocab"] = lambda p: self.vocab.to_disk(p)
316+
serialize["vocab"] = lambda p: self.vocab.to_disk(p, exclude=exclude)
317317
serialize["model"] = lambda p: self.model.to_disk(p)
318318
util.to_disk(path, serialize, exclude)
319319
@@ -338,7 +338,7 @@ cdef class TrainablePipe(Pipe):
338338
deserialize = {}
339339
if hasattr(self, "cfg") and self.cfg is not None:
340340
deserialize["cfg"] = lambda p: self.cfg.update(deserialize_config(p))
341-
deserialize["vocab"] = lambda p: self.vocab.from_disk(p)
341+
deserialize["vocab"] = lambda p: self.vocab.from_disk(p, exclude=exclude)
342342
deserialize["model"] = load_model
343343
util.from_disk(path, deserialize, exclude)
344344
return self

spacy/pipeline/transition_parser.pyx

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -569,15 +569,15 @@ cdef class Parser(TrainablePipe):
569569
def to_disk(self, path, exclude=tuple()):
570570
serializers = {
571571
"model": lambda p: (self.model.to_disk(p) if self.model is not True else True),
572-
"vocab": lambda p: self.vocab.to_disk(p),
572+
"vocab": lambda p: self.vocab.to_disk(p, exclude=exclude),
573573
"moves": lambda p: self.moves.to_disk(p, exclude=["strings"]),
574574
"cfg": lambda p: srsly.write_json(p, self.cfg)
575575
}
576576
util.to_disk(path, serializers, exclude)
577577

578578
def from_disk(self, path, exclude=tuple()):
579579
deserializers = {
580-
"vocab": lambda p: self.vocab.from_disk(p),
580+
"vocab": lambda p: self.vocab.from_disk(p, exclude=exclude),
581581
"moves": lambda p: self.moves.from_disk(p, exclude=["strings"]),
582582
"cfg": lambda p: self.cfg.update(srsly.read_json(p)),
583583
"model": lambda p: None,
@@ -597,15 +597,15 @@ cdef class Parser(TrainablePipe):
597597
def to_bytes(self, exclude=tuple()):
598598
serializers = {
599599
"model": lambda: (self.model.to_bytes()),
600-
"vocab": lambda: self.vocab.to_bytes(),
600+
"vocab": lambda: self.vocab.to_bytes(exclude=exclude),
601601
"moves": lambda: self.moves.to_bytes(exclude=["strings"]),
602602
"cfg": lambda: srsly.json_dumps(self.cfg, indent=2, sort_keys=True)
603603
}
604604
return util.to_bytes(serializers, exclude)
605605

606606
def from_bytes(self, bytes_data, exclude=tuple()):
607607
deserializers = {
608-
"vocab": lambda b: self.vocab.from_bytes(b),
608+
"vocab": lambda b: self.vocab.from_bytes(b, exclude=exclude),
609609
"moves": lambda b: self.moves.from_bytes(b, exclude=["strings"]),
610610
"cfg": lambda b: self.cfg.update(srsly.json_loads(b)),
611611
"model": lambda b: None,

spacy/tests/serialize/test_serialize_pipeline.py

Lines changed: 19 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
import pytest
2-
from spacy import registry, Vocab
2+
from spacy import registry, Vocab, load
33
from spacy.pipeline import Tagger, DependencyParser, EntityRecognizer
44
from spacy.pipeline import TextCategorizer, SentenceRecognizer, TrainablePipe
55
from spacy.pipeline.dep_parser import DEFAULT_PARSER_MODEL
@@ -268,3 +268,21 @@ def __init__(self, vocab, model):
268268
pipe.to_disk(d)
269269
new_pipe = CustomPipe(Vocab(), Linear()).from_disk(d)
270270
assert new_pipe.to_bytes() == pipe_bytes
271+
272+
273+
def test_load_without_strings():
274+
nlp = spacy.blank("en")
275+
orig_strings_length = len(nlp.vocab.strings)
276+
word = "unlikely_word_" * 20
277+
nlp.vocab.strings.add(word)
278+
assert len(nlp.vocab.strings) == orig_strings_length + 1
279+
with make_tempdir() as d:
280+
nlp.to_disk(d)
281+
# reload with strings
282+
reloaded_nlp = load(d)
283+
assert len(nlp.vocab.strings) == len(reloaded_nlp.vocab.strings)
284+
assert word in reloaded_nlp.vocab.strings
285+
# reload without strings
286+
reloaded_nlp = load(d, exclude=["strings"])
287+
assert orig_strings_length == len(reloaded_nlp.vocab.strings)
288+
assert word not in reloaded_nlp.vocab.strings

spacy/tokenizer.pyx

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -765,7 +765,7 @@ cdef class Tokenizer:
765765
DOCS: https://spacy.io/api/tokenizer#to_bytes
766766
"""
767767
serializers = {
768-
"vocab": lambda: self.vocab.to_bytes(),
768+
"vocab": lambda: self.vocab.to_bytes(exclude=exclude),
769769
"prefix_search": lambda: _get_regex_pattern(self.prefix_search),
770770
"suffix_search": lambda: _get_regex_pattern(self.suffix_search),
771771
"infix_finditer": lambda: _get_regex_pattern(self.infix_finditer),
@@ -786,7 +786,7 @@ cdef class Tokenizer:
786786
"""
787787
data = {}
788788
deserializers = {
789-
"vocab": lambda b: self.vocab.from_bytes(b),
789+
"vocab": lambda b: self.vocab.from_bytes(b, exclude=exclude),
790790
"prefix_search": lambda b: data.setdefault("prefix_search", b),
791791
"suffix_search": lambda b: data.setdefault("suffix_search", b),
792792
"infix_finditer": lambda b: data.setdefault("infix_finditer", b),

0 commit comments

Comments
 (0)