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 }