[ExecuTorch][WebGPU] Add q8ta_linear op (int8 quantized linear)#21197
[ExecuTorch][WebGPU] Add q8ta_linear op (int8 quantized linear)#21197JCNTH wants to merge 1 commit into
Conversation
🔗 Helpful Links🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/21197
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: ❌ 33 New Failures, 3 Unrelated FailuresAs of commit 4a0ab21 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 linear — the core C2 GEMM. A plain
nn.Linearthrough XNNPACK static PT2E lowers to a delegatedquantize_per_tensor -> q8ta_linear -> dequantize_per_tensorsubgraph; the quantize/dequantize are the landed C0 ops, soq8ta_linearis the one missing piece to run a quantized linear end-to-end.Solution: Port
et_vk.q8ta_linear(int8 activation x int8 per-channel weight -> int8), pluset_vk.q8ta_linear_gemv(identical schema, the M==1 case) on the same handler. Register-tiled i32 GEMM:acc = Σ_k (x_int8[m,k] - input_zero_point) * weight_int8[n,k], dequantizeacc * input_scale * weight_scales[n] + bias, requantizeclamp(round(v * inv_output_scale) + output_zero_point, -128, 127), packed 4 int8 per word. Mirrors Vulkanq8ta_linear.glsl+ the landed q4gsw register tiling.Implementation:
Q8taLinear.cppregistersq8ta_linear.default+q8ta_linear_gemv.default,out = args.back(), guards int8 x/weight/out +N % 4 == 0(output word alignment) +M*K/N*K% 4 == 0(array binding), fail-loud.activationis read via a newWebGPUGraphstring accessor (strings_/get_string, mirroringints_/doubles_) and throws unless"none"— a relu-fused variant would otherwise silently emit non-relu output.Constraints / divergences from the Vulkan reference (both numerically equivalent, byte-exact-verified): (1) the
weight_sumsarg is unused — the input zero-point correction is folded per-element (Σ(x-zp)·w == Σx·w - zp·Σw) instead of Vulkan's precomputedweight_sums. (2) x/weight are read row-major[M,K]/[N,K](WebGPU's always-buffer convention; C0quantize_per_tensoremits row-major int8 and no prepack runs), not Vulkan's 4H4W-packed/block-transposed layout — consistent with the landed q4gsw.activationscoped to"none"(the XNNPACK-static default; relu-fused is fail-loud, a follow-up).Differential Revision: D112257601