lib/simd/src/sort/heap.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const key = @import("key.zig");
3
4 pub fn sort(
5 comptime Key: type,
6 comptime direction: key.Direction,
7 values: []Key,
8 ) void {
9 if (comptime !key.supported(Key)) @compileError("VQSort key type is not supported");
10 if (values.len < 2) return;
11
12 var start = values.len / 2;
13 while (start != 0) {
14 start -= 1;
15 siftDown(Key, direction, values, start);
16 }
17
18 var end = values.len;
19 while (end > 1) {
20 end -= 1;
21 std.mem.swap(Key, &values[0], &values[end]);
22 siftDown(Key, direction, values[0..end], 0);
23 }
24 }
25
26 fn siftDown(
27 comptime Key: type,
28 comptime direction: key.Direction,
29 values: []Key,
30 start: usize,
31 ) void {
32 std.debug.assert(start < values.len);
33 var root = start;
34 var iterations: usize = 0;
35 while (root < values.len / 2) : (iterations += 1) {
36 std.debug.assert(iterations < @bitSizeOf(usize));
37 const left = root * 2 + 1;
38 const right = left + 1;
39 var later = left;
40 if (right < values.len and key.before(Key, direction, values[left], values[right])) {
41 later = right;
42 }
43 if (!key.before(Key, direction, values[root], values[later])) return;
44 std.mem.swap(Key, &values[root], &values[later]);
45 root = later;
46 }
47 }
48
49 fn Order(comptime Key: type, comptime direction: key.Direction) type {
50 return struct {
51 fn lessThan(_: void, a: Key, b: Key) bool {
52 return key.before(Key, direction, a, b);
53 }
54 };
55 }
56
57 test "Highway VQSort heap fallback matches both scalar orders" {
58 var ascending = [_]i64{ 9, -4, 7, 7, 0, std.math.minInt(i64), 11, -4, std.math.maxInt(i64) };
59 var expected_ascending = ascending;
60 std.mem.sort(i64, &expected_ascending, {}, Order(i64, .ascending).lessThan);
61 sort(i64, .ascending, &ascending);
62 try std.testing.expectEqualSlices(i64, &expected_ascending, &ascending);
63
64 var descending = ascending;
65 var expected_descending = descending;
66 std.mem.sort(i64, &expected_descending, {}, Order(i64, .descending).lessThan);
67 sort(i64, .descending, &descending);
68 try std.testing.expectEqualSlices(i64, &expected_descending, &descending);
69 }
70
71 test "Highway VQSort heap fallback orders 128-bit records by key" {
72 var values = [_]key.K64V64{
73 key.K64V64.init(8, 1),
74 key.K64V64.init(2, 7),
75 key.K64V64.init(9, 4),
76 key.K64V64.init(5, 3),
77 };
78 sort(key.K64V64, .descending, &values);
79 for (values[1..], 1..) |value, index| try std.testing.expect(values[index - 1].key >= value.key);
80 }