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 }