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 }