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 }