lib/simd/src/dot.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const arithmetic = @import("arithmetic.zig");
3 const bfloat = @import("bfloat.zig");
4 const construct = @import("construct.zig");
5 const memory = @import("memory.zig");
6 const multiply = @import("multiply.zig");
7 const reduce = @import("reduce.zig");
8
9 pub const Assumptions = packed struct(u3) {
10 at_least_one_vector: bool = false,
11 multiple_of_vector: bool = false,
12 padded_to_vector: bool = false,
13 };
14
15 pub fn compute(
16 comptime D: type,
17 a: []const D.Lane,
18 b: []const D.Lane,
19 ) resultType(D.Lane) {
20 std.debug.assert(a.len == b.len);
21 return computeAssume(D, a, b, a.len, .{});
22 }
23
24 pub fn computeAssume(
25 comptime D: type,
26 a: []const D.Lane,
27 b: []const D.Lane,
28 count: usize,
29 comptime assumptions: Assumptions,
30 ) resultType(D.Lane) {
31 validateSameLane(D.Lane);
32 validateInputs(D.lane_count, a.len, b.len, count, assumptions);
33 if (D.Lane == i16) return computeI16(D, a, b, count, assumptions);
34 return computeFloat(D, a, b, count, assumptions);
35 }
36
37 pub fn computeBFloat(
38 comptime D: type,
39 a: []const bfloat.BFloat16,
40 b: []const bfloat.BFloat16,
41 ) f32 {
42 std.debug.assert(a.len == b.len);
43 return computeBFloatAssume(D, a, b, a.len, .{});
44 }
45
46 pub fn computeBFloatAssume(
47 comptime D: type,
48 a: []const bfloat.BFloat16,
49 b: []const bfloat.BFloat16,
50 count: usize,
51 comptime assumptions: Assumptions,
52 ) f32 {
53 requireBFloatTag(D);
54 validateInputs(D.lane_count, a.len, b.len, count, assumptions);
55 if (D.lane_count < 2) return computeBFloatScalar(a, b, count);
56 const DF = D.repartition(f32);
57 var sum0: DF.Vector = @splat(0);
58 var sum1: DF.Vector = @splat(0);
59 var sum2: DF.Vector = @splat(0);
60 var sum3: DF.Vector = @splat(0);
61 var index: usize = 0;
62 while (index + 2 * D.lane_count <= count) : (index += 2 * D.lane_count) {
63 const a0 = bfloat.load(D, a[index..]);
64 const b0 = bfloat.load(D, b[index..]);
65 sum0 = bfloat.reorderWidenMulAccumulate(DF, a0, b0, sum0, &sum1);
66 const a1 = bfloat.load(D, a[index + D.lane_count ..]);
67 const b1 = bfloat.load(D, b[index + D.lane_count ..]);
68 sum2 = bfloat.reorderWidenMulAccumulate(DF, a1, b1, sum2, &sum3);
69 }
70 if (index + D.lane_count <= count) {
71 const av = bfloat.load(D, a[index..]);
72 const bv = bfloat.load(D, b[index..]);
73 sum0 = bfloat.reorderWidenMulAccumulate(DF, av, bv, sum0, &sum1);
74 index += D.lane_count;
75 }
76 if (!assumptions.multiple_of_vector and index != count) {
77 const remaining = count - index;
78 const av = loadBFloatTail(D, a[index..], remaining, assumptions.padded_to_vector);
79 const bv = loadBFloatTail(D, b[index..], remaining, assumptions.padded_to_vector);
80 sum2 = bfloat.reorderWidenMulAccumulate(DF, av, bv, sum2, &sum3);
81 }
82 return reduce.sum(DF, (sum0 + sum1) + (sum2 + sum3));
83 }
84
85 pub fn computeF32BFloat(
86 comptime D: type,
87 a: []const f32,
88 b: []const bfloat.BFloat16,
89 ) f32 {
90 std.debug.assert(a.len == b.len);
91 return computeF32BFloatAssume(D, a, b, a.len, .{});
92 }
93
94 pub fn computeF32BFloatAssume(
95 comptime D: type,
96 a: []const f32,
97 b: []const bfloat.BFloat16,
98 count: usize,
99 comptime assumptions: Assumptions,
100 ) f32 {
101 if (comptime D.Lane != f32) @compileError("mixed dot requires an f32 descriptor");
102 validateInputs(D.lane_count, a.len, b.len, count, assumptions);
103 var sum0: D.Vector = @splat(0);
104 var sum1: D.Vector = @splat(0);
105 var sum2: D.Vector = @splat(0);
106 var sum3: D.Vector = @splat(0);
107 var index: usize = 0;
108 while (index + 4 * D.lane_count <= count) : (index += 4 * D.lane_count) {
109 sum0 = mixedMulAdd(D, a[index..], b[index..], sum0);
110 const index1 = index + D.lane_count;
111 sum1 = mixedMulAdd(D, a[index1..], b[index1..], sum1);
112 const index2 = index + 2 * D.lane_count;
113 sum2 = mixedMulAdd(D, a[index2..], b[index2..], sum2);
114 const index3 = index + 3 * D.lane_count;
115 sum3 = mixedMulAdd(D, a[index3..], b[index3..], sum3);
116 }
117 while (index + D.lane_count <= count) : (index += D.lane_count) {
118 sum0 = mixedMulAdd(D, a[index..], b[index..], sum0);
119 }
120 if (!assumptions.multiple_of_vector and index != count) {
121 const remaining = count - index;
122 const av = loadTail(D, a[index..], remaining, assumptions.padded_to_vector);
123 const DB = bfloat.Tag(D.lane_count);
124 const bits = loadBFloatTail(DB, b[index..], remaining, assumptions.padded_to_vector);
125 sum1 = arithmetic.mulAdd(D, av, bfloat.promoteF32(D, bits), sum1);
126 }
127 return reduce.sum(D, (sum0 + sum1) + (sum2 + sum3));
128 }
129
130 fn computeFloat(
131 comptime D: type,
132 a: []const D.Lane,
133 b: []const D.Lane,
134 count: usize,
135 comptime assumptions: Assumptions,
136 ) D.Lane {
137 var sum0: D.Vector = @splat(0);
138 var sum1: D.Vector = @splat(0);
139 var sum2: D.Vector = @splat(0);
140 var sum3: D.Vector = @splat(0);
141 var index: usize = 0;
142 while (index + 4 * D.lane_count <= count) : (index += 4 * D.lane_count) {
143 sum0 = arithmetic.mulAdd(D, memory.load(D, a[index..]), memory.load(D, b[index..]), sum0);
144 const index1 = index + D.lane_count;
145 sum1 = arithmetic.mulAdd(D, memory.load(D, a[index1..]), memory.load(D, b[index1..]), sum1);
146 const index2 = index + 2 * D.lane_count;
147 sum2 = arithmetic.mulAdd(D, memory.load(D, a[index2..]), memory.load(D, b[index2..]), sum2);
148 const index3 = index + 3 * D.lane_count;
149 sum3 = arithmetic.mulAdd(D, memory.load(D, a[index3..]), memory.load(D, b[index3..]), sum3);
150 }
151 while (index + D.lane_count <= count) : (index += D.lane_count) {
152 sum0 = arithmetic.mulAdd(D, memory.load(D, a[index..]), memory.load(D, b[index..]), sum0);
153 }
154 if (!assumptions.multiple_of_vector and index != count) {
155 const remaining = count - index;
156 const av = loadTail(D, a[index..], remaining, assumptions.padded_to_vector);
157 const bv = loadTail(D, b[index..], remaining, assumptions.padded_to_vector);
158 sum1 = arithmetic.mulAdd(D, av, bv, sum1);
159 }
160 return reduce.sum(D, (sum0 + sum1) + (sum2 + sum3));
161 }
162
163 fn computeI16(
164 comptime D: type,
165 a: []const i16,
166 b: []const i16,
167 count: usize,
168 comptime assumptions: Assumptions,
169 ) i32 {
170 if (D.lane_count < 2) return computeI16Scalar(a, b, count);
171 const DW = D.repartition(i32);
172 var sum0: DW.Vector = @splat(0);
173 var sum1: DW.Vector = @splat(0);
174 var sum2: DW.Vector = @splat(0);
175 var sum3: DW.Vector = @splat(0);
176 var index: usize = 0;
177 while (index + 2 * D.lane_count <= count) : (index += 2 * D.lane_count) {
178 const a0 = memory.load(D, a[index..]);
179 const b0 = memory.load(D, b[index..]);
180 sum0 = multiply.reorderWidenMulAccumulate(DW, a0, b0, sum0, &sum1);
181 const index1 = index + D.lane_count;
182 const a1 = memory.load(D, a[index1..]);
183 const b1 = memory.load(D, b[index1..]);
184 sum2 = multiply.reorderWidenMulAccumulate(DW, a1, b1, sum2, &sum3);
185 }
186 if (index + D.lane_count <= count) {
187 const av = memory.load(D, a[index..]);
188 const bv = memory.load(D, b[index..]);
189 sum0 = multiply.reorderWidenMulAccumulate(DW, av, bv, sum0, &sum1);
190 index += D.lane_count;
191 }
192 if (!assumptions.multiple_of_vector and index != count) {
193 const remaining = count - index;
194 const av = loadTail(D, a[index..], remaining, assumptions.padded_to_vector);
195 const bv = loadTail(D, b[index..], remaining, assumptions.padded_to_vector);
196 sum2 = multiply.reorderWidenMulAccumulate(DW, av, bv, sum2, &sum3);
197 }
198 return reduce.sum(DW, (sum0 +% sum1) +% (sum2 +% sum3));
199 }
200
201 fn mixedMulAdd(
202 comptime D: type,
203 a: []const f32,
204 b: []const bfloat.BFloat16,
205 sum: D.Vector,
206 ) D.Vector {
207 const DB = bfloat.Tag(D.lane_count);
208 return arithmetic.mulAdd(D, memory.load(D, a), bfloat.promoteF32(D, bfloat.load(DB, b)), sum);
209 }
210
211 fn loadTail(
212 comptime D: type,
213 input: []const D.Lane,
214 count: usize,
215 comptime padded: bool,
216 ) D.Vector {
217 if (padded) {
218 const value = memory.load(D, input);
219 return @select(D.Lane, construct.firstN(D, count), value, @as(D.Vector, @splat(0)));
220 }
221 return memory.loadN(D, input, count);
222 }
223
224 fn loadBFloatTail(
225 comptime D: type,
226 input: []const bfloat.BFloat16,
227 count: usize,
228 comptime padded: bool,
229 ) D.Vector {
230 if (padded) {
231 const value = bfloat.load(D, input);
232 return @select(u16, construct.firstN(D, count), value, @as(D.Vector, @splat(0)));
233 }
234 var result: D.Vector = @splat(0);
235 inline for (0..D.lane_count) |index| {
236 if (index < count) result[index] = input[index].bits;
237 }
238 return result;
239 }
240
241 fn computeI16Scalar(a: []const i16, b: []const i16, count: usize) i32 {
242 var sum: i32 = 0;
243 for (a[0..count], b[0..count]) |av, bv| sum +%= @as(i32, av) * @as(i32, bv);
244 return sum;
245 }
246
247 fn computeBFloatScalar(
248 a: []const bfloat.BFloat16,
249 b: []const bfloat.BFloat16,
250 count: usize,
251 ) f32 {
252 var sum: f32 = 0;
253 for (a[0..count], b[0..count]) |av, bv| sum = @mulAdd(f32, av.toF32(), bv.toF32(), sum);
254 return sum;
255 }
256
257 fn validateInputs(
258 comptime lanes: usize,
259 a_len: usize,
260 b_len: usize,
261 count: usize,
262 comptime assumptions: Assumptions,
263 ) void {
264 std.debug.assert(a_len >= count);
265 std.debug.assert(b_len >= count);
266 if (assumptions.at_least_one_vector) std.debug.assert(count >= lanes);
267 if (assumptions.multiple_of_vector) std.debug.assert(count % lanes == 0);
268 if (assumptions.padded_to_vector and count % lanes != 0) {
269 const padded_count = std.mem.alignForward(usize, count, lanes);
270 std.debug.assert(a_len >= padded_count);
271 std.debug.assert(b_len >= padded_count);
272 }
273 }
274
275 fn validateSameLane(comptime T: type) void {
276 if (T != f16 and T != f32 and T != f64 and T != i16) {
277 @compileError("dot requires f16/f32/f64 or i16 lanes");
278 }
279 }
280
281 fn requireBFloatTag(comptime D: type) void {
282 if (comptime !@hasDecl(D, "is_bfloat16") or !D.is_bfloat16) {
283 @compileError("bfloat16 dot requires a bfloat16 descriptor");
284 }
285 }
286
287 fn resultType(comptime T: type) type {
288 validateSameLane(T);
289 return if (T == i16) i32 else T;
290 }
291
292 fn close(comptime T: type, expected: T, actual: T, scale: T) bool {
293 const tolerance = scale * @max(@abs(expected), @as(T, 1));
294 return @abs(expected - actual) <= tolerance;
295 }
296
297 test "Highway floating dot handles every assumption and awkward alignment" {
298 const simd = @import("root.zig");
299 const D = simd.FixedTag(f32, 8);
300 var a_storage: [96]f32 = @splat(std.math.nan(f32));
301 var b_storage: [96]f32 = @splat(std.math.nan(f32));
302 const a = a_storage[1..];
303 const b = b_storage[3..];
304 for (0..75) |index| {
305 a[index] = @as(f32, @floatFromInt(@as(i32, @intCast(index % 17)) - 8)) * 0.25;
306 b[index] = @as(f32, @floatFromInt(@as(i32, @intCast(index % 13)) - 6)) * 0.5;
307 }
308 const count: usize = 67;
309 var expected: f32 = 0;
310 for (a[0..count], b[0..count]) |av, bv| expected = @mulAdd(f32, av, bv, expected);
311 inline for (.{
312 Assumptions{},
313 Assumptions{ .at_least_one_vector = true },
314 Assumptions{ .padded_to_vector = true },
315 Assumptions{ .at_least_one_vector = true, .padded_to_vector = true },
316 }) |assumptions| {
317 const actual = computeAssume(D, a, b, count, assumptions);
318 try std.testing.expect(close(f32, expected, actual, 32 * std.math.floatEps(f32)));
319 }
320 const multiple_count: usize = 64;
321 inline for (.{
322 Assumptions{ .multiple_of_vector = true },
323 Assumptions{ .at_least_one_vector = true, .multiple_of_vector = true },
324 Assumptions{ .multiple_of_vector = true, .padded_to_vector = true },
325 Assumptions{
326 .at_least_one_vector = true,
327 .multiple_of_vector = true,
328 .padded_to_vector = true,
329 },
330 }) |assumptions| {
331 _ = computeAssume(D, a, b, multiple_count, assumptions);
332 }
333 }
334
335 test "Highway dot supports every same-input lane class" {
336 const simd = @import("root.zig");
337 inline for (.{ f16, f32, f64 }) |T| {
338 const D = simd.FixedTag(T, 4);
339 var a: [13]T = undefined;
340 var b: [13]T = undefined;
341 var expected: T = 0;
342 for (&a, &b, 0..) |*av, *bv, index| {
343 av.* = @floatFromInt(@as(i32, @intCast(index % 7)) - 3);
344 bv.* = @floatFromInt(@as(i32, @intCast(index % 5)) - 2);
345 expected = @mulAdd(T, av.*, bv.*, expected);
346 }
347 const actual = compute(D, &a, &b);
348 try std.testing.expect(close(T, expected, actual, 32 * std.math.floatEps(T)));
349 }
350 const DI = simd.FixedTag(i16, 8);
351 const ai = [_]i16{ 7, -3, 12, 9, -8, 4, 6, -11, 5, 2, -1 };
352 const bi = [_]i16{ -2, 8, 3, -7, 5, 9, -4, 6, 10, -3, 12 };
353 var expected_i16: i32 = 0;
354 for (ai, bi) |av, bv| expected_i16 +%= @as(i32, av) * @as(i32, bv);
355 try std.testing.expectEqual(expected_i16, compute(DI, &ai, &bi));
356 const D1 = simd.FixedTag(i16, 1);
357 try std.testing.expectEqual(expected_i16, compute(D1, &ai, &bi));
358 }
359
360 test "Highway bfloat dot widens both same and mixed inputs" {
361 const simd = @import("root.zig");
362 const DB = simd.BFloat16Tag(8);
363 const DF = simd.FixedTag(f32, 8);
364 var a: [19]bfloat.BFloat16 = undefined;
365 var b: [19]bfloat.BFloat16 = undefined;
366 var af: [19]f32 = undefined;
367 var expected_bf: f32 = 0;
368 var expected_mixed: f32 = 0;
369 for (&a, &b, &af, 0..) |*av, *bv, *fv, index| {
370 const ai = @as(i32, @intCast(index % 11)) - 5;
371 const bi = @as(i32, @intCast(index % 7)) - 3;
372 av.* = bfloat.BFloat16.fromF32(@as(f32, @floatFromInt(ai)) * 0.5);
373 bv.* = bfloat.BFloat16.fromF32(@as(f32, @floatFromInt(bi)) * 0.25);
374 fv.* = @as(f32, @floatFromInt(ai)) * 0.125;
375 expected_bf = @mulAdd(f32, av.toF32(), bv.toF32(), expected_bf);
376 expected_mixed = @mulAdd(f32, fv.*, bv.toF32(), expected_mixed);
377 }
378 try std.testing.expect(close(
379 f32,
380 expected_bf,
381 computeBFloat(DB, &a, &b),
382 32 * std.math.floatEps(f32),
383 ));
384 try std.testing.expect(close(
385 f32,
386 expected_mixed,
387 computeF32BFloat(DF, &af, &b),
388 32 * std.math.floatEps(f32),
389 ));
390 }
391
392 test "Highway AVX2 dot oracle matches awkward tails" {
393 const simd = @import("root.zig");
394 const DF32 = simd.FixedTag(f32, 8);
395 const DF64 = simd.FixedTag(f64, 4);
396 const DI16 = simd.FixedTag(i16, 16);
397 const DBF16 = simd.BFloat16Tag(16);
398 var a32: [67]f32 = undefined;
399 var b32: [67]f32 = undefined;
400 for (&a32, &b32, 0..) |*av, *bv, index| {
401 const ai = @as(i32, @intCast(index % 19)) - 9;
402 const bi = @as(i32, @intCast(index % 13)) - 6;
403 av.* = @as(f32, @floatFromInt(ai)) * 0.1375 + 0.03125;
404 bv.* = @as(f32, @floatFromInt(bi)) * -0.2125 + 0.015625;
405 }
406 try std.testing.expect(close(
407 f32,
408 @bitCast(@as(u32, 0xbf29_b714)),
409 compute(DF32, &a32, &b32),
410 96 * std.math.floatEps(f32),
411 ));
412
413 var a64: [37]f64 = undefined;
414 var b64: [37]f64 = undefined;
415 for (&a64, &b64, 0..) |*av, *bv, index| {
416 const ai = @as(i32, @intCast(index % 11)) - 5;
417 const bi = @as(i32, @intCast(index % 7)) - 3;
418 av.* = @as(f64, @floatFromInt(ai)) * 0.1375 + 0.03125;
419 bv.* = @as(f64, @floatFromInt(bi)) * -0.2125 + 0.015625;
420 }
421 try std.testing.expect(close(
422 f64,
423 @bitCast(@as(u64, 0x3fa9_cf5c_28f5_c270)),
424 compute(DF64, &a64, &b64),
425 96 * std.math.floatEps(f64),
426 ));
427
428 var ai16: [53]i16 = undefined;
429 var bi16: [53]i16 = undefined;
430 for (&ai16, &bi16, 0..) |*av, *bv, index| {
431 av.* = @intCast(@as(i32, @intCast(index % 31)) - 15);
432 bv.* = @intCast(@as(i32, @intCast(index % 23)) - 11);
433 }
434 try std.testing.expectEqual(@as(i32, 24), compute(DI16, &ai16, &bi16));
435
436 var abf: [35]bfloat.BFloat16 = undefined;
437 var bbf: [35]bfloat.BFloat16 = undefined;
438 var mixed: [35]f32 = undefined;
439 for (&abf, &bbf, &mixed, 0..) |*av, *bv, *mv, index| {
440 const ai = @as(i32, @intCast(index % 17)) - 8;
441 const bi = @as(i32, @intCast(index % 9)) - 4;
442 const af = @as(f32, @floatFromInt(ai)) * 0.1375 + 0.03125;
443 const bf = @as(f32, @floatFromInt(bi)) * -0.2125 + 0.015625;
444 av.* = bfloat.BFloat16.fromF32(af);
445 bv.* = bfloat.BFloat16.fromF32(bf);
446 const offset = @as(i32, @intCast(index % 5)) - 2;
447 mv.* = af + @as(f32, @floatFromInt(offset)) * 0.003;
448 }
449 try std.testing.expect(close(
450 f32,
451 @bitCast(@as(u32, 0xc015_a4a0)),
452 computeBFloat(DBF16, &abf, &bbf),
453 96 * std.math.floatEps(f32),
454 ));
455 try std.testing.expect(close(
456 f32,
457 @bitCast(@as(u32, 0xc015_d992)),
458 computeF32BFloat(DF32, &mixed, &bbf),
459 96 * std.math.floatEps(f32),
460 ));
461 }