lib/accy/src/kernel/library/test.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const library = @import("root.zig");
4
5 const artifact = @import("../../artifact/root.zig");
6 const catalog = library.catalog;
7 const attention = library.attention;
8 const compaction = library.compaction;
9 const entry = library.entry;
10 const elementwise = library.elementwise;
11 const extent = library.extent;
12 const fused = library.fused;
13 const geometry = library.geometry;
14 const histogram = library.histogram;
15 const image = library.image;
16 const sort = library.sort;
17 const factor = library.factor;
18 const sparse = library.sparse;
19 const spatial = library.spatial;
20 const indexing = library.indexing;
21 const layout = library.layout;
22 const linalg = library.linalg;
23 const loss = library.loss;
24 const normalization = library.normalization;
25 const random = library.random;
26 const reduction = library.reduction;
27 const scan = library.scan;
28 const sdf = library.sdf;
29 const segmented = library.segmented;
30 const stencil = library.stencil;
31 const tuning = library.tuning;
32 const CatalogDescriptor = library.CatalogDescriptor;
33 const OwnedCatalogDescriptor = library.OwnedCatalogDescriptor;
34 const OwnedKernelCallArtifactRegistry = library.OwnedKernelCallArtifactRegistry;
35 const OwnedKernelCallPipelinePackage = library.OwnedKernelCallPipelinePackage;
36 const CatalogQuery = library.CatalogQuery;
37 const ActivationKind = library.ActivationKind;
38 const ActivationQuery = library.ActivationQuery;
39 const AttentionKind = library.AttentionKind;
40 const AttentionSchedule = library.AttentionSchedule;
41 const AttentionQuery = library.AttentionQuery;
42 const FusedQuery = library.FusedQuery;
43 const FusedVectorKind = library.FusedVectorKind;
44 const FusedVectorQuery = library.FusedVectorQuery;
45 const FusedLinalgEpilogue = library.FusedLinalgEpilogue;
46 const FusedMatrixProductQuery = library.FusedMatrixProductQuery;
47 const FusedMatrixVectorProductQuery = library.FusedMatrixVectorProductQuery;
48 const FusedRowNormalizationKind = library.FusedRowNormalizationKind;
49 const FusedRowNormalizationQuery = library.FusedRowNormalizationQuery;
50 const LayoutQuery = library.LayoutQuery;
51 const LayoutKind = library.LayoutKind;
52 const TransposeQuery = library.TransposeQuery;
53 const BatchedMatrixProductSchedule = library.BatchedMatrixProductSchedule;
54 const BatchedMatrixProductQuery = library.BatchedMatrixProductQuery;
55 const MatrixProductSchedule = library.MatrixProductSchedule;
56 const MatrixProductQuery = library.MatrixProductQuery;
57 const MatrixProductCandidateDescriptors = library.MatrixProductCandidateDescriptors;
58 const MatrixVectorProductSchedule = library.MatrixVectorProductSchedule;
59 const MatrixVectorProductQuery = library.MatrixVectorProductQuery;
60 const OuterProductSchedule = library.OuterProductSchedule;
61 const OuterProductQuery = library.OuterProductQuery;
62 const ReductionKind = library.ReductionKind;
63 const ReductionOperand = library.ReductionOperand;
64 const ReductionQuery = library.ReductionQuery;
65 const RowNormalizationParameterization = library.RowNormalizationParameterization;
66 const RowNormalizationKind = library.RowNormalizationKind;
67 const RowNormalizationSchedule = library.RowNormalizationSchedule;
68 const RowSparseCrossEntropyQuery = library.RowSparseCrossEntropyQuery;
69 const RowSparseCrossEntropySchedule = library.RowSparseCrossEntropySchedule;
70 const RowNormalizationQuery = library.RowNormalizationQuery;
71 const GatherSchedule = library.GatherSchedule;
72 const GatherQuery = library.GatherQuery;
73 const GatherCandidateDescriptors = library.GatherCandidateDescriptors;
74 const ScanKind = library.ScanKind;
75 const ScanSchedule = library.ScanSchedule;
76 const ScanQuery = library.ScanQuery;
77 const PrefixSumCandidateDescriptors = library.PrefixSumCandidateDescriptors;
78 const CompactionPredicateKind = library.CompactionPredicateKind;
79 const FilterSchedule = library.FilterSchedule;
80 const FilterQuery = library.FilterQuery;
81 const FilterCandidateDescriptors = library.FilterCandidateDescriptors;
82 const RandomAlgorithm = library.RandomAlgorithm;
83 const RandomSchedule = library.RandomSchedule;
84 const RandomQuery = library.RandomQuery;
85 const RandomCandidateDescriptors = library.RandomCandidateDescriptors;
86 const ScatterSchedule = library.ScatterSchedule;
87 const ScatterQuery = library.ScatterQuery;
88 const ScatterAddSchedule = library.ScatterAddSchedule;
89 const ScatterAddQuery = library.ScatterAddQuery;
90 const ScatterAddCandidateDescriptors = library.ScatterAddCandidateDescriptors;
91 const HistogramSchedule = library.HistogramSchedule;
92 const HistogramBinningPolicy = library.HistogramBinningPolicy;
93 const HistogramQuery = library.HistogramQuery;
94 const FactorQuery = library.FactorQuery;
95 const FactorKind = library.FactorKind;
96 const FactorTileLayout = library.FactorTileLayout;
97 const FactorSchedule = library.FactorSchedule;
98 const SparseQuery = library.SparseQuery;
99 const SparseKind = library.SparseKind;
100 const SparseStructure = library.SparseStructure;
101 const SparseSchedule = library.SparseSchedule;
102 const SparseCandidateDescriptors = library.SparseCandidateDescriptors;
103 const SpatialQuery = library.SpatialQuery;
104 const SpatialKind = library.SpatialKind;
105 const SpatialSchedule = library.SpatialSchedule;
106 const ImageQuery = library.ImageQuery;
107 const ImageKind = library.ImageKind;
108 const ImageSchedule = library.ImageSchedule;
109 const ImageCandidateDescriptors = library.ImageCandidateDescriptors;
110 const HistogramCandidateDescriptors = library.HistogramCandidateDescriptors;
111 const ScatterCandidateDescriptors = library.ScatterCandidateDescriptors;
112 const SegmentedKind = library.SegmentedKind;
113 const SegmentedSchedule = library.SegmentedSchedule;
114 const SegmentedQuery = library.SegmentedQuery;
115 const SegmentSumCandidateDescriptors = library.SegmentSumCandidateDescriptors;
116 const StencilKind = library.StencilKind;
117 const StencilSchedule = library.StencilSchedule;
118 const StencilQuery = library.StencilQuery;
119 const StencilWindowCandidateDescriptors = library.StencilWindowCandidateDescriptors;
120 const CatalogArtifactRequest = library.CatalogArtifactRequest;
121 const catalog_entries = library.catalog_entries;
122 const catalog_descriptors = library.catalog_descriptors;
123 const findEntry = library.findEntry;
124 const select = library.select;
125 const selectOwned = library.selectOwned;
126 const selectOwnedMatrixProductCandidates = library.selectOwnedMatrixProductCandidates;
127 const selectOwnedStencilWindowCandidates = library.selectOwnedStencilWindowCandidates;
128 const selectOwnedGatherCandidates = library.selectOwnedGatherCandidates;
129 const selectOwnedScatterCandidates = library.selectOwnedScatterCandidates;
130 const selectOwnedScatterAddCandidates = library.selectOwnedScatterAddCandidates;
131 const selectOwnedHistogramCandidates = library.selectOwnedHistogramCandidates;
132 const selectOwnedFactor = library.selectOwnedFactor;
133 const selectOwnedSparse = library.selectOwnedSparse;
134 const selectOwnedSparseCandidates = library.selectOwnedSparseCandidates;
135 const selectOwnedSpatial = library.selectOwnedSpatial;
136 const selectOwnedSpatialCandidates = library.selectOwnedSpatialCandidates;
137 const selectOwnedImageCandidates = library.selectOwnedImageCandidates;
138 const selectOwnedRandomCandidates = library.selectOwnedRandomCandidates;
139 const selectOwnedFilterCandidates = library.selectOwnedFilterCandidates;
140 const selectOwnedPrefixSumCandidates = library.selectOwnedPrefixSumCandidates;
141 const selectOwnedSegmentSumCandidates = library.selectOwnedSegmentSumCandidates;
142 const createKernelCallArtifact = library.createKernelCallArtifact;
143 const createOwnedKernelCallArtifact = library.createOwnedKernelCallArtifact;
144 const createOwnedKernelCallPipelinePackage = library.createOwnedKernelCallPipelinePackage;
145 const createOwnedKernelCallArtifactRegistry = library.createOwnedKernelCallArtifactRegistry;
146 const Entry = library.Entry;
147 const Metadata = library.Metadata;
148 const Layer = library.Layer;
149 const Category = library.Category;
150 const Operation = library.Operation;
151 const operationFingerprint = library.operationFingerprint;
152 const AttentionOperator = library.AttentionOperator;
153 const ElementwiseOperator = library.ElementwiseOperator;
154 const Axis = library.Axis;
155 const Vector1D = library.Vector1D;
156 const Matrix2D = library.Matrix2D;
157 const Threads2D = library.Threads2D;
158 const Threads3D = library.Threads3D;
159 const Shape = library.Shape;
160 const Reduction = library.Reduction;
161 const ReductionContract = library.ReductionContract;
162 const ReductionReuse = library.ReductionReuse;
163 const ReductionReuseContract = library.ReductionReuseContract;
164 const ReductionOperator = library.ReductionOperator;
165 const StaticParameter = library.StaticParameter;
166 const RowNormalizationOperator = library.RowNormalizationOperator;
167 const LayoutOperator = library.LayoutOperator;
168 const LinalgOperator = library.LinalgOperator;
169 const StencilOperator = library.StencilOperator;
170 const ImageOperator = library.ImageOperator;
171 const IndexingOperator = library.IndexingOperator;
172 const SegmentedOperator = library.SegmentedOperator;
173 const ScanOperator = library.ScanOperator;
174 const RandomOperator = library.RandomOperator;
175 const CompactionOperator = library.CompactionOperator;
176 const CompactionPredicate = library.CompactionPredicate;
177 const InputTransform = library.InputTransform;
178 const InputTransformContract = library.InputTransformContract;
179 const InputTransformOperator = library.InputTransformOperator;
180 const EpilogueStep = library.EpilogueStep;
181 const EpilogueContract = library.EpilogueContract;
182 const EpilogueOperator = library.EpilogueOperator;
183 const Launch = library.Launch;
184 const Schedule = library.Schedule;
185 const ScheduleBinding = library.ScheduleBinding;
186 const Specialization = library.Specialization;
187 const OwnedSpecialization = library.OwnedSpecialization;
188 const ArtifactOptions = library.ArtifactOptions;
189 const shape1D = library.shape1D;
190 const runtimeShape1D = library.runtimeShape1D;
191 const shapeScalar = library.shapeScalar;
192 const runtimeShapeScalar = library.runtimeShapeScalar;
193 const shape2D = library.shape2D;
194 const runtimeShape2D = library.runtimeShape2D;
195 const shape3D = library.shape3D;
196 const runtimeShape3D = library.runtimeShape3D;
197 const describeReduction = library.describeReduction;
198 const describeRuntimeReduction = library.describeRuntimeReduction;
199 const describeDependentReduction = library.describeDependentReduction;
200 const describeRuntimeDependentReduction = library.describeRuntimeDependentReduction;
201 const describeReductionReuse = library.describeReductionReuse;
202 const describeRuntimeReductionReuse = library.describeRuntimeReductionReuse;
203 const describeStaticParameter = library.describeStaticParameter;
204 const describeRuntimeStaticParameter = library.describeRuntimeStaticParameter;
205 const describeInputTransform = library.describeInputTransform;
206 const describeEpilogue = library.describeEpilogue;
207 const describeInputEpilogue = library.describeInputEpilogue;
208 const grid1D = library.grid1D;
209 const runtimeGrid1D = library.runtimeGrid1D;
210 const launch1D = library.launch1D;
211 const runtimeLaunch1D = library.runtimeLaunch1D;
212 const launch2D = library.launch2D;
213 const runtimeLaunch2D = library.runtimeLaunch2D;
214 const launch3D = library.launch3D;
215 const runtimeLaunch3D = library.runtimeLaunch3D;
216 const threadBlocks1D = library.threadBlocks1D;
217 const runtimeThreadBlocks1D = library.runtimeThreadBlocks1D;
218 const threadBlocks2D = library.threadBlocks2D;
219 const runtimeThreadBlocks2D = library.runtimeThreadBlocks2D;
220 const threadBlocks3D = library.threadBlocks3D;
221 const runtimeThreadBlocks3D = library.runtimeThreadBlocks3D;
222
223 test {
224 _ = @import("catalog/test.zig");
225 _ = @import("histogram/test.zig");
226 _ = @import("random/test.zig");
227 @import("test_discovery").discover(library);
228 }
229
230 test "kernel library family artifacts round-trip through the wire registry" {
231 const allocator = std.testing.allocator;
232 var state = gpu.recording.BackendState{
233 .allocator = allocator,
234 .kind = .cuda,
235 .format = .cuda_ptx,
236 };
237
238 const instance = linalg.MatrixProduct{ .m = 5, .n = 7, .k = 3, .threads = .{ .x = 4, .y = 2 } };
239 var family_artifact = try linalg.createMatrixProductFamilyArtifact(allocator, state.handle(), instance, .{ .limits = .testing });
240 defer family_artifact.deinit();
241 var fixed_artifact = try linalg.MatrixProduct2x3x4F32.createKernelCallArtifact(allocator, state.handle(), .{ .limits = linalg.MatrixProduct2x3x4F32.Limits.testing });
242 defer fixed_artifact.deinit();
243
244 const entries = [_]artifact.KernelCallArtifact{ family_artifact.entry(), fixed_artifact.entry() };
245 const encoded = try artifact.wire.encode(allocator, entries[0..], &.{});
246 defer allocator.free(encoded);
247
248 var decoded = try artifact.wire.decode(allocator, encoded);
249 defer decoded.deinit();
250
251 const family_entry = decoded.registry().find(
252 "accy.kernel.linalg.matmul_family_4x2_f32",
253 linalg.matrix_product_family_version,
254 .cuda_ptx,
255 ) orelse return error.TestExpectedFamilyArtifact;
256 try std.testing.expectEqualDeep(family_artifact.entry(), family_entry);
257
258 const fixed_entry = decoded.registry().find(
259 linalg.MatrixProduct2x3x4F32.target,
260 linalg.MatrixProduct2x3x4F32.version,
261 .cuda_ptx,
262 ) orelse return error.TestExpectedFixedArtifact;
263 try std.testing.expectEqualDeep(fixed_artifact.entry(), fixed_entry);
264 }
265
266 test "accy kernel library declaration coverage" {
267 std.testing.refAllDecls(library);
268 }