Skip to content
This repository was archived by the owner on Dec 16, 2022. It is now read-only.

Commit a330876

Browse files
authored
fix seq2seq reader bug (#139)
1 parent dd36890 commit a330876

3 files changed

Lines changed: 27 additions & 4 deletions

File tree

CHANGELOG.md

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
1515
### Fixed
1616

1717
- Fixed BART for latest `transformers` version.
18+
- Fixed a bug with `Seq2SeqDatasetReader` that would cause an exception when
19+
the desired behavior is to not add start or end symbols to either the source or the target
20+
and the default `start_symbol` or `end_symbol` are not part of the tokenizer's vocabulary.
1821

1922
## [v1.1.0](https://github.com/allenai/allennlp-models/releases/tag/v1.1.0) - 2020-09-08
2023

allennlp_models/generation/dataset_readers/seq2seq.py

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@
1111
from allennlp.data.dataset_readers.dataset_reader import DatasetReader
1212
from allennlp.data.fields import TextField
1313
from allennlp.data.instance import Instance
14-
from allennlp.data.tokenizers import Tokenizer, SpacyTokenizer
14+
from allennlp.data.tokenizers import Tokenizer, SpacyTokenizer, Token
1515
from allennlp.data.token_indexers import TokenIndexer, SingleIdTokenIndexer
1616

1717
logger = logging.getLogger(__name__)
@@ -78,13 +78,29 @@ def __init__(
7878
self._target_tokenizer = target_tokenizer or self._source_tokenizer
7979
self._source_token_indexers = source_token_indexers or {"tokens": SingleIdTokenIndexer()}
8080
self._target_token_indexers = target_token_indexers or self._source_token_indexers
81+
8182
self._source_add_start_token = source_add_start_token
8283
self._source_add_end_token = source_add_end_token
8384
self._target_add_start_token = target_add_start_token
8485
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+
88104
self._delimiter = delimiter
89105
self._source_max_tokens = source_max_tokens
90106
self._target_max_tokens = target_max_tokens

tests/generation/dataset_readers/seq2seq_test.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -216,3 +216,7 @@ def test_correct_quote_handling(self, line):
216216
"d",
217217
"@end@",
218218
]
219+
220+
def test_bad_start_or_end_symbol(self):
221+
with pytest.raises(ValueError, match="Bad start or end symbol"):
222+
Seq2SeqDatasetReader(start_symbol="BAD SYMBOL")

0 commit comments

Comments
 (0)