[ExecuTorch][WebGPU] Add q8ta_conv2d op (int8 general conv)#21203
Conversation
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/21203
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 5ac2d5c 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 general (ungrouped) convolution. A standard
nn.Conv2d(groups == 1) through XNNPACK static PT2E lowers to a delegatedquantize_per_tensor -> q8ta_conv2d -> dequantize_per_tensorsubgraph; the quantize/dequantize are the landed C0 ops, soq8ta_conv2dis the one missing piece to run a quantized conv end-to-end. Together with the stagedq8ta_conv2d_pw(1x1) andq8ta_conv2d_dw(depthwise), this covers the int8 conv2d family.Solution: Port
et_vk.q8ta_conv2d(int8 activation x int8 per-channel weight -> int8) as a direct windowed convolution with full input-channel reduction. Per output element:acc = Σ_{ic,kh,kw} (x_int8[n,ic,ih,iw] - input_zero_point) * weight_int8[oc, (kh*Kw+kw)*IC + ic]withih = oh*stride_h - pad_h + kh*dil_h,iw = ow*stride_w - pad_w + kw*dil_w(out-of-bounds taps skipped = zero padding); 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(Q8taConv2d.cpp+q8ta_conv2d.glsl, the direct general-shader path). Vulkan'sq8ta_conv2ddispatches im2col-vs-general; the WebGPU port implements the direct-windowed general path only.Implementation:
Q8taConv2d.cppregistersq8ta_conv2d.default,out = args.back(), one thread per output word = 4 consecutiveW_outpositions of a fixed(n, oc, oh)(W_out % 4 == 0for output word alignment), 2D dispatch fold to lift the 65535 cap. Reads the weight row stride fromweight.dims[1]to tolerate the AOT's align-width padding ofKh*Kw*IC, and reads scales/bias as>= [OC]since 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, andactivation == "none"(all fail-loud).Q8taConvParamsis 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, including at zero-padded taps. (2) weight arrives raw[OC, Kh*Kw*IC]int8 (the AOT pattern's im2col reshapepermute(0,2,3,1); the WebGPU prepack is a passthrough), not Vulkan's int8x4-block texture layout.groups > 1andactivation != "none"are fail-loud (follow-ups);im2colis not ported (the direct general path covers all shapes).Differential Revision: D112257589