Do not remove half seq length in generation tests (#30016)
* remove seq length from generation tests * style and quality * [test_all] & PR suggestion Co-authored-by: Joao Gante <joaofranciscocardosogante@gmail.com> * Update tests/generation/test_utils.py Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com> * [test all] remove unused variables --------- Co-authored-by: Joao Gante <joaofranciscocardosogante@gmail.com> Co-authored-by: Arthur <48595927+ArthurZucker@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
parent
b4fd49b6c5
commit
b1cd48740e
@@ -686,6 +686,18 @@ class ReformerLocalAttnModelTest(ReformerTesterMixin, GenerationTesterMixin, Mod
|
||||
def test_left_padding_compatibility(self):
|
||||
pass
|
||||
|
||||
def _get_input_ids_and_config(self, batch_size=2):
|
||||
# override because overwise we hit max possible seq length for model (4*8=32)
|
||||
# decreasing the seq_length in tester causes errors for "training_tests", those need exactly max seq length
|
||||
# NOTE: seq_length has to be multiple of 4, otherwise it fails for other tests
|
||||
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
|
||||
input_ids = inputs_dict[self.input_name]
|
||||
input_ids = input_ids[:batch_size, :16]
|
||||
attention_mask = torch.ones_like(input_ids, dtype=torch.long)[:batch_size, :16]
|
||||
config.eos_token_id = None
|
||||
config.forced_eos_token_id = None
|
||||
return config, input_ids, attention_mask
|
||||
|
||||
|
||||
@require_torch
|
||||
class ReformerLSHAttnModelTest(
|
||||
|
||||
Reference in New Issue
Block a user