lib/simd/src/test.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const simd = @import("root.zig");
  3 
  4 const parity = simd.parity;
  5 const FixedTag = simd.FixedTag;
  6 const allTrue = simd.allTrue;
  7 const eq = simd.eq;
  8 const shl = simd.shl;
  9 const shr = simd.shr;
 10 const roundingShr = simd.roundingShr;
 11 const shiftLeft = simd.shiftLeft;
 12 const shiftRight = simd.shiftRight;
 13 const rol = simd.rol;
 14 const ror = simd.ror;
 15 const rotateLeftSame = simd.rotateLeftSame;
 16 const rotateRightSame = simd.rotateRightSame;
 17 const multiRotateRight = simd.multiRotateRight;
 18 const populationCount = simd.populationCount;
 19 const leadingZeroCount = simd.leadingZeroCount;
 20 const trailingZeroCount = simd.trailingZeroCount;
 21 const highestSetBitIndex = simd.highestSetBitIndex;
 22 const copySign = simd.copySign;
 23 const copySignToAbs = simd.copySignToAbs;
 24 const bitsFromMask = simd.bitsFromMask;
 25 const isNegative = simd.isNegative;
 26 const isNaN = simd.isNaN;
 27 const broadcastSignBit = simd.broadcastSignBit;
 28 const lowerHalf = simd.lowerHalf;
 29 const upperHalf = simd.upperHalf;
 30 const zeroExtendVector = simd.zeroExtendVector;
 31 const combine = simd.combine;
 32 const concatLowerLower = simd.concatLowerLower;
 33 const concatUpperUpper = simd.concatUpperUpper;
 34 const concatLowerUpper = simd.concatLowerUpper;
 35 const concatUpperLower = simd.concatUpperLower;
 36 const concatOdd = simd.concatOdd;
 37 const concatEven = simd.concatEven;
 38 const interleaveWholeLower = simd.interleaveWholeLower;
 39 const interleaveWholeUpper = simd.interleaveWholeUpper;
 40 const lowerHalfOfMask = simd.lowerHalfOfMask;
 41 const upperHalfOfMask = simd.upperHalfOfMask;
 42 const combineMasks = simd.combineMasks;
 43 const resizeBitCast = simd.resizeBitCast;
 44 const zeroExtendResizeBitCast = simd.zeroExtendResizeBitCast;
 45 const convertTo = simd.convertTo;
 46 const demoteTo = simd.demoteTo;
 47 const promoteLowerTo = simd.promoteLowerTo;
 48 const promoteUpperTo = simd.promoteUpperTo;
 49 const promoteEvenTo = simd.promoteEvenTo;
 50 const promoteOddTo = simd.promoteOddTo;
 51 const orderedDemote2To = simd.orderedDemote2To;
 52 const orderedTruncate2To = simd.orderedTruncate2To;
 53 const interleaveLower = simd.interleaveLower;
 54 const interleaveUpper = simd.interleaveUpper;
 55 const interleaveEven = simd.interleaveEven;
 56 const interleaveOdd = simd.interleaveOdd;
 57 const shiftLeftLanes = simd.shiftLeftLanes;
 58 const shiftRightLanes = simd.shiftRightLanes;
 59 const combineShiftRightLanes = simd.combineShiftRightLanes;
 60 const reverse = simd.reverse;
 61 const reverseBlocks = simd.reverseBlocks;
 62 const compress = simd.compress;
 63 const expand = simd.expand;
 64 const slideUpLanes = simd.slideUpLanes;
 65 const slideDownLanes = simd.slideDownLanes;
 66 const sumOfLanes = simd.sumOfLanes;
 67 const maskedReduceSum = simd.maskedReduceSum;
 68 const tableLookupLanes = simd.tableLookupLanes;
 69 const per4LaneBlockShuffle = simd.per4LaneBlockShuffle;
 70 const storeInterleaved2 = simd.storeInterleaved2;
 71 const load = simd.load;
 72 const store = simd.store;
 73 const add = simd.add;
 74 const iota = simd.iota;
 75 const lt = simd.lt;
 76 const sub = simd.sub;
 77 const mul = simd.mul;
 78 const select = simd.select;
 79 const reduceSum = simd.reduceSum;
 80 const countTrue = simd.countTrue;
 81 const Divisor = simd.Divisor;
 82 const Divisor64 = simd.Divisor64;
 83 const lemireMod = simd.lemireMod;
 84 
 85 fn expectOracleVector(comptime D: type, expected: D.Vector, actual: D.Vector) !void {
 86     try std.testing.expect(allTrue(D, eq(D, expected, actual)));
 87 }
 88 
 89 fn verifySignedBitOracle() !void {
 90     const D = FixedTag(i16, 4);
 91     const U = FixedTag(u16, 4);
 92     const value: D.Vector = @bitCast(@as(U.Vector, .{ 0x8001, 0xfffd, 0x1234, 0x7fff }));
 93     const amounts: D.Vector = .{ 0, 1, 4, 15 };
 94     try expectOracleVector(D, @bitCast(@as(U.Vector, .{ 0x8001, 0xfffa, 0x2340, 0x8000 })), shl(D, value, amounts));
 95     try expectOracleVector(D, @bitCast(@as(U.Vector, .{ 0x8001, 0xfffe, 0x0123, 0x0000 })), shr(D, value, amounts));
 96     try expectOracleVector(D, @bitCast(@as(U.Vector, .{ 0x8001, 0xffff, 0x0123, 0x0001 })), roundingShr(D, value, amounts));
 97     try expectOracleVector(D, @bitCast(@as(U.Vector, .{ 0x0008, 0xffe8, 0x91a0, 0xfff8 })), shiftLeft(D, 3, value));
 98     try expectOracleVector(D, @bitCast(@as(U.Vector, .{ 0xf000, 0xffff, 0x0246, 0x0fff })), shiftRight(D, 3, value));
 99 }
