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 }