lib/accy/src/kernel/model/plan/model.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const testing = std.testing;
3
4 const core = @import("../core/root.zig");
5
6 const builder = core.builder;
7 const schedule_mod = core.schedule;
8
9 pub const Error = schedule_mod.ScheduleError ||
10 builder.BackendError ||
11 builder.FingerprintError ||
12 std.mem.Allocator.Error ||
13 error{
14 ArgumentCountOverflow,
15 MissingKernelName,
16 ParameterSchemaMismatch,
17 };
18
19 pub const Options = struct {
20 diagnostic_id: ?[]const u8 = null,
21 diagnostic: ?*builder.VerificationDiagnostic = null,
22 };
23
24 pub const plan_version: u32 = 1;
25
26 pub const Plan = struct {
27 allocator: std.mem.Allocator,
28 version: u32,
29 schedule: schedule_mod.Snapshot,
30 schedule_version: u32,
31 launch: schedule_mod.Launch,
32 body_fingerprint: u64,
33 schedule_fingerprint: u64,
34 entry_name: []const u8,
35 diagnostic_id: []const u8,
36 argument_count: u32,
37 params: []builder.Param,
38
39 pub fn deinit(self: *Plan) void {
40 if (self.params.len != 0) self.allocator.free(self.params);
41 self.allocator.free(self.diagnostic_id);
42 self.allocator.free(self.entry_name);
43 self.schedule.deinit(self.allocator);
44 self.* = undefined;
45 }
46 };
47
48 pub fn create(
49 allocator: std.mem.Allocator,
50 kernel: *builder.Kernel,
51 schedule: *const schedule_mod.Schedule,
52 snapshot_limits: schedule_mod.Snapshot.Limits,
53 options: Options,
54 ) Error!Plan {
55 var local_diagnostic: builder.VerificationDiagnostic = .{};
56 try kernel.verifyWithDiagnostic(options.diagnostic orelse &local_diagnostic);
57
58 const launch = try schedule.launch();
59 var snapshot = try schedule_mod.createSnapshot(allocator, snapshot_limits, schedule);
60 errdefer snapshot.deinit(allocator);
61 const body_fingerprint = try kernel.fingerprint(allocator);
62 const schedule_fingerprint = snapshot.fingerprint();
63
64 const entry_name = kernel.func().getName() orelse return error.MissingKernelName;
65 const raw_argument_count = kernel.func().getNumArguments();
66 if (raw_argument_count != kernel.params().len) return error.ParameterSchemaMismatch;
67 const argument_count = std.math.cast(u32, raw_argument_count) orelse return error.ArgumentCountOverflow;
68
69 const owned_entry_name = try allocator.dupe(u8, entry_name);
70 errdefer allocator.free(owned_entry_name);
71 const owned_diagnostic_id = try allocator.dupe(u8, options.diagnostic_id orelse entry_name);
72 errdefer allocator.free(owned_diagnostic_id);
73 const owned_params = try allocator.dupe(builder.Param, kernel.params());
74 errdefer allocator.free(owned_params);
75
76 return .{
77 .allocator = allocator,
78 .version = plan_version,
79 .schedule = snapshot,
80 .schedule_version = snapshot.version,
81 .launch = launch,
82 .body_fingerprint = body_fingerprint,
83 .schedule_fingerprint = schedule_fingerprint,
84 .entry_name = owned_entry_name,
85 .diagnostic_id = owned_diagnostic_id,
86 .argument_count = argument_count,
87 .params = owned_params,
88 };
89 }
90
91 test "Plan captures authored kernel schedule without backend artifact payload" {
92 var schedule = try schedule_mod.Schedule.init(testing.allocator, schedule_mod.Schedule.Limits.testing);
93 defer schedule.deinit(testing.allocator);
94
95 const i = try schedule.addAxis("i", 4);
96 try schedule.bind(i, .thread_x);
97
98 var b = try builder.Builder.init(testing.allocator, builder.Builder.Limits.testing, "copy_i32_plan", &.{
99 builder.dynamicBuffer(.i32),
100 builder.dynamicBuffer(.i32),
101 }, .{});
102 errdefer b.deinit();
103
104 const src = b.argument(0);
105 const dst = b.argument(1);
106 const index = try b.globalId(.x);
107 const value = try b.load(src, index);
108 try b.store(value, dst, index);
109 try b.return_();
110
111 var kernel = try b.finish();
112 defer kernel.deinit();
113
114 var out = try create(
115 testing.allocator,
116 &kernel,
117 &schedule,
118 schedule_mod.Snapshot.Limits.testing,
119 .{},
120 );
121 defer out.deinit();
122
123 try testing.expectEqual(plan_version, out.version);
124 try testing.expectEqual(schedule_mod.snapshot_version, out.schedule_version);
125 try testing.expectEqualStrings("copy_i32_plan", out.entry_name);
126 try testing.expectEqualStrings("copy_i32_plan", out.diagnostic_id);
127 try testing.expectEqual(@as(u32, 2), out.argument_count);
128 try testing.expectEqual(@as(usize, 2), out.params.len);
129 try testing.expect((builder.Param{ .buffer = .{ .dtype = .i32 } }).eql(out.params[0]));
130 try testing.expect((builder.Param{ .buffer = .{ .dtype = .i32 } }).eql(out.params[1]));
131
132 try testing.expectEqual(@as(u32, 1), out.launch.grid[0]);
133 try testing.expectEqual(@as(u32, 4), out.launch.block[0]);
134 try testing.expectEqual(@as(usize, 1), out.schedule.allAxes().len);
135 try testing.expectEqual(try kernel.fingerprint(testing.allocator), out.body_fingerprint);
136 try testing.expectEqual(schedule.fingerprint(), out.schedule_fingerprint);
137 }