onnxruntime
2a720f92 - [WebGPU] Move several elements per thread in Concat, Slice, Transpose and Pool (#32497)

Commit
11 days ago
[WebGPU] Move several elements per thread in Concat, Slice, Transpose and Pool (#32497) ### Description Four WebGPU kernels — Concat, Slice, Transpose and Pool — are bound by memory traffic, not arithmetic, and each currently gives a thread a single element to move. On an RTX 3060 that leaves most of the available bandwidth unused. Each now works in units of four (or two) elements when the shape allows it. Every offset the shader deals with is expressed in those units, so the index arithmetic inside the shaders is unchanged; only the host-side setup picks the component count and reduces the shapes it passes down. **Concat** — when the concatenated axis is the innermost one, a four-element group must not straddle two inputs, so *every input* has to be a multiple of four there; checking only the output would be wrong. On any other axis the innermost dimension is untouched and the offsets along the concat axis stay in single elements. (Split already does the equivalent on `main`.) **Slice** — vectorized only when the innermost axis is a contiguous forward run starting on a four-element boundary. A negative step or an unaligned start breaks the correspondence between consecutive output and input elements. **Transpose** — the untiled path vectorizes when the permutation leaves the innermost dimension innermost. Separately, the shared-memory path now uses a 32-wide tile with eight workgroup rows instead of a 16×16 tile, which doubles the width of each coalesced global access. A full 32×32 workgroup (1024 threads) was tried and measured slower than the 16×16 it would have replaced. **Pool** — vectorizes across channels in NHWC, where the channel is innermost and pooling never crosses it, so the window arithmetic (which is what this kernel spends its time on) is paid once per group instead of once per element. The parallel-reduction path stays scalar; it reduces one window across a workgroup and has nothing to fold the per-channel grouping into. `int64` stays scalar throughout: it is already stored as `vec2<u32>`. ### Motivation and Context Part of a WebGPU optimization pass on a document-layout model, measured on an RTX 3060 with the Dawn/Vulkan backend. These four kernels together accounted for about 1.2 ms of a 16 ms inference. ### Testing Added cases to the existing Concat, Slice, Transpose and Pool op tests covering both the vectorized shapes and the shapes that must fall back to one element per thread: unaligned Concat inputs, an unaligned Slice start, a negative Slice step, a permutation that moves the innermost dimension, and a channel count that is a multiple of two but not four. The full `PoolTest`, `ConcatOpTest`, `SliceTest` and `TransposeOpTest` suites pass on the WebGPU EP — 198 tests, 191 passed, 7 pre-existing skips, 0 failures. Verified on Windows / MSVC / NVIDIA (Dawn Vulkan backend) only; I do not have other vendors or backends to hand, so CI is the first run on those.
Author
Parents
Loading