// Pass 1 of parallel DynamicQuantizeLinear. Regular mode reduces one contiguous // WG * EPT block per workgroup. Grid-stride mode caps the partial count and has // every lane revisit the tensor at grid-sized strides. Both modes use the same // subgroup/tree combine and write one min/max pair per workgroup. The subgroup // combine publishes one shared-memory slot per invocation (its subgroup pair // from the elected lane, the neutral 0.0 from every other lane) and every // subgroup folds all of them. // // The local min/max start at 0.0 because the ONNX DynamicQuantizeLinear range // always includes zero; out-of-range lanes contribute the same neutral value. // f32 min/max is order-independent, so subgroup and workgroup-tree combinations // produce the same extrema. Without subgroups, a full workgroup tree combines // the per-invocation partials. {% set gridStride = gridStride if gridStride is defined else false %} {% if useSubgroups %} enable subgroups; {% endif %} {{ env.wgsl.resourceDeclarations }} const WG: u32 = {{ workgroupSize }}u; {% if not gridStride and not vec4 %} const EPT: u32 = {{ elemsPerThread }}u; {% endif %} var wgMin: array; var wgMax: array; @compute @workgroup_size(WG, 1, 1) fn main(@builtin(workgroup_id) wg: vec3, @builtin(local_invocation_id) lid: vec3 {%- if gridStride %}, @builtin(num_workgroups) nwg: vec3 {%- endif %} {%- if useSubgroups %}, @builtin(subgroup_size) sgSize: u32 {%- endif %}) { let tid = lid.x; {% if gridStride %} let blk = wg.x; // The capped x-only dispatch is the full logical grid. Each workgroup writes // exactly one partial, matching finalize's numPartials uniform. let gridThreads = nwg.x * WG; {% else %} // 2D-folded block index: wg.y carries the high bits past the device's // per-dimension dispatch limit. The whole over-dispatched workgroup returns // before barriers or scratch stores. // blk/numBlocks are workgroup-uniform, so the return can't split the barrier. let blk = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u; {% if vec4 %} let numBlocks = (params.count / 4u + WG - 1u) / WG; {% else %} let numBlocks = (params.count + WG * EPT - 1u) / (WG * EPT); {% endif %} if (blk >= numBlocks) { return; } {% endif %} var localMin = 0.0; var localMax = 0.0; {% if gridStride %} {% if vec4 %} // Adjacent lanes read adjacent vec4s on each grid-stride iteration. let count4 = params.count / 4u; for (var idx4 = blk * WG + tid; idx4 < count4; idx4 = idx4 + gridThreads) { let v = x[idx4]; localMin = min(localMin, min(min(v.x, v.y), min(v.z, v.w))); localMax = max(localMax, max(max(v.x, v.y), max(v.z, v.w))); } {% else %} for (var idx = blk * WG + tid; idx < params.count; idx = idx + gridThreads) { let value = x[idx]; localMin = min(localMin, value); localMax = max(localMax, value); } {% endif %} {% else %} {% if vec4 %} // Each thread reads one contiguous vec4; this path assumes params.count is // divisible by 4. let count4 = params.count / 4u; let idx4 = blk * WG + tid; if (idx4 < count4) { let v = x[idx4]; localMin = min(localMin, min(min(v.x, v.y), min(v.z, v.w))); localMax = max(localMax, max(max(v.x, v.y), max(v.z, v.w))); } {% else %} let base = blk * WG * EPT; for (var e = 0u; e < EPT; e = e + 1u) { let idx = base + e * WG + tid; if (idx < params.count) { let v = x[idx]; localMin = min(localMin, v); localMax = max(localMax, v); } } {% endif %} {% endif %} {% if useSubgroups %} let sgMin = subgroupMin(localMin); let sgMax = subgroupMax(localMax); // Cross-subgroup fold that assumes nothing about which invocations share a // subgroup, how many subgroups there are, or which of a subgroup's lanes are // active: every invocation owns the slot at its own index, the elected lane // publishes its subgroup pair there and every other lane publishes 0.0, which // is an exact identity here because every lane's local range already includes // zero (so every published minimum is <= 0 and every maximum >= 0). Each // subgroup then folds all WG slots — lane `rank`, its dense position among // the active lanes, walks slots rank, rank + count, ... — and one more // collective merges the lane partials, so every slot is merged exactly once // at any legal width and partition. min/max is commutative and associative, // so the fold order does not change the result. var totalMin = sgMin; var totalMax = sgMax; // A one-subgroup workgroup is already fully reduced by the collectives above. // The test reads the `subgroup_size` builtin, which is uniform; a collective's // result is not uniform to WGSL's analysis and may not guard a barrier. if (sgSize != WG) { let rank = subgroupExclusiveAdd(1u); let count = subgroupAdd(1u); let leader = rank == 0u; wgMin[tid] = select(0.0, sgMin, leader); wgMax[tid] = select(0.0, sgMax, leader); workgroupBarrier(); var foldMin = 0.0; var foldMax = 0.0; for (var i = rank; i < WG; i = i + count) { foldMin = min(foldMin, wgMin[i]); foldMax = max(foldMax, wgMax[i]); } totalMin = subgroupMin(foldMin); totalMax = subgroupMax(foldMax); } if (tid == 0u) { partial_min[blk] = totalMin; partial_max[blk] = totalMax; } {% else %} // No-subgroups fallback: full workgroup-tree min/max over per-thread partials. wgMin[tid] = localMin; wgMax[tid] = localMax; workgroupBarrier(); for (var step = WG >> 1u; step > 0u; step = step >> 1u) { if (tid < step) { wgMin[tid] = min(wgMin[tid], wgMin[tid + step]); wgMax[tid] = max(wgMax[tid], wgMax[tid + step]); } workgroupBarrier(); } if (tid == 0u) { partial_min[blk] = wgMin[0]; partial_max[blk] = wgMax[0]; } {% endif %} }