Compare commits

...
Author SHA1 Message Date
Morgan Funtowicz a9aae9e3dd Increase the default max_length parameter when using TransfoXL & XLnet.
This is needed because the prepended PADDING_TEXT constant is already bigger than the default max_length parameter on the generate method, thus leading to no token being generated.

Signed-off-by: Morgan Funtowicz <morgan@huggingface.co>
2020-06-23 11:10:04 +02:00
+5
View File
@@ -624,6 +624,11 @@ class TextGenerationPipeline(Pipeline):
inputs = self._parse_and_tokenize(
self.PADDING_TEXT + prompt_text, pad_to_max_length=False, add_special_tokens=False
)
# PADDING_TEXT will prepend ~~ 160 tokens to the input which will be bigger
# than the default max_len parameter of the model.generate() function.
# This line ensure we are generating at least 20 tokens when max_length is not defined
generate_kwargs["max_length"] = inputs.input_ids.shape[-1] + generate_kwargs.get("max_length", 20)
else:
inputs = self._parse_and_tokenize(prompt_text, pad_to_max_length=False, add_special_tokens=False)