lib/accy/src/artifact/model/pipeline.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const gpu = @import("gpu");
3 const choir_abi = @import("choir_abi");
4
5 const plan = @import("registry.zig");
6
7 const DType = choir_abi.DType;
8
9 pub const product_name = "accy.kernel_call_pipeline";
10
11 pub const PipelineValueRef = union(enum) {
12 operand: u32,
13 result: u32,
14 intermediate: u32,
15 };
16
17 pub const PipelineScalarDerivation = union(enum) {
18 forward: u32,
19 ceil_div: plan.KernelCallDerivedLaunchAxis.RuntimeU32CeilDiv,
20 ceil_div_scaled: CeilDivScaled,
21 ceil_div_scaled_by_arg: CeilDivScaledByArg,
22
23 pub const CeilDivScaled = struct {
24 argument_index: u32,
25 divisor: u32 = 1,
26 scale: u32 = 1,
27 };
28
29 pub const CeilDivScaledByArg = struct {
30 argument_index: u32,
31 divisor: u32 = 1,
32 scale_argument_index: u32,
33 };
34
35 pub fn resolveScalar(
36 self: PipelineScalarDerivation,
37 runtime_scalar_arguments: []const choir_abi.ScalarArgument,
38 ) gpu.BackendError!choir_abi.ScalarArgument {
39 return switch (self) {
40 .forward => |index| if (index < runtime_scalar_arguments.len)
41 runtime_scalar_arguments[index]
42 else
43 error.LaunchArgumentMismatch,
44 .ceil_div => |axis| .{
45 .u32 = try (plan.KernelCallDerivedLaunchAxis{ .runtime_u32_ceil_div = axis }).extent(runtime_scalar_arguments),
46 },
47 .ceil_div_scaled => |scaled| blk: {
48 const base = try (plan.KernelCallDerivedLaunchAxis{ .runtime_u32_ceil_div = .{
49 .argument_index = scaled.argument_index,
50 .divisor = scaled.divisor,
51 } }).extent(runtime_scalar_arguments);
52 const value = std.math.mul(u32, base, scaled.scale) catch return error.LaunchArgumentMismatch;
53 break :blk .{ .u32 = value };
54 },
55 .ceil_div_scaled_by_arg => |scaled| blk: {
56 const base = try (plan.KernelCallDerivedLaunchAxis{ .runtime_u32_ceil_div = .{
57 .argument_index = scaled.argument_index,
58 .divisor = scaled.divisor,
59 } }).extent(runtime_scalar_arguments);
60 if (scaled.scale_argument_index >= runtime_scalar_arguments.len) return error.LaunchArgumentMismatch;
61 const scale = switch (runtime_scalar_arguments[scaled.scale_argument_index]) {
62 .u32 => |value| value,
63 else => return error.LaunchArgumentMismatch,
64 };
65 const value = std.math.mul(u32, base, scale) catch return error.LaunchArgumentMismatch;
66 break :blk .{ .u32 = value };
67 },
68 };
69 }
70
71 pub fn resolveExtent(
72 self: PipelineScalarDerivation,
73 runtime_scalar_arguments: []const choir_abi.ScalarArgument,
74 ) gpu.BackendError!u32 {
75 return switch (try self.resolveScalar(runtime_scalar_arguments)) {
76 .u32 => |value| value,
77 else => error.LaunchArgumentMismatch,
78 };
79 }
80 };
81
82 pub const PipelineIntermediate = struct {
83 dtype: DType,
84 extent: PipelineScalarDerivation,
85 };
86
87 pub const PipelineRuntimeScalarBound = struct {
88 argument_index: u32,
89 max_u32: u32,
90 };
91
92 pub const PipelineStage = struct {
93 target: []const u8,
94 version: u32,
95 buffers: []const PipelineValueRef = &.{},
96 scalars: []const PipelineScalarDerivation = &.{},
97 };
98
99 /// An ordered list of stages under one target name and version, so a caller
100 /// describes a multi-kernel computation and hands it to the loader and
101 /// launcher. Each stage names a registry entry, a prebuilt kernel, by target
102 /// and version and so holds no copy of its code. Each stage lists its buffers
103 /// as caller operands, caller results or scratch buffers allocated between
104 /// stages, and derives its integer arguments from the runtime scalar arguments
105 /// by passing one through, dividing and rounding up, or scaling first. The same
106 /// runtime scalar arguments size the intermediates and the launch grids, and
107 /// each runtime scalar argument may carry an upper bound. `validate` checks a
108 /// nonempty target, a nonzero version and at least one stage, then finds every
109 /// stage's entry for the format and checks each stage's buffer count, argument
110 /// count, references and divisors before anything is launched, and every
111 /// failure is `error.InvalidArtifact`. `validateRuntimeScalarArguments` checks
112 /// the count and the declared bounds of the runtime scalar arguments and
113 /// returns `error.LaunchArgumentMismatch` on a mismatch.
114 pub const KernelCallPipeline = struct {
115 target: []const u8,
116 version: u32,
117 operand_count: u32 = 0,
118 result_count: u32 = 0,
119 runtime_scalar_argument_count: u32 = 0,
120 runtime_scalar_bounds: []const PipelineRuntimeScalarBound = &.{},
121 intermediates: []const PipelineIntermediate = &.{},
122 stages: []const PipelineStage = &.{},
123
124 pub fn validate(
125 self: KernelCallPipeline,
126 registry: plan.KernelCallRegistry,
127 format: gpu.ArtifactFormat,
128 ) gpu.BackendError!void {
129 if (self.target.len == 0) return error.InvalidArtifact;
130 if (self.version == 0) return error.InvalidArtifact;
131 if (self.stages.len == 0) return error.InvalidArtifact;
132 for (self.runtime_scalar_bounds) |bound| {
133 if (bound.argument_index >= self.runtime_scalar_argument_count) return error.InvalidArtifact;
134 }
135 for (self.intermediates) |intermediate| {
136 try self.validateDerivation(intermediate.extent);
137 }
138 for (self.stages) |stage| {
139 const entry = registry.find(stage.target, stage.version, format) orelse return error.InvalidArtifact;
140 if (entry.argument_count < entry.runtime_scalar_argument_count) return error.InvalidArtifact;
141 const buffer_count = entry.argument_count - entry.runtime_scalar_argument_count;
142 if (stage.buffers.len != buffer_count) return error.InvalidArtifact;
143 if (stage.scalars.len != entry.runtime_scalar_argument_count) return error.InvalidArtifact;
144 for (stage.buffers) |ref| try self.validateValueRef(ref);
145 for (stage.scalars) |derivation| try self.validateDerivation(derivation);
146 }
147 }
148
149 pub fn validateRuntimeScalarArguments(
150 self: KernelCallPipeline,
151 runtime_scalar_arguments: []const choir_abi.ScalarArgument,
152 ) gpu.BackendError!void {
153 if (runtime_scalar_arguments.len != self.runtime_scalar_argument_count) return error.LaunchArgumentMismatch;
154 for (self.runtime_scalar_bounds) |bound| {
155 if (bound.argument_index >= runtime_scalar_arguments.len) return error.LaunchArgumentMismatch;
156 const value = switch (runtime_scalar_arguments[bound.argument_index]) {
157 .u32 => |value| value,
158 else => return error.LaunchArgumentMismatch,
159 };
160 if (value > bound.max_u32) return error.LaunchArgumentMismatch;
161 }
162 }
163
164 fn validateValueRef(self: KernelCallPipeline, ref: PipelineValueRef) gpu.BackendError!void {
165 switch (ref) {
166 .operand => |index| if (index >= self.operand_count) return error.InvalidArtifact,
167 .result => |index| if (index >= self.result_count) return error.InvalidArtifact,
168 .intermediate => |index| if (index >= self.intermediates.len) return error.InvalidArtifact,
169 }
170 }
171
172 fn validateDerivation(self: KernelCallPipeline, derivation: PipelineScalarDerivation) gpu.BackendError!void {
173 switch (derivation) {
174 .forward => |index| if (index >= self.runtime_scalar_argument_count) return error.InvalidArtifact,
175 .ceil_div => |axis| {
176 if (axis.divisor == 0) return error.InvalidArtifact;
177 if (axis.argument_index >= self.runtime_scalar_argument_count) return error.InvalidArtifact;
178 },
179 .ceil_div_scaled => |scaled| {
180 if (scaled.divisor == 0 or scaled.scale == 0) return error.InvalidArtifact;
181 if (scaled.argument_index >= self.runtime_scalar_argument_count) return error.InvalidArtifact;
182 },
183 .ceil_div_scaled_by_arg => |scaled| {
184 if (scaled.divisor == 0) return error.InvalidArtifact;
185 if (scaled.argument_index >= self.runtime_scalar_argument_count) return error.InvalidArtifact;
186 if (scaled.scale_argument_index >= self.runtime_scalar_argument_count) return error.InvalidArtifact;
187 },
188 }
189 }
190 };
191
192 pub fn findPipeline(
193 pipelines: []const KernelCallPipeline,
194 target: []const u8,
195 version: u32,
196 ) ?KernelCallPipeline {
197 for (pipelines) |entry| {
198 if (entry.version != version) continue;
199 if (!std.mem.eql(u8, entry.target, target)) continue;
200 return entry;
201 }
202 return null;
203 }
204
205 pub const OwnedKernelCallPipeline = struct {
206 arena: std.heap.ArenaAllocator,
207 value: KernelCallPipeline = .{ .target = &.{}, .version = 0 },
208
209 pub fn init(backing_allocator: std.mem.Allocator) OwnedKernelCallPipeline {
210 return .{ .arena = std.heap.ArenaAllocator.init(backing_allocator) };
211 }
212
213 pub fn allocator(self: *OwnedKernelCallPipeline) std.mem.Allocator {
214 return self.arena.allocator();
215 }
216
217 pub fn deinit(self: *OwnedKernelCallPipeline) void {
218 self.arena.deinit();
219 self.* = undefined;
220 }
221 };
222
223 const testing = std.testing;
224
225 fn testEntry(comptime target: []const u8, argument_count: u32, runtime_scalars: u32) plan.KernelCallArtifact {
226 return .{
227 .target = target,
228 .version = 1,
229 .format = .cuda_ptx,
230 .entry_name = target,
231 .argument_count = argument_count,
232 .payload = .{ .text = "// " ++ target },
233 .runtime_scalar_argument_count = runtime_scalars,
234 };
235 }
236
237 const device_scan_entries = [_]plan.KernelCallArtifact{
238 testEntry("accy.kernel.scan.device_prefix_sum_block_scan_family_64_f32", 4, 1),
239 testEntry("accy.kernel.scan.prefix_sum_exclusive_family_32_f32", 3, 1),
240 testEntry("accy.kernel.scan.device_add_base_family_64_f32", 3, 1),
241 };
242
243 const device_scan_pipeline = KernelCallPipeline{
244 .target = "accy.kernel.scan.device_prefix_sum_family_64_f32",
245 .version = 1,
246 .operand_count = 1,
247 .result_count = 1,
248 .runtime_scalar_argument_count = 1,
249 .intermediates = &.{
250 .{ .dtype = .f32, .extent = .{ .ceil_div = .{ .argument_index = 0, .divisor = 64 } } },
251 .{ .dtype = .f32, .extent = .{ .ceil_div = .{ .argument_index = 0, .divisor = 64 } } },
252 },
253 .stages = &.{
254 .{
255 .target = "accy.kernel.scan.device_prefix_sum_block_scan_family_64_f32",
256 .version = 1,
257 .buffers = &.{ .{ .result = 0 }, .{ .operand = 0 }, .{ .intermediate = 0 } },
258 .scalars = &.{.{ .forward = 0 }},
259 },
260 .{
261 .target = "accy.kernel.scan.prefix_sum_exclusive_family_32_f32",
262 .version = 1,
263 .buffers = &.{ .{ .intermediate = 1 }, .{ .intermediate = 0 } },
264 .scalars = &.{.{ .ceil_div = .{ .argument_index = 0, .divisor = 64 } }},
265 },
266 .{
267 .target = "accy.kernel.scan.device_add_base_family_64_f32",
268 .version = 1,
269 .buffers = &.{ .{ .result = 0 }, .{ .intermediate = 1 } },
270 .scalars = &.{.{ .forward = 0 }},
271 },
272 },
273 };
274
275 test "pipeline validates the device scan shape against its stage registry" {
276 const registry = plan.KernelCallRegistry{ .entries = device_scan_entries[0..] };
277 try device_scan_pipeline.validate(registry, .cuda_ptx);
278 }
279
280 test "pipeline validation names every violated contract" {
281 const registry = plan.KernelCallRegistry{ .entries = device_scan_entries[0..] };
282
283 var missing_stage = device_scan_pipeline;
284 missing_stage.stages = &.{.{
285 .target = "accy.kernel.scan.absent",
286 .version = 1,
287 .buffers = &.{ .{ .result = 0 }, .{ .operand = 0 }, .{ .intermediate = 0 } },
288 .scalars = &.{.{ .forward = 0 }},
289 }};
290 try testing.expectError(error.InvalidArtifact, missing_stage.validate(registry, .cuda_ptx));
291
292 var wrong_format = device_scan_pipeline;
293 try testing.expectError(error.InvalidArtifact, wrong_format.validate(registry, .vulkan_spirv));
294
295 var bad_version = device_scan_pipeline;
296 bad_version.version = 0;
297 try testing.expectError(error.InvalidArtifact, bad_version.validate(registry, .cuda_ptx));
298
299 var no_stages = device_scan_pipeline;
300 no_stages.stages = &.{};
301 try testing.expectError(error.InvalidArtifact, no_stages.validate(registry, .cuda_ptx));
302
303 var bad_operand = device_scan_pipeline;
304 bad_operand.operand_count = 0;
305 try testing.expectError(error.InvalidArtifact, bad_operand.validate(registry, .cuda_ptx));
306
307 var bad_intermediate = device_scan_pipeline;
308 bad_intermediate.intermediates = device_scan_pipeline.intermediates[0..1];
309 try testing.expectError(error.InvalidArtifact, bad_intermediate.validate(registry, .cuda_ptx));
310
311 var bad_scalar = device_scan_pipeline;
312 bad_scalar.runtime_scalar_argument_count = 0;
313 try testing.expectError(error.InvalidArtifact, bad_scalar.validate(registry, .cuda_ptx));
314
315 var bad_runtime_bound = device_scan_pipeline;
316 bad_runtime_bound.runtime_scalar_bounds = &.{.{ .argument_index = 1, .max_u32 = 5000 }};
317 try testing.expectError(error.InvalidArtifact, bad_runtime_bound.validate(registry, .cuda_ptx));
318
319 const short_buffers = [_]PipelineStage{.{
320 .target = "accy.kernel.scan.device_add_base_family_64_f32",
321 .version = 1,
322 .buffers = &.{.{ .result = 0 }},
323 .scalars = &.{.{ .forward = 0 }},
324 }};
325 var wrong_arity = device_scan_pipeline;
326 wrong_arity.stages = short_buffers[0..];
327 try testing.expectError(error.InvalidArtifact, wrong_arity.validate(registry, .cuda_ptx));
328 }
329
330 test "pipeline runtime scalar bounds reject out-of-capacity launches" {
331 const registry = plan.KernelCallRegistry{ .entries = device_scan_entries[0..] };
332 var bounded = device_scan_pipeline;
333 bounded.runtime_scalar_bounds = &.{.{ .argument_index = 0, .max_u32 = 5000 }};
334 try bounded.validate(registry, .cuda_ptx);
335
336 const in_range = [_]choir_abi.ScalarArgument{.{ .u32 = 5000 }};
337 try bounded.validateRuntimeScalarArguments(in_range[0..]);
338
339 const too_large = [_]choir_abi.ScalarArgument{.{ .u32 = 5001 }};
340 try testing.expectError(error.LaunchArgumentMismatch, bounded.validateRuntimeScalarArguments(too_large[0..]));
341
342 const wrong_type = [_]choir_abi.ScalarArgument{.{ .f32 = 5000.0 }};
343 try testing.expectError(error.LaunchArgumentMismatch, bounded.validateRuntimeScalarArguments(wrong_type[0..]));
344
345 try testing.expectError(error.LaunchArgumentMismatch, bounded.validateRuntimeScalarArguments(&.{}));
346 }
347
348 test "pipeline scalar derivations resolve forward and ceil-div values" {
349 const args = [_]choir_abi.ScalarArgument{.{ .u32 = 5000 }};
350
351 const forward = PipelineScalarDerivation{ .forward = 0 };
352 try testing.expectEqual(@as(u32, 5000), try forward.resolveExtent(args[0..]));
353
354 const blocks = PipelineScalarDerivation{ .ceil_div = .{ .argument_index = 0, .divisor = 64 } };
355 try testing.expectEqual(@as(u32, 79), try blocks.resolveExtent(args[0..]));
356
357 const scalar = try blocks.resolveScalar(args[0..]);
358 try testing.expectEqual(@as(u32, 79), scalar.u32);
359
360 const float_args = [_]choir_abi.ScalarArgument{.{ .f32 = 1.0 }};
361 try testing.expectError(error.LaunchArgumentMismatch, forward.resolveExtent(float_args[0..]));
362 try testing.expectError(error.LaunchArgumentMismatch, blocks.resolveExtent(float_args[0..]));
363
364 const out_of_range = PipelineScalarDerivation{ .forward = 3 };
365 try testing.expectError(error.LaunchArgumentMismatch, out_of_range.resolveScalar(args[0..]));
366 }
367
368 test "pipeline lookup finds by target and version" {
369 const pipelines = [_]KernelCallPipeline{device_scan_pipeline};
370 try testing.expect(findPipeline(pipelines[0..], device_scan_pipeline.target, 1) != null);
371 try testing.expect(findPipeline(pipelines[0..], device_scan_pipeline.target, 2) == null);
372 try testing.expect(findPipeline(pipelines[0..], "missing", 1) == null);
373 }
374
375 test "pipeline scaled ceil-div derivations resolve and validate" {
376 const args = [_]choir_abi.ScalarArgument{.{ .u32 = 5000 }};
377 const counts_extent = PipelineScalarDerivation{
378 .ceil_div_scaled = .{ .argument_index = 0, .divisor = 64, .scale = 16 },
379 };
380 try testing.expectEqual(@as(u32, 79 * 16), try counts_extent.resolveExtent(args[0..]));
381
382 const identity_scale = PipelineScalarDerivation{
383 .ceil_div_scaled = .{ .argument_index = 0, .divisor = 64, .scale = 1 },
384 };
385 try testing.expectEqual(@as(u32, 79), try identity_scale.resolveExtent(args[0..]));
386
387 const overflow = PipelineScalarDerivation{
388 .ceil_div_scaled = .{ .argument_index = 0, .divisor = 1, .scale = std.math.maxInt(u32) },
389 };
390 try testing.expectError(error.LaunchArgumentMismatch, overflow.resolveExtent(args[0..]));
391
392 var scaled_pipeline = device_scan_pipeline;
393 const scaled_intermediates = [_]PipelineIntermediate{
394 .{ .dtype = .f32, .extent = .{ .ceil_div_scaled = .{ .argument_index = 0, .divisor = 64, .scale = 0 } } },
395 };
396 scaled_pipeline.intermediates = scaled_intermediates[0..];
397 const registry = plan.KernelCallRegistry{ .entries = device_scan_entries[0..] };
398 try testing.expectError(error.InvalidArtifact, scaled_pipeline.validate(registry, .cuda_ptx));
399
400 var bad_index = device_scan_pipeline;
401 const bad_intermediates = [_]PipelineIntermediate{
402 .{ .dtype = .f32, .extent = .{ .ceil_div_scaled = .{ .argument_index = 7, .divisor = 64, .scale = 16 } } },
403 };
404 bad_index.intermediates = bad_intermediates[0..];
405 try testing.expectError(error.InvalidArtifact, bad_index.validate(registry, .cuda_ptx));
406 }
407
408 test "pipeline runtime-scaled derivations resolve and validate" {
409 const args = [_]choir_abi.ScalarArgument{ .{ .u32 = 5000 }, .{ .u32 = 64 } };
410 const cells_buffer = PipelineScalarDerivation{
411 .ceil_div_scaled_by_arg = .{ .argument_index = 0, .divisor = 64, .scale_argument_index = 1 },
412 };
413 try testing.expectEqual(@as(u32, 79 * 64), try cells_buffer.resolveExtent(args[0..]));
414
415 const bad_index = PipelineScalarDerivation{
416 .ceil_div_scaled_by_arg = .{ .argument_index = 0, .divisor = 64, .scale_argument_index = 7 },
417 };
418 try testing.expectError(error.LaunchArgumentMismatch, bad_index.resolveExtent(args[0..]));
419
420 const float_scale_args = [_]choir_abi.ScalarArgument{ .{ .u32 = 5000 }, .{ .f32 = 2.0 } };
421 try testing.expectError(error.LaunchArgumentMismatch, cells_buffer.resolveExtent(float_scale_args[0..]));
422
423 var scaled_pipeline = device_scan_pipeline;
424 const bad_intermediates = [_]PipelineIntermediate{
425 .{ .dtype = .f32, .extent = .{ .ceil_div_scaled_by_arg = .{ .argument_index = 0, .divisor = 0, .scale_argument_index = 0 } } },
426 };
427 scaled_pipeline.intermediates = bad_intermediates[0..];
428 const registry = plan.KernelCallRegistry{ .entries = device_scan_entries[0..] };
429 try testing.expectError(error.InvalidArtifact, scaled_pipeline.validate(registry, .cuda_ptx));
430 }