File size: 6,843 Bytes
bbc3807
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
{% if source.fromPartials %}
// Parallel range finalization folds the guarded min/max partials and computes
// the output scale and zero point. Each partial already includes zero in its range.
{% else %}
// The serial path computes the complete range, scale, zero point, and quantized
// output in one invocation.
{% endif %}
{{ 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) {
    // DynamicQuantizeLinear's GPU-supported ranges are normal and finite.
    // Preserve the implementation-defined subnormal/overflow behavior outside
    // that contract; the integer quantizer never takes this branch near a
    // half-integer boundary.
    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));
}


@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
  if (gid.x != 0u) { return; }
{% if source.fromPartials %}
  var min_value = partial_min[0];
  var max_value = partial_max[0];
  for (var i = 1u; i < params.numPartials; i = i + 1u) {
    min_value = min(min_value, partial_min[i]);
    max_value = max(max_value, partial_max[i]);
  }
{% else %}
  var min_value = 0.0;
  var max_value = 0.0;
  for (var i = 0u; i < params.count; i = i + 1u) {
    min_value = min(min_value, x[i]);
    max_value = max(max_value, x[i]);
  }
{% endif %}
  var scale = dynamic_quantize_exact_div_normal(max_value - min_value, 255.0);
  if (scale == 0.0) {
    scale = 1.0;
  }
  let zp_i32 = clamp(round_dynamic_half_to_even(-min_value, scale), 0, 255);
  y_scale[0] = scale;
  y_zero_point[0] = u32(zp_i32);
{% if not source.fromPartials %}
  for (var i = 0u; i < params.count; i = i + 1u) {
    let q = clamp(round_dynamic_half_to_even(x[i], scale) + zp_i32, 0, 255);
    y[i] = u32(q);
  }
{% endif %}
}