[ExecuTorch][WebGPU] Add q8ta_conv2d_transposed op (int8 transposed conv)#21205
[ExecuTorch][WebGPU] Add q8ta_conv2d_transposed op (int8 transposed conv)#21205JCNTH wants to merge 1 commit into
Conversation
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/21205
Note: Links to docs will display an error until the docs builds have been completed. ❗ 1 Active SEVsThere are 1 currently active SEVs. If your PR is affected, please view them below: ❌ 46 New Failures, 3 Unrelated FailuresAs of commit 75b891d with merge base 266e0dc ( NEW FAILURES - The following jobs have failed:
FLAKY - The following jobs failed but were likely due to flakiness present on trunk:
BROKEN TRUNK - The following job failed but were present on the merge base:👉 Rebase onto the `viable/strict` branch to avoid these failures
This comment was automatically generated by Dr. CI and updates every 15 minutes. |
This PR needs a
|
psiddh
left a comment
There was a problem hiding this comment.
Approving full WebGPU stack
Stack from ghstack (oldest at bottom):
Problem: The WebGPU delegate has no quantized transposed convolution (deconvolution). An
nn.ConvTranspose2dthrough XNNPACK static PT2E lowers to a delegatedquantize_per_tensor -> q8ta_conv2d_transposed -> dequantize_per_tensorsubgraph; the quantize/dequantize are the landed C0 ops, soq8ta_conv2d_transposedis the one missing piece to run a quantized deconv end-to-end. This completes the int8 conv2d family (pointwise, depthwise, general, transposed).Solution: Port
et_vk.q8ta_conv2d_transposed(int8 activation x int8 per-channel weight -> int8) as a direct gather. Per output element:acc = Σ_{ic,kh,kw} (x_int8[n,ic,ih,iw] - input_zero_point) * weight_int8[oc, (kh*Kw+kw)*IC + ic]whereih = (oh + pad_h - kh*dil_h) / stride_his taken only when the numerator is non-negative and divisible bystride_handih < H_in(iwanalogously); dequantizeacc * input_scale * weight_scales[oc] + bias, requantizeclamp(round(v * inv_output_scale) + output_zero_point, -128, 127), packed 4 int8 per word. Mirrors Vulkanq8ta_conv2d_transposed(Q8taConv2dTransposed.cpp+q8ta_conv2d_transposed.glsl); reuses the generalq8ta_conv2dint8 scaffold +[OC, Kh*Kw*IC]weight layout, changing only the index relation to the transposed gather.Implementation:
Q8taConv2dTransposed.cppregistersq8ta_conv2d_transposed.default,out = args.back(), one thread per output word = 4 consecutiveW_outpositions of a fixed(n, oc, oh)(W_out % 4 == 0), 2D dispatch fold. The transposed schema insertsoutput_paddingat arg 12, sodilationis read from arg 13 andgroupsfrom arg 14 (output_paddingitself is unused — the output H/W come from the serialized output dims). The weight row stride is read fromweight.dims[1]to tolerate align-width padding ofKh*Kw*IC; scales/bias are read as>= [OC](the AOT padsOCto a multiple of 4; the shader reads only[0, OC)). Guards[N,IC,H_in,W_in]/[OC,Kh*Kw*IC]/[N,OC,H_out,W_out]ranks, int8 dtypes,groups == 1,dilation == 1,stride >= 1, andactivation == "none"(all fail-loud).Q8taConvTParamsis a 96-byte (16-aligned) uniform.Constraints / divergences from the Vulkan reference (numerically equivalent, golden-verified): (1) the
weight_sumsarg is unused — the input zero-point correction is folded per-element (Σ(x-zp)·w == Σx·w - zp·Σw), matching Vulkan'sweight_sumscompensation exactly (Vulkan fills out-of-bounds taps with the zero-point, which contribute(zp-zp)·w = 0= the skipped taps here). (2) weight arrives raw[OC, Kh*Kw*IC]int8 (the AOT reshape mapsweight[oc,(kh*Kw+kw)*IC+ic] = original[ic,oc,kh,kw]; the WebGPU prepack is a passthrough), not Vulkan's int8x4-block layout.dilation != 1(matching Vulkan's own restriction),groups > 1, andactivation != "none"are fail-loud follow-ups.Differential Revision: D112257642