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 }