Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6cdaff706e | ||
|
|
904119b9e0 | ||
|
|
353deb4f6a | ||
|
|
6580dbdda4 | ||
|
|
70947826c5 | ||
|
|
b1530b5145 | ||
|
|
664a04e90d | ||
|
|
626727bdee |
No files matched your search
@@ -98,7 +98,7 @@ setup(
|
|||||||
packages=find_packages("src"),
|
packages=find_packages("src"),
|
||||||
install_requires=[
|
install_requires=[
|
||||||
"numpy",
|
"numpy",
|
||||||
"tokenizers == 0.7.0",
|
"tokenizers == 0.8.0.dev1",
|
||||||
# dataclasses for Python versions that don't have it
|
# dataclasses for Python versions that don't have it
|
||||||
"dataclasses;python_version<'3.7'",
|
"dataclasses;python_version<'3.7'",
|
||||||
# filesystem locks e.g. to prevent parallel downloads
|
# filesystem locks e.g. to prevent parallel downloads
|
||||||
|
|||||||
@@ -185,6 +185,15 @@ class BatchEncoding(UserDict):
|
|||||||
|
|
||||||
self._encodings = encoding
|
self._encodings = encoding
|
||||||
|
|
||||||
|
@property
|
||||||
|
def is_fast(self):
|
||||||
|
"""
|
||||||
|
Indicate if this BatchEncoding was generated from the result of a PreTrainedTokenizerFast
|
||||||
|
Returns: True if generated from subclasses of PreTrainedTokenizerFast, else otherwise
|
||||||
|
|
||||||
|
"""
|
||||||
|
return self._encodings is not None
|
||||||
|
|
||||||
def __getitem__(self, item: Union[int, str]) -> EncodingFast:
|
def __getitem__(self, item: Union[int, str]) -> EncodingFast:
|
||||||
""" If the key is a string, get the value of the dict associated to `key` ('input_ids', 'attention_mask'...)
|
""" If the key is a string, get the value of the dict associated to `key` ('input_ids', 'attention_mask'...)
|
||||||
If the key is an integer, get the EncodingFast for batch item with index `key`
|
If the key is an integer, get the EncodingFast for batch item with index `key`
|
||||||
@@ -202,6 +211,16 @@ class BatchEncoding(UserDict):
|
|||||||
def __getattr__(self, item: str):
|
def __getattr__(self, item: str):
|
||||||
return self.data[item]
|
return self.data[item]
|
||||||
|
|
||||||
|
def __getstate__(self):
|
||||||
|
return {"data": self.data, "encodings": self._encodings}
|
||||||
|
|
||||||
|
def __setstate__(self, state):
|
||||||
|
if "data" in state:
|
||||||
|
self.data = state["data"]
|
||||||
|
|
||||||
|
if "encodings" in state:
|
||||||
|
self._encodings = state["encodings"]
|
||||||
|
|
||||||
def keys(self):
|
def keys(self):
|
||||||
return self.data.keys()
|
return self.data.keys()
|
||||||
|
|
||||||
@@ -224,7 +243,7 @@ class BatchEncoding(UserDict):
|
|||||||
"""
|
"""
|
||||||
return self._encodings
|
return self._encodings
|
||||||
|
|
||||||
def tokens(self, batch_index: int = 0) -> List[int]:
|
def tokens(self, batch_index: int = 0) -> List[str]:
|
||||||
if not self._encodings:
|
if not self._encodings:
|
||||||
raise ValueError("tokens() is not available when using Python based tokenizers")
|
raise ValueError("tokens() is not available when using Python based tokenizers")
|
||||||
return self._encodings[batch_index].tokens
|
return self._encodings[batch_index].tokens
|
||||||
|
|||||||
@@ -16,7 +16,7 @@
|
|||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
from transformers import PreTrainedTokenizer
|
from transformers import BertTokenizer, BertTokenizerFast, PreTrainedTokenizer
|
||||||
from transformers.tokenization_gpt2 import GPT2Tokenizer
|
from transformers.tokenization_gpt2 import GPT2Tokenizer
|
||||||
|
|
||||||
from .utils import slow
|
from .utils import slow
|
||||||
@@ -39,3 +39,34 @@ class TokenizerUtilsTest(unittest.TestCase):
|
|||||||
@slow
|
@slow
|
||||||
def test_pretrained_tokenizers(self):
|
def test_pretrained_tokenizers(self):
|
||||||
self.check_tokenizer_from_pretrained(GPT2Tokenizer)
|
self.check_tokenizer_from_pretrained(GPT2Tokenizer)
|
||||||
|
|
||||||
|
def test_batch_encoding_pickle(self):
|
||||||
|
from pickle import loads, dumps
|
||||||
|
|
||||||
|
# Get a slow & a fast tokenizer
|
||||||
|
tok_slow = BertTokenizer.from_pretrained("bert-base-cased")
|
||||||
|
tok_fast = BertTokenizerFast.from_pretrained("bert-base-cased")
|
||||||
|
|
||||||
|
# Encode a sentence
|
||||||
|
be_slow = tok_slow.encode_plus("This is a dummy input sentence")
|
||||||
|
be_fast = tok_fast.encode_plus("This is a dummy input sentence")
|
||||||
|
|
||||||
|
# Make sure both are pickable
|
||||||
|
be_slow_data = dumps(be_slow)
|
||||||
|
be_fast_data = dumps(be_fast)
|
||||||
|
|
||||||
|
# Try to restore
|
||||||
|
be_slow_pickled = loads(be_slow_data)
|
||||||
|
be_fast_pickled = loads(be_fast_data)
|
||||||
|
|
||||||
|
# Ensure pickled objects keeps the is_fast attribute
|
||||||
|
self.assertFalse(be_slow_pickled.is_fast)
|
||||||
|
self.assertTrue(be_fast_pickled.is_fast)
|
||||||
|
|
||||||
|
# Ensure .data match
|
||||||
|
self.assertDictEqual(be_slow_pickled.data, be_slow.data)
|
||||||
|
self.assertDictEqual(be_fast_pickled.data, be_fast.data)
|
||||||
|
|
||||||
|
# Ensure .encodings match
|
||||||
|
self.assertIsNone(be_slow_pickled.encodings)
|
||||||
|
self.assertEqual(len(be_fast_pickled.encodings), len(be_fast.encodings))
|
||||||
Reference in new issue
Block a user