|
11 | 11 | from allennlp.data.dataset_readers.dataset_reader import DatasetReader |
12 | 12 | from allennlp.data.fields import TextField |
13 | 13 | from allennlp.data.instance import Instance |
14 | | -from allennlp.data.tokenizers import Tokenizer, SpacyTokenizer |
| 14 | +from allennlp.data.tokenizers import Tokenizer, SpacyTokenizer, Token |
15 | 15 | from allennlp.data.token_indexers import TokenIndexer, SingleIdTokenIndexer |
16 | 16 |
|
17 | 17 | logger = logging.getLogger(__name__) |
@@ -78,13 +78,29 @@ def __init__( |
78 | 78 | self._target_tokenizer = target_tokenizer or self._source_tokenizer |
79 | 79 | self._source_token_indexers = source_token_indexers or {"tokens": SingleIdTokenIndexer()} |
80 | 80 | self._target_token_indexers = target_token_indexers or self._source_token_indexers |
| 81 | + |
81 | 82 | self._source_add_start_token = source_add_start_token |
82 | 83 | self._source_add_end_token = source_add_end_token |
83 | 84 | self._target_add_start_token = target_add_start_token |
84 | 85 | self._target_add_end_token = target_add_end_token |
85 | | - self._start_token, self._end_token = self._source_tokenizer.tokenize( |
86 | | - start_symbol + " " + end_symbol |
87 | | - ) |
| 86 | + self._start_token: Optional[Token] = None |
| 87 | + self._end_token: Optional[Token] = None |
| 88 | + if ( |
| 89 | + source_add_start_token |
| 90 | + or source_add_end_token |
| 91 | + or target_add_start_token |
| 92 | + or target_add_end_token |
| 93 | + ): |
| 94 | + try: |
| 95 | + self._start_token, self._end_token = self._source_tokenizer.tokenize( |
| 96 | + start_symbol + " " + end_symbol |
| 97 | + ) |
| 98 | + except ValueError: |
| 99 | + raise ValueError( |
| 100 | + f"Bad start or end symbol ({'start_symbol', 'end_symbol'}) " |
| 101 | + f"for tokenizer {self._source_tokenizer}" |
| 102 | + ) |
| 103 | + |
88 | 104 | self._delimiter = delimiter |
89 | 105 | self._source_max_tokens = source_max_tokens |
90 | 106 | self._target_max_tokens = target_max_tokens |
|
0 commit comments