lib/simd/src/compare.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 pub fn eq(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
4 return a == b;
5 }
6
7 pub fn ne(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
8 return a != b;
9 }
10
11 pub fn lt(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
12 return a < b;
13 }
14
15 pub fn le(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
16 return a <= b;
17 }
18
19 pub fn gt(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
20 return a > b;
21 }
22
23 pub fn ge(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
24 return a >= b;
25 }
26
27 pub fn lt128(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
28 return compare128(D, a, b, .less, false);
29 }
30
31 pub fn lt128Upper(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
32 return compare128(D, a, b, .less, true);
33 }
34
35 pub fn eq128(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
36 return compare128(D, a, b, .equal, false);
37 }
38
39 pub fn ne128(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
40 return compare128(D, a, b, .not_equal, false);
41 }
42
43 pub fn eq128Upper(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
44 return compare128(D, a, b, .equal, true);
45 }
46
47 pub fn ne128Upper(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
48 return compare128(D, a, b, .not_equal, true);
49 }
50
51 pub fn isNegative(comptime D: type, value: D.Vector) D.Mask {
52 const info = @typeInfo(D.Lane);
53 if (info != .float and (info != .int or info.int.signedness != .signed)) {
54 @compileError("isNegative requires signed integer or floating-point lanes");
55 }
56 const U = @Int(.unsigned, @bitSizeOf(D.Lane));
57 const UV = @Vector(D.lane_count, U);
58 const bits: UV = @bitCast(value);
59 const sign: UV = @splat(@as(U, 1) << (@bitSizeOf(D.Lane) - 1));
60 return bits & sign != @as(UV, @splat(0));
61 }
62
63 pub fn isNaN(comptime D: type, value: D.Vector) D.Mask {
64 validateFloat(D);
65 return value != value;
66 }
67
68 pub fn isEitherNaN(comptime D: type, a: D.Vector, b: D.Vector) D.Mask {
69 return isNaN(D, a) | isNaN(D, b);
70 }
71
72 pub fn isInf(comptime D: type, value: D.Vector) D.Mask {
73 validateFloat(D);
74 return @abs(value) == @as(D.Vector, @splat(std.math.inf(D.Lane)));
75 }
76
77 pub fn isFinite(comptime D: type, value: D.Vector) D.Mask {
78 validateFloat(D);
79 return @abs(value) < @as(D.Vector, @splat(std.math.inf(D.Lane)));
80 }
81
82 pub fn maskedEq(comptime D: type, mask: D.Mask, a: D.Vector, b: D.Vector) D.Mask {
83 return mask & eq(D, a, b);
84 }
85
86 pub fn maskedNe(comptime D: type, mask: D.Mask, a: D.Vector, b: D.Vector) D.Mask {
87 return mask & ne(D, a, b);
88 }
89
90 pub fn maskedLt(comptime D: type, mask: D.Mask, a: D.Vector, b: D.Vector) D.Mask {
91 return mask & lt(D, a, b);
92 }
93
94 pub fn maskedLe(comptime D: type, mask: D.Mask, a: D.Vector, b: D.Vector) D.Mask {
95 return mask & le(D, a, b);
96 }
97
98 pub fn maskedGt(comptime D: type, mask: D.Mask, a: D.Vector, b: D.Vector) D.Mask {
99 return mask & gt(D, a, b);
100 }
101
102 pub fn maskedGe(comptime D: type, mask: D.Mask, a: D.Vector, b: D.Vector) D.Mask {
103 return mask & ge(D, a, b);
104 }
105
106 pub fn maskedIsNaN(comptime D: type, mask: D.Mask, value: D.Vector) D.Mask {
107 return mask & isNaN(D, value);
108 }
109
110 pub fn select(
111 comptime D: type,
112 mask: D.Mask,
113 yes: D.Vector,
114 no: D.Vector,
115 ) D.Vector {
116 return @select(D.Lane, mask, yes, no);
117 }
118
119 pub fn maskNot(comptime D: type, mask: D.Mask) D.Mask {
120 return @select(bool, mask, @as(D.Mask, @splat(false)), @as(D.Mask, @splat(true)));
121 }
122
123 pub fn maskAnd(comptime D: type, a: D.Mask, b: D.Mask) D.Mask {
124 return a & b;
125 }
126
127 pub fn maskAndNot(comptime D: type, not_a: D.Mask, b: D.Mask) D.Mask {
128 return maskNot(D, not_a) & b;
129 }
130
131 pub fn maskOr(comptime D: type, a: D.Mask, b: D.Mask) D.Mask {
132 return a | b;
133 }
134
135 pub fn maskXor(comptime D: type, a: D.Mask, b: D.Mask) D.Mask {
136 return a ^ b;
137 }
138
139 pub fn exclusiveNeither(comptime D: type, a: D.Mask, b: D.Mask) D.Mask {
140 return maskNot(D, a) & maskNot(D, b);
141 }
142
143 pub fn allTrue(comptime D: type, mask: D.Mask) bool {
144 return @reduce(.And, mask);
145 }
146
147 pub fn anyTrue(comptime D: type, mask: D.Mask) bool {
148 return @reduce(.Or, mask);
149 }
150
151 pub fn allFalse(comptime D: type, mask: D.Mask) bool {
152 return !anyTrue(D, mask);
153 }
154
155 pub fn countTrue(comptime D: type, mask: D.Mask) usize {
156 var count: usize = 0;
157 inline for (0..D.lane_count) |index| {
158 count += @intFromBool(mask[index]);
159 }
160 std.debug.assert(count <= D.lane_count);
161 return count;
162 }
163
164 /// Returns the lowest true lane, or -1 when no lane is true. The mask selects
165 /// lane indices and a minimum reduction picks the first, which compiles to a
166 /// few vector instructions. Walking the lanes of an array copy, or bitcasting
167 /// the mask to an integer, scalarizes under the pinned compiler.
168 pub fn findFirstTrue(comptime D: type, mask: D.Mask) isize {
169 const index = std.simd.firstTrue(mask) orelse return -1;
170 return @intCast(index);
171 }
172
173 pub fn findKnownFirstTrue(comptime D: type, mask: D.Mask) usize {
174 std.debug.assert(anyTrue(D, mask));
175 return @intCast(findFirstTrue(D, mask));
176 }
177
178 /// Returns the highest true lane, or -1 when no lane is true, by a maximum
179 /// reduction over the selected lane indices.
180 pub fn findLastTrue(comptime D: type, mask: D.Mask) isize {
181 const index = std.simd.lastTrue(mask) orelse return -1;
182 return @intCast(index);
183 }
184
185 pub fn findKnownLastTrue(comptime D: type, mask: D.Mask) usize {
186 std.debug.assert(anyTrue(D, mask));
187 return @intCast(findLastTrue(D, mask));
188 }
189
190 pub fn setOnlyFirst(comptime D: type, mask: D.Mask) D.Mask {
191 var result: [D.lane_count]bool = @splat(false);
192 const first = findFirstTrue(D, mask);
193 if (first >= 0) result[@intCast(first)] = true;
194 return result;
195 }
196
197 pub fn setBeforeFirst(comptime D: type, mask: D.Mask) D.Mask {
198 const first = findFirstTrue(D, mask);
199 const boundary: usize = if (first < 0) D.lane_count else @intCast(first);
200 var result: D.Mask = @splat(false);
201 inline for (0..D.lane_count) |index| result[index] = index < boundary;
202 return result;
203 }
204
205 pub fn setAtOrBeforeFirst(comptime D: type, mask: D.Mask) D.Mask {
206 return maskOr(D, setBeforeFirst(D, mask), setOnlyFirst(D, mask));
207 }
208
209 pub fn setAtOrAfterFirst(comptime D: type, mask: D.Mask) D.Mask {
210 return maskNot(D, setBeforeFirst(D, mask));
211 }
212
213 pub fn bitsFromMask(comptime D: type, mask: D.Mask) u64 {
214 if (D.lane_count > 64) @compileError("bitsFromMask supports at most 64 lanes");
215 var bits: u64 = 0;
216 inline for (0..D.lane_count) |index| {
217 if (mask[index]) bits |= @as(u64, 1) << @intCast(index);
218 }
219 return bits;
220 }
221
222 pub fn maskFromBits(comptime D: type, bits: u64) D.Mask {
223 if (D.lane_count > 64) @compileError("maskFromBits supports at most 64 lanes");
224 var mask: D.Mask = @splat(false);
225 inline for (0..D.lane_count) |index| {
226 mask[index] = bits & (@as(u64, 1) << @intCast(index)) != 0;
227 }
228 return mask;
229 }
230
231 pub fn rebindMask(comptime D: type, mask: anytype) D.Mask {
232 if (@TypeOf(mask) != D.Mask) @compileError("rebindMask requires equal lane counts");
233 return mask;
234 }
235
236 pub fn vecFromMask(comptime D: type, mask: D.Mask) D.Vector {
237 const U = @Int(.unsigned, @bitSizeOf(D.Lane));
238 const UV = @Vector(D.lane_count, U);
239 const bits: UV = @select(
240 U,
241 mask,
242 @as(UV, @splat(std.math.maxInt(U))),
243 @as(UV, @splat(0)),
244 );
245 return @bitCast(bits);
246 }
247
248 pub fn maskFromVec(comptime D: type, value: D.Vector) D.Mask {
249 const U = @Int(.unsigned, @bitSizeOf(D.Lane));
250 const UV = @Vector(D.lane_count, U);
251 const bits: UV = @bitCast(value);
252 return bits != @as(UV, @splat(0));
253 }
254
255 pub fn loadMaskBits(comptime D: type, input: []const u8) D.Mask {
256 const byte_count = (D.lane_count + 7) / 8;
257 std.debug.assert(input.len >= byte_count);
258 var result: D.Mask = @splat(false);
259 inline for (0..D.lane_count) |index| {
260 result[index] = input[index / 8] & (@as(u8, 1) << @intCast(index % 8)) != 0;
261 }
262 return result;
263 }
264
265 pub fn storeMaskBits(comptime D: type, mask: D.Mask, output: []u8) usize {
266 const byte_count = (D.lane_count + 7) / 8;
267 std.debug.assert(output.len >= byte_count);
268 @memset(output[0..byte_count], 0);
269 inline for (0..D.lane_count) |index| {
270 if (mask[index]) output[index / 8] |= @as(u8, 1) << @intCast(index % 8);
271 }
272 return byte_count;
273 }
274
275 pub fn dup128MaskFromMaskBits(comptime D: type, bits: u64) D.Mask {
276 const period = 16 / @sizeOf(D.Lane);
277 var result: D.Mask = @splat(false);
278 inline for (0..D.lane_count) |index| {
279 result[index] = bits & (@as(u64, 1) << @intCast(index % period)) != 0;
280 }
281 return result;
282 }
283
284 pub fn maskFalse(comptime D: type) D.Mask {
285 return @splat(false);
286 }
287
288 pub fn setMask(comptime D: type, value: bool) D.Mask {
289 return @splat(value);
290 }
291
292 fn validateFloat(comptime D: type) void {
293 if (@typeInfo(D.Lane) != .float) @compileError("classification requires floating-point lanes");
294 }
295
296 const PairComparison = enum { less, equal, not_equal };
297
298 fn compare128(
299 comptime D: type,
300 a: D.Vector,
301 b: D.Vector,
302 comptime comparison: PairComparison,
303 comptime upper_only: bool,
304 ) D.Mask {
305 if (comptime D.Lane != u64 or D.lane_count < 2 or D.lane_count & 1 != 0) {
306 @compileError("128-bit comparison requires an even number of u64 lanes");
307 }
308 var result: D.Mask = @splat(false);
309 inline for (0..D.lane_count / 2) |pair| {
310 const low = pair * 2;
311 const high = low + 1;
312 const matches = switch (comparison) {
313 .less => if (upper_only)
314 a[high] < b[high]
315 else
316 a[high] < b[high] or (a[high] == b[high] and a[low] < b[low]),
317 .equal => a[high] == b[high] and (upper_only or a[low] == b[low]),
318 .not_equal => a[high] != b[high] or (!upper_only and a[low] != b[low]),
319 };
320 result[low] = matches;
321 result[high] = matches;
322 }
323 return result;
324 }
325
326 test "comparisons and selection operate independently per lane" {
327 const simd = @import("root.zig");
328 const D = simd.FixedTag(i32, 4);
329 const a: D.Vector = .{ 1, 5, -3, 8 };
330 const b: D.Vector = .{ 2, 5, -4, 9 };
331 const expected_mask: D.Mask = .{ true, false, false, true };
332 try std.testing.expect(allTrue(D, lt(D, a, b) == expected_mask));
333 const expected: D.Vector = .{ 1, 5, -4, 8 };
334 try std.testing.expect(allTrue(D, eq(D, select(D, expected_mask, a, b), expected)));
335 }
336
337 test "mask bit round trips retain lane order" {
338 const simd = @import("root.zig");
339 const D = simd.FixedTag(u8, 8);
340 const mask = maskFromBits(D, 0xa5);
341 try std.testing.expectEqual(@as(u64, 0xa5), bitsFromMask(D, mask));
342 try std.testing.expectEqual(@as(usize, 4), countTrue(D, mask));
343 try std.testing.expect(allTrue(D, maskOr(D, mask, maskNot(D, mask))));
344 try std.testing.expect(allFalse(D, maskAnd(D, mask, maskNot(D, mask))));
345 }
346
347 test "Highway mask vectors and byte storage retain canonical bits" {
348 const simd = @import("root.zig");
349 inline for (.{ u8, i8, u16, i16, u32, i32, u64, i64, f16, f32, f64 }) |T| {
350 const E = simd.FixedTag(T, 4);
351 const expected: E.Mask = .{ true, false, true, false };
352 try std.testing.expect(@reduce(.And, maskFromVec(E, vecFromMask(E, expected)) == expected));
353 }
354 const D = simd.FixedTag(u16, 16);
355 const mask: D.Mask = .{
356 true, false, true, false, false, true, false, true,
357 false, true, false, false, true, false, true, false,
358 };
359 try std.testing.expect(@reduce(.And, maskFromVec(D, vecFromMask(D, mask)) == mask));
360 var bytes: [2]u8 = undefined;
361 try std.testing.expectEqual(@as(usize, 2), storeMaskBits(D, mask, &bytes));
362 try std.testing.expectEqualSlices(u8, &.{ 0xa5, 0x52 }, &bytes);
363 try std.testing.expect(@reduce(.And, loadMaskBits(D, &bytes) == mask));
364 try std.testing.expect(@reduce(.And, dup128MaskFromMaskBits(D, 0xa5) ==
365 @as(D.Mask, .{
366 true, false, true, false, false, true, false, true,
367 true, false, true, false, false, true, false, true,
368 })));
369 try std.testing.expect(allFalse(D, maskFalse(D)));
370 try std.testing.expect(allTrue(D, setMask(D, true)));
371 try std.testing.expect(@reduce(.And, rebindMask(simd.FixedTag(f16, 16), mask) == mask));
372 }
373
374 test "Highway sign classification reads bits for integers and floats" {
375 const simd = @import("root.zig");
376 const DI = simd.FixedTag(i32, 4);
377 const DF = simd.FixedTag(f32, 4);
378 const DU = simd.FixedTag(u32, 4);
379 const integers: DI.Vector = .{ 0, -1, 1, std.math.minInt(i32) };
380 const floats: DF.Vector = @bitCast(@as(DU.Vector, .{
381 0, 0x8000_0000, 0x7fc0_0001, 0xffc0_0001,
382 }));
383 try std.testing.expect(allTrue(DI, isNegative(DI, integers) ==
384 @as(DI.Mask, .{ false, true, false, true })));
385 try std.testing.expect(allTrue(DF, isNegative(DF, floats) ==
386 @as(DF.Mask, .{ false, true, false, true })));
387 try std.testing.expect(allTrue(DF, isNaN(DF, floats) ==
388 @as(DF.Mask, .{ false, false, true, true })));
389 try std.testing.expect(allTrue(DF, isEitherNaN(DF, floats, @as(DF.Vector, @splat(0))) ==
390 @as(DF.Mask, .{ false, false, true, true })));
391 }
392
393 test "Highway floating classification distinguishes infinity and finiteness" {
394 const simd = @import("root.zig");
395 inline for (.{ f16, f32, f64 }) |T| {
396 const D = simd.FixedTag(T, 4);
397 const value: D.Vector = .{ 0, -1, std.math.inf(T), std.math.nan(T) };
398 try std.testing.expect(allTrue(D, isInf(D, value) ==
399 @as(D.Mask, .{ false, false, true, false })));
400 try std.testing.expect(allTrue(D, isFinite(D, value) ==
401 @as(D.Mask, .{ true, true, false, false })));
402 }
403 }
404
405 test "Highway masked comparisons clear inactive lanes" {
406 const simd = @import("root.zig");
407 const D = simd.FixedTag(u16, 4);
408 const a: D.Vector = .{ 1, 2, 3, 4 };
409 const b: D.Vector = .{ 1, 3, 2, 4 };
410 const mask: D.Mask = .{ true, true, false, false };
411 try std.testing.expect(allTrue(D, maskedEq(D, mask, a, b) ==
412 @as(D.Mask, .{ true, false, false, false })));
413 try std.testing.expect(allTrue(D, maskedLt(D, mask, a, b) ==
414 @as(D.Mask, .{ false, true, false, false })));
415 try std.testing.expect(allTrue(D, maskedNe(D, mask, a, b) ==
416 @as(D.Mask, .{ false, true, false, false })));
417 try std.testing.expect(allTrue(D, maskedLe(D, mask, a, b) ==
418 @as(D.Mask, .{ true, true, false, false })));
419 try std.testing.expect(allTrue(D, maskedGt(D, mask, a, b) ==
420 @as(D.Mask, .{ false, false, false, false })));
421 try std.testing.expect(allTrue(D, maskedGe(D, mask, a, b) ==
422 @as(D.Mask, .{ true, false, false, false })));
423 const F = simd.FixedTag(f32, 4);
424 const floats: F.Vector = .{ std.math.nan(f32), 0, std.math.nan(f32), 1 };
425 try std.testing.expect(allTrue(F, maskedIsNaN(F, @as(F.Mask, mask), floats) ==
426 @as(F.Mask, .{ true, false, false, false })));
427 }
428
429 test "Highway pair comparisons broadcast full and upper-key results" {
430 const simd = @import("root.zig");
431 const D = simd.FixedTag(u64, 8);
432 const a: D.Vector = .{ 9, 1, 7, 4, 10, 6, 20, 8 };
433 const b: D.Vector = .{ 10, 1, 8, 3, 10, 6, 19, 8 };
434 try std.testing.expect(@reduce(.And, lt128(D, a, b) ==
435 @as(D.Mask, .{ true, true, false, false, false, false, false, false })));
436 try std.testing.expect(@reduce(.And, lt128Upper(D, a, b) ==
437 @as(D.Mask, .{ false, false, false, false, false, false, false, false })));
438 try std.testing.expect(@reduce(.And, eq128(D, a, b) ==
439 @as(D.Mask, .{ false, false, false, false, true, true, false, false })));
440 try std.testing.expect(@reduce(.And, ne128(D, a, b) ==
441 @as(D.Mask, .{ true, true, true, true, false, false, true, true })));
442 try std.testing.expect(@reduce(.And, eq128Upper(D, a, b) ==
443 @as(D.Mask, .{ true, true, false, false, true, true, true, true })));
444 try std.testing.expect(@reduce(.And, ne128Upper(D, a, b) ==
445 @as(D.Mask, .{ false, false, true, true, false, false, false, false })));
446 }
447
448 test "Highway mask searches and first-boundary transforms cover empty and populated masks" {
449 const simd = @import("root.zig");
450 const D = simd.FixedTag(u8, 8);
451 const mask: D.Mask = .{ false, false, true, false, true, false, false, true };
452 try std.testing.expectEqual(@as(isize, 2), findFirstTrue(D, mask));
453 try std.testing.expectEqual(@as(usize, 2), findKnownFirstTrue(D, mask));
454 try std.testing.expectEqual(@as(isize, 7), findLastTrue(D, mask));
455 try std.testing.expectEqual(@as(usize, 7), findKnownLastTrue(D, mask));
456 try std.testing.expect(@reduce(.And, setOnlyFirst(D, mask) ==
457 @as(D.Mask, .{ false, false, true, false, false, false, false, false })));
458 try std.testing.expect(@reduce(.And, setBeforeFirst(D, mask) ==
459 @as(D.Mask, .{ true, true, false, false, false, false, false, false })));
460 try std.testing.expect(@reduce(.And, setAtOrBeforeFirst(D, mask) ==
461 @as(D.Mask, .{ true, true, true, false, false, false, false, false })));
462 try std.testing.expect(@reduce(.And, setAtOrAfterFirst(D, mask) ==
463 @as(D.Mask, .{ false, false, true, true, true, true, true, true })));
464 const empty: D.Mask = @splat(false);
465 try std.testing.expectEqual(@as(isize, -1), findFirstTrue(D, empty));
466 try std.testing.expectEqual(@as(isize, -1), findLastTrue(D, empty));
467 try std.testing.expect(allTrue(D, setBeforeFirst(D, empty)));
468 try std.testing.expect(allFalse(D, setOnlyFirst(D, empty)));
469 try std.testing.expect(@reduce(.And, maskAndNot(D, mask, setBeforeFirst(D, mask)) ==
470 @as(D.Mask, .{ true, true, false, false, false, false, false, false })));
471 try std.testing.expect(@reduce(.And, exclusiveNeither(D, mask, setBeforeFirst(D, mask)) ==
472 @as(D.Mask, .{ false, false, false, true, false, true, true, false })));
473 }