lib/accy/src/kernel/library/catalog/family/factor.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const descriptor_mod = @import("../root.zig");
4 const factor_mod = @import("../../root.zig").factor;
5 const match_mod = @import("../match/root.zig");
6 const query_mod = descriptor_mod;
7
8 const Descriptor = descriptor_mod.Descriptor;
9 const OwnedDescriptor = descriptor_mod.OwnedDescriptor;
10 const FactorQuery = query_mod.FactorQuery;
11
12 pub fn selectFactor(backing_allocator: std.mem.Allocator, query: FactorQuery) !?OwnedDescriptor {
13 if (query.dtype != .f32) return null;
14 const threads = factorQueryThreads(query) orelse return null;
15 switch (query.kind) {
16 .batched_cholesky => |facts| {
17 const instance = factor_mod.BatchedCholesky{
18 .batch = facts.batch,
19 .n = facts.n,
20 .threads = threads,
21 .layout = factorInstanceLayout(query.layout),
22 };
23 if (!factor_mod.batchedCholeskyInstanceValid(instance)) return null;
24 var specialization = try factor_mod.batchedCholeskyFamilySpecialization(backing_allocator, instance);
25 errdefer specialization.deinit();
26 const lifetime_allocator = specialization.allocator();
27 const descriptor = Descriptor{
28 .name = try factor_mod.batchedCholeskyFamilyEntryName(lifetime_allocator, instance),
29 .metadata = .{
30 .target = try factor_mod.batchedCholeskyFamilyTarget(lifetime_allocator, instance),
31 .version = factor_mod.batched_cholesky_family_version,
32 .layer = .logical,
33 .category = .linalg,
34 .specialization = specialization.value,
35 },
36 };
37 if (!match_mod.factorDescriptorMatches(descriptor, query)) {
38 specialization.deinit();
39 return null;
40 }
41 return .{ .descriptor = descriptor, .specialization = specialization };
42 },
43 .batched_cholesky_solve => |facts| {
44 const instance = factor_mod.BatchedCholeskySolve{
45 .batch = facts.batch,
46 .n = facts.n,
47 .threads = threads,
48 .layout = factorInstanceLayout(query.layout),
49 };
50 if (!factor_mod.batchedCholeskySolveInstanceValid(instance)) return null;
51 var specialization = try factor_mod.batchedCholeskySolveFamilySpecialization(backing_allocator, instance);
52 errdefer specialization.deinit();
53 const lifetime_allocator = specialization.allocator();
54 const descriptor = Descriptor{
55 .name = try factor_mod.batchedCholeskySolveFamilyEntryName(lifetime_allocator, instance),
56 .metadata = .{
57 .target = try factor_mod.batchedCholeskySolveFamilyTarget(lifetime_allocator, instance),
58 .version = factor_mod.batched_cholesky_solve_family_version,
59 .layer = .logical,
60 .category = .linalg,
61 .specialization = specialization.value,
62 },
63 };
64 if (!match_mod.factorDescriptorMatches(descriptor, query)) {
65 specialization.deinit();
66 return null;
67 }
68 return .{ .descriptor = descriptor, .specialization = specialization };
69 },
70 .batched_inverse => |facts| {
71 const instance = factor_mod.BatchedInverse{
72 .batch = facts.batch,
73 .n = facts.n,
74 .threads = threads,
75 .layout = factorInstanceLayout(query.layout),
76 };
77 if (!factor_mod.batchedInverseInstanceValid(instance)) return null;
78 var specialization = try factor_mod.batchedInverseFamilySpecialization(backing_allocator, instance);
79 errdefer specialization.deinit();
80 const lifetime_allocator = specialization.allocator();
81 const descriptor = Descriptor{
82 .name = try factor_mod.batchedInverseFamilyEntryName(lifetime_allocator, instance),
83 .metadata = .{
84 .target = try factor_mod.batchedInverseFamilyTarget(lifetime_allocator, instance),
85 .version = factor_mod.batched_inverse_family_version,
86 .layer = .logical,
87 .category = .linalg,
88 .specialization = specialization.value,
89 },
90 };
91 if (!match_mod.factorDescriptorMatches(descriptor, query)) {
92 specialization.deinit();
93 return null;
94 }
95 return .{ .descriptor = descriptor, .specialization = specialization };
96 },
97 }
98 }
99
100 fn factorInstanceLayout(layout: query_mod.FactorTileLayout) factor_mod.TileLayout {
101 return switch (layout) {
102 .row_major => .row_major,
103 .interleaved => .interleaved,
104 };
105 }
106
107 fn factorQueryThreads(query: FactorQuery) ?u32 {
108 if (query.schedule) |requested| {
109 switch (requested) {
110 .thread_blocks => |threads| {
111 if (threads == 0) return null;
112 return threads;
113 },
114 }
115 }
116 return 256;
117 }
118
119 test "catalog factor selection covers both kernels" {
120 const allocator = std.testing.allocator;
121
122 var cholesky = (try selectFactor(allocator, .{
123 .dtype = .f32,
124 .kind = .{ .batched_cholesky = .{ .n = 3, .batch = 5000 } },
125 })) orelse return error.TestExpectedFactorDescriptor;
126 defer cholesky.deinit();
127 try std.testing.expectEqualStrings(
128 "accy.kernel.linalg.batched_cholesky_family_3_256_f32",
129 cholesky.descriptor.metadata.target,
130 );
131
132 var solve = (try selectFactor(allocator, .{
133 .dtype = .f32,
134 .kind = .{ .batched_cholesky_solve = .{ .n = 4, .batch = 5000 } },
135 .schedule = .{ .thread_blocks = 64 },
136 })) orelse return error.TestExpectedFactorDescriptor;
137 defer solve.deinit();
138 try std.testing.expectEqualStrings(
139 "accy.kernel.linalg.batched_cholesky_solve_family_4_64_f32",
140 solve.descriptor.metadata.target,
141 );
142
143 var interleaved = (try selectFactor(allocator, .{
144 .dtype = .f32,
145 .kind = .{ .batched_cholesky = .{ .n = 3, .batch = 5000 } },
146 .layout = .interleaved,
147 .schedule = .{ .thread_blocks = 64 },
148 })) orelse return error.TestExpectedFactorDescriptor;
149 defer interleaved.deinit();
150 try std.testing.expectEqualStrings(
151 "accy.kernel.linalg.batched_cholesky_family_3_64_il_f32",
152 interleaved.descriptor.metadata.target,
153 );
154 try std.testing.expect(interleaved.descriptor.metadata.specialization.layoutIs("interleaved"));
155
156 var interleaved_solve = (try selectFactor(allocator, .{
157 .dtype = .f32,
158 .kind = .{ .batched_cholesky_solve = .{ .n = 3, .batch = 5000 } },
159 .layout = .interleaved,
160 .schedule = .{ .thread_blocks = 64 },
161 })) orelse return error.TestExpectedFactorDescriptor;
162 defer interleaved_solve.deinit();
163 try std.testing.expectEqualStrings(
164 "accy.kernel.linalg.batched_cholesky_solve_family_3_64_il_f32",
165 interleaved_solve.descriptor.metadata.target,
166 );
167 try std.testing.expect(interleaved_solve.descriptor.metadata.specialization.layoutIs("interleaved"));
168
169 var inverse = (try selectFactor(allocator, .{
170 .dtype = .f32,
171 .kind = .{ .batched_inverse = .{ .n = 3, .batch = 5000 } },
172 .schedule = .{ .thread_blocks = 128 },
173 })) orelse return error.TestExpectedFactorDescriptor;
174 defer inverse.deinit();
175 try std.testing.expectEqualStrings(
176 "accy.kernel.linalg.batched_inverse_family_3_128_f32",
177 inverse.descriptor.metadata.target,
178 );
179 try std.testing.expect(inverse.descriptor.metadata.specialization.operationIs(.{ .linalg = .batched_inverse }));
180
181 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectFactor(allocator, .{
182 .dtype = .i32,
183 .kind = .{ .batched_cholesky = .{ .n = 3, .batch = 5000 } },
184 }));
185 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectFactor(allocator, .{
186 .dtype = .f32,
187 .kind = .{ .batched_cholesky = .{ .n = 5, .batch = 5000 } },
188 }));
189 try std.testing.expectEqual(@as(?OwnedDescriptor, null), try selectFactor(allocator, .{
190 .dtype = .f32,
191 .kind = .{ .batched_cholesky = .{ .n = 3, .batch = 0 } },
192 }));
193 }