diff --git a/tests/kernels/test_conv3d_implicit.py b/tests/kernels/test_conv3d_implicit.py index a2b5dcd69..cd3595152 100644 --- a/tests/kernels/test_conv3d_implicit.py +++ b/tests/kernels/test_conv3d_implicit.py @@ -254,3 +254,95 @@ def test_conv1d_vs_torch(s, stride, padding): assert y.shape == y_ref.shape assert torch.allclose(y, y_ref, rtol=2e-2, atol=2e-2) + + +# ---- Qwen-Image VAE classic conv (T2I T=1) --------------------------------- +# CausalConv3d 3x3x3 degenerates to conv2d with weight[:, :, 2, :, :]. +# Shapes below are the 1024x1024 spatial ladder plus the two hottest layers of +# the 1328x1328 default resolution, taken from forward-hook traces of +# AutoencoderKLQwenImage rather than from the config alone: the decoder halves +# its channel count inside the UpBlock loop (in_dim // 2) before each stage, so +# its ResBlock channel pairs coincide with the encoder ones instead of +# continuing 384 -> 192 -> 96 -> 48. + + +_QWENIMAGE_T1_RES3 = [ + pytest.param(3, 96, 1024, 1024, id="enc_conv_in"), + pytest.param(96, 96, 1024, 1024, id="enc_e0_res__dec_d3_res"), + pytest.param(96, 192, 512, 512, id="enc_e1_res1"), + pytest.param(192, 192, 512, 512, id="enc_e1_res2__dec_d2_res"), + pytest.param(192, 384, 256, 256, id="enc_e2_res1__dec_d1_res1"), + pytest.param(384, 384, 256, 256, id="enc_e2_res2__dec_d1_res"), + pytest.param(384, 384, 128, 128, id="enc_e3_mid__dec_mid_d0"), + pytest.param(384, 32, 128, 128, id="enc_conv_out"), + pytest.param(16, 384, 128, 128, id="dec_conv_in"), + pytest.param(96, 3, 1024, 1024, id="dec_conv_out"), + pytest.param(384, 384, 166, 166, id="dec_bottleneck_1328"), + pytest.param(96, 96, 1328, 1328, id="dec_d3_res_hot_1328"), +] + +# Resample downsample2d: ZeroPad2d((0, 1, 0, 1)) then Conv2d(k=3, s=2, p=0). +_QWENIMAGE_DOWN2D = [ + pytest.param(96, 1024, 1024, id="enc_e0_downsample"), + pytest.param(192, 512, 512, id="enc_e1_downsample_spatial"), + pytest.param(384, 256, 256, id="enc_e2_downsample_spatial"), +] + +# Resample upsample2d/3d: nearest-exact x2 happens outside the kernel, so the +# conv runs at the already-doubled resolution with Conv2d(dim, dim // 2, k=3, +# s=1, p=1). +_QWENIMAGE_UP2D = [ + pytest.param(384, 192, 256, 256, id="dec_d0_upsample"), + pytest.param(384, 192, 512, 512, id="dec_d1_upsample"), + pytest.param(192, 96, 1024, 1024, id="dec_d2_upsample"), +] + + +@_skip_non_cdna4 +@pytest.mark.parametrize("c_in,c_out,h,w", _QWENIMAGE_T1_RES3) +def test_qwenimage_vae_t1_res3_bf16(c_in, c_out, h, w): + torch.manual_seed(8800 + c_in + c_out + h + w) + x2 = torch.randn((1, c_in, h, w), device="cuda", dtype=torch.bfloat16) + weight5 = torch.randn((c_out, c_in, 3, 3, 3), device="cuda", dtype=torch.bfloat16) + weight2 = weight5[:, :, 2, :, :] + bias = torch.randn((c_out,), device="cuda", dtype=torch.float32) + + y = conv3d_implicit(x2, weight2, bias=bias, stride=1, padding=1) + y_ref = F.conv2d(x2, weight2, bias=bias.to(torch.bfloat16), stride=1, padding=1) + torch.cuda.synchronize() + + assert y.shape == y_ref.shape == (1, c_out, h, w) + assert torch.allclose(y, y_ref, rtol=2e-2, atol=2e-2) + + +@_skip_non_cdna4 +@pytest.mark.parametrize("c,h,w", _QWENIMAGE_DOWN2D) +def test_qwenimage_vae_downsample2d_bf16(c, h, w): + torch.manual_seed(9100 + c + h + w) + x2 = torch.randn((1, c, h, w), device="cuda", dtype=torch.bfloat16) + x_pad = F.pad(x2, (0, 1, 0, 1)) + weight = torch.randn((c, c, 3, 3), device="cuda", dtype=torch.bfloat16) + bias = torch.randn((c,), device="cuda", dtype=torch.float32) + + y = conv3d_implicit(x_pad, weight, bias=bias, stride=2, padding=0) + y_ref = F.conv2d(x_pad, weight, bias=bias.to(torch.bfloat16), stride=2, padding=0) + torch.cuda.synchronize() + + assert y.shape == y_ref.shape + assert torch.allclose(y, y_ref, rtol=2e-2, atol=2e-2) + + +@_skip_non_cdna4 +@pytest.mark.parametrize("c_in,c_out,h,w", _QWENIMAGE_UP2D) +def test_qwenimage_vae_upsample2d_bf16(c_in, c_out, h, w): + torch.manual_seed(9400 + c_in + c_out + h + w) + x2 = torch.randn((1, c_in, h, w), device="cuda", dtype=torch.bfloat16) + weight = torch.randn((c_out, c_in, 3, 3), device="cuda", dtype=torch.bfloat16) + bias = torch.randn((c_out,), device="cuda", dtype=torch.float32) + + y = conv3d_implicit(x2, weight, bias=bias, stride=1, padding=1) + y_ref = F.conv2d(x2, weight, bias=bias.to(torch.bfloat16), stride=1, padding=1) + torch.cuda.synchronize() + + assert y.shape == y_ref.shape == (1, c_out, h, w) + assert torch.allclose(y, y_ref, rtol=2e-2, atol=2e-2) diff --git a/tests/kernels/test_conv3d_implicit_fp8.py b/tests/kernels/test_conv3d_implicit_fp8.py index a43ba2864..a9891fdb2 100644 --- a/tests/kernels/test_conv3d_implicit_fp8.py +++ b/tests/kernels/test_conv3d_implicit_fp8.py @@ -202,3 +202,108 @@ def test_fp8_transpose_dword_aligned_spatial(c, t, h, w): ref = x.permute(0, 2, 3, 4, 1).contiguous().view(torch.int8).view(-1) assert torch.equal(got, ref) + + +# ---- Qwen-Image VAE classic conv (T2I T=1) --------------------------------- +# Same 2D-degenerate shapes as test_conv3d_implicit.test_qwenimage_vae_*. + + +_QWENIMAGE_T1_RES3 = [ + pytest.param(3, 96, 1024, 1024, id="enc_conv_in"), + pytest.param(96, 96, 1024, 1024, id="enc_e0_res__dec_d3_res"), + pytest.param(96, 192, 512, 512, id="enc_e1_res1"), + pytest.param(192, 192, 512, 512, id="enc_e1_res2__dec_d2_res"), + pytest.param(192, 384, 256, 256, id="enc_e2_res1__dec_d1_res1"), + pytest.param(384, 384, 256, 256, id="enc_e2_res2__dec_d1_res"), + pytest.param(384, 384, 128, 128, id="enc_e3_mid__dec_mid_d0"), + pytest.param(384, 32, 128, 128, id="enc_conv_out"), + pytest.param(16, 384, 128, 128, id="dec_conv_in"), + pytest.param(96, 3, 1024, 1024, id="dec_conv_out"), + pytest.param(384, 384, 166, 166, id="dec_bottleneck_1328"), + pytest.param(96, 96, 1328, 1328, id="dec_d3_res_hot_1328"), +] + +_QWENIMAGE_DOWN2D = [ + pytest.param(96, 1024, 1024, id="enc_e0_downsample"), + pytest.param(192, 512, 512, id="enc_e1_downsample_spatial"), + pytest.param(384, 256, 256, id="enc_e2_downsample_spatial"), +] + +_QWENIMAGE_UP2D = [ + pytest.param(384, 192, 256, 256, id="dec_d0_upsample"), + pytest.param(384, 192, 512, 512, id="dec_d1_upsample"), + pytest.param(192, 96, 1024, 1024, id="dec_d2_upsample"), +] + + +@_skip_no_fp8 +@pytest.mark.parametrize("c_in,c_out,h,w", _QWENIMAGE_T1_RES3) +def test_qwenimage_vae_t1_res3_fp8(c_in, c_out, h, w): + torch.manual_seed(8900 + c_in + c_out + h + w) + x2 = _fp8(torch.randn((1, c_in, h, w), device="cuda", dtype=torch.bfloat16)) + weight5 = _fp8(torch.randn((c_out, c_in, 3, 3, 3), device="cuda", dtype=torch.bfloat16)) + weight2 = weight5[:, :, 2, :, :] + bias = torch.randn((c_out,), device="cuda", dtype=torch.float32) + + y = conv3d_implicit_fp8(x2, weight2, bias=bias, stride=1, padding=1) + y_ref = F.conv2d( + x2.to(torch.bfloat16), + weight2.to(torch.bfloat16), + bias=bias.to(torch.bfloat16), + stride=1, + padding=1, + ) + torch.cuda.synchronize() + + assert y.shape == y_ref.shape + assert y.dtype == torch.bfloat16 + rel = (y.float() - y_ref.float()).abs().mean() / y_ref.float().abs().mean().clamp_min(1e-6) + assert rel.item() < 2e-2, f"Qwen-Image FP8 rel_err {rel.item():.3e}" + + +@_skip_no_fp8 +@pytest.mark.parametrize("c,h,w", _QWENIMAGE_DOWN2D) +def test_qwenimage_vae_downsample2d_fp8(c, h, w): + torch.manual_seed(9200 + c + h + w) + x2 = _fp8(torch.randn((1, c, h, w), device="cuda", dtype=torch.bfloat16)) + x_pad = F.pad(x2, (0, 1, 0, 1)) + weight = _fp8(torch.randn((c, c, 3, 3), device="cuda", dtype=torch.bfloat16)) + bias = torch.randn((c,), device="cuda", dtype=torch.float32) + + y = conv3d_implicit_fp8(x_pad, weight, bias=bias, stride=2, padding=0) + y_ref = F.conv2d( + x_pad.to(torch.bfloat16), + weight.to(torch.bfloat16), + bias=bias.to(torch.bfloat16), + stride=2, + padding=0, + ) + torch.cuda.synchronize() + + assert y.shape == y_ref.shape + rel = (y.float() - y_ref.float()).abs().mean() / y_ref.float().abs().mean().clamp_min(1e-6) + assert rel.item() < 2e-2, f"Qwen-Image FP8 downsample rel_err {rel.item():.3e}" + + +@_skip_no_fp8 +@pytest.mark.parametrize("c_in,c_out,h,w", _QWENIMAGE_UP2D) +def test_qwenimage_vae_upsample2d_fp8(c_in, c_out, h, w): + torch.manual_seed(9500 + c_in + c_out + h + w) + x2 = _fp8(torch.randn((1, c_in, h, w), device="cuda", dtype=torch.bfloat16)) + weight = _fp8(torch.randn((c_out, c_in, 3, 3), device="cuda", dtype=torch.bfloat16)) + bias = torch.randn((c_out,), device="cuda", dtype=torch.float32) + + y = conv3d_implicit_fp8(x2, weight, bias=bias, stride=1, padding=1) + y_ref = F.conv2d( + x2.to(torch.bfloat16), + weight.to(torch.bfloat16), + bias=bias.to(torch.bfloat16), + stride=1, + padding=1, + ) + torch.cuda.synchronize() + + assert y.shape == y_ref.shape + assert y.dtype == torch.bfloat16 + rel = (y.float() - y_ref.float()).abs().mean() / y_ref.float().abs().mean().clamp_min(1e-6) + assert rel.item() < 2e-2, f"Qwen-Image FP8 upsample rel_err {rel.item():.3e}"