lib/simd/src/precision.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 
  3 pub fn twoSumScalar(comptime T: type, a: T, b: T, err: *T) T {
  4     requireFloat(T);
  5     const sum = a + b;
  6     const a2 = sum - b;
  7     const b2 = sum - a2;
  8     err.* = (a - a2) + (b - b2);
  9     return sum;
 10 }
 11 
 12 pub fn twoProducts(comptime D: type, a: D.Vector, b: D.Vector, err: *D.Vector) D.Vector {
 13     requireFloat(D.Lane);
 14     const product = a * b;
 15     err.* = @mulAdd(D.Vector, a, b, -product);
 16     return product;
 17 }
 18 
 19 pub fn twoSums(comptime D: type, a: D.Vector, b: D.Vector, err: *D.Vector) D.Vector {
 20     requireFloat(D.Lane);
 21     const sum = a + b;
 22     const a2 = sum - b;
 23     const b2 = sum - a2;
 24     err.* = (a - a2) + (b - b2);
 25     return sum;
 26 }
 27 
 28 pub fn fastTwoSums(comptime D: type, a: D.Vector, b: D.Vector, err: *D.Vector) D.Vector {
 29     requireFloat(D.Lane);
 30     const sum = a + b;
 31     err.* = b - (sum - a);
 32     return sum;
 33 }
 34 
 35 pub fn updateCascadedSums(
 36     comptime D: type,
 37     value: D.Vector,
 38     sum: *D.Vector,
 39     sum_err: *D.Vector,
 40 ) void {
 41     var err: D.Vector = undefined;
 42     sum.* = twoSums(D, sum.*, value, &err);
 43     sum_err.* += err;
 44 }
 45 
 46 pub fn assimilateCascadedSums(
 47     comptime D: type,
 48     other_sum: D.Vector,
 49     other_sum_err: D.Vector,
 50     sum: *D.Vector,
 51     sum_err: *D.Vector,
 52 ) void {
 53     sum_err.* += other_sum_err;
 54     updateCascadedSums(D, other_sum, sum, sum_err);
 55 }
 56 
 57 pub fn reduceCascadedSums(comptime D: type, sum: D.Vector, sum_err: D.Vector) D.Lane {
 58     requireFloat(D.Lane);
 59     var total: D.Lane = 0;
 60     var total_err: D.Lane = 0;
 61     inline for (0..D.lane_count) |index| {
 62         var err: D.Lane = undefined;
 63         total_err += sum_err[index];
 64         total = twoSumScalar(D.Lane, total, sum[index], &err);
 65         total_err += err;
 66     }
 67     return total + total_err;
 68 }
 69 
 70 pub fn ddAdd(
 71     comptime D: type,
 72     a_hi: D.Vector,
 73     a_lo: D.Vector,
 74     b_hi: D.Vector,
 75     b_lo: D.Vector,
 76     result_lo: *D.Vector,
 77 ) D.Vector {
 78     var error_value: D.Vector = undefined;
 79     const sum = twoSums(D, a_hi, b_hi, &error_value);
 80     error_value += a_lo + b_lo;
 81     return fastTwoSums(D, sum, error_value, result_lo);
 82 }
 83 
 84 pub fn ddMul1(
 85     comptime D: type,
 86     a_hi: D.Vector,
 87     a_lo: D.Vector,
 88     b: D.Vector,
 89     result_lo: *D.Vector,
 90 ) D.Vector {
 91     var product_lo: D.Vector = undefined;
 92     const product_hi = twoProducts(D, a_hi, b, &product_lo);
 93     product_lo = @mulAdd(D.Vector, a_lo, b, product_lo);
 94     return fastTwoSums(D, product_hi, product_lo, result_lo);
 95 }
 96 
 97 pub fn ddMul2(
 98     comptime D: type,
 99     a_hi: D.Vector,
100     a_lo: D.Vector,
101     b_hi: D.Vector,
102     b_lo: D.Vector,
103     result_lo: *D.Vector,
104 ) D.Vector {
105     var product_lo: D.Vector = undefined;
106     const product_hi = twoProducts(D, a_hi, b_hi, &product_lo);
107     product_lo = @mulAdd(D.Vector, a_hi, b_lo, @mulAdd(D.Vector, a_lo, b_hi, product_lo));
108     return fastTwoSums(D, product_hi, product_lo, result_lo);
109 }
110 
111 pub fn ddDiv(
112     comptime D: type,
113     a_hi: D.Vector,
114     a_lo: D.Vector,
115     b_hi: D.Vector,
116     b_lo: D.Vector,
117     result_lo: *D.Vector,
118 ) D.Vector {
119     const q1 = a_hi / b_hi;
120     var product_lo: D.Vector = undefined;
121     const product_hi = ddMul1(D, b_hi, b_lo, q1, &product_lo);
122     var remainder_lo: D.Vector = undefined;
123     const remainder_hi = ddAdd(D, a_hi, a_lo, -product_hi, -product_lo, &remainder_lo);
124     const q2 = remainder_hi / b_hi;
125     return fastTwoSums(D, q1, q2, result_lo);
126 }
127 
128 fn requireFloat(comptime T: type) void {
129     if (T != f32 and T != f64) @compileError("compensated arithmetic requires f32 or f64 lanes");
130 }
131 
132 test "Highway error-free sums and products reconstruct in wider precision" {
133     const simd = @import("root.zig");
134     const D = simd.FixedTag(f32, 4);
135     const a: D.Vector = .{ 1.0e20, 1.0, 1.25, -7.0 };
136     const b: D.Vector = .{ 1.0, 0x1p-24, 3.5, 0.3 };
137 
138     var sum_err: D.Vector = undefined;
139     const sum = twoSums(D, a, b, &sum_err);
140     var product_err: D.Vector = undefined;
141     const product = twoProducts(D, a, b, &product_err);
142     inline for (0..D.lane_count) |index| {
143         try std.testing.expectEqual(
144             @as(f64, a[index]) + @as(f64, b[index]),
145             @as(f64, sum[index]) + @as(f64, sum_err[index]),
146         );
147         try std.testing.expectEqual(
148             @as(f64, a[index]) * @as(f64, b[index]),
149             @as(f64, product[index]) + @as(f64, product_err[index]),
150         );
151     }
152 }
153 
154 test "Highway cascaded sums retain low-order values" {
155     const simd = @import("root.zig");
156     const D = simd.FixedTag(f64, 4);
157     var sum: D.Vector = @splat(0);
158     var sum_err: D.Vector = @splat(0);
159     updateCascadedSums(D, @as(D.Vector, .{ 1.0e16, 1, -1.0e16, 1 }), &sum, &sum_err);
160     try std.testing.expectEqual(@as(f64, 2), reduceCascadedSums(D, sum, sum_err));
161 
162     var other_sum: D.Vector = .{ 2, 3, 4, 5 };
163     var other_err: D.Vector = @splat(0);
164     assimilateCascadedSums(D, other_sum, other_err, &sum, &sum_err);
165     try std.testing.expectEqual(@as(f64, 16), reduceCascadedSums(D, sum, sum_err));
166     _ = &other_sum;
167     _ = &other_err;
168 }
169 
170 test "Highway double-double helpers preserve low components" {
171     const simd = @import("root.zig");
172     const D = simd.FixedTag(f64, 2);
173     const a_hi: D.Vector = .{ 1.0e16, 3 };
174     const a_lo: D.Vector = .{ 1, 0x1p-52 };
175     const b_hi: D.Vector = .{ 2, 7 };
176     const b_lo: D.Vector = .{ 0, -0x1p-51 };
177 
178     var result_lo: D.Vector = undefined;
179     const add_hi = ddAdd(D, a_hi, a_lo, b_hi, b_lo, &result_lo);
180     try std.testing.expectEqual(@as(f128, a_hi[0]) + a_lo[0] + b_hi[0] + b_lo[0], @as(f128, add_hi[0]) + result_lo[0]);
181 
182     const mul_hi = ddMul2(D, a_hi, a_lo, b_hi, b_lo, &result_lo);
183     try std.testing.expectApproxEqRel(
184         (@as(f128, a_hi[1]) + a_lo[1]) * (@as(f128, b_hi[1]) + b_lo[1]),
185         @as(f128, mul_hi[1]) + result_lo[1],
186         0x1p-100,
187     );
188 
189     const div_hi = ddDiv(D, a_hi, a_lo, b_hi, b_lo, &result_lo);
190     try std.testing.expectApproxEqRel(
191         (@as(f128, a_hi[1]) + a_lo[1]) / (@as(f128, b_hi[1]) + b_lo[1]),
192         @as(f128, div_hi[1]) + result_lo[1],
193         0x1p-100,
194     );
195 }