Compare commits

...
Author SHA1 Message Date
Morgan Funtowicz e6db36d60e Ensure from_pretrained takes a str in unittest.
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-05-28 17:00:27 +02:00
Morgan Funtowicz 719b9fb6fd Added unittest to save_pretrained when using pathlib.Path
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-05-28 16:45:46 +02:00
Morgan Funtowicz d453872e22 Allow pathlib.Path to be used on save_pretrained and save_vocabulary
Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-05-28 16:35:09 +02:00
2 changed files with 20 additions and 12 deletions
+10 -5
View File
@@ -25,6 +25,7 @@ import re
import warnings
from collections import UserDict, defaultdict
from contextlib import contextmanager
from pathlib import Path
from typing import Any, Dict, List, NamedTuple, Optional, Sequence, Tuple, Union
from tokenizers import AddedToken as AddedTokenFast
@@ -1086,7 +1087,7 @@ class PreTrainedTokenizer(SpecialTokensMixin):
return tokenizer
def save_pretrained(self, save_directory):
def save_pretrained(self, save_directory: Union[str, Path]):
""" Save the tokenizer vocabulary files together with:
- added tokens,
- special-tokens-to-class-attributes-mapping,
@@ -1098,6 +1099,10 @@ class PreTrainedTokenizer(SpecialTokensMixin):
This method make sure the full tokenizer can then be re-loaded using the
:func:`~transformers.PreTrainedTokenizer.from_pretrained` class method.
"""
# Ensure save_directory is a str
save_directory = str(save_directory)
if not os.path.isdir(save_directory):
logger.error("Saving directory ({}) should be a directory".format(save_directory))
return
@@ -1127,7 +1132,7 @@ class PreTrainedTokenizer(SpecialTokensMixin):
return vocab_files + (special_tokens_map_file, added_tokens_file)
def save_vocabulary(self, save_directory) -> Tuple[str]:
def save_vocabulary(self, save_directory: Union[str, Path]) -> Tuple[str]:
""" Save the tokenizer vocabulary to a directory. This method does *NOT* save added tokens
and special token mappings.
@@ -2661,12 +2666,12 @@ class PreTrainedTokenizerFast(PreTrainedTokenizer):
else:
return text
def save_vocabulary(self, save_directory: str) -> Tuple[str]:
def save_vocabulary(self, save_directory: Union[str, Path]) -> Tuple[str]:
if os.path.isdir(save_directory):
files = self._tokenizer.save(save_directory)
files = self._tokenizer.save(str(save_directory))
else:
folder, file = os.path.split(os.path.abspath(save_directory))
files = self._tokenizer.save(folder, name=file)
files = self._tokenizer.save(str(folder), name=file)
return tuple(files)
+10 -7
View File
@@ -19,6 +19,7 @@ import pickle
import shutil
import tempfile
from collections import OrderedDict
from pathlib import Path
from typing import TYPE_CHECKING, Dict, Tuple, Union
from tests.utils import require_tf, require_torch
@@ -122,15 +123,17 @@ class TokenizerTesterMixin:
sample_text = "He is very happy, UNwant\u00E9d,running"
before_tokens = tokenizer.encode(sample_text, add_special_tokens=False)
tokenizer.save_pretrained(self.tmpdirname)
tokenizer = self.tokenizer_class.from_pretrained(self.tmpdirname)
# Test for str and pathlib.Path
for path in [self.tmpdirname, Path(self.tmpdirname)]:
tokenizer.save_pretrained(path)
tokenizer = self.tokenizer_class.from_pretrained(str(path))
after_tokens = tokenizer.encode(sample_text, add_special_tokens=False)
self.assertListEqual(before_tokens, after_tokens)
after_tokens = tokenizer.encode(sample_text, add_special_tokens=False)
self.assertListEqual(before_tokens, after_tokens)
self.assertEqual(tokenizer.max_len, 42)
tokenizer = self.tokenizer_class.from_pretrained(self.tmpdirname, max_len=43)
self.assertEqual(tokenizer.max_len, 43)
self.assertEqual(tokenizer.max_len, 42)
tokenizer = self.tokenizer_class.from_pretrained(str(path), max_len=43)
self.assertEqual(tokenizer.max_len, 43)
def test_pickle_tokenizer(self):
"""Google pickle __getstate__ __setstate__ if you are struggling with this."""