lib/accy/src/kernel/library/catalog/select/owned.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2
3 const descriptor_mod = @import("../root.zig");
4 const family_mod = @import("../family/root.zig");
5 const query_mod = descriptor_mod;
6
7 const static = @import("static/root.zig");
8
9 const OwnedDescriptor = descriptor_mod.OwnedDescriptor;
10 const Query = query_mod.Query;
11
12 pub fn selectOwned(backing_allocator: std.mem.Allocator, query: Query) !?OwnedDescriptor {
13 if (static.select(query)) |descriptor| return .{ .descriptor = descriptor };
14 return switch (query) {
15 .batched_matrix_product => |batched_matrix_product| family_mod.selectBatchedMatrixProduct(backing_allocator, batched_matrix_product),
16 .matrix_product => |matrix_product| family_mod.selectMatrixProduct(backing_allocator, matrix_product),
17 .matrix_vector_product => |matrix_vector_product| family_mod.selectMatrixVectorProduct(backing_allocator, matrix_vector_product),
18 .outer_product => |outer_product| family_mod.selectOuterProduct(backing_allocator, outer_product),
19 .stencil => |stencil_query| family_mod.selectStencilWindow(backing_allocator, stencil_query),
20 .gather => |gather_query| family_mod.selectGather(backing_allocator, gather_query),
21 .scan => |scan_query| family_mod.selectPrefixSum(backing_allocator, scan_query),
22 .sort => |sort_query| family_mod.selectSort(backing_allocator, sort_query),
23 .sparse => |sparse_query| family_mod.selectSparse(backing_allocator, sparse_query),
24 .factor => |factor_query| family_mod.selectFactor(backing_allocator, factor_query),
25 .spatial => |spatial_query| family_mod.selectSpatial(backing_allocator, spatial_query),
26 .image => |image_query| family_mod.selectImage(backing_allocator, image_query),
27 .scatter => |scatter_query| family_mod.selectScatter(backing_allocator, scatter_query),
28 .scatter_add => |scatter_add_query| family_mod.selectScatterAdd(backing_allocator, scatter_add_query),
29 .row_sparse_cross_entropy => |loss_query| family_mod.selectRowSparseCrossEntropy(backing_allocator, loss_query),
30 .histogram => |histogram_query| family_mod.selectHistogram(backing_allocator, histogram_query),
31 .segmented => |segmented_query| family_mod.selectSegmentSum(backing_allocator, segmented_query),
32 .random => |random_query| family_mod.selectRandom(backing_allocator, random_query),
33 .filter => |filter_query| family_mod.selectFilter(backing_allocator, filter_query),
34 else => null,
35 };
36 }