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 }