|
22 | 22 | import torch |
23 | 23 |
|
24 | 24 | from diffusers.models.embeddings import get_timestep_embedding |
25 | | -from diffusers.models.resnet import Upsample |
| 25 | +from diffusers.models.resnet import Downsample, Upsample |
26 | 26 | from diffusers.testing_utils import floats_tensor, slow, torch_device |
27 | 27 |
|
28 | 28 |
|
@@ -164,3 +164,58 @@ def test_upsample_with_transpose(self): |
164 | 164 | output_slice = upsampled[0, -1, -3:, -3:] |
165 | 165 | expected_slice = torch.tensor([-0.3028, -0.1582, 0.0071, 0.0350, -0.4799, -0.1139, 0.1056, -0.1153, -0.1046]) |
166 | 166 | assert torch.allclose(output_slice.flatten(), expected_slice, atol=1e-3) |
| 167 | + |
| 168 | + |
| 169 | +class DownsampleBlockTests(unittest.TestCase): |
| 170 | + def test_downsample_default(self): |
| 171 | + torch.manual_seed(0) |
| 172 | + sample = torch.randn(1, 32, 64, 64) |
| 173 | + downsample = Downsample(channels=32, use_conv=False) |
| 174 | + with torch.no_grad(): |
| 175 | + downsampled = downsample(sample) |
| 176 | + |
| 177 | + assert downsampled.shape == (1, 32, 32, 32) |
| 178 | + output_slice = downsampled[0, -1, -3:, -3:] |
| 179 | + expected_slice = torch.tensor([-0.0513, -0.3889, 0.0640, 0.0836, -0.5460, -0.0341, -0.0169, -0.6967, 0.1179]) |
| 180 | + max_diff = (output_slice.flatten() - expected_slice).abs().sum().item() |
| 181 | + assert max_diff <= 1e-3 |
| 182 | + # assert torch.allclose(output_slice.flatten(), expected_slice, atol=1e-1) |
| 183 | + |
| 184 | + def test_downsample_with_conv(self): |
| 185 | + torch.manual_seed(0) |
| 186 | + sample = torch.randn(1, 32, 64, 64) |
| 187 | + downsample = Downsample(channels=32, use_conv=True) |
| 188 | + with torch.no_grad(): |
| 189 | + downsampled = downsample(sample) |
| 190 | + |
| 191 | + assert downsampled.shape == (1, 32, 32, 32) |
| 192 | + output_slice = downsampled[0, -1, -3:, -3:] |
| 193 | + |
| 194 | + expected_slice = torch.tensor( |
| 195 | + [0.9267, 0.5878, 0.3337, 1.2321, -0.1191, -0.3984, -0.7532, -0.0715, -0.3913], |
| 196 | + ) |
| 197 | + assert torch.allclose(output_slice.flatten(), expected_slice, atol=1e-3) |
| 198 | + |
| 199 | + def test_downsample_with_conv_pad1(self): |
| 200 | + torch.manual_seed(0) |
| 201 | + sample = torch.randn(1, 32, 64, 64) |
| 202 | + downsample = Downsample(channels=32, use_conv=True, padding=1) |
| 203 | + with torch.no_grad(): |
| 204 | + downsampled = downsample(sample) |
| 205 | + |
| 206 | + assert downsampled.shape == (1, 32, 32, 32) |
| 207 | + output_slice = downsampled[0, -1, -3:, -3:] |
| 208 | + expected_slice = torch.tensor([0.9267, 0.5878, 0.3337, 1.2321, -0.1191, -0.3984, -0.7532, -0.0715, -0.3913]) |
| 209 | + assert torch.allclose(output_slice.flatten(), expected_slice, atol=1e-3) |
| 210 | + |
| 211 | + def test_downsample_with_conv_out_dim(self): |
| 212 | + torch.manual_seed(0) |
| 213 | + sample = torch.randn(1, 32, 64, 64) |
| 214 | + downsample = Downsample(channels=32, use_conv=True, out_channels=16) |
| 215 | + with torch.no_grad(): |
| 216 | + downsampled = downsample(sample) |
| 217 | + |
| 218 | + assert downsampled.shape == (1, 16, 32, 32) |
| 219 | + output_slice = downsampled[0, -1, -3:, -3:] |
| 220 | + expected_slice = torch.tensor([-0.6586, 0.5985, 0.0721, 0.1256, -0.1492, 0.4436, -0.2544, 0.5021, 1.1522]) |
| 221 | + assert torch.allclose(output_slice.flatten(), expected_slice, atol=1e-3) |
0 commit comments