lib/accy/src/executable/fixture.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const gpu = @import("gpu");
  3 const choir = @import("choir");
  4 const accy_root = @import("../root.zig");
  5 const accy_choir = @import("../choir/root.zig");
  6 const artifact_product = @import("../artifact/root.zig");
  7 const preparation = @import("../preparation/root.zig");
  8 const binding_mod = @import("binding.zig");
  9 
 10 const ir = choir.ir;
 11 const passes = choir.passes;
 12 const semantic = accy_choir.semantic;
 13 
 14 pub const OwnedSemanticChoirModule = struct {
 15     module: *semantic.SemanticModule,
 16     choir_module: *ir.Operation,
 17     ctx: *ir.Context,
 18 
 19     pub fn init(module: *semantic.SemanticModule) OwnedSemanticChoirModule {
 20         return .{
 21             .module = module,
 22             .choir_module = module.choir_module,
 23             .ctx = module.context(),
 24         };
 25     }
 26 
 27     pub fn deinit(self: *OwnedSemanticChoirModule) void {
 28         self.module.deinit();
 29         self.* = undefined;
 30     }
 31 };
 32 
 33 pub fn addChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
 34     return OwnedSemanticChoirModule.init(try addSemanticModule(allocator, name));
 35 }
 36 
 37 pub fn addU32ChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
 38     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
 39     errdefer builder.deinit();
 40     const u32_8 = try builder.tensor(.u32, &.{8});
 41     var fb = try builder.beginFunction(name, &.{ u32_8, u32_8 }, &.{u32_8});
 42     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
 43     try fb.return_(&.{sum});
 44     try fb.finish();
 45     return OwnedSemanticChoirModule.init(try builder.finish());
 46 }
 47 
 48 pub fn addSemanticModule(allocator: std.mem.Allocator, name: []const u8) !*semantic.SemanticModule {
 49     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
 50     errdefer builder.deinit();
 51     const f32_8 = try builder.tensor(.f32, &.{8});
 52     var fb = try builder.beginFunction(name, &.{ f32_8, f32_8 }, &.{f32_8});
 53     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
 54     try fb.return_(&.{sum});
 55     try fb.finish();
 56     return try builder.finish();
 57 }
 58 
 59 pub fn fusedAddMulChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
 60     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
 61     errdefer builder.deinit();
 62     const f32_8 = try builder.tensor(.f32, &.{8});
 63     var fb = try builder.beginFunction(name, &.{ f32_8, f32_8, f32_8 }, &.{f32_8});
 64     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
 65     const product = try fb.mul(sum, fb.parameter(2));
 66     try fb.return_(&.{product});
 67     try fb.finish();
 68     return OwnedSemanticChoirModule.init(try builder.finish());
 69 }
 70 
 71 pub fn kernelCallChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
 72     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
 73     errdefer builder.deinit();
 74     const f32_8 = try builder.tensor(.f32, &.{8});
 75     var fb = try builder.beginFunction(name, &.{ f32_8, f32_8 }, &.{f32_8});
 76     const call = try fb.kernelCall(
 77         &.{ fb.parameter(0), fb.parameter(1) },
 78         &.{f32_8},
 79         .{
 80             .target = "accy.custom.scale",
 81             .operand_effects = &.{ .read, .write },
 82             .result_aliases = &.{null},
 83         },
 84     );
 85     try fb.return_(&.{call.getFirstResult()});
 86     try fb.finish();
 87     return OwnedSemanticChoirModule.init(try builder.finish());
 88 }
 89 
 90 pub fn aliasedKernelCallChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
 91     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
 92     errdefer builder.deinit();
 93     const f32_8 = try builder.tensor(.f32, &.{8});
 94     var fb = try builder.beginFunction(name, &.{f32_8}, &.{f32_8});
 95     const call = try fb.kernelCall(
 96         &.{fb.parameter(0)},
 97         &.{f32_8},
 98         .{
 99             .target = "accy.custom.update",
100             .operand_effects = &.{.read_write},
101             .result_aliases = &.{0},
102         },
103     );
104     try fb.return_(&.{call.getFirstResult()});
105     try fb.finish();
106     return OwnedSemanticChoirModule.init(try builder.finish());
107 }
108 
109 pub fn constantAddChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
110     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
111     errdefer builder.deinit();
112     const f32_8 = try builder.tensor(.f32, &.{8});
113     var fb = try builder.beginFunction(name, &.{f32_8}, &.{f32_8});
114     const values = @as([8]f32, @splat(2.0));
115     const constant = try fb.constant(f32_8, std.mem.sliceAsBytes(values[0..]));
116     const sum = try fb.add(fb.parameter(0), constant);
117     try fb.return_(&.{sum});
118     try fb.finish();
119     return OwnedSemanticChoirModule.init(try builder.finish());
120 }
121 
122 pub fn dotGeneralChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
123     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
124     errdefer builder.deinit();
125     const f32_16x16 = try builder.tensor(.f32, &.{ 16, 16 });
126     var fb = try builder.beginFunction(name, &.{ f32_16x16, f32_16x16 }, &.{f32_16x16});
127     const product = try fb.dotGeneral(
128         fb.parameter(0),
129         fb.parameter(1),
130         f32_16x16,
131         &.{1},
132         &.{0},
133         &.{},
134         &.{},
135     );
136     try fb.return_(&.{product});
137     try fb.finish();
138     return OwnedSemanticChoirModule.init(try builder.finish());
139 }
140 
141 pub fn dotGeneralF16ChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
142     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
143     errdefer builder.deinit();
144     const f16_16x16 = try builder.tensor(.f16, &.{ 16, 16 });
145     const f32_16x16 = try builder.tensor(.f32, &.{ 16, 16 });
146     var fb = try builder.beginFunction(name, &.{ f16_16x16, f16_16x16 }, &.{f32_16x16});
147     const product = try fb.dotGeneral(
148         fb.parameter(0),
149         fb.parameter(1),
150         f32_16x16,
151         &.{1},
152         &.{0},
153         &.{},
154         &.{},
155     );
156     try fb.return_(&.{product});
157     try fb.finish();
158     return OwnedSemanticChoirModule.init(try builder.finish());
159 }
160 
161 pub fn reduceChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
162     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
163     errdefer builder.deinit();
164     const f32_256 = try builder.tensor(.f32, &.{256});
165     const f32_scalar = try builder.tensor(.f32, &.{});
166     var fb = try builder.beginFunction(name, &.{f32_256}, &.{f32_scalar});
167     const zero = try fb.constant(f32_scalar, std.mem.asBytes(&@as(f32, 0.0)));
168     const reduced = try fb.reduce(fb.parameter(0), zero, f32_scalar, "sum", &.{0});
169     try fb.return_(&.{reduced});
170     try fb.finish();
171     return OwnedSemanticChoirModule.init(try builder.finish());
172 }
173 
174 pub fn reduceI32ChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
175     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
176     errdefer builder.deinit();
177     const i32_256 = try builder.tensor(.i32, &.{256});
178     const i32_scalar = try builder.tensor(.i32, &.{});
179     var fb = try builder.beginFunction(name, &.{i32_256}, &.{i32_scalar});
180     const zero = try fb.constant(i32_scalar, std.mem.asBytes(&@as(i32, 0)));
181     const reduced = try fb.reduce(fb.parameter(0), zero, i32_scalar, "sum", &.{0});
182     try fb.return_(&.{reduced});
183     try fb.finish();
184     return OwnedSemanticChoirModule.init(try builder.finish());
185 }
186 
187 pub fn escapedTwoKernelChoirModule(allocator: std.mem.Allocator, name: []const u8) !OwnedSemanticChoirModule {
188     var builder = try semantic.Builder.init(allocator, semantic.Builder.ContextLimits.standard);
189     errdefer builder.deinit();
190     const f32_8 = try builder.tensor(.f32, &.{8});
191     var fb = try builder.beginFunction(name, &.{ f32_8, f32_8, f32_8 }, &.{ f32_8, f32_8 });
192     const sum = try fb.add(fb.parameter(0), fb.parameter(1));
193     const product = try fb.mul(sum, fb.parameter(2));
194     try fb.return_(&.{ sum, product });
195     try fb.finish();
196     return OwnedSemanticChoirModule.init(try builder.finish());
197 }
198 
199 pub fn createTestBackendArtifactPlan(
200     allocator: std.mem.Allocator,
201     handle: gpu.BackendHandle,
202     pass_ctx: *passes.PassContext,
203     choir_module: *ir.Operation,
204     options: artifact_product.ArtifactPlanOptions,
205 ) !artifact_product.BackendArtifactPlan {
206     const lowered_kernels = try preparation.kernelization.getKernelizationAnalysis(pass_ctx, choir_module);
207     return try artifact_product.createBackendArtifactPlan(
208         allocator,
209         handle,
210         .{
211             .pass_ctx = pass_ctx,
212             .choir_module = choir_module,
213             .lowered_kernels = lowered_kernels,
214         },
215         options,
216     );
217 }
218 
219 pub fn bufferBinding(
220     id: gpu.BackendObjectId,
221     kind: gpu.BackendKind,
222     bytes: usize,
223 ) gpu.BufferBinding {
224     return .{
225         .handle = .{
226             .id = id,
227             .backend = kind,
228             .byte_size = bytes,
229             .ownership = .backend,
230         },
231         .access = .read_write,
232         .ownership = .backend,
233         .byte_size = bytes,
234     };
235 }
236 
237 pub fn slotBindingsForKernel(
238     allocator: std.mem.Allocator,
239     kernel: artifact_product.PlannedKernel,
240     kind: gpu.BackendKind,
241 ) ![]binding_mod.SlotBinding {
242     const bindings = try allocator.alloc(binding_mod.SlotBinding, 1 + kernel.input_slot_ids.len);
243     bindings[0] = .{
244         .slot_id = kernel.output_slot_id,
245         .binding = bufferBinding(100, kind, 32),
246     };
247     for (kernel.input_slot_ids, 0..) |slot_id, index| {
248         bindings[1 + index] = .{
249             .slot_id = slot_id,
250             .binding = bufferBinding(@intCast(101 + index), kind, 32),
251         };
252     }
253     return bindings;
254 }
255 
256 pub fn slotBindingsForPlan(
257     allocator: std.mem.Allocator,
258     artifact_plan: *const artifact_product.BackendArtifactPlan,
259     kind: gpu.BackendKind,
260 ) ![]binding_mod.SlotBinding {
261     const bindings = try allocator.alloc(binding_mod.SlotBinding, artifact_plan.slots.len);
262     for (artifact_plan.slots, 0..) |slot, index| {
263         const bytes = std.math.cast(usize, slot.byte_size orelse 32) orelse return error.InvalidArtifact;
264         bindings[index] = .{
265             .slot_id = slot.slot_id,
266             .binding = bufferBinding(@intCast(100 + index), kind, bytes),
267         };
268     }
269     return bindings;
270 }
271 
272 pub fn elementCountBindingsForPlan(
273     allocator: std.mem.Allocator,
274     artifact_plan: *const artifact_product.BackendArtifactPlan,
275     kind: gpu.BackendKind,
276 ) ![]binding_mod.ElementCountBufferBinding {
277     var bindings = std.ArrayListUnmanaged(binding_mod.ElementCountBufferBinding).empty;
278     errdefer bindings.deinit(allocator);
279     for (artifact_plan.kernels.items) |kernel| {
280         if (kernel.element_count_argument != .device_buffer_u32) continue;
281         try bindings.append(allocator, .{
282             .kernel_id = kernel.kernel_id,
283             .binding = bufferBinding(@intCast(900 + kernel.kernel_id), kind, 4),
284         });
285     }
286     return try bindings.toOwnedSlice(allocator);
287 }
288 
289 pub fn firstElementCountBinding(bindings: []const binding_mod.ElementCountBufferBinding) ?gpu.BufferBinding {
290     if (bindings.len == 0) return null;
291     return bindings[0].binding;
292 }