100 
101 fn verifyWideRotateOracle() !void {
102     const D = FixedTag(u64, 4);
103     const value: D.Vector = .{ 0, 1, 0x8000_0000_0000_0001, 0xf0f0_0000_0000_0001 };
104     const amounts: D.Vector = .{ 0, 1, @bitCast(@as(i64, -4)), 65 };
105     try expectOracleVector(D, .{ 0, 2, 0x1800_0000_0000_0000, 0xe1e0_0000_0000_0003 }, rol(D, value, amounts));
106     try expectOracleVector(D, .{ 0, 0x8000_0000_0000_0000, 0x18, 0xf878_0000_0000_0000 }, ror(D, value, amounts));
107     try expectOracleVector(D, .{ 0, 0x0200_0000_0000_0000, 0x0300_0000_0000_0000, 0x03e1_e000_0000_0000 }, rotateLeftSame(D, value, -7));
108     try expectOracleVector(D, .{ 0, 0x80, 0xc0, 0x7800_0000_0000_00f8 }, rotateRightSame(D, value, -7));
109     const DI = D.repartition(u8);
110     const indices: DI.Vector = .{
111         0,  8,  16, 24, 32, 40, 48, 56,
112         0,  9,  18, 27, 36, 45, 54, 63,
113         4,  12, 20, 28, 36, 44, 52, 60,
114         60, 52, 44, 36, 28, 20, 12, 4,
115     };
116     try expectOracleVector(D, .{ 0, 0x0200_0000_0000_0001, 0x1800_0000_0000_0000, 0x0f1f }, multiRotateRight(D, value, indices));
117 }
118 
119 fn verifyWideCountOracle() !void {
120     const D = FixedTag(u64, 4);
121     const value: D.Vector = .{ 0, 1, 0x8000_0000_0000_0001, 0xf0f0_0000_0000_0001 };
122     try expectOracleVector(D, .{ 0, 1, 2, 9 }, populationCount(D, value));
123     try expectOracleVector(D, .{ 64, 63, 0, 0 }, leadingZeroCount(D, value));
124     try expectOracleVector(D, .{ 64, 0, 0, 0 }, trailingZeroCount(D, value));
125     try expectOracleVector(D, .{ std.math.maxInt(u64), 0, 63, 63 }, highestSetBitIndex(D, value));
126 }
127 
128 fn verifySignOracle() !void {
129     const D = FixedTag(f32, 4);
130     const U = FixedTag(u32, 4);
131     const magnitude: D.Vector = @bitCast(@as(U.Vector, .{ 0, 0x3f80_0000, 0xc000_0000, 0x7fc0_1234 }));
132     const absolute: D.Vector = @bitCast(@as(U.Vector, .{ 0, 0x3f80_0000, 0x4000_0000, 0x7fc0_1234 }));
133     const signs: D.Vector = @bitCast(@as(U.Vector, .{ 0x8000_0000, 0, 0x8000_0000, 0 }));
134     const expected: U.Vector = .{ 0x8000_0000, 0x3f80_0000, 0xc000_0000, 0x7fc0_1234 };
135     try expectOracleVector(U, expected, @bitCast(copySign(D, magnitude, signs)));
136     try expectOracleVector(U, expected, @bitCast(copySignToAbs(D, absolute, signs)));
137     try std.testing.expectEqual(@as(u64, 4), bitsFromMask(D, isNegative(D, magnitude)));
138     try std.testing.expectEqual(@as(u64, 8), bitsFromMask(D, isNaN(D, magnitude)));
139     const DI = FixedTag(i32, 4);
140     const sign_values: DI.Vector = .{ 0, 1, -1, std.math.minInt(i32) };
141     try expectOracleVector(DI, .{ 0, 0, -1, -1 }, broadcastSignBit(DI, sign_values));
142 }
143 
144 fn verifyLayoutOracle() !void {
145     const D = FixedTag(u32, 8);
146     const H = D.half();
147     const low: D.Vector = .{ 0, 1, 2, 3, 4, 5, 6, 7 };
148     const high: D.Vector = .{ 10, 11, 12, 13, 14, 15, 16, 17 };
149     try expectOracleVector(H, .{ 0, 1, 2, 3 }, lowerHalf(H, low));
150     try expectOracleVector(H, .{ 4, 5, 6, 7 }, upperHalf(H, low));
151     try expectOracleVector(D, .{ 0, 1, 2, 3, 0, 0, 0, 0 }, zeroExtendVector(D, lowerHalf(H, low)));
152     try expectOracleVector(D, low, combine(D, upperHalf(H, low), lowerHalf(H, low)));
153     try expectOracleVector(D, .{ 0, 1, 2, 3, 10, 11, 12, 13 }, concatLowerLower(D, high, low));
154     try expectOracleVector(D, .{ 4, 5, 6, 7, 14, 15, 16, 17 }, concatUpperUpper(D, high, low));
155     try expectOracleVector(D, .{ 4, 5, 6, 7, 10, 11, 12, 13 }, concatLowerUpper(D, high, low));
156     try expectOracleVector(D, .{ 0, 1, 2, 3, 14, 15, 16, 17 }, concatUpperLower(D, high, low));
157     try expectOracleVector(D, .{ 1, 3, 5, 7, 11, 13, 15, 17 }, concatOdd(D, high, low));
158     try expectOracleVector(D, .{ 0, 2, 4, 6, 10, 12, 14, 16 }, concatEven(D, high, low));
159     try expectOracleVector(D, .{ 0, 10, 1, 11, 2, 12, 3, 13 }, interleaveWholeLower(D, low, high));
160     try expectOracleVector(D, .{ 4, 14, 5, 15, 6, 16, 7, 17 }, interleaveWholeUpper(D, low, high));
161     const mask = lt(D, low, @as(D.Vector, @splat(5)));
162     try std.testing.expectEqual(@as(u64, 0x0f), bitsFromMask(H, lowerHalfOfMask(H, mask)));
163     try std.testing.expectEqual(@as(u64, 0x01), bitsFromMask(H, upperHalfOfMask(H, mask)));
164     try std.testing.expectEqual(@as(u64, 0x1f), bitsFromMask(D, combineMasks(D, upperHalfOfMask(H, mask), lowerHalfOfMask(H, mask))));
165 }
166 
167 fn verifyResizeOracle() !void {
168     const D = FixedTag(u16, 8);
169     const H = D.half();
170     const T = D.twice();
171     const value: D.Vector = .{ 0x1001, 0x2002, 0x3003, 0x4004, 0x5005, 0x6006, 0x7007, 0x8008 };
172     try expectOracleVector(H, .{ 0x1001, 0x2002, 0x3003, 0x4004 }, resizeBitCast(H, value));
173     try expectOracleVector(T, .{
174         0x1001, 0x2002, 0x3003, 0x4004, 0x5005, 0x6006, 0x7007, 0x8008,
175         0,      0,      0,      0,      0,      0,      0,      0,
176     }, zeroExtendResizeBitCast(T, D, value));
177 }
178 
179 fn verifyConversionOracle() !void {
180     const F = FixedTag(f32, 8);
181     const I = FixedTag(i32, 8);
182     const U = FixedTag(u32, 8);
183     const float_bits: U.Vector = .{
184         0xc060_0000,
185         0xbf00_0000,
186         0,
187         0x409c_cccd,
188         0x4f00_0000,
189         0xcf00_0000,
190         0x7f80_0000,
191         0xff80_0000,
192     };
193     try expectOracleVector(I, .{
194         -3,
195         0,
196         0,
197         4,
198         std.math.maxInt(i32),
199         std.math.minInt(i32),
200         std.math.maxInt(i32),
201         std.math.minInt(i32),
202     }, convertTo(I, @as(F.Vector, @bitCast(float_bits))));
203     try expectOracleVector(U, .{
204         0,
205         0,
206         0,
207         4,
208         2_147_483_648,
209         0,
210         std.math.maxInt(u32),
211         0,
212     }, convertTo(U, @as(F.Vector, @bitCast(float_bits))));
213 
214     const I16 = FixedTag(i16, 8);
215     const U16 = FixedTag(u16, 8);
216     const integers: I.Vector = .{
217         std.math.minInt(i32), -65_536, -1, 0, 1, 65_535, 65_536, std.math.maxInt(i32),
218     };
219     try expectOracleVector(I16, .{
220         -32_768, -32_768, -1, 0, 1, 32_767, 32_767, 32_767,
221     }, demoteTo(I16, integers));
222     try expectOracleVector(U16, .{
223         0, 0, 0, 0, 1, 65_535, 65_535, 65_535,
224     }, demoteTo(U16, integers));
225 
226     const B = FixedTag(u8, 8);
227     const W = FixedTag(u16, 4);
228     const bytes: B.Vector = .{ 0, 1, 2, 3, 4, 5, 6, 7 };
229     try expectOracleVector(W, .{ 0, 1, 2, 3 }, promoteLowerTo(W, bytes));
230     try expectOracleVector(W, .{ 4, 5, 6, 7 }, promoteUpperTo(W, bytes));
231     try expectOracleVector(W, .{ 0, 2, 4, 6 }, promoteEvenTo(W, bytes));
232     try expectOracleVector(W, .{ 1, 3, 5, 7 }, promoteOddTo(W, bytes));
233     const a: W.Vector = .{ 0x0102, 0x03ff, 0x0400, 0xffff };
234     const b: W.Vector = .{ 0x1005, 0x2006, 0x3007, 0x4008 };
235     try expectOracleVector(B, .{ 255, 255, 255, 255, 255, 255, 255, 255 }, orderedDemote2To(B, a, b));
236     try expectOracleVector(B, .{ 2, 255, 0, 255, 5, 6, 7, 8 }, orderedTruncate2To(B, a, b));
237 }
238 
239 fn verifyArrangeOracle() !void {
240     const D = FixedTag(u32, 8);
241     const a: D.Vector = .{ 0, 1, 2, 3, 4, 5, 6, 7 };
242     const b: D.Vector = .{ 10, 11, 12, 13, 14, 15, 16, 17 };
243     const mask: D.Mask = .{ false, true, true, false, true, false, false, true };
244     try expectOracleVector(D, .{ 0, 10, 1, 11, 4, 14, 5, 15 }, interleaveLower(D, a, b));
245     try expectOracleVector(D, .{ 2, 12, 3, 13, 6, 16, 7, 17 }, interleaveUpper(D, a, b));
246     try expectOracleVector(D, .{ 0, 10, 2, 12, 4, 14, 6, 16 }, interleaveEven(D, a, b));
247     try expectOracleVector(D, .{ 1, 11, 3, 13, 5, 15, 7, 17 }, interleaveOdd(D, a, b));
248     try expectOracleVector(D, .{ 0, 0, 1, 2, 0, 4, 5, 6 }, shiftLeftLanes(D, 1, a));
249     try expectOracleVector(D, .{ 1, 2, 3, 0, 5, 6, 7, 0 }, shiftRightLanes(D, 1, a));
250     try expectOracleVector(D, .{ 2, 3, 10, 11, 6, 7, 14, 15 }, combineShiftRightLanes(D, 2, b, a));
251     try expectOracleVector(D, .{ 7, 6, 5, 4, 3, 2, 1, 0 }, reverse(D, a));
252     try expectOracleVector(D, .{ 4, 5, 6, 7, 0, 1, 2, 3 }, reverseBlocks(D, a));
253     const compacted = compress(D, a, mask);
254     try expectOracleVector(D, .{ 1, 2, 4, 7, 0, 3, 5, 6 }, compacted);
255     try expectOracleVector(D, .{ 0, 1, 2, 0, 4, 0, 0, 7 }, expand(D, compacted, mask));
256     try expectOracleVector(D, .{ 0, 0, 0, 0, 1, 2, 3, 4 }, slideUpLanes(D, a, 3));
257     try expectOracleVector(D, .{ 3, 4, 5, 6, 7, 0, 0, 0 }, slideDownLanes(D, a, 3));
258     try expectOracleVector(D, @splat(28), sumOfLanes(D, a));
259     try std.testing.expectEqual(@as(u32, 14), maskedReduceSum(D, mask, a));
260     const indices: D.Vector = .{ 7, 0, 5, 2, 4, 1, 6, 3 };
261     try expectOracleVector(D, .{ 7, 0, 5, 2, 4, 1, 6, 3 }, tableLookupLanes(D, a, indices));
262     try expectOracleVector(D, .{ 3, 2, 1, 0, 7, 6, 5, 4 }, per4LaneBlockShuffle(D, 0, 1, 2, 3, a));
263     var interleaved: [16]u32 = undefined;
264     storeInterleaved2(D, a, b, &interleaved);
265     try std.testing.expectEqualSlices(u32, &.{
266         0, 10, 1, 11, 2, 12, 3, 13, 4, 14, 5, 15, 6, 16, 7, 17,
267     }, &interleaved);
268 }
269 
270 test "simd root composes a vector load compute store pipeline" {
271     const D = FixedTag(u32, 4);
272     const left = [_]u32{ 1, 2, 3, 4 };
273     const right = [_]u32{ 10, 20, 30, 40 };
274     var output: [4]u32 = undefined;
275     store(D, add(D, load(D, &left), load(D, &right)), &output);
276     try std.testing.expectEqualSlices(u32, &.{ 11, 22, 33, 44 }, &output);
277 }
278 
279 test "pinned Highway AVX2 oracle agrees on wrapping lanes" {
280     const D = FixedTag(u32, 8);
281     const a = iota(D, 0xffff_fffc);
282     const b = iota(D, 3);
283     const mask = lt(D, a, b);
284     try std.testing.expect(allTrue(D, eq(D, add(D, a, b), @as(D.Vector, .{ 0xffff_ffff, 1, 3, 5, 7, 9, 11, 13 }))));
285     try std.testing.expect(allTrue(D, eq(D, sub(D, a, b), @as(D.Vector, @splat(0xffff_fff9)))));
286     try std.testing.expect(allTrue(D, eq(D, mul(D, a, b), @as(D.Vector, .{ 0xffff_fff4, 0xffff_fff4, 0xffff_fff6, 0xffff_fffa, 0, 8, 18, 30 }))));
287     try std.testing.expect(allTrue(D, eq(D, select(D, mask, a, b), @as(D.Vector, .{ 3, 4, 5, 6, 0, 1, 2, 3 }))));
288     try std.testing.expectEqual(@as(u32, 48), reduceSum(D, add(D, a, b)));
289     try std.testing.expectEqual(@as(usize, 4), countTrue(D, mask));
290 }
291 
292 test "pinned Highway AVX2 oracle agrees on bit operations" {
293     try verifySignedBitOracle();
294     try verifyWideRotateOracle();
295     try verifyWideCountOracle();
296     try verifySignOracle();
297 }
298 
299 test "pinned Highway AVX2 oracle agrees on lane geometry" {
300     try verifyLayoutOracle();
301     try verifyResizeOracle();
302 }
303 
304 test "pinned Highway AVX2 oracle agrees on numeric conversions" {
305     try verifyConversionOracle();
306 }
307 
308 test "pinned Highway AVX2 oracle agrees on swizzle compaction and reduction" {
309     try verifyArrangeOracle();
310 }
311 
312 test "simd package namespace" {
313     std.testing.refAllDecls(simd);
314 }
315 
316 test "pinned Highway public declarations stay present" {
317     comptime {
318         for (parity.core_callables) |callable| {
319             if (!@hasDecl(simd, callable.zig)) {
320                 @compileError("missing pinned Highway callable: " ++ callable.upstream);
321             }
322         }
323         for (parity.supplemental_callables) |callable| {
324             if (!@hasDecl(simd, callable.zig)) {
325                 @compileError("missing pinned Highway declaration: " ++ callable.upstream);
326             }
327         }
328     }
329 }
330 
331 test "pinned Highway base scalar differential digest" {
332     var digest: u64 = 0xcbf2_9ce4_8422_2325;
333     const divisors32 = [_]u32{ 1, 2, 3, 7, 31, 65_537, 0x8000_0001, 0xffff_ffff };
334     const dividends32 = [_]u32{ 0, 1, 2, 7, 0x7fff_ffff, 0x8000_0000, 0xffff_fffe, 0xffff_ffff };
335     for (divisors32) |divisor| {
336         const prepared = Divisor.init(divisor);
337         digest = baseDigestStep(digest, prepared.getDivisor());
338         for (dividends32) |dividend| {
339             digest = baseDigestStep(digest, prepared.divide(dividend));
340             digest = baseDigestStep(digest, prepared.remainder(dividend));
341             digest = baseDigestStep(digest, lemireMod(dividend, divisor));
342         }
343     }
344 
345     const divisors64 = [_]u64{
346         1,
347         2,
348         3,
349         7,
350         0x0000_0001_0000_0001,
351         0x7fff_ffff_ffff_ffff,
352         0x8000_0000_0000_0001,
353         0xffff_ffff_ffff_ffff,
354     };
355     const dividends64 = [_]u64{
356         0,
357         1,
358         2,
359         7,
360         0x7fff_ffff_ffff_ffff,
361         0x8000_0000_0000_0000,
362         0xffff_ffff_ffff_fffe,
363         0xffff_ffff_ffff_ffff,
364     };
365     for (divisors64) |divisor| {
366         const prepared = Divisor64.init(divisor);
367         digest = baseDigestStep(digest, prepared.getDivisor());
368         for (dividends64) |dividend| {
369             digest = baseDigestStep(digest, prepared.divide(dividend));
370             digest = baseDigestStep(digest, prepared.remainder(dividend));
371         }
372     }
373     try std.testing.expectEqual(@as(u64, 300_788_720_801_098_909), digest);
374 }
375 
376 fn baseDigestStep(digest: u64, value: u64) u64 {
377     return (digest ^ value) *% 0x0000_0100_0000_01b3;
378 }
379 
380 test {
381     _ = parity;
382     _ = @import("topology/test.zig");
383     _ = @import("thread/test.zig");
384     _ = @import("image/test.zig");
385     _ = @import("sort/test.zig");
386     _ = @import("phast/test.zig");
387     _ = @import("cuckoo/test.zig");
388 }