From 098cd3b86b3beafefed92686429938bd6c12d87f Mon Sep 17 00:00:00 2001 From: Dong Wang Date: Wed, 23 Sep 2026 17:20:15 +0200 Subject: [PATCH] perf: run one-frame Wan VAE convolutions as 2D convolutions Image VAEs built on the Wan VAE (Qwen Image 2.1 among them) keep Conv3d weights one frame deep, and an image is a single frame, so each such convolution is a 2D convolution. Metal has no IM2COL_3D, so these fall back to GGML_OP_CONV_3D (#1731), whose Metal kernel accumulates every output element in a scalar loop: a 512x512 Qwen Image 2.1 decode takes over a minute on an M1 Pro. Run them through ggml_ext_conv_2d (im2col + mul_mat) instead. Backends that implement IM2COL_3D compute the same products through the 2D path. --- src/model/vae/wan_vae.hpp | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/src/model/vae/wan_vae.hpp b/src/model/vae/wan_vae.hpp index 035deae324..8d14cf9570 100644 --- a/src/model/vae/wan_vae.hpp +++ b/src/model/vae/wan_vae.hpp @@ -82,6 +82,20 @@ namespace WAN { } x = ggml_ext_pad_ext(ctx->ggml_ctx, ctx->backend, x, lp0, rp0, lp1, rp1, lp2, rp2, 0, 0, ctx->circular_x_enabled, ctx->circular_y_enabled); + if (w->ne[2] == 1 && x->ne[2] == 1 && x->ne[3] == in_channels) { + // One frame through a one-frame-deep kernel is a 2D conv; backends without + // im2col_3d (Metal) otherwise fall back to a much slower direct conv_3d. + if (!ggml_is_contiguous(x)) { + x = ggml_cont(ctx->ggml_ctx, x); + } + ggml_tensor* x2 = ggml_reshape_4d(ctx->ggml_ctx, x, x->ne[0], x->ne[1], in_channels, 1); + ggml_tensor* w2 = ggml_reshape_4d(ctx->ggml_ctx, w, w->ne[0], w->ne[1], in_channels, out_channels); + x2 = ggml_ext_conv_2d(ctx->ggml_ctx, x2, w2, b, + std::get<2>(stride), std::get<1>(stride), 0, 0, + std::get<2>(dilation), std::get<1>(dilation), + ctx->conv2d_direct_enabled); + return ggml_reshape_4d(ctx->ggml_ctx, x2, x2->ne[0], x2->ne[1], 1, out_channels); + } return ggml_ext_conv_3d(ctx->ggml_ctx, ctx->backend, x, w, b, in_channels, std::get<2>(stride), std::get<1>(stride), std::get<0>(stride), 0, 0, 0,