From 6aff79ac2220c13d58a9468a8c7722a5db74f389 Mon Sep 17 00:00:00 2001 From: Akshan Krithick Date: Tue, 9 Jun 2026 08:21:31 -0700 Subject: [PATCH 1/2] refactor unet_1d tests --- tests/models/unets/test_models_unet_1d.py | 257 ++++++---------------- 1 file changed, 66 insertions(+), 191 deletions(-) diff --git a/tests/models/unets/test_models_unet_1d.py b/tests/models/unets/test_models_unet_1d.py index bac017e7e7d3..f058212f7cb9 100644 --- a/tests/models/unets/test_models_unet_1d.py +++ b/tests/models/unets/test_models_unet_1d.py @@ -13,77 +13,37 @@ # See the License for the specific language governing permissions and # limitations under the License. -import unittest - -import pytest import torch from diffusers import UNet1DModel +from diffusers.utils.torch_utils import randn_tensor -from ...testing_utils import ( - backend_manual_seed, - floats_tensor, - slow, - torch_device, -) -from ..test_modeling_common import ModelTesterMixin, UNetTesterMixin - +from ...testing_utils import backend_manual_seed, enable_full_determinism, slow, torch_device +from ..testing_utils import BaseModelTesterConfig, ModelTesterMixin -class UNet1DModelTests(ModelTesterMixin, UNetTesterMixin, unittest.TestCase): - model_class = UNet1DModel - main_input_name = "sample" - @property - def dummy_input(self): - batch_size = 4 - num_features = 14 - seq_len = 16 +enable_full_determinism() - noise = floats_tensor((batch_size, num_features, seq_len)).to(torch_device) - time_step = torch.tensor([10] * batch_size).to(torch_device) - return {"sample": noise, "timestep": time_step} +class UNet1DModelTesterConfig(BaseModelTesterConfig): + @property + def model_class(self): + return UNet1DModel @property - def input_shape(self): - return (4, 14, 16) + def main_input_name(self) -> str: + return "sample" @property - def output_shape(self): + def output_shape(self) -> tuple: return (4, 14, 16) - @unittest.skip("Test not supported.") - def test_ema_training(self): - pass - - @unittest.skip("Test not supported.") - def test_training(self): - pass - - @unittest.skip("Test not supported.") - def test_layerwise_casting_training(self): - pass - - def test_determinism(self): - super().test_determinism() - - def test_outputs_equivalence(self): - super().test_outputs_equivalence() - - def test_from_save_pretrained(self): - super().test_from_save_pretrained() - - def test_from_save_pretrained_variant(self): - super().test_from_save_pretrained_variant() - - def test_model_from_pretrained(self): - super().test_model_from_pretrained() - - def test_output(self): - super().test_output() + @property + def generator(self): + return torch.Generator("cpu").manual_seed(0) - def prepare_init_args_and_inputs_for_common(self): - init_dict = { + def get_init_dict(self) -> dict: + return { "block_out_channels": (8, 8, 16, 16), "in_channels": 14, "out_channels": 14, @@ -97,19 +57,26 @@ def prepare_init_args_and_inputs_for_common(self): "up_block_types": ("UpResnetBlock1D", "UpResnetBlock1D", "UpResnetBlock1D"), "act_fn": "swish", } - inputs_dict = self.dummy_input - return init_dict, inputs_dict + def get_dummy_inputs(self) -> dict: + batch_size = 4 + num_features = 14 + seq_len = 16 + noise = randn_tensor((batch_size, num_features, seq_len), generator=self.generator, device=torch_device) + timestep = torch.tensor([10] * batch_size, device=torch_device) + return {"sample": noise, "timestep": timestep} + + +class TestUNet1DModel(UNet1DModelTesterConfig, ModelTesterMixin): def test_from_pretrained_hub(self): model, loading_info = UNet1DModel.from_pretrained( "bglick13/hopper-medium-v2-value-function-hor32", output_loading_info=True, subfolder="unet" ) - self.assertIsNotNone(model) - self.assertEqual(len(loading_info["missing_keys"]), 0) + assert model is not None + assert len(loading_info["missing_keys"]) == 0 model.to(torch_device) - image = model(**self.dummy_input) - + image = model(**self.get_dummy_inputs()) assert image is not None, "Make sure output is not None" def test_output_pretrained(self): @@ -119,9 +86,7 @@ def test_output_pretrained(self): num_features = model.config.in_channels seq_len = 16 - noise = torch.randn((1, seq_len, num_features)).permute( - 0, 2, 1 - ) # match original, we can update values and remove + noise = torch.randn((1, seq_len, num_features)).permute(0, 2, 1) time_step = torch.full((num_features,), 0) with torch.no_grad(): @@ -131,12 +96,7 @@ def test_output_pretrained(self): # fmt: off expected_output_slice = torch.tensor([-2.137172, 1.1426016, 0.3688687, -0.766922, 0.7303146, 0.11038864, -0.4760633, 0.13270172, 0.02591348]) # fmt: on - self.assertTrue(torch.allclose(output_slice, expected_output_slice, rtol=1e-3)) - - @unittest.skip("Test not supported.") - def test_forward_with_norm_groups(self): - # Not implemented yet for this UNet - pass + assert torch.allclose(output_slice, expected_output_slice, rtol=1e-3) @slow def test_unet_1d_maestro(self): @@ -157,98 +117,26 @@ def test_unet_1d_maestro(self): assert (output_sum - 224.0896).abs() < 0.5 assert (output_max - 0.0607).abs() < 4e-4 - @pytest.mark.xfail( - reason=( - "RuntimeError: 'fill_out' not implemented for 'Float8_e4m3fn'. The error is caused due to certain torch.float8_e4m3fn and torch.float8_e5m2 operations " - "not being supported when using deterministic algorithms (which is what the tests run with). To fix:\n" - "1. Wait for next PyTorch release: https://github.com/pytorch/pytorch/issues/137160.\n" - "2. Unskip this test." - ), - ) - def test_layerwise_casting_inference(self): - super().test_layerwise_casting_inference() - - @pytest.mark.xfail( - reason=( - "RuntimeError: 'fill_out' not implemented for 'Float8_e4m3fn'. The error is caused due to certain torch.float8_e4m3fn and torch.float8_e5m2 operations " - "not being supported when using deterministic algorithms (which is what the tests run with). To fix:\n" - "1. Wait for next PyTorch release: https://github.com/pytorch/pytorch/issues/137160.\n" - "2. Unskip this test." - ), - ) - def test_layerwise_casting_memory(self): - pass - - -class UNetRLModelTests(ModelTesterMixin, UNetTesterMixin, unittest.TestCase): - model_class = UNet1DModel - main_input_name = "sample" +class UNetRLModelTesterConfig(BaseModelTesterConfig): @property - def dummy_input(self): - batch_size = 4 - num_features = 14 - seq_len = 16 - - noise = floats_tensor((batch_size, num_features, seq_len)).to(torch_device) - time_step = torch.tensor([10] * batch_size).to(torch_device) - - return {"sample": noise, "timestep": time_step} + def model_class(self): + return UNet1DModel @property - def input_shape(self): - return (4, 14, 16) + def main_input_name(self) -> str: + return "sample" @property - def output_shape(self): + def output_shape(self) -> tuple: return (4, 14, 1) - def test_determinism(self): - super().test_determinism() - - def test_outputs_equivalence(self): - super().test_outputs_equivalence() - - def test_from_save_pretrained(self): - super().test_from_save_pretrained() - - def test_from_save_pretrained_variant(self): - super().test_from_save_pretrained_variant() - - def test_model_from_pretrained(self): - super().test_model_from_pretrained() - - def test_output(self): - # UNetRL is a value-function is different output shape - init_dict, inputs_dict = self.prepare_init_args_and_inputs_for_common() - model = self.model_class(**init_dict) - model.to(torch_device) - model.eval() - - with torch.no_grad(): - output = model(**inputs_dict) - - if isinstance(output, dict): - output = output.sample - - self.assertIsNotNone(output) - expected_shape = torch.Size((inputs_dict["sample"].shape[0], 1)) - self.assertEqual(output.shape, expected_shape, "Input and output shapes do not match") - - @unittest.skip("Test not supported.") - def test_ema_training(self): - pass - - @unittest.skip("Test not supported.") - def test_training(self): - pass - - @unittest.skip("Test not supported.") - def test_layerwise_casting_training(self): - pass + @property + def generator(self): + return torch.Generator("cpu").manual_seed(0) - def prepare_init_args_and_inputs_for_common(self): - init_dict = { + def get_init_dict(self) -> dict: + return { "in_channels": 14, "out_channels": 14, "down_block_types": ["DownResnetBlock1D", "DownResnetBlock1D", "DownResnetBlock1D", "DownResnetBlock1D"], @@ -264,19 +152,35 @@ def prepare_init_args_and_inputs_for_common(self): "time_embedding_type": "positional", "act_fn": "mish", } - inputs_dict = self.dummy_input - return init_dict, inputs_dict + + def get_dummy_inputs(self) -> dict: + batch_size = 4 + num_features = 14 + seq_len = 16 + noise = randn_tensor((batch_size, num_features, seq_len), generator=self.generator, device=torch_device) + timestep = torch.tensor([10] * batch_size, device=torch_device) + return {"sample": noise, "timestep": timestep} + + +class TestUNetRLModel(UNetRLModelTesterConfig, ModelTesterMixin): + # UNetRL is a value function, so it has a different output shape. + def test_output(self): + model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + + with torch.no_grad(): + output = model(**self.get_dummy_inputs()).sample + + assert output.shape == (self.get_dummy_inputs()["sample"].shape[0], 1), "Input and output shapes do not match" def test_from_pretrained_hub(self): value_function, vf_loading_info = UNet1DModel.from_pretrained( "bglick13/hopper-medium-v2-value-function-hor32", output_loading_info=True, subfolder="value_function" ) - self.assertIsNotNone(value_function) - self.assertEqual(len(vf_loading_info["missing_keys"]), 0) + assert value_function is not None + assert len(vf_loading_info["missing_keys"]) == 0 value_function.to(torch_device) - image = value_function(**self.dummy_input) - + image = value_function(**self.get_dummy_inputs()) assert image is not None, "Make sure output is not None" def test_output_pretrained(self): @@ -288,9 +192,7 @@ def test_output_pretrained(self): num_features = value_function.config.in_channels seq_len = 14 - noise = torch.randn((1, seq_len, num_features)).permute( - 0, 2, 1 - ) # match original, we can update values and remove + noise = torch.randn((1, seq_len, num_features)).permute(0, 2, 1) time_step = torch.full((num_features,), 0) with torch.no_grad(): @@ -299,31 +201,4 @@ def test_output_pretrained(self): # fmt: off expected_output_slice = torch.tensor([165.25] * seq_len) # fmt: on - self.assertTrue(torch.allclose(output, expected_output_slice, rtol=1e-3)) - - @unittest.skip("Test not supported.") - def test_forward_with_norm_groups(self): - # Not implemented yet for this UNet - pass - - @pytest.mark.xfail( - reason=( - "RuntimeError: 'fill_out' not implemented for 'Float8_e4m3fn'. The error is caused due to certain torch.float8_e4m3fn and torch.float8_e5m2 operations " - "not being supported when using deterministic algorithms (which is what the tests run with). To fix:\n" - "1. Wait for next PyTorch release: https://github.com/pytorch/pytorch/issues/137160.\n" - "2. Unskip this test." - ), - ) - def test_layerwise_casting_inference(self): - pass - - @pytest.mark.xfail( - reason=( - "RuntimeError: 'fill_out' not implemented for 'Float8_e4m3fn'. The error is caused due to certain torch.float8_e4m3fn and torch.float8_e5m2 operations " - "not being supported when using deterministic algorithms (which is what the tests run with). To fix:\n" - "1. Wait for next PyTorch release: https://github.com/pytorch/pytorch/issues/137160.\n" - "2. Unskip this test." - ), - ) - def test_layerwise_casting_memory(self): - pass + assert torch.allclose(output, expected_output_slice, rtol=1e-3) From 39f8abb973ed3638600ad0021d6549360d05cd7d Mon Sep 17 00:00:00 2001 From: Akshan Krithick Date: Tue, 9 Jun 2026 20:16:37 -0700 Subject: [PATCH 2/2] use per-sample output_shape for unet_1d tests --- tests/models/unets/test_models_unet_1d.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/tests/models/unets/test_models_unet_1d.py b/tests/models/unets/test_models_unet_1d.py index f058212f7cb9..ee11cba3dd2c 100644 --- a/tests/models/unets/test_models_unet_1d.py +++ b/tests/models/unets/test_models_unet_1d.py @@ -36,7 +36,7 @@ def main_input_name(self) -> str: @property def output_shape(self) -> tuple: - return (4, 14, 16) + return (14, 16) @property def generator(self): @@ -129,7 +129,7 @@ def main_input_name(self) -> str: @property def output_shape(self) -> tuple: - return (4, 14, 1) + return (1,) @property def generator(self): @@ -167,10 +167,11 @@ class TestUNetRLModel(UNetRLModelTesterConfig, ModelTesterMixin): def test_output(self): model = self.model_class(**self.get_init_dict()).to(torch_device).eval() + inputs = self.get_dummy_inputs() with torch.no_grad(): - output = model(**self.get_dummy_inputs()).sample + output = model(**inputs).sample - assert output.shape == (self.get_dummy_inputs()["sample"].shape[0], 1), "Input and output shapes do not match" + assert output.shape == (inputs["sample"].shape[0], 1), "Input and output shapes do not match" def test_from_pretrained_hub(self): value_function, vf_loading_info = UNet1DModel.from_pretrained(