Download build/webgpu/dynamic-quantize-linear-quantize.wgsl.jinja from webgpu-kernels/ai.onnx.DynamicQuantizeLinear: direct link, hf CLI and curl.
- Browser
- Download file 7.25 kB
-
https://huggingface.co/kernels/webgpu-kernels/ai.onnx.DynamicQuantizeLinear/resolve/v1/build/webgpu/dynamic-quantize-linear-quantize.wgsl.jinja
- Command line
-
hf download hf://webgpu-kernels/ai.onnx.DynamicQuantizeLinear@v1/build/webgpu/dynamic-quantize-linear-quantize.wgsl.jinja
-
curl -L -o dynamic-quantize-linear-quantize.wgsl.jinja https://huggingface.co/kernels/webgpu-kernels/ai.onnx.DynamicQuantizeLinear/resolve/v1/build/webgpu/dynamic-quantize-linear-quantize.wgsl.jinja
7.25 kB
| // Pass 3 of parallel DynamicQuantizeLinear: elementwise quantization using the | |
| // read-only y_scale/y_zero_point produced by the finalize pass. Each workgroup | |
| // covers the same WG * EPT contiguous chunk as the reduction pass, so their | |
| // dispatch counts match. Quantization applies round-to-even and saturates the | |
| // result to the uint8 range using the common scale and zero point. | |
| {{ env.wgsl.resourceDeclarations }} | |
| // ONNX DynamicQuantizeLinear uses correctly rounded f32 division followed by | |
| // round-half-to-even. Every execution path uses these helpers so their numerical | |
| // behavior cannot drift apart. | |
| // WGSL permits f32 division to differ from the correctly-rounded result by | |
| // 2.5 ULP, and fma() inherits separate multiply/add accuracy rather than | |
| // promising a fused residual. Reconstruct the correctly-rounded normal result | |
| // with integer significand division when the quotient can affect an integer | |
| // rounding boundary. This is backend-independent and uses only exact u32 ops. | |
| fn dynamic_quantize_exact_div_normal(numerator: f32, denominator: f32) -> f32 { | |
| if (numerator == 0.0) { | |
| return numerator; | |
| } | |
| let numerator_bits = bitcast<u32>(numerator); | |
| let denominator_bits = bitcast<u32>(denominator); | |
| let sign_bits = (numerator_bits ^ denominator_bits) & 0x80000000u; | |
| let numerator_abs = numerator_bits & 0x7fffffffu; | |
| let denominator_abs = denominator_bits & 0x7fffffffu; | |
| if (denominator_abs == 0u | |
| || (numerator_abs & 0x7f800000u) == 0x7f800000u | |
| || (denominator_abs & 0x7f800000u) == 0x7f800000u) { | |
| return numerator / denominator; | |
| } | |
| var numerator_mantissa = numerator_abs & 0x007fffffu; | |
| var denominator_mantissa = denominator_abs & 0x007fffffu; | |
| let numerator_biased_exponent = (numerator_abs >> 23u) & 0xffu; | |
| let denominator_biased_exponent = (denominator_abs >> 23u) & 0xffu; | |
| var numerator_exponent: i32; | |
| var denominator_exponent: i32; | |
| if (numerator_biased_exponent == 0u) { | |
| numerator_exponent = -126; | |
| // Zero returned above. A non-zero subnormal reaches the implicit-bit | |
| // position in at most 23 exact shifts. | |
| while ((numerator_mantissa & 0x00800000u) == 0u) { | |
| numerator_mantissa = numerator_mantissa << 1u; | |
| numerator_exponent = numerator_exponent - 1; | |
| } | |
| } else { | |
| numerator_mantissa = numerator_mantissa | 0x00800000u; | |
| numerator_exponent = i32(numerator_biased_exponent) - 127; | |
| } | |
| if (denominator_biased_exponent == 0u) { | |
| denominator_exponent = -126; | |
| while ((denominator_mantissa & 0x00800000u) == 0u) { | |
| denominator_mantissa = denominator_mantissa << 1u; | |
| denominator_exponent = denominator_exponent - 1; | |
| } | |
| } else { | |
| denominator_mantissa = denominator_mantissa | 0x00800000u; | |
| denominator_exponent = i32(denominator_biased_exponent) - 127; | |
| } | |
| var quotient_exponent = numerator_exponent - denominator_exponent; | |
| var remainder = numerator_mantissa; | |
| if (remainder < denominator_mantissa) { | |
| remainder = remainder << 1u; | |
| quotient_exponent = quotient_exponent - 1; | |
| } | |
| // The normalized ratio is now in [1, 2). Emit its implicit bit followed by | |
| // all 23 stored significand bits using exact binary long division. | |
| var quotient_mantissa = 0x00800000u; | |
| remainder = remainder - denominator_mantissa; | |
| for (var digit = 0u; digit < 23u; digit = digit + 1u) { | |
| remainder = remainder << 1u; | |
| if (remainder >= denominator_mantissa) { | |
| remainder = remainder - denominator_mantissa; | |
| quotient_mantissa = quotient_mantissa | (1u << (22u - digit)); | |
| } | |
| } | |
| // Round the 24-bit significand to nearest, ties to even. remainder and its | |
| // doubled value are below 2^25, so no u32 overflow is possible. | |
| let twice_remainder = remainder << 1u; | |
| if (twice_remainder > denominator_mantissa | |
| || (twice_remainder == denominator_mantissa && (quotient_mantissa & 1u) != 0u)) { | |
| quotient_mantissa = quotient_mantissa + 1u; | |
| } | |
| if (quotient_mantissa == 0x01000000u) { | |
| quotient_mantissa = quotient_mantissa >> 1u; | |
| quotient_exponent = quotient_exponent + 1; | |
| } | |
| let biased_exponent = quotient_exponent + 127; | |
| if (biased_exponent <= 0 || biased_exponent >= 255) { | |
| // The exact normal-range reconstruction below cannot encode a subnormal or | |
| // overflowed quotient. Use WGSL division for those exponent ranges; this | |
| // branch is disjoint from the finite halfway cases handled below. | |
| return numerator / denominator; | |
| } | |
| let result_bits = sign_bits | |
| | (u32(biased_exponent) << 23u) | |
| | (quotient_mantissa & 0x007fffffu); | |
| return bitcast<f32>(result_bits); | |
| } | |
| fn dynamic_quantize_division_may_cross_half(estimate: f32) -> bool { | |
| let lower = floor(estimate); | |
| let fraction = estimate - lower; | |
| let magnitude = abs(estimate); | |
| let magnitude_bits = bitcast<u32>(magnitude); | |
| let adjacent = bitcast<f32>(magnitude_bits + 1u); | |
| let ulp = adjacent - magnitude; | |
| // Division is allowed 2.5 ULP error. Eight ULP also covers the factor-of-two | |
| // ULP change when an estimate straddles the 0.5 exponent boundary. | |
| return abs(fraction - 0.5) <= ulp * 8.0; | |
| } | |
| fn round_dynamic_half_to_even(value: f32, scale: f32) -> i32 { | |
| let estimate = value / scale; | |
| var scaled = estimate; | |
| // `select` evaluates both value operands in WGSL; use control flow so the | |
| // 23-bit software divide remains a rare boundary fallback, not O(23) work | |
| // for every quantized element. | |
| if (dynamic_quantize_division_may_cross_half(estimate)) { | |
| scaled = dynamic_quantize_exact_div_normal(value, scale); | |
| } | |
| let lower = floor(scaled); | |
| let fraction = scaled - lower; | |
| if (fraction < 0.5) { | |
| return i32(lower); | |
| } | |
| if (fraction > 0.5) { | |
| return i32(lower + 1.0); | |
| } | |
| let upper = lower + 1.0; | |
| let half_lower = floor(lower * 0.5); | |
| let lower_is_even = (lower - half_lower * 2.0) == 0.0; | |
| return i32(select(upper, lower, lower_is_even)); | |
| } | |
| const WG: u32 = {{ workgroupSize }}u; | |
| {% if not vec4 %} | |
| const EPT: u32 = {{ elemsPerThread }}u; | |
| {% endif %} | |
| @compute @workgroup_size(WG, 1, 1) | |
| fn main(@builtin(workgroup_id) wg: vec3<u32>, | |
| @builtin(local_invocation_id) lid: vec3<u32>) { | |
| let tid = lid.x; | |
| let scale = y_scale[0]; | |
| let zp_i32 = i32(y_zero_point[0]); | |
| // Fold the block grid across x/y at a fixed per-axis workgroup width. | |
| // Per-element guards discard the over-dispatched tail. | |
| let blk = wg.x + wg.y * {{ DISPATCH_FOLD_WIDTH }}u; | |
| {% if vec4 %} | |
| // The vec4 input load reads four scalars at once. Output storage still uses | |
| // one u32 element for each quantized value. | |
| let count4 = params.count / 4u; | |
| let idx4 = blk * WG + tid; | |
| if (idx4 < count4) { | |
| let v = x[idx4]; | |
| let o = idx4 * 4u; | |
| y[o + 0u] = u32(clamp(round_dynamic_half_to_even(v.x, scale) + zp_i32, 0, 255)); | |
| y[o + 1u] = u32(clamp(round_dynamic_half_to_even(v.y, scale) + zp_i32, 0, 255)); | |
| y[o + 2u] = u32(clamp(round_dynamic_half_to_even(v.z, scale) + zp_i32, 0, 255)); | |
| y[o + 3u] = u32(clamp(round_dynamic_half_to_even(v.w, scale) + zp_i32, 0, 255)); | |
| } | |
| {% 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 q = clamp(round_dynamic_half_to_even(x[idx], scale) + zp_i32, 0, 255); | |
| y[idx] = u32(q); | |
| } | |
| } | |
| {% endif %} | |
| } | |