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 }