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 }