From a85958435ac6f4a506e2fc4d0e2f3df706203fb0 Mon Sep 17 00:00:00 2001 From: Will Berman Date: Thu, 18 May 2023 02:20:50 +0000 Subject: [PATCH] parameterize pass single args through tuple --- tests/models/test_models_vae.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tests/models/test_models_vae.py b/tests/models/test_models_vae.py index fd4cf0114f51..9a3e49cdfbc0 100644 --- a/tests/models/test_models_vae.py +++ b/tests/models/test_models_vae.py @@ -321,7 +321,7 @@ def test_stable_diffusion_decode_fp16(self, seed, expected_slice): assert torch_all_close(output_slice, expected_output_slice, atol=5e-3) - @parameterized.expand([13, 16, 27]) + @parameterized.expand([(13,), (16,), (27,)]) @require_torch_gpu @unittest.skipIf(not is_xformers_available(), reason="xformers is not required when using PyTorch 2.0.") def test_stable_diffusion_decode_xformers_vs_2_0_fp16(self, seed): @@ -339,7 +339,7 @@ def test_stable_diffusion_decode_xformers_vs_2_0_fp16(self, seed): assert torch_all_close(sample, sample_2, atol=1e-1) - @parameterized.expand([13, 16, 37]) + @parameterized.expand([(13,), (16,), (37,)]) @require_torch_gpu @unittest.skipIf(not is_xformers_available(), reason="xformers is not required when using PyTorch 2.0.") def test_stable_diffusion_decode_xformers_vs_2_0(self, seed):