lib/choir/src/backends/gpu/spirv/emitter/codegen.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const abi = @import("choir_abi");
3 const sys = @import("sys");
4 const choir = @import("../../../../root.zig");
5
6 const ir = choir.ir;
7 const dialects = choir.dialects;
8 const spirv_target = @import("../root.zig");
9 const calls = @import("../../calls.zig");
10 const gpu_target = @import("../../../../dialects/gpu/root.zig");
11 const binary = @import("module.zig");
12 const catalog = @import("catalog.zig");
13 const dialect_writer = @import("dialect.zig");
14 const gpu = @import("gpu.zig");
15 const memory = @import("memory.zig");
16 const scalar = @import("scalar.zig");
17 const spec = @import("spec.zig");
18 const stage_emit = @import("stage.zig");
19 const validation = @import("validation.zig");
20 const spirv_ops = @import("ops.zig");
21
22 const GpuDialect = gpu_target.GpuDialect;
23 const Stage = gpu_target.Stage;
24 const ArithDialect = dialects.arith.ArithDialect;
25 const arith = dialects.arith;
26 const MemrefDialect = dialects.memref.MemrefDialect;
27 const FuncDialect = dialects.func.FuncDialect;
28 const ScfDialect = dialects.scf.ScfDialect;
29 const SpirvDialect = spirv_target.SpirvDialect;
30 const BuiltinDialect = dialects.builtin.BuiltinDialect;
31 pub const SpirvOp = spirv_ops.SpirvOp;
32 const GLSLstd450 = scalar.GLSLstd450;
33 const ScalarKind = scalar.Kind;
34 const SpvAddressingModel = spec.AddressingModel;
35 const SpvCapability = spec.Capability;
36 const SpvDecoration = spec.Decoration;
37 const SpvExecutionMode = spec.ExecutionMode;
38 const SpvExecutionModel = spec.ExecutionModel;
39 const SpvMemoryModel = spec.MemoryModel;
40 const SpvStorageClass = spec.StorageClass;
41 const ModuleBuilder = binary.Builder;
42 const Section = binary.Section;
43 const SpirvHeader = binary.Header;
44 const SpirvVersion = binary.Version;
45
46 pub const SpirvCodegenError = error{
47 OutOfMemory,
48 InvalidModule,
49 MissingFunctionName,
50 MissingKernelAttribute,
51 UnsupportedFunctionSignature,
52 UnsupportedOperation,
53 UnsupportedType,
54 UnsupportedMask,
55 UnsupportedAddressSpace,
56 UnsupportedControlFlow,
57 MissingValue,
58 MissingAttribute,
59 InvalidMemrefType,
60 } || gpu_target.stage.MemoryError || calls.SignatureError;
61
62 pub const FloatWidths = struct {
63 f16: bool = false,
64 f32: bool = false,
65 f64: bool = false,
66 };
67
68 /// Modes selected from the properties of the device that will run the module.
69 pub const FloatControls = struct {
70 denorm_preserve: FloatWidths = .{},
71 signed_zero_inf_nan_preserve: FloatWidths = .{},
72 };
73
74 const VectorKey = struct {
75 elem: ScalarKind,
76 len: u32,
77 };
78
79 const PointerKey = struct {
80 storage_class: u32,
81 base_type: u32,
82 };
83
84 const RuntimeArrayKey = struct {
85 elem_type: u32,
86 stride: u32,
87 };
88
89 const ArrayKey = struct {
90 elem_type: u32,
91 length: u32,
92 };
93
94 const ConstKey = struct {
95 type_id: u32,
96 word0: u32,
97 word1: u32,
98 word_count: u32,
99 };
100
101 const ScalarArgument = struct {
102 arg: *ir.Value,
103 elem_type: u32,
104 ptr_elem_type: u32,
105 member_index: u32,
106 byte_offset: u32,
107 };
108
109 const MemrefBinding = memory.Binding;
110 pub const SpirvCodegen = struct {
111 allocator: std.mem.Allocator,
112 builder: ModuleBuilder,
113 value_ids: std.AutoHashMapUnmanaged(*ir.Value, u32),
114 memref_bindings: std.AutoHashMapUnmanaged(*ir.Value, MemrefBinding),
115 scalar_types: std.AutoHashMapUnmanaged(ScalarKind, u32),
116 vector_types: std.AutoHashMapUnmanaged(VectorKey, u32),
117 pointer_types: std.AutoHashMapUnmanaged(PointerKey, u32),
118 array_types: std.AutoHashMapUnmanaged(ArrayKey, u32),
119 runtime_array_types: std.AutoHashMapUnmanaged(RuntimeArrayKey, u32),
120 struct_types: std.AutoHashMapUnmanaged(u32, u32),
121 pair_struct_types: std.AutoHashMapUnmanaged(u32, u32),
122 constants: std.AutoHashMapUnmanaged(ConstKey, u32),
123 builtin_vars: std.AutoHashMapUnmanaged(gpu.BuiltinKind, u32),
124 interface_vars: std.ArrayListUnmanaged(u32),
125 interface_set: std.AutoHashMapUnmanaged(u32, void),
126 capability_set: std.AutoHashMapUnmanaged(u32, void),
127 extension_set: std.StringHashMapUnmanaged(void),
128 void_function_type: ?u32,
129 workgroup_size_constant: ?u32,
130 binding_counter: u32,
131 current_block: ?u32,
132 /// Target bounds on the push-constant block.
133 limits: abi.Limits = .{},
134 float_controls: FloatControls = .{},
135 /// Kernel whose push-constant layout the caller wants back.
136 entry_name: ?[]const u8 = null,
137 /// Layout of `entry_name`'s push-constant block once it is emitted.
138 entry_push_constants: ?abi.PushConstants = null,
139 /// Interface of the stage function being emitted.
140 stage_interface: ?stage_emit.Interface = null,
141 sampled_image_type: ?u32 = null,
142 /// Combined image sampler variables by group and binding.
143 texture_vars: std.AutoHashMapUnmanaged(stage_emit.ResourceKey, u32) = .{},
144 /// Bytes of the push-constant block that stage functions read.
145 push_extent: u32 = 0,
146 push_var: ?u32 = null,
147 /// Uniform blocks by group and binding.
148 uniform_blocks: std.AutoHashMapUnmanaged(stage_emit.ResourceKey, stage_emit.UniformBlock) = .{},
149 call_ids: std.StringHashMapUnmanaged(u32) = .{},
150
151 pub fn init(allocator: std.mem.Allocator) SpirvCodegen {
152 return .{
153 .allocator = allocator,
154 .builder = ModuleBuilder.init(allocator),
155 .value_ids = .{},
156 .memref_bindings = .{},
157 .scalar_types = .{},
158 .vector_types = .{},
159 .pointer_types = .{},
160 .array_types = .{},
161 .runtime_array_types = .{},
162 .struct_types = .{},
163 .pair_struct_types = .{},
164 .constants = .{},
165 .builtin_vars = .{},
166 .interface_vars = .empty,
167 .interface_set = .{},
168 .capability_set = .{},
169 .extension_set = .{},
170 .void_function_type = null,
171 .workgroup_size_constant = null,
172 .binding_counter = 0,
173 .current_block = null,
174 };
175 }
176
177 pub fn deinit(self: *SpirvCodegen) void {
178 self.builder.deinit();
179 self.value_ids.deinit(self.allocator);
180 self.memref_bindings.deinit(self.allocator);
181 self.scalar_types.deinit(self.allocator);
182 self.vector_types.deinit(self.allocator);
183 self.pointer_types.deinit(self.allocator);
184 self.array_types.deinit(self.allocator);
185 self.runtime_array_types.deinit(self.allocator);
186 self.struct_types.deinit(self.allocator);
187 self.pair_struct_types.deinit(self.allocator);
188 self.constants.deinit(self.allocator);
189 self.builtin_vars.deinit(self.allocator);
190 self.interface_vars.deinit(self.allocator);
191 self.interface_set.deinit(self.allocator);
192 self.capability_set.deinit(self.allocator);
193 self.extension_set.deinit(self.allocator);
194 self.texture_vars.deinit(self.allocator);
195 self.uniform_blocks.deinit(self.allocator);
196 self.call_ids.deinit(self.allocator);
197 }
198
199 pub fn getWorkgroupSizeConstant(self: *SpirvCodegen) SpirvCodegenError!u32 {
200 if (self.workgroup_size_constant) |id| return id;
201
202 const u32_type = try self.getScalarType(.u32);
203 const vec_type = try self.getVectorType(.u32, 3);
204 var member_ids: [3]u32 = undefined;
205 for (&member_ids, 0..) |*member_id, spec_index| {
206 member_id.* = self.builder.newId();
207 try self.builder.emit(&self.builder.types, SpirvOp.SpecConstant, &.{
208 u32_type,
209 member_id.*,
210 1,
211 });
212 try self.builder.emit(&self.builder.annotations, SpirvOp.Decorate, &.{
213 member_id.*,
214 SpvDecoration.SpecId,
215 @intCast(spec_index),
216 });
217 }
218 const composite_id = self.builder.newId();
219 try self.builder.emit(&self.builder.types, SpirvOp.SpecConstantComposite, &.{
220 vec_type,
221 composite_id,
222 member_ids[0],
223 member_ids[1],
224 member_ids[2],
225 });
226 try self.builder.emit(&self.builder.annotations, SpirvOp.Decorate, &.{
227 composite_id,
228 SpvDecoration.BuiltIn,
229 gpu.SpvBuiltIn.WorkgroupSize,
230 });
231
232 self.workgroup_size_constant = composite_id;
233 return composite_id;
234 }
235
236 pub fn requireCapability(self: *SpirvCodegen, capability: u32) !void {
237 const entry = try self.capability_set.getOrPut(self.allocator, capability);
238 if (entry.found_existing) return;
239 try self.builder.emitCapability(capability);
240 }
241
242 pub fn requireExtension(self: *SpirvCodegen, extension: []const u8) !void {
243 const entry = try self.extension_set.getOrPut(self.allocator, extension);
244 if (entry.found_existing) return;
245 try self.builder.emitExtension(extension);
246 }
247
248 pub fn emitFloatExecutionModes(self: *SpirvCodegen, func_id: u32) SpirvCodegenError!void {
249 const widths = [_]u32{ 16, 32, 64 };
250 const denorm = self.float_controls.denorm_preserve;
251 const signed = self.float_controls.signed_zero_inf_nan_preserve;
252 const denorm_flags = [_]bool{ denorm.f16, denorm.f32, denorm.f64 };
253 const signed_flags = [_]bool{ signed.f16, signed.f32, signed.f64 };
254 for (widths, 0..) |width, index| {
255 if (!denorm_flags[index] and !signed_flags[index]) continue;
256 self.builder.requireVersion(SpirvVersion.v13);
257 try self.requireExtension("SPV_KHR_float_controls");
258 if (denorm_flags[index]) {
259 try self.requireCapability(SpvCapability.DenormPreserve);
260 try self.builder.emitFloatExecutionMode(func_id, SpvExecutionMode.DenormPreserve, width);
261 }
262 if (signed_flags[index]) {
263 try self.requireCapability(SpvCapability.SignedZeroInfNanPreserve);
264 try self.builder.emitFloatExecutionMode(func_id, SpvExecutionMode.SignedZeroInfNanPreserve, width);
265 }
266 }
267 }
268
269 pub fn emitModuleWords(self: *SpirvCodegen, module: *ir.Operation) SpirvCodegenError![]u32 {
270 self.entry_push_constants = null;
271 if (std.mem.eql(u8, module.name.name, SpirvDialect.ModuleOp.operation_name)) {
272 return dialect_writer.emitModuleWords(self, module);
273 }
274 var call_plan = try calls.Plan.init(self.allocator, module, null);
275 defer call_plan.deinit();
276 try validation.validateModule(module);
277 try stage_emit.scanBlocks(self, module);
278 self.call_ids.clearRetainingCapacity();
279
280 try self.requireCapability(SpvCapability.Shader);
281 try self.builder.emitMemoryModel(SpvAddressingModel.Logical, SpvMemoryModel.GLSL450);
282
283 const region = module.getRegion(0) orelse return error.InvalidModule;
284 const block = region.getEntryBlock() orelse return error.InvalidModule;
285
286 for (call_plan.helpers.items) |helper| try self.emitHelperFunction(helper);
287
288 var op_iter = block.operations.head;
289 while (op_iter) |op_ptr| {
290 const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
291 if (validation.isFunctionOp(op)) {
292 if (op.getAttr("kernel") != null) {
293 try self.emitKernelFunction(op);
294 } else if (gpu_target.stage.stageOf(op)) |stage| {
295 try self.emitStageFunction(op, stage);
296 }
297 } else {
298 return error.UnsupportedOperation;
299 }
300 op_iter = op.next_op;
301 }
302
303 return self.builder.toWords(self.allocator);
304 }
305
306 pub fn emitModuleBytes(self: *SpirvCodegen, module: *ir.Operation) SpirvCodegenError![]u8 {
307 const words = try self.emitModuleWords(module);
308 defer self.allocator.free(words);
309
310 const byte_len = words.len * @sizeOf(u32);
311 const bytes = try self.allocator.alloc(u8, byte_len);
312 var offset: usize = 0;
313 for (words) |word| {
314 std.mem.writeInt(u32, bytes[offset..][0..4], word, .little);
315 offset += 4;
316 }
317 return bytes;
318 }
319
320 pub fn emitLabel(self: *SpirvCodegen, label_id: u32) SpirvCodegenError!void {
321 try self.builder.emit(&self.builder.functions, SpirvOp.Label, &.{label_id});
322 self.current_block = label_id;
323 }
324
325 fn emitKernelFunction(self: *SpirvCodegen, func_op: *ir.Operation) SpirvCodegenError!void {
326 const name = dialect_writer.functionName(func_op) orelse return error.MissingFunctionName;
327
328 self.value_ids.clearRetainingCapacity();
329 self.memref_bindings.clearRetainingCapacity();
330 self.interface_vars.clearRetainingCapacity();
331 self.interface_set.clearRetainingCapacity();
332
333 const region = func_op.getRegion(0) orelse return error.InvalidModule;
334 const entry = region.getEntryBlock() orelse return error.InvalidModule;
335
336 var scalar_args = std.ArrayListUnmanaged(ScalarArgument).empty;
337 defer scalar_args.deinit(self.allocator);
338 var push_constants: abi.PushConstants = .{};
339
340 for (entry.arguments.items) |arg| {
341 if (memory.parseType(arg.type)) |memref| {
342 const binding = self.binding_counter;
343 self.binding_counter += 1;
344
345 if (memref.addr_space == .shared) return error.UnsupportedAddressSpace;
346 const storage_class = memory.storageClassForAddressSpace(memref.addr_space) orelse return error.UnsupportedAddressSpace;
347 const elem_kind = scalar.kindFromName(memref.element_type_name) orelse return error.UnsupportedType;
348
349 try self.requireBufferStorageCapability(memory.storageElementKind(elem_kind), storage_class);
350 const elem_type_id = try self.getScalarType(elem_kind);
351 const storage_elem_kind = memory.storageElementKind(elem_kind);
352 const storage_elem_type_id = try self.getScalarType(storage_elem_kind);
353 const layout = try memory.getBufferLayout(self, storage_elem_type_id, storage_elem_kind, storage_class);
354
355 const var_id = self.builder.newId();
356 try self.builder.emit(&self.builder.globals, SpirvOp.Variable, &.{
357 layout.ptr_struct_type,
358 var_id,
359 storage_class,
360 });
361
362 try memory.decorateBufferVariable(self, var_id, @intCast(binding), storage_class, memref.addr_space);
363 try self.bindValue(arg, var_id);
364 try self.memref_bindings.put(self.allocator, arg, .{
365 .storage_class = storage_class,
366 .elem_kind = elem_kind,
367 .elem_type = elem_type_id,
368 .storage_elem_kind = storage_elem_kind,
369 .storage_elem_type = storage_elem_type_id,
370 .ptr_elem_type = layout.ptr_elem_type,
371 .layout = .buffer,
372 });
373
374 try self.addInterfaceVar(var_id);
375 continue;
376 }
377
378 const scalar_kind = scalar.kindFromType(arg.type) orelse return error.UnsupportedType;
379 const elem_type_id = try self.getScalarType(scalar_kind);
380 const byte_size: u32 = @intCast(scalar.elementByteSize(scalar_kind) orelse return error.UnsupportedType);
381 const byte_offset = push_constants.append(byte_size, self.limits) catch
382 return error.UnsupportedFunctionSignature;
383 const ptr_elem_type = try self.getPointerType(SpvStorageClass.PushConstant, elem_type_id);
384 try scalar_args.append(self.allocator, .{
385 .arg = arg,
386 .elem_type = elem_type_id,
387 .ptr_elem_type = ptr_elem_type,
388 .member_index = @intCast(scalar_args.items.len),
389 .byte_offset = byte_offset,
390 });
391 }
392
393 if (self.entry_name) |entry_name| {
394 if (std.mem.eql(u8, name, entry_name)) self.entry_push_constants = push_constants;
395 }
396
397 var push_var: u32 = 0;
398 if (scalar_args.items.len > 0) {
399 var struct_operands = std.ArrayListUnmanaged(u32).empty;
400 defer struct_operands.deinit(self.allocator);
401 const struct_id = self.builder.newId();
402 try struct_operands.append(self.allocator, struct_id);
403 for (scalar_args.items) |scalar_arg| {
404 try struct_operands.append(self.allocator, scalar_arg.elem_type);
405 }
406 try self.builder.emit(&self.builder.types, SpirvOp.TypeStruct, struct_operands.items);
407 try self.builder.emit(&self.builder.annotations, SpirvOp.Decorate, &.{
408 struct_id,
409 SpvDecoration.Block,
410 });
411 for (scalar_args.items) |scalar_arg| {
412 try self.builder.emit(&self.builder.annotations, SpirvOp.MemberDecorate, &.{
413 struct_id,
414 scalar_arg.member_index,
415 SpvDecoration.Offset,
416 scalar_arg.byte_offset,
417 });
418 }
419 const ptr_struct = try self.getPointerType(SpvStorageClass.PushConstant, struct_id);
420 push_var = self.builder.newId();
421 try self.builder.emit(&self.builder.globals, SpirvOp.Variable, &.{
422 ptr_struct,
423 push_var,
424 SpvStorageClass.PushConstant,
425 });
426 try self.addInterfaceVar(push_var);
427 }
428
429 const void_type = try self.getScalarType(.void);
430 const func_type = try self.getVoidFunctionType(void_type);
431 const func_id = self.builder.newId();
432
433 try self.builder.emit(&self.builder.functions, SpirvOp.Function, &.{
434 void_type,
435 func_id,
436 0,
437 func_type,
438 });
439
440 const label_id = self.builder.newId();
441 try self.emitLabel(label_id);
442 try self.declareFunctionAllocas(entry);
443
444 if (scalar_args.items.len > 0) {
445 const index_type = try self.getScalarType(.u32);
446 for (scalar_args.items) |scalar_arg| {
447 const member_index = try self.getIntConstant(index_type, .u32, scalar_arg.member_index);
448 const access_id = self.builder.newId();
449 try self.builder.emit(&self.builder.functions, SpirvOp.AccessChain, &.{
450 scalar_arg.ptr_elem_type,
451 access_id,
452 push_var,
453 member_index,
454 });
455
456 const load_id = self.builder.newId();
457 try self.builder.emit(&self.builder.functions, SpirvOp.Load, &.{
458 scalar_arg.elem_type,
459 load_id,
460 access_id,
461 });
462 try self.bindValue(scalar_arg.arg, load_id);
463 }
464 }
465
466 var op_iter = entry.operations.head;
467 while (op_iter) |op_ptr| {
468 const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
469 try self.emitOperation(op);
470 op_iter = op.next_op;
471 }
472
473 try self.builder.emit(&self.builder.functions, SpirvOp.FunctionEnd, &.{});
474 self.current_block = null;
475
476 try self.builder.emitEntryPoint(
477 SpvExecutionModel.GLCompute,
478 func_id,
479 name,
480 self.interface_vars.items,
481 );
482
483 try self.builder.emitExecutionModeLocalSize(func_id, 1, 1, 1);
484 try self.emitFloatExecutionModes(func_id);
485 _ = try self.getWorkgroupSizeConstant();
486 }
487
488 /// Emits a vertex or fragment entry. Its interface variables enter the
489 /// entry point; a fragment entry takes an upper-left origin, as Vulkan
490 /// requires.
491 fn emitStageFunction(
492 self: *SpirvCodegen,
493 func_op: *ir.Operation,
494 stage: Stage,
495 ) SpirvCodegenError!void {
496 const name = dialect_writer.functionName(func_op) orelse return error.MissingFunctionName;
497 const region = func_op.getRegion(0) orelse return error.InvalidModule;
498 const entry = region.getEntryBlock() orelse return error.InvalidModule;
499 if (entry.arguments.items.len != 0) return error.UnsupportedFunctionSignature;
500
501 self.value_ids.clearRetainingCapacity();
502 self.interface_vars.clearRetainingCapacity();
503 self.interface_set.clearRetainingCapacity();
504 self.stage_interface = .{ .stage = stage };
505 defer self.stage_interface = null;
506
507 const void_type = try self.getScalarType(.void);
508 const func_type = try self.getVoidFunctionType(void_type);
509 const func_id = self.builder.newId();
510 try self.builder.emit(&self.builder.functions, SpirvOp.Function, &.{
511 void_type,
512 func_id,
513 0,
514 func_type,
515 });
516 try self.emitLabel(self.builder.newId());
517 self.memref_bindings.clearRetainingCapacity();
518 try self.declareFunctionAllocas(entry);
519
520 var op_iter = entry.operations.head;
521 while (op_iter) |op_ptr| {
522 const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
523 try self.emitOperation(op);
524 op_iter = op.next_op;
525 }
526
527 try self.builder.emit(&self.builder.functions, SpirvOp.FunctionEnd, &.{});
528 self.current_block = null;
529
530 try self.builder.emitEntryPoint(
531 stage_emit.executionModel(stage),
532 func_id,
533 name,
534 self.interface_vars.items,
535 );
536 if (stage == .fragment) {
537 try self.builder.emitExecutionMode(func_id, SpvExecutionMode.OriginUpperLeft);
538 }
539 try self.emitFloatExecutionModes(func_id);
540 }
541
542 fn emitOperation(self: *SpirvCodegen, op: *ir.Operation) SpirvCodegenError!void {
543 return catalog.emit(SpirvCodegenError, catalog_methods, self, op);
544 }
545
546 fn emitHelperFunction(self: *SpirvCodegen, func_op: *ir.Operation) SpirvCodegenError!void {
547 try validation.validateHelperOp(func_op);
548 const func = FuncDialect.FuncOp{ .op = func_op };
549 const name = func.getName() orelse return error.MissingFunctionName;
550 const results = func.getResultTypes();
551 if (results.len > 1) return error.UnsupportedFunctionSignature;
552 const return_type = if (results.len == 0)
553 try self.getScalarType(.void)
554 else
555 try self.getScalarType(scalar.kindFromType(results[0]) orelse return error.UnsupportedType);
556 var param_types = std.ArrayListUnmanaged(u32).empty;
557 defer param_types.deinit(self.allocator);
558 for (func.getArguments()) |arg| {
559 const kind = scalar.kindFromType(arg.type) orelse return error.UnsupportedFunctionSignature;
560 try param_types.append(self.allocator, try self.getScalarType(kind));
561 }
562 const func_type = try dialect_writer.emitFunctionType(self, return_type, param_types.items);
563 const func_id = self.builder.newId();
564 try self.call_ids.put(self.allocator, name, func_id);
565 self.value_ids.clearRetainingCapacity();
566 self.memref_bindings.clearRetainingCapacity();
567 try self.builder.emit(&self.builder.functions, SpirvOp.Function, &.{ return_type, func_id, 0, func_type });
568 for (func.getArguments(), param_types.items) |arg, param_type| {
569 const param_id = self.builder.newId();
570 try self.builder.emit(&self.builder.functions, SpirvOp.FunctionParameter, &.{ param_type, param_id });
571 try self.bindValue(arg, param_id);
572 }
573 try self.emitLabel(self.builder.newId());
574 const entry = func.getEntryBlock();
575 try self.declareFunctionAllocas(entry);
576 var ops = entry.getOperations();
577 while (ops.next()) |op| try self.emitOperation(op);
578 try self.builder.emit(&self.builder.functions, SpirvOp.FunctionEnd, &.{});
579 self.current_block = null;
580 }
581
582 fn emitCall(self: *SpirvCodegen, op: *ir.Operation) SpirvCodegenError!void {
583 const call = FuncDialect.CallOp{ .op = op };
584 const name = call.getCallee() orelse return error.MissingAttribute;
585 const callee_id = self.call_ids.get(name) orelse return error.UnsupportedOperation;
586 if (call.getNumResults() > 1) return error.UnsupportedFunctionSignature;
587 const return_type = if (call.getNumResults() == 0)
588 try self.getScalarType(.void)
589 else
590 try self.getScalarType(scalar.kindFromType(call.getResult(0).?.type) orelse return error.UnsupportedType);
591 const result_id = self.builder.newId();
592 var operands = std.ArrayListUnmanaged(u32).empty;
593 defer operands.deinit(self.allocator);
594 try operands.appendSlice(self.allocator, &.{ return_type, result_id, callee_id });
595 for (call.getOperands()) |arg| try operands.append(self.allocator, try self.getValue(arg));
596 try self.builder.emit(&self.builder.functions, SpirvOp.FunctionCall, operands.items);
597 if (call.getResult(0)) |result| try self.bindValue(result, result_id);
598 }
599
600 fn declareFunctionAllocas(self: *SpirvCodegen, block: *ir.Block) SpirvCodegenError!void {
601 var ops = block.getOperations();
602 while (ops.next()) |op| {
603 if (std.mem.eql(u8, op.name.name, MemrefDialect.AllocaOp.operation_name)) {
604 try memory.declareAlloca(self, op);
605 }
606 for (0..op.getNumRegions()) |index| {
607 const region = op.getRegion(index) orelse continue;
608 var blocks = region.getBlocks();
609 while (blocks.next()) |nested| try self.declareFunctionAllocas(nested);
610 }
611 }
612 }
613
614 fn emitReturn(self: *SpirvCodegen, op: *ir.Operation) SpirvCodegenError!void {
615 if (op.operands.items.len == 0) {
616 try self.builder.emit(&self.builder.functions, SpirvOp.Return, &.{});
617 return;
618 }
619 if (op.operands.items.len == 1) {
620 const value_id = try self.getValue(op.operands.items[0].value);
621 try self.builder.emit(&self.builder.functions, SpirvOp.ReturnValue, &.{value_id});
622 return;
623 }
624 return error.UnsupportedFunctionSignature;
625 }
626
627 fn emitBlockUntilYield(
628 self: *SpirvCodegen,
629 block: *ir.Block,
630 yields: *std.ArrayListUnmanaged(u32),
631 ) SpirvCodegenError!void {
632 yields.clearRetainingCapacity();
633 var op_iter = block.operations.head;
634 while (op_iter) |op_ptr| {
635 const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
636 if (std.mem.eql(u8, op.name.name, ScfDialect.YieldOp.operation_name)) {
637 if (op.next_op != null) return error.UnsupportedControlFlow;
638 const yield_op = ScfDialect.YieldOp{ .op = op };
639 for (yield_op.getOperands()) |operand| {
640 const value_id = try self.getValue(operand);
641 try yields.append(self.allocator, value_id);
642 }
643 return;
644 }
645 try self.emitOperation(op);
646 op_iter = op.next_op;
647 }
648 return error.UnsupportedControlFlow;
649 }
650
651 fn emitBlockUntilCondition(
652 self: *SpirvCodegen,
653 block: *ir.Block,
654 cond_id: *u32,
655 args: *std.ArrayListUnmanaged(u32),
656 ) SpirvCodegenError!void {
657 args.clearRetainingCapacity();
658 var op_iter = block.operations.head;
659 while (op_iter) |op_ptr| {
660 const op: *ir.Operation = @ptrCast(@alignCast(op_ptr));
661 if (std.mem.eql(u8, op.name.name, ScfDialect.ConditionOp.operation_name)) {
662 if (op.next_op != null) return error.UnsupportedControlFlow;
663 const cond_op = ScfDialect.ConditionOp{ .op = op };
664 cond_id.* = try self.getValue(cond_op.getCondition());
665 for (cond_op.getArgs()) |operand| {
666 const value_id = try self.getValue(operand);
667 try args.append(self.allocator, value_id);
668 }
669 return;
670 }
671 try self.emitOperation(op);
672 op_iter = op.next_op;
673 }
674 return error.UnsupportedControlFlow;
675 }
676
677 fn bindBlockArguments(self: *SpirvCodegen, block: *ir.Block, values: []const u32) SpirvCodegenError!void {
678 if (block.arguments.items.len != values.len) return error.UnsupportedControlFlow;
679 for (block.arguments.items, values) |arg, value_id| {
680 try self.bindValue(arg, value_id);
681 }
682 }
683
684 fn emitCompareLess(
685 self: *SpirvCodegen,
686 kind: ScalarKind,
687 lhs_id: u32,
688 rhs_id: u32,
689 ) SpirvCodegenError!u32 {
690 const bool_type_id = try self.getScalarType(.bool);
691 const opcode: u16 = if (scalar.isFloat(kind))
692 SpirvOp.FOrdLessThan
693 else if (scalar.isSignedInt(kind))
694 SpirvOp.SLessThan
695 else
696 SpirvOp.ULessThan;
697
698 const result_id = self.builder.newId();
699 try self.builder.emit(&self.builder.functions, opcode, &.{
700 bool_type_id,
701 result_id,
702 lhs_id,
703 rhs_id,
704 });
705 return result_id;
706 }
707
708 fn emitPhiWithPatch(
709 self: *SpirvCodegen,
710 result_type_id: u32,
711 result_id: u32,
712 init_value_id: u32,
713 init_label: u32,
714 backedge_label: u32,
715 ) SpirvCodegenError!usize {
716 const start_index = self.builder.functions.items.len;
717 try self.builder.emit(&self.builder.functions, SpirvOp.Phi, &.{
718 result_type_id,
719 result_id,
720 init_value_id,
721 init_label,
722 0,
723 backedge_label,
724 });
725 return start_index + 1 + 4;
726 }
727
728 fn emitScfIf(self: *SpirvCodegen, op: *ir.Operation) SpirvCodegenError!void {
729 const if_op = ScfDialect.IfOp{ .op = op };
730 const cond_id = try self.getValue(if_op.getCondition());
731
732 const then_block = if_op.getThenBlock();
733 const else_block = if_op.getElseBlock();
734
735 const merge_label = self.builder.newId();
736 const then_label = self.builder.newId();
737 const else_label = if (else_block != null) self.builder.newId() else merge_label;
738
739 try self.builder.emit(&self.builder.functions, SpirvOp.SelectionMerge, &.{ merge_label, 0 });
740 try self.builder.emit(&self.builder.functions, SpirvOp.BranchConditional, &.{ cond_id, then_label, else_label });
741
742 var then_yields: std.ArrayListUnmanaged(u32) = .empty;
743 defer then_yields.deinit(self.allocator);
744 try self.emitLabel(then_label);
745 try self.emitBlockUntilYield(then_block, &then_yields);
746 try self.builder.emit(&self.builder.functions, SpirvOp.Branch, &.{merge_label});
747
748 var else_yields: std.ArrayListUnmanaged(u32) = .empty;
749 defer else_yields.deinit(self.allocator);
750 if (else_block) |block| {
751 try self.emitLabel(else_label);
752 try self.emitBlockUntilYield(block, &else_yields);
753 try self.builder.emit(&self.builder.functions, SpirvOp.Branch, &.{merge_label});
754 }
755
756 try self.emitLabel(merge_label);
757
758 const result_count = if_op.getNumResults();
759 if (result_count == 0) return;
760 if (then_yields.items.len != result_count) return error.UnsupportedControlFlow;
761 if (else_block == null) return error.UnsupportedControlFlow;
762 if (else_yields.items.len != result_count) return error.UnsupportedControlFlow;
763
764 for (0..result_count) |i| {
765 const result_val = if_op.getResult(i) orelse return error.UnsupportedOperation;
766 const result_type_id = try self.getTypeForValue(result_val);
767 const result_id = self.builder.newId();
768 try self.builder.emit(&self.builder.functions, SpirvOp.Phi, &.{
769 result_type_id,
770 result_id,
771 then_yields.items[i],
772 then_label,
773 else_yields.items[i],
774 else_label,
775 });
776 try self.bindValue(result_val, result_id);
777 }
778 }
779
780 fn emitScfFor(self: *SpirvCodegen, op: *ir.Operation) SpirvCodegenError!void {
781 const for_op = ScfDialect.ForOp{ .op = op };
782 const preheader = self.current_block orelse return error.UnsupportedControlFlow;
783
784 const header_label = self.builder.newId();
785 const body_label = self.builder.newId();
786 const merge_label = self.builder.newId();
787
788 try self.builder.emit(&self.builder.functions, SpirvOp.Branch, &.{header_label});
789
790 try self.emitLabel(header_label);
791
792 const lower_id = try self.getValue(for_op.getLowerBound());
793 const upper_id = try self.getValue(for_op.getUpperBound());
794 const step_id = try self.getValue(for_op.getStep());
795 const iv_kind = scalar.kindFromType(for_op.getLowerBound().type) orelse return error.UnsupportedType;
796 const iv_type_id = try self.getScalarType(iv_kind);
797
798 var phi_ids: std.ArrayListUnmanaged(u32) = .empty;
799 defer phi_ids.deinit(self.allocator);
800 var phi_patches: std.ArrayListUnmanaged(usize) = .empty;
801 defer phi_patches.deinit(self.allocator);
802
803 const iv_phi_id = self.builder.newId();
804 const iv_patch = try self.emitPhiWithPatch(iv_type_id, iv_phi_id, lower_id, preheader, body_label);
805 try phi_ids.append(self.allocator, iv_phi_id);
806 try phi_patches.append(self.allocator, iv_patch);
807
808 const init_args = for_op.getInitArgs();
809 for (init_args) |init_arg| {
810 const init_id = try self.getValue(init_arg);
811 const init_type_id = try self.getTypeForValue(init_arg);
812 const phi_id = self.builder.newId();
813 const patch = try self.emitPhiWithPatch(init_type_id, phi_id, init_id, preheader, body_label);
814 try phi_ids.append(self.allocator, phi_id);
815 try phi_patches.append(self.allocator, patch);
816 }
817
818 const cond_id = try self.emitCompareLess(iv_kind, iv_phi_id, upper_id);
819 try self.builder.emit(&self.builder.functions, SpirvOp.LoopMerge, &.{ merge_label, body_label, 0 });
820 try self.builder.emit(&self.builder.functions, SpirvOp.BranchConditional, &.{ cond_id, body_label, merge_label });
821
822 try self.emitLabel(body_label);
823
824 var body_args: std.ArrayListUnmanaged(u32) = .empty;
825 defer body_args.deinit(self.allocator);
826 try body_args.append(self.allocator, iv_phi_id);
827 for (phi_ids.items[1..]) |phi_id| {
828 try body_args.append(self.allocator, phi_id);
829 }
830 try self.bindBlockArguments(for_op.getBodyBlock(), body_args.items);
831
832 var yields: std.ArrayListUnmanaged(u32) = .empty;
833 defer yields.deinit(self.allocator);
834 try self.emitBlockUntilYield(for_op.getBodyBlock(), &yields);
835 if (yields.items.len != init_args.len) return error.UnsupportedControlFlow;
836
837 const next_iv_id = try scalar.emitBinaryArithIds(self, iv_kind, iv_phi_id, step_id, .add);
838 try self.builder.emit(&self.builder.functions, SpirvOp.Branch, &.{header_label});
839
840 self.builder.functions.items[phi_patches.items[0]] = next_iv_id;
841 for (yields.items, 0..) |yield_id, idx| {
842 const patch_index = phi_patches.items[idx + 1];
843 self.builder.functions.items[patch_index] = yield_id;
844 }
845
846 try self.emitLabel(merge_label);
847
848 for (0..init_args.len) |i| {
849 const result_val = for_op.getResult(i) orelse return error.UnsupportedOperation;
850 try self.bindValue(result_val, phi_ids.items[i + 1]);
851 }
852 }
853
854 fn emitScfWhile(self: *SpirvCodegen, op: *ir.Operation) SpirvCodegenError!void {
855 const while_op = ScfDialect.WhileOp{ .op = op };
856 const preheader = self.current_block orelse return error.UnsupportedControlFlow;
857
858 const header_label = self.builder.newId();
859 const body_label = self.builder.newId();
860 const merge_label = self.builder.newId();
861
862 try self.builder.emit(&self.builder.functions, SpirvOp.Branch, &.{header_label});
863
864 try self.emitLabel(header_label);
865
866 var phi_ids: std.ArrayListUnmanaged(u32) = .empty;
867 defer phi_ids.deinit(self.allocator);
868 var phi_patches: std.ArrayListUnmanaged(usize) = .empty;
869 defer phi_patches.deinit(self.allocator);
870
871 for (op.operands.items) |operand| {
872 const init_val = operand.value;
873 const init_id = try self.getValue(init_val);
874 const type_id = try self.getTypeForValue(init_val);
875 const phi_id = self.builder.newId();
876 const patch = try self.emitPhiWithPatch(type_id, phi_id, init_id, preheader, body_label);
877 try phi_ids.append(self.allocator, phi_id);
878 try phi_patches.append(self.allocator, patch);
879 }
880
881 try self.bindBlockArguments(while_op.getBeforeBlock(), phi_ids.items);
882
883 var cond_args: std.ArrayListUnmanaged(u32) = .empty;
884 defer cond_args.deinit(self.allocator);
885 var cond_id: u32 = 0;
886 try self.emitBlockUntilCondition(while_op.getBeforeBlock(), &cond_id, &cond_args);
887 if (cond_args.items.len != phi_ids.items.len) return error.UnsupportedControlFlow;
888
889 try self.builder.emit(&self.builder.functions, SpirvOp.LoopMerge, &.{ merge_label, body_label, 0 });
890 try self.builder.emit(&self.builder.functions, SpirvOp.BranchConditional, &.{ cond_id, body_label, merge_label });
891
892 try self.emitLabel(body_label);
893 try self.bindBlockArguments(while_op.getAfterBlock(), cond_args.items);
894
895 var yields: std.ArrayListUnmanaged(u32) = .empty;
896 defer yields.deinit(self.allocator);
897 try self.emitBlockUntilYield(while_op.getAfterBlock(), &yields);
898 if (yields.items.len != phi_ids.items.len) return error.UnsupportedControlFlow;
899
900 try self.builder.emit(&self.builder.functions, SpirvOp.Branch, &.{header_label});
901
902 for (yields.items, 0..) |yield_id, idx| {
903 const patch_index = phi_patches.items[idx];
904 self.builder.functions.items[patch_index] = yield_id;
905 }
906
907 try self.emitLabel(merge_label);
908
909 for (0..cond_args.items.len) |i| {
910 const result_val = op.getResult(i) orelse return error.UnsupportedOperation;
911 try self.bindValue(result_val, cond_args.items[i]);
912 }
913 }
914
915 pub fn addInterfaceVar(self: *SpirvCodegen, var_id: u32) SpirvCodegenError!void {
916 const gop = try self.interface_set.getOrPut(self.allocator, var_id);
917 if (!gop.found_existing) {
918 try self.interface_vars.append(self.allocator, var_id);
919 }
920 }
921
922 pub fn bindValue(self: *SpirvCodegen, value: *ir.Value, id: u32) SpirvCodegenError!void {
923 try self.value_ids.put(self.allocator, value, id);
924 }
925
926 pub fn getValue(self: *SpirvCodegen, value: *ir.Value) SpirvCodegenError!u32 {
927 return self.value_ids.get(value) orelse error.MissingValue;
928 }
929
930 pub fn getTypeForValue(self: *SpirvCodegen, value: *ir.Value) SpirvCodegenError!u32 {
931 const kind = scalar.kindFromType(value.type) orelse return error.UnsupportedType;
932 return self.getScalarType(kind);
933 }
934
935 pub fn getScalarType(self: *SpirvCodegen, kind: ScalarKind) SpirvCodegenError!u32 {
936 if (self.scalar_types.get(kind)) |id| return id;
937
938 try self.requireScalarCapability(kind);
939
940 const id = self.builder.newId();
941 switch (kind) {
942 .void => try self.builder.emit(&self.builder.types, SpirvOp.TypeVoid, &.{id}),
943 .bool => try self.builder.emit(&self.builder.types, SpirvOp.TypeBool, &.{id}),
944 .i8 => try self.builder.emit(&self.builder.types, SpirvOp.TypeInt, &.{ id, 8, 1 }),
945 .i16 => try self.builder.emit(&self.builder.types, SpirvOp.TypeInt, &.{ id, 16, 1 }),
946 .i32 => try self.builder.emit(&self.builder.types, SpirvOp.TypeInt, &.{ id, 32, 1 }),
947 .i64 => try self.builder.emit(&self.builder.types, SpirvOp.TypeInt, &.{ id, 64, 1 }),
948 .u8 => try self.builder.emit(&self.builder.types, SpirvOp.TypeInt, &.{ id, 8, 0 }),
949 .u16 => try self.builder.emit(&self.builder.types, SpirvOp.TypeInt, &.{ id, 16, 0 }),
950 .u32 => try self.builder.emit(&self.builder.types, SpirvOp.TypeInt, &.{ id, 32, 0 }),
951 .u64 => try self.builder.emit(&self.builder.types, SpirvOp.TypeInt, &.{ id, 64, 0 }),
952 .f16 => try self.builder.emit(&self.builder.types, SpirvOp.TypeFloat, &.{ id, 16 }),
953 .f32 => try self.builder.emit(&self.builder.types, SpirvOp.TypeFloat, &.{ id, 32 }),
954 .f64 => try self.builder.emit(&self.builder.types, SpirvOp.TypeFloat, &.{ id, 64 }),
955 }
956
957 try self.scalar_types.put(self.allocator, kind, id);
958 return id;
959 }
960
961 fn requireScalarCapability(self: *SpirvCodegen, kind: ScalarKind) SpirvCodegenError!void {
962 const cap: ?u32 = switch (kind) {
963 .i8, .u8 => SpvCapability.Int8,
964 .i16, .u16 => SpvCapability.Int16,
965 .i64, .u64 => SpvCapability.Int64,
966 .f16 => SpvCapability.Float16,
967 .f64 => SpvCapability.Float64,
968 .void, .bool, .i32, .u32, .f32 => null,
969 };
970 if (cap) |c| try self.requireCapability(c);
971 }
972
973 fn requireBufferStorageCapability(self: *SpirvCodegen, kind: ScalarKind, storage_class: u32) SpirvCodegenError!void {
974 if (storage_class == SpvStorageClass.StorageBuffer) {
975 try self.requireExtension("SPV_KHR_storage_buffer_storage_class");
976 }
977 const width: ?u8 = switch (kind) {
978 .i8, .u8 => 8,
979 .i16, .u16, .f16 => 16,
980 else => null,
981 };
982 const bits = width orelse return;
983 const capability = switch (storage_class) {
984 SpvStorageClass.StorageBuffer => if (bits == 8)
985 SpvCapability.StorageBuffer8BitAccess
986 else
987 SpvCapability.StorageBuffer16BitAccess,
988 SpvStorageClass.Uniform => if (bits == 8)
989 SpvCapability.UniformAndStorageBuffer8BitAccess
990 else
991 SpvCapability.UniformAndStorageBuffer16BitAccess,
992 else => return,
993 };
994 try self.requireCapability(capability);
995 try self.requireExtension(if (bits == 8) "SPV_KHR_8bit_storage" else "SPV_KHR_16bit_storage");
996 }
997
998 pub fn getVectorType(self: *SpirvCodegen, elem: ScalarKind, len: u32) SpirvCodegenError!u32 {
999 const key = VectorKey{ .elem = elem, .len = len };
1000 if (self.vector_types.get(key)) |id| return id;
1001
1002 const elem_type = try self.getScalarType(elem);
1003 const id = self.builder.newId();
1004 try self.builder.emit(&self.builder.types, SpirvOp.TypeVector, &.{ id, elem_type, len });
1005 try self.vector_types.put(self.allocator, key, id);
1006 return id;
1007 }
1008
1009 pub fn getPointerType(self: *SpirvCodegen, storage_class: u32, base_type: u32) SpirvCodegenError!u32 {
1010 const key = PointerKey{ .storage_class = storage_class, .base_type = base_type };
1011 if (self.pointer_types.get(key)) |id| return id;
1012
1013 const id = self.builder.newId();
1014 try self.builder.emit(&self.builder.types, SpirvOp.TypePointer, &.{ id, storage_class, base_type });
1015 try self.pointer_types.put(self.allocator, key, id);
1016 return id;
1017 }
1018
1019 pub fn getArrayType(self: *SpirvCodegen, elem_type: u32, length: u32) SpirvCodegenError!u32 {
1020 const key = ArrayKey{ .elem_type = elem_type, .length = length };
1021 if (self.array_types.get(key)) |id| return id;
1022
1023 const length_const = try self.getIntConstant(
1024 try self.getScalarType(.u32),
1025 .u32,
1026 @intCast(length),
1027 );
1028 const id = self.builder.newId();
1029 try self.builder.emit(&self.builder.types, SpirvOp.TypeArray, &.{ id, elem_type, length_const });
1030 try self.array_types.put(self.allocator, key, id);
1031 return id;
1032 }
1033
1034 pub fn getRuntimeArrayType(self: *SpirvCodegen, elem_type: u32, stride: u32) SpirvCodegenError!u32 {
1035 const key = RuntimeArrayKey{ .elem_type = elem_type, .stride = stride };
1036 if (self.runtime_array_types.get(key)) |id| return id;
1037
1038 const id = self.builder.newId();
1039 try self.builder.emit(&self.builder.types, SpirvOp.TypeRuntimeArray, &.{ id, elem_type });
1040 try self.builder.emit(&self.builder.annotations, SpirvOp.Decorate, &.{
1041 id,
1042 SpvDecoration.ArrayStride,
1043 stride,
1044 });
1045
1046 try self.runtime_array_types.put(self.allocator, key, id);
1047 return id;
1048 }
1049
1050 pub fn getStructType(self: *SpirvCodegen, member_type: u32) SpirvCodegenError!u32 {
1051 if (self.struct_types.get(member_type)) |id| return id;
1052
1053 const id = self.builder.newId();
1054 try self.builder.emit(&self.builder.types, SpirvOp.TypeStruct, &.{ id, member_type });
1055 try self.builder.emit(&self.builder.annotations, SpirvOp.Decorate, &.{
1056 id,
1057 SpvDecoration.Block,
1058 });
1059 try self.builder.emit(&self.builder.annotations, SpirvOp.MemberDecorate, &.{
1060 id,
1061 0,
1062 SpvDecoration.Offset,
1063 0,
1064 });
1065
1066 try self.struct_types.put(self.allocator, member_type, id);
1067 return id;
1068 }
1069
1070 pub fn getPairStructType(self: *SpirvCodegen, member_type: u32) SpirvCodegenError!u32 {
1071 if (self.pair_struct_types.get(member_type)) |id| return id;
1072
1073 const id = self.builder.newId();
1074 try self.builder.emit(&self.builder.types, SpirvOp.TypeStruct, &.{ id, member_type, member_type });
1075 try self.pair_struct_types.put(self.allocator, member_type, id);
1076 return id;
1077 }
1078
1079 fn getVoidFunctionType(self: *SpirvCodegen, void_type: u32) SpirvCodegenError!u32 {
1080 if (self.void_function_type) |id| return id;
1081
1082 const id = self.builder.newId();
1083 try self.builder.emit(&self.builder.types, SpirvOp.TypeFunction, &.{ id, void_type });
1084 self.void_function_type = id;
1085 return id;
1086 }
1087
1088 pub fn getIntConstant(
1089 self: *SpirvCodegen,
1090 type_id: u32,
1091 kind: ScalarKind,
1092 value: i64,
1093 ) SpirvCodegenError!u32 {
1094 const bits = scalar.integerConstantBits(kind, value) orelse return error.UnsupportedType;
1095 const key = ConstKey{
1096 .type_id = type_id,
1097 .word0 = bits.word0,
1098 .word1 = bits.word1,
1099 .word_count = bits.word_count,
1100 };
1101 if (self.constants.get(key)) |id| return id;
1102
1103 const id = self.builder.newId();
1104 var operands = Section.empty;
1105 defer operands.deinit(self.allocator);
1106 try operands.append(self.allocator, type_id);
1107 try operands.append(self.allocator, id);
1108 try operands.append(self.allocator, bits.word0);
1109 if (bits.word_count == 2) {
1110 try operands.append(self.allocator, bits.word1);
1111 }
1112
1113 try self.builder.emit(&self.builder.types, SpirvOp.Constant, operands.items);
1114 try self.constants.put(self.allocator, key, id);
1115 return id;
1116 }
1117
1118 pub fn getFloatConstant(
1119 self: *SpirvCodegen,
1120 type_id: u32,
1121 kind: ScalarKind,
1122 value: f64,
1123 ) SpirvCodegenError!u32 {
1124 const bits = scalar.floatConstantBits(kind, value) orelse return error.UnsupportedType;
1125 const key = ConstKey{
1126 .type_id = type_id,
1127 .word0 = bits.word0,
1128 .word1 = bits.word1,
1129 .word_count = bits.word_count,
1130 };
1131 if (self.constants.get(key)) |id| return id;
1132
1133 const id = self.builder.newId();
1134 var operands = Section.empty;
1135 defer operands.deinit(self.allocator);
1136 try operands.append(self.allocator, type_id);
1137 try operands.append(self.allocator, id);
1138 try operands.append(self.allocator, bits.word0);
1139 if (bits.word_count == 2) {
1140 try operands.append(self.allocator, bits.word1);
1141 }
1142
1143 try self.builder.emit(&self.builder.types, SpirvOp.Constant, operands.items);
1144 try self.constants.put(self.allocator, key, id);
1145 return id;
1146 }
1147
1148 pub fn getBoolConstant(self: *SpirvCodegen, value: bool) SpirvCodegenError!u32 {
1149 const type_id = try self.getScalarType(.bool);
1150 const key = ConstKey{
1151 .type_id = type_id,
1152 .word0 = if (value) 1 else 0,
1153 .word1 = 0,
1154 .word_count = 1,
1155 };
1156 if (self.constants.get(key)) |id| return id;
1157
1158 const id = self.builder.newId();
1159 const opcode = if (value) SpirvOp.ConstantTrue else SpirvOp.ConstantFalse;
1160 try self.builder.emit(&self.builder.types, opcode, &.{ type_id, id });
1161 try self.constants.put(self.allocator, key, id);
1162 return id;
1163 }
1164 };
1165
1166 const catalog_methods = .{
1167 .emitCall = SpirvCodegen.emitCall,
1168 .emitReturn = SpirvCodegen.emitReturn,
1169 .emitScfIf = SpirvCodegen.emitScfIf,
1170 .emitScfFor = SpirvCodegen.emitScfFor,
1171 .emitScfWhile = SpirvCodegen.emitScfWhile,
1172 };
1173
1174 test "spirv catalog drives supported and unsupported emission" {
1175 var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
1176 defer ctx.deinit(std.testing.allocator);
1177
1178 const loc = ir.Location.getUnknown();
1179 const i32_type = try ArithDialect.getI32Type(&ctx);
1180 const constant = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 7);
1181
1182 var codegen = SpirvCodegen.init(std.testing.allocator);
1183 defer codegen.deinit();
1184
1185 try codegen.emitOperation(constant.op);
1186 try std.testing.expect(codegen.builder.types.items.len > 0);
1187
1188 const unsupported = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1189 try std.testing.expectError(error.UnsupportedOperation, codegen.emitOperation(unsupported.op));
1190 }
1191
1192 test "spirv codegen emits minimal header and entry point" {
1193 const testing = std.testing;
1194 const allocator = testing.allocator;
1195
1196 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1197 defer ctx.deinit(allocator);
1198
1199 const loc = ir.Location.getUnknown();
1200 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1201
1202 const memref_elem = try ArithDialect.getScalarType(&ctx, .f32);
1203 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, memref_elem, .device);
1204
1205 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
1206 try module.getBodyBlock().addOperation(func.op);
1207
1208 const entry = func.getEntryBlock();
1209 const idx_type = try ArithDialect.getIndexType(&ctx);
1210 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
1211 try entry.addOperation(gid.op);
1212
1213 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), memref_elem);
1214 try entry.addOperation(load.op);
1215
1216 const store = try MemrefDialect.StoreOp.create(&ctx, loc, load.getResult(), func.getArgument(0), gid.getResult());
1217 try entry.addOperation(store.op);
1218
1219 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1220 try entry.addOperation(ret.op);
1221
1222 var codegen = SpirvCodegen.init(allocator);
1223 defer codegen.deinit();
1224
1225 const words = try codegen.emitModuleWords(module.op);
1226 defer allocator.free(words);
1227
1228 try testing.expect(words.len > 5);
1229 try testing.expectEqual(SpirvHeader.magic, words[0]);
1230 try testing.expectEqual(SpirvHeader.version, words[1]);
1231
1232 try testing.expect(containsOpcode(words, SpirvOp.EntryPoint));
1233 try testing.expect(containsOpcode(words, SpirvOp.ExecutionMode));
1234 try testing.expect(containsOpcode(words, SpirvOp.Variable));
1235 try testing.expect(containsOpcode(words, SpirvOp.Load));
1236 try testing.expect(containsOpcode(words, SpirvOp.Store));
1237 _ = idx_type;
1238 }
1239
1240 test "spirv codegen stores bool memrefs as byte runtime arrays" {
1241 const testing = std.testing;
1242 const allocator = testing.allocator;
1243
1244 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1245 defer ctx.deinit(allocator);
1246
1247 const loc = ir.Location.getUnknown();
1248 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1249
1250 const bool_type = try ArithDialect.getScalarType(&ctx, .bool);
1251 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, bool_type, .device);
1252
1253 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel_bool_copy", &.{memref_type});
1254 try module.getBodyBlock().addOperation(func.op);
1255
1256 const entry = func.getEntryBlock();
1257 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
1258 try entry.addOperation(gid.op);
1259 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), bool_type);
1260 try entry.addOperation(load.op);
1261 const store = try MemrefDialect.StoreOp.create(&ctx, loc, load.getResult(), func.getArgument(0), gid.getResult());
1262 try entry.addOperation(store.op);
1263 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1264 try entry.addOperation(ret.op);
1265
1266 var codegen = SpirvCodegen.init(allocator);
1267 defer codegen.deinit();
1268
1269 const words = try codegen.emitModuleWords(module.op);
1270 defer allocator.free(words);
1271
1272 const bool_type_id = findTypeBool(words) orelse return error.TestExpectedBoolType;
1273 const u8_type_id = findTypeInt(words, 8, 0) orelse return error.TestExpectedU8Type;
1274 try testing.expect(containsCapability(words, SpvCapability.Int8));
1275 try testing.expect(containsCapability(words, SpvCapability.StorageBuffer8BitAccess));
1276 try testing.expect(containsExtension(words, "SPV_KHR_storage_buffer_storage_class"));
1277 try testing.expect(containsExtension(words, "SPV_KHR_8bit_storage"));
1278 try testing.expect(containsRuntimeArrayElement(words, u8_type_id));
1279 try testing.expect(!containsRuntimeArrayElement(words, bool_type_id));
1280 try testing.expect(containsOpcode(words, SpirvOp.Load));
1281 try testing.expect(containsOpcode(words, SpirvOp.INotEqual));
1282 try testing.expect(containsOpcode(words, SpirvOp.Select));
1283 try testing.expect(containsOpcode(words, SpirvOp.Store));
1284 }
1285
1286 test "spirv codegen emits memref integer atomics" {
1287 const testing = std.testing;
1288 const allocator = testing.allocator;
1289
1290 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1291 defer ctx.deinit(allocator);
1292
1293 const loc = ir.Location.getUnknown();
1294 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1295
1296 const i32_type = try ArithDialect.getI32Type(&ctx);
1297 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, i32_type, .device);
1298
1299 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel_atomic", &.{memref_type});
1300 try module.getBodyBlock().addOperation(func.op);
1301
1302 const entry = func.getEntryBlock();
1303 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
1304 try entry.addOperation(gid.op);
1305
1306 const one = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 1);
1307 try entry.addOperation(one.op);
1308 const two = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 2);
1309 try entry.addOperation(two.op);
1310
1311 const add = try MemrefDialect.AtomicRmwOp.create(
1312 &ctx,
1313 loc,
1314 .add,
1315 one.getResult(),
1316 func.getArgument(0),
1317 gid.getResult(),
1318 i32_type,
1319 );
1320 try entry.addOperation(add.op);
1321
1322 const cas = try MemrefDialect.AtomicCasOp.create(
1323 &ctx,
1324 loc,
1325 add.getResult(),
1326 two.getResult(),
1327 func.getArgument(0),
1328 gid.getResult(),
1329 i32_type,
1330 );
1331 try entry.addOperation(cas.op);
1332
1333 const store = try MemrefDialect.StoreOp.create(&ctx, loc, cas.getResult(), func.getArgument(0), gid.getResult());
1334 try entry.addOperation(store.op);
1335
1336 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1337 try entry.addOperation(ret.op);
1338
1339 var codegen = SpirvCodegen.init(allocator);
1340 defer codegen.deinit();
1341
1342 const words = try codegen.emitModuleWords(module.op);
1343 defer allocator.free(words);
1344
1345 try testing.expect(containsOpcode(words, SpirvOp.AccessChain));
1346 try testing.expect(containsOpcode(words, SpirvOp.AtomicIAdd));
1347 try testing.expect(containsOpcode(words, SpirvOp.AtomicCompareExchange));
1348 try testing.expect(containsOpcode(words, SpirvOp.Store));
1349
1350 const bytes = try wordsToBytes(allocator, words);
1351 defer allocator.free(bytes);
1352 try maybeRunSpirvVal(bytes);
1353 }
1354
1355 test "spirv codegen emits shared integer atomic rmw" {
1356 const testing = std.testing;
1357 const allocator = testing.allocator;
1358
1359 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1360 defer ctx.deinit(allocator);
1361
1362 const loc = ir.Location.getUnknown();
1363 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1364
1365 const i32_type = try ArithDialect.getI32Type(&ctx);
1366 const index_type = try ArithDialect.getIndexType(&ctx);
1367 const shared_type = try MemrefDialect.getMemrefType1D(&ctx, 4, i32_type, .shared);
1368
1369 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel_shared_atomic", &.{});
1370 try module.getBodyBlock().addOperation(func.op);
1371
1372 const entry = func.getEntryBlock();
1373 const shared_alloc = try MemrefDialect.AllocOp.createStatic(&ctx, loc, shared_type);
1374 try entry.addOperation(shared_alloc.op);
1375 const zero = try ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 0);
1376 try entry.addOperation(zero.op);
1377 const one = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 1);
1378 try entry.addOperation(one.op);
1379 const max = try MemrefDialect.AtomicRmwOp.create(
1380 &ctx,
1381 loc,
1382 .max,
1383 one.getResult(),
1384 shared_alloc.getResult(),
1385 zero.getResult(),
1386 i32_type,
1387 );
1388 try entry.addOperation(max.op);
1389
1390 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1391 try entry.addOperation(ret.op);
1392
1393 var codegen = SpirvCodegen.init(allocator);
1394 defer codegen.deinit();
1395
1396 const words = try codegen.emitModuleWords(module.op);
1397 defer allocator.free(words);
1398
1399 try testing.expect(containsOpcode(words, SpirvOp.Variable));
1400 try testing.expect(containsOpcode(words, SpirvOp.AtomicSMax));
1401
1402 const bytes = try wordsToBytes(allocator, words);
1403 defer allocator.free(bytes);
1404 try maybeRunSpirvVal(bytes);
1405 }
1406
1407 test "spirv codegen rejects memref f32 atomic add without extension support" {
1408 const testing = std.testing;
1409 const allocator = testing.allocator;
1410
1411 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1412 defer ctx.deinit(allocator);
1413
1414 const loc = ir.Location.getUnknown();
1415 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1416
1417 const f32_type = try ArithDialect.getScalarType(&ctx, .f32);
1418 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
1419
1420 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel_atomic_f32", &.{memref_type});
1421 try module.getBodyBlock().addOperation(func.op);
1422
1423 const entry = func.getEntryBlock();
1424 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
1425 try entry.addOperation(gid.op);
1426 const one = try ArithDialect.ConstantOp.createFloat(&ctx, loc, f32_type, 1.0);
1427 try entry.addOperation(one.op);
1428 const add = try MemrefDialect.AtomicRmwOp.create(
1429 &ctx,
1430 loc,
1431 .add,
1432 one.getResult(),
1433 func.getArgument(0),
1434 gid.getResult(),
1435 f32_type,
1436 );
1437 try entry.addOperation(add.op);
1438
1439 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1440 try entry.addOperation(ret.op);
1441
1442 var codegen = SpirvCodegen.init(allocator);
1443 defer codegen.deinit();
1444
1445 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
1446 }
1447
1448 test "spirv dialect serialization emits entry point and arithmetic ops" {
1449 const testing = std.testing;
1450 const allocator = testing.allocator;
1451
1452 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1453 defer ctx.deinit(allocator);
1454
1455 const loc = ir.Location.getUnknown();
1456 const module = try SpirvDialect.ModuleOp.create(
1457 &ctx,
1458 loc,
1459 .logical,
1460 .glsl450,
1461 .shader,
1462 "GLSL.std.450",
1463 );
1464 const module_block = module.getBodyBlock();
1465
1466 const i32_type = try ArithDialect.getScalarType(&ctx, .i32);
1467
1468 const global_var = try SpirvDialect.VariableOp.create(&ctx, loc, i32_type, .workgroup, null);
1469 try module_block.addOperation(global_var.op);
1470
1471 var func = try SpirvDialect.FuncOp.create(&ctx, loc, "spirv_entry", &.{i32_type}, &.{});
1472 try func.setEntryPoint(&ctx, .gl_compute);
1473 try module_block.addOperation(func.op);
1474
1475 const entry = func.getEntryBlock();
1476 const arg0 = entry.arguments.items[0];
1477
1478 const const_op = try SpirvDialect.ConstantOp.createInt(&ctx, loc, i32_type, 7);
1479 try entry.addOperation(const_op.op);
1480
1481 const add = try SpirvDialect.IAddOp.create(&ctx, loc, arg0, const_op.getResult());
1482 try entry.addOperation(add.op);
1483
1484 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1485 try entry.addOperation(ret.op);
1486
1487 var codegen = SpirvCodegen.init(allocator);
1488 defer codegen.deinit();
1489
1490 const words = try codegen.emitModuleWords(module.op);
1491 defer allocator.free(words);
1492
1493 try testing.expectEqual(SpirvHeader.magic, words[0]);
1494 try testing.expect(containsOpcode(words, SpirvOp.EntryPoint));
1495 try testing.expect(containsOpcode(words, SpirvOp.ExecutionMode));
1496 try testing.expect(containsOpcode(words, SpirvOp.TypeFunction));
1497 try testing.expect(containsOpcode(words, SpirvOp.FunctionParameter));
1498 try testing.expect(containsOpcode(words, SpirvOp.Variable));
1499 try testing.expect(containsOpcode(words, SpirvOp.Constant));
1500 try testing.expect(containsOpcode(words, SpirvOp.IAdd));
1501 try testing.expect(containsOpcode(words, SpirvOp.Return));
1502 try testing.expect(containsOpcode(words, SpirvOp.FunctionEnd));
1503 }
1504
1505 test "spirv codegen handles control flow, shared alloc, and warp ops" {
1506 const testing = std.testing;
1507 const allocator = testing.allocator;
1508
1509 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1510 defer ctx.deinit(allocator);
1511
1512 const loc = ir.Location.getUnknown();
1513 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1514
1515 const i32_type = try ArithDialect.getI32Type(&ctx);
1516 const index_type = try ArithDialect.getIndexType(&ctx);
1517 const memref_device = try MemrefDialect.getMemrefType1D(&ctx, 4, i32_type, .device);
1518 const memref_shared = try MemrefDialect.getMemrefType1D(&ctx, 4, i32_type, .shared);
1519
1520 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel_cf", &.{ memref_device, i32_type });
1521 try module.getBodyBlock().addOperation(func.op);
1522
1523 const entry = func.getEntryBlock();
1524 const buf_arg = func.getArgument(0);
1525 const scalar_arg = func.getArgument(1);
1526
1527 const tid = try GpuDialect.ThreadIdxOp.create(&ctx, loc, .x);
1528 try entry.addOperation(tid.op);
1529
1530 const load = try MemrefDialect.LoadOp.create(&ctx, loc, buf_arg, tid.getResult(), i32_type);
1531 try entry.addOperation(load.op);
1532
1533 const sum = try ArithDialect.AddOp.create(&ctx, loc, load.getResult(), scalar_arg);
1534 try entry.addOperation(sum.op);
1535
1536 const shared_alloc = try MemrefDialect.AllocOp.createStatic(&ctx, loc, memref_shared);
1537 try entry.addOperation(shared_alloc.op);
1538 const store_shared = try MemrefDialect.StoreOp.create(&ctx, loc, sum.getResult(), shared_alloc.getResult(), tid.getResult());
1539 try entry.addOperation(store_shared.op);
1540
1541 const barrier = try GpuDialect.BarrierOp.create(&ctx, loc, .block);
1542 try entry.addOperation(barrier.op);
1543
1544 const mask_const = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, -1);
1545 try entry.addOperation(mask_const.op);
1546 const reduce = try GpuDialect.WarpReduceOp.create(&ctx, loc, .add, mask_const.getResult(), sum.getResult());
1547 try entry.addOperation(reduce.op);
1548
1549 const zero = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 0);
1550 try entry.addOperation(zero.op);
1551 const cmp = try ArithDialect.CmpOp.create(&ctx, loc, .gt, reduce.getResult(), zero.getResult());
1552 try entry.addOperation(cmp.op);
1553
1554 const all_sync = try GpuDialect.AllSyncOp.create(&ctx, loc, mask_const.getResult(), cmp.getResult());
1555 try entry.addOperation(all_sync.op);
1556
1557 const ballot = try GpuDialect.BallotSyncOp.create(&ctx, loc, mask_const.getResult(), cmp.getResult());
1558 try entry.addOperation(ballot.op);
1559
1560 const shfl_delta = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 1);
1561 try entry.addOperation(shfl_delta.op);
1562 const shfl_down = try GpuDialect.ShflSyncOp.create(&ctx, loc, .down, mask_const.getResult(), sum.getResult(), shfl_delta.getResult());
1563 try entry.addOperation(shfl_down.op);
1564
1565 const shfl_xor = try GpuDialect.ShflSyncOp.create(&ctx, loc, .xor, mask_const.getResult(), sum.getResult(), shfl_delta.getResult());
1566 try entry.addOperation(shfl_xor.op);
1567
1568 var if_op = try ScfDialect.IfOp.create(&ctx, loc, cmp.getResult(), &.{i32_type});
1569 try entry.addOperation(if_op.op);
1570 const then_block = if_op.getThenBlock();
1571 const one = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 1);
1572 try then_block.addOperation(one.op);
1573 const then_yield = try ScfDialect.YieldOp.create(&ctx, loc, &.{one.getResult()});
1574 try then_block.addOperation(then_yield.op);
1575 const else_block = if_op.getElseBlock().?;
1576 const two = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 2);
1577 try else_block.addOperation(two.op);
1578 const else_yield = try ScfDialect.YieldOp.create(&ctx, loc, &.{two.getResult()});
1579 try else_block.addOperation(else_yield.op);
1580
1581 const if_result = if_op.getResult(0).?;
1582 const lo = try ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 0);
1583 try entry.addOperation(lo.op);
1584 const hi = try ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 4);
1585 try entry.addOperation(hi.op);
1586 const step = try ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 1);
1587 try entry.addOperation(step.op);
1588
1589 var for_op = try ScfDialect.ForOp.create(
1590 &ctx,
1591 loc,
1592 lo.getResult(),
1593 hi.getResult(),
1594 step.getResult(),
1595 &.{if_result},
1596 &.{i32_type},
1597 );
1598 try entry.addOperation(for_op.op);
1599
1600 const body = for_op.getBodyBlock();
1601 const iv = body.arguments.items[0];
1602 const acc = body.arguments.items[1];
1603 const iv_cast = try ArithDialect.CastOp.create(&ctx, loc, iv, i32_type);
1604 try body.addOperation(iv_cast.op);
1605 const acc_add = try ArithDialect.AddOp.create(&ctx, loc, acc, iv_cast.getResult());
1606 try body.addOperation(acc_add.op);
1607 const for_yield = try ScfDialect.YieldOp.create(&ctx, loc, &.{acc_add.getResult()});
1608 try body.addOperation(for_yield.op);
1609
1610 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1611 try entry.addOperation(ret.op);
1612
1613 var codegen = SpirvCodegen.init(allocator);
1614 defer codegen.deinit();
1615
1616 const words = try codegen.emitModuleWords(module.op);
1617 defer allocator.free(words);
1618
1619 try testing.expectEqual(SpirvVersion.v13, words[1]);
1620 try testing.expect(containsCapability(words, SpvCapability.GroupNonUniform));
1621 try testing.expect(containsCapability(words, SpvCapability.GroupNonUniformVote));
1622 try testing.expect(containsCapability(words, SpvCapability.GroupNonUniformBallot));
1623 try testing.expect(containsCapability(words, SpvCapability.GroupNonUniformShuffle));
1624 try testing.expect(containsCapability(words, SpvCapability.GroupNonUniformShuffleRelative));
1625 try testing.expect(containsCapability(words, SpvCapability.GroupNonUniformArithmetic));
1626 try testing.expect(containsOpcode(words, SpirvOp.LoopMerge));
1627 try testing.expect(loopMergesImmediatelyPrecedeBranches(words));
1628 try testing.expect(containsOpcode(words, SpirvOp.SelectionMerge));
1629 try testing.expect(containsOpcode(words, SpirvOp.ControlBarrier));
1630 try testing.expect(containsOpcode(words, SpirvOp.GroupNonUniformIAdd));
1631 try testing.expect(containsOpcode(words, SpirvOp.ULessThan));
1632 }
1633
1634 test "spirv codegen emits arith comparison opcode families" {
1635 const testing = std.testing;
1636 const allocator = testing.allocator;
1637
1638 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1639 defer ctx.deinit(allocator);
1640
1641 const loc = ir.Location.getUnknown();
1642 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1643 const bool_type = try ArithDialect.getScalarType(&ctx, .bool);
1644 const f32_type = try ArithDialect.getScalarType(&ctx, .f32);
1645 const i32_type = try ArithDialect.getScalarType(&ctx, .i32);
1646 const index_type = try ArithDialect.getScalarType(&ctx, .index);
1647 const bool_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, bool_type, .device);
1648 const f32_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
1649 const i32_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, i32_type, .device);
1650 const index_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, index_type, .device);
1651
1652 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel_cmp", &.{ bool_memref, f32_memref, i32_memref, index_memref });
1653 try module.getBodyBlock().addOperation(func.op);
1654
1655 const entry = func.getEntryBlock();
1656 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
1657 try entry.addOperation(gid.op);
1658
1659 const f32_value = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(1), gid.getResult(), f32_type);
1660 try entry.addOperation(f32_value.op);
1661 const f32_limit = try ArithDialect.ConstantOp.createFloat(&ctx, loc, f32_type, 0.0);
1662 try entry.addOperation(f32_limit.op);
1663 const f32_cmp = try ArithDialect.CmpOp.create(&ctx, loc, .lt, f32_value.getResult(), f32_limit.getResult());
1664 try entry.addOperation(f32_cmp.op);
1665 const store_f32 = try MemrefDialect.StoreOp.create(&ctx, loc, f32_cmp.getResult(), func.getArgument(0), gid.getResult());
1666 try entry.addOperation(store_f32.op);
1667
1668 const i32_value = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(2), gid.getResult(), i32_type);
1669 try entry.addOperation(i32_value.op);
1670 const i32_limit = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 7);
1671 try entry.addOperation(i32_limit.op);
1672 const i32_cmp = try ArithDialect.CmpOp.create(&ctx, loc, .lt, i32_value.getResult(), i32_limit.getResult());
1673 try entry.addOperation(i32_cmp.op);
1674 const store_i32 = try MemrefDialect.StoreOp.create(&ctx, loc, i32_cmp.getResult(), func.getArgument(0), gid.getResult());
1675 try entry.addOperation(store_i32.op);
1676
1677 const index_value = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(3), gid.getResult(), index_type);
1678 try entry.addOperation(index_value.op);
1679 const index_limit = try ArithDialect.ConstantOp.createInt(&ctx, loc, index_type, 9);
1680 try entry.addOperation(index_limit.op);
1681 const index_cmp = try ArithDialect.CmpOp.create(&ctx, loc, .lt, index_value.getResult(), index_limit.getResult());
1682 try entry.addOperation(index_cmp.op);
1683 const store_index = try MemrefDialect.StoreOp.create(&ctx, loc, index_cmp.getResult(), func.getArgument(0), gid.getResult());
1684 try entry.addOperation(store_index.op);
1685
1686 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1687 try entry.addOperation(ret.op);
1688
1689 var codegen = SpirvCodegen.init(allocator);
1690 defer codegen.deinit();
1691 const words = try codegen.emitModuleWords(module.op);
1692 defer allocator.free(words);
1693
1694 try testing.expect(containsOpcode(words, SpirvOp.FOrdLessThan));
1695 try testing.expect(containsOpcode(words, SpirvOp.SLessThan));
1696 try testing.expect(containsOpcode(words, SpirvOp.ULessThan));
1697 try testing.expect(!containsOpcode(words, SpirvOp.LogicalOr));
1698 try testing.expect(!containsOpcode(words, SpirvOp.LogicalAnd));
1699 }
1700
1701 test "spirv codegen keeps subgroup barrier headers minimal" {
1702 const testing = std.testing;
1703 const allocator = testing.allocator;
1704
1705 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1706 defer ctx.deinit(allocator);
1707
1708 const loc = ir.Location.getUnknown();
1709 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1710
1711 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "sync_warp_only", &.{});
1712 try module.getBodyBlock().addOperation(func.op);
1713
1714 const entry = func.getEntryBlock();
1715 const i32_type = try ArithDialect.getI32Type(&ctx);
1716 const mask_const = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, -1);
1717 try entry.addOperation(mask_const.op);
1718
1719 const sync = try GpuDialect.SyncWarpOp.create(&ctx, loc, mask_const.getResult());
1720 try entry.addOperation(sync.op);
1721
1722 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1723 try entry.addOperation(ret.op);
1724
1725 var codegen = SpirvCodegen.init(allocator);
1726 defer codegen.deinit();
1727
1728 const words = try codegen.emitModuleWords(module.op);
1729 defer allocator.free(words);
1730
1731 try testing.expectEqual(SpirvHeader.version, words[1]);
1732 try testing.expect(!containsCapability(words, SpvCapability.GroupNonUniform));
1733 try testing.expect(containsOpcode(words, SpirvOp.ControlBarrier));
1734 }
1735
1736 test "spirv codegen decorates bindings for vector add kernel" {
1737 const testing = std.testing;
1738 const allocator = testing.allocator;
1739
1740 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1741 defer ctx.deinit(allocator);
1742
1743 const module = try buildVecAddKernelJob(&ctx);
1744
1745 var codegen = SpirvCodegen.init(allocator);
1746 defer codegen.deinit();
1747
1748 const words = try codegen.emitModuleWords(module);
1749 defer allocator.free(words);
1750
1751 try testing.expectEqual(@as(usize, 3), countDecorations(words, SpvDecoration.Binding, null));
1752 try testing.expect(containsDecoration(words, SpvDecoration.Binding, 0));
1753 try testing.expect(containsDecoration(words, SpvDecoration.Binding, 1));
1754 try testing.expect(containsDecoration(words, SpvDecoration.Binding, 2));
1755 try testing.expectEqual(@as(usize, 3), countDecorations(words, SpvDecoration.DescriptorSet, 0));
1756 try testing.expect(containsOpcode(words, SpirvOp.IAdd));
1757 try testing.expect(containsOpcode(words, SpirvOp.Load));
1758 try testing.expect(containsOpcode(words, SpirvOp.Store));
1759
1760 const bytes = try wordsToBytes(allocator, words);
1761 defer allocator.free(bytes);
1762 try maybeRunSpirvVal(bytes);
1763 }
1764
1765 test "spirv codegen decorates bindings for reduction kernel" {
1766 const testing = std.testing;
1767 const allocator = testing.allocator;
1768
1769 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1770 defer ctx.deinit(allocator);
1771
1772 const module = try buildReductionKernelJob(&ctx);
1773
1774 var codegen = SpirvCodegen.init(allocator);
1775 defer codegen.deinit();
1776
1777 const words = try codegen.emitModuleWords(module);
1778 defer allocator.free(words);
1779
1780 try testing.expectEqual(SpirvVersion.v13, words[1]);
1781 try testing.expectEqual(@as(usize, 2), countDecorations(words, SpvDecoration.Binding, null));
1782 try testing.expect(containsDecoration(words, SpvDecoration.Binding, 0));
1783 try testing.expect(containsDecoration(words, SpvDecoration.Binding, 1));
1784 try testing.expectEqual(@as(usize, 2), countDecorations(words, SpvDecoration.DescriptorSet, 0));
1785 try testing.expect(containsOpcode(words, SpirvOp.GroupNonUniformIAdd));
1786
1787 const bytes = try wordsToBytes(allocator, words);
1788 defer allocator.free(bytes);
1789 try maybeRunSpirvVal(bytes);
1790 }
1791
1792 test "spirv codegen rejects non-full warp masks" {
1793 const testing = std.testing;
1794 const allocator = testing.allocator;
1795
1796 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1797 defer ctx.deinit(allocator);
1798
1799 const loc = ir.Location.getUnknown();
1800 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1801
1802 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "mask_reject", &.{});
1803 try module.getBodyBlock().addOperation(func.op);
1804
1805 const entry = func.getEntryBlock();
1806 const i32_type = try ArithDialect.getI32Type(&ctx);
1807 const mask = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 0);
1808 try entry.addOperation(mask.op);
1809 const pred = try ArithDialect.ConstantOp.createBool(&ctx, loc, true);
1810 try entry.addOperation(pred.op);
1811
1812 const all_sync = try GpuDialect.AllSyncOp.create(&ctx, loc, mask.getResult(), pred.getResult());
1813 try entry.addOperation(all_sync.op);
1814
1815 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1816 try entry.addOperation(ret.op);
1817
1818 var codegen = SpirvCodegen.init(allocator);
1819 defer codegen.deinit();
1820
1821 try testing.expectError(error.UnsupportedMask, codegen.emitModuleWords(module.op));
1822 }
1823
1824 const UnaryExtInstCase = struct {
1825 name: []const u8,
1826 opcode: u32,
1827 emit: *const fn (ctx: *ir.Context, entry: *ir.Block, x: *ir.Value) anyerror!*ir.Value,
1828 };
1829
1830 const ScalarCapabilityCase = struct {
1831 elem: ArithDialect.ScalarTypeKind,
1832 cap: u32,
1833 };
1834
1835 fn unary_test_body(comptime Op: type) type {
1836 return struct {
1837 fn emit(ctx: *ir.Context, entry: *ir.Block, x: *ir.Value) !*ir.Value {
1838 const loc = ir.Location.getUnknown();
1839 const op = try Op.create(ctx, loc, x);
1840 try entry.addOperation(op.op);
1841 return op.getResult();
1842 }
1843 };
1844 }
1845
1846 fn integer_constant_test_body(comptime Op: type, comptime value: i64) type {
1847 return struct {
1848 fn emit(ctx: *ir.Context, entry: *ir.Block, x: *ir.Value) !*ir.Value {
1849 const loc = ir.Location.getUnknown();
1850 const i32_type = try ArithDialect.getScalarType(ctx, .i32);
1851 const constant = try ArithDialect.ConstantOp.createInt(ctx, loc, i32_type, value);
1852 try entry.addOperation(constant.op);
1853 const op = try Op.create(ctx, loc, x, constant.getResult());
1854 try entry.addOperation(op.op);
1855 return op.getResult();
1856 }
1857 };
1858 }
1859
1860 fn emit_test_add(ctx: *ir.Context, entry: *ir.Block, x: *ir.Value) !*ir.Value {
1861 const loc = ir.Location.getUnknown();
1862 const add = try ArithDialect.AddOp.create(ctx, loc, x, x);
1863 try entry.addOperation(add.op);
1864 return add.getResult();
1865 }
1866
1867 fn emit_test_fma_f32(ctx: *ir.Context, entry: *ir.Block, x: *ir.Value) !*ir.Value {
1868 const loc = ir.Location.getUnknown();
1869 const f32_type = try ArithDialect.getScalarType(ctx, .f32);
1870 const a = try ArithDialect.ConstantOp.createFloat(ctx, loc, f32_type, 2.0);
1871 try entry.addOperation(a.op);
1872 const b = try ArithDialect.ConstantOp.createFloat(ctx, loc, f32_type, 3.0);
1873 try entry.addOperation(b.op);
1874 const fma = try ArithDialect.FmaOp.create(ctx, loc, a.getResult(), b.getResult(), x);
1875 try entry.addOperation(fma.op);
1876 return fma.getResult();
1877 }
1878
1879 fn emit_test_fma_f16(ctx: *ir.Context, entry: *ir.Block, x: *ir.Value) !*ir.Value {
1880 const loc = ir.Location.getUnknown();
1881 const fma = try ArithDialect.FmaOp.create(ctx, loc, x, x, x);
1882 try entry.addOperation(fma.op);
1883 return fma.getResult();
1884 }
1885
1886 fn emit_test_fma_f64(ctx: *ir.Context, entry: *ir.Block, x: *ir.Value) !*ir.Value {
1887 const loc = ir.Location.getUnknown();
1888 const f64_type = try ArithDialect.getScalarType(ctx, .f64);
1889 const a = try ArithDialect.ConstantOp.createFloat(ctx, loc, f64_type, 1.0);
1890 try entry.addOperation(a.op);
1891 const b = try ArithDialect.ConstantOp.createFloat(ctx, loc, f64_type, 2.0);
1892 try entry.addOperation(b.op);
1893 const fma = try ArithDialect.FmaOp.create(ctx, loc, a.getResult(), b.getResult(), x);
1894 try entry.addOperation(fma.op);
1895 return fma.getResult();
1896 }
1897
1898 fn emitArithKernelWords(
1899 allocator: std.mem.Allocator,
1900 elem_kind: ArithDialect.ScalarTypeKind,
1901 body: anytype,
1902 ) ![]u32 {
1903 return emitArithKernelWordsWithControls(allocator, elem_kind, body, .{});
1904 }
1905
1906 fn emitArithKernelWordsWithControls(
1907 allocator: std.mem.Allocator,
1908 elem_kind: ArithDialect.ScalarTypeKind,
1909 body: anytype,
1910 controls: FloatControls,
1911 ) ![]u32 {
1912 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1913 defer ctx.deinit(allocator);
1914
1915 const loc = ir.Location.getUnknown();
1916 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
1917
1918 const elem_type = try ArithDialect.getScalarType(&ctx, elem_kind);
1919 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, elem_type, .device);
1920
1921 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
1922 try module.getBodyBlock().addOperation(func.op);
1923
1924 const entry = func.getEntryBlock();
1925 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
1926 try entry.addOperation(gid.op);
1927
1928 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), elem_type);
1929 try entry.addOperation(load.op);
1930
1931 const stored = try body(&ctx, entry, load.getResult());
1932
1933 const store = try MemrefDialect.StoreOp.create(&ctx, loc, stored, func.getArgument(0), gid.getResult());
1934 try entry.addOperation(store.op);
1935
1936 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
1937 try entry.addOperation(ret.op);
1938
1939 var codegen = SpirvCodegen.init(allocator);
1940 defer codegen.deinit();
1941 codegen.float_controls = controls;
1942 return codegen.emitModuleWords(module.op);
1943 }
1944
1945 fn isDecorated(words: []const u32, result_id: u32, decoration: u32) bool {
1946 var i: usize = 5;
1947 while (i < words.len) {
1948 const count: usize = @intCast(words[i] >> 16);
1949 if (count == 0 or i + count > words.len) return false;
1950 if (@as(u16, @truncate(words[i])) == SpirvOp.Decorate and count >= 3 and
1951 words[i + 1] == result_id and words[i + 2] == decoration) return true;
1952 i += count;
1953 }
1954 return false;
1955 }
1956
1957 fn hasFloatExecutionMode(words: []const u32, mode: u32, width: u32) bool {
1958 var i: usize = 5;
1959 while (i < words.len) {
1960 const count: usize = @intCast(words[i] >> 16);
1961 if (count == 0 or i + count > words.len) return false;
1962 if (@as(u16, @truncate(words[i])) == SpirvOp.ExecutionMode and count == 4 and
1963 words[i + 2] == mode and words[i + 3] == width) return true;
1964 i += count;
1965 }
1966 return false;
1967 }
1968
1969 test "spirv float arithmetic carries NoContraction and selected execution modes" {
1970 const testing = std.testing;
1971 const allocator = testing.allocator;
1972 const controls = FloatControls{
1973 .denorm_preserve = .{ .f32 = true },
1974 .signed_zero_inf_nan_preserve = .{ .f32 = true },
1975 };
1976 const words = try emitArithKernelWordsWithControls(allocator, .f32, emit_test_add, controls);
1977 defer allocator.free(words);
1978 const add_at = firstOpcodeIndex(words, SpirvOp.FAdd) orelse return error.TestExpectedFloatAdd;
1979 try testing.expect(isDecorated(words, words[add_at + 2], SpvDecoration.NoContraction));
1980 try testing.expect(hasFloatExecutionMode(words, SpvExecutionMode.DenormPreserve, 32));
1981 try testing.expect(hasFloatExecutionMode(words, SpvExecutionMode.SignedZeroInfNanPreserve, 32));
1982 try testing.expect(!hasFloatExecutionMode(words, SpvExecutionMode.DenormPreserve, 16));
1983 try testing.expect(containsCapability(words, SpvCapability.DenormPreserve));
1984 try testing.expect(containsCapability(words, SpvCapability.SignedZeroInfNanPreserve));
1985
1986 const fma = try emitArithKernelWords(allocator, .f32, emit_test_fma_f32);
1987 defer allocator.free(fma);
1988 const ext_at = firstOpcodeIndex(fma, SpirvOp.ExtInst) orelse return error.TestExpectedExtInst;
1989 try testing.expect(isDecorated(fma, fma[ext_at + 2], SpvDecoration.NoContraction));
1990
1991 const ints = try emitArithKernelWords(allocator, .i32, emit_test_add);
1992 defer allocator.free(ints);
1993 const int_at = firstOpcodeIndex(ints, SpirvOp.IAdd) orelse return error.TestExpectedIntegerAdd;
1994 try testing.expect(!isDecorated(ints, ints[int_at + 2], SpvDecoration.NoContraction));
1995 }
1996
1997 fn findExtInstImportId(words: []const u32, set_name: []const u8) ?u32 {
1998 if (words.len < 6) return null;
1999 var i: usize = 5;
2000 while (i < words.len) {
2001 const word = words[i];
2002 const word_count: usize = @intCast(word >> 16);
2003 const op = @as(u16, @truncate(word));
2004 if (op == SpirvOp.ExtInstImport and word_count >= 3 and i + word_count <= words.len) {
2005 const id = words[i + 1];
2006 const string_words = word_count - 2;
2007 var matches = true;
2008 var name_idx: usize = 0;
2009 var w: usize = 0;
2010 outer: while (w < string_words) : (w += 1) {
2011 const word_value = words[i + 2 + w];
2012 var byte_idx: u5 = 0;
2013 while (byte_idx < 4) : (byte_idx += 1) {
2014 const byte: u8 = @truncate(word_value >> (8 * @as(u5, byte_idx)));
2015 if (byte == 0) {
2016 if (name_idx != set_name.len) matches = false;
2017 break :outer;
2018 }
2019 if (name_idx >= set_name.len or set_name[name_idx] != byte) {
2020 matches = false;
2021 break :outer;
2022 }
2023 name_idx += 1;
2024 }
2025 }
2026 if (matches) return id;
2027 }
2028 if (word_count == 0) break;
2029 i += word_count;
2030 }
2031 return null;
2032 }
2033
2034 fn containsExtInst(words: []const u32, set_id: u32, ext_opcode: u32) bool {
2035 if (words.len < 6) return false;
2036 var i: usize = 5;
2037 while (i < words.len) {
2038 const word = words[i];
2039 const word_count: usize = @intCast(word >> 16);
2040 const op = @as(u16, @truncate(word));
2041 if (op == SpirvOp.ExtInst and word_count >= 5 and i + word_count <= words.len) {
2042 if (words[i + 3] == set_id and words[i + 4] == ext_opcode) return true;
2043 }
2044 if (word_count == 0) break;
2045 i += word_count;
2046 }
2047 return false;
2048 }
2049
2050 fn firstOpcodeIndex(words: []const u32, opcode: u16) ?usize {
2051 if (words.len < 5) return null;
2052 var i: usize = 5;
2053 while (i < words.len) {
2054 const word = words[i];
2055 const word_count = word >> 16;
2056 const op = @as(u16, @truncate(word));
2057 if (op == opcode) return i;
2058 if (word_count == 0) break;
2059 i += word_count;
2060 }
2061 return null;
2062 }
2063
2064 test "spirv codegen emits arith.neg via core SNegate / FNegate" {
2065 const testing = std.testing;
2066 const allocator = testing.allocator;
2067
2068 const f32_words = try emitArithKernelWords(allocator, .f32, unary_test_body(ArithDialect.NegOp).emit);
2069 defer allocator.free(f32_words);
2070 try testing.expect(containsOpcode(f32_words, SpirvOp.FNegate));
2071
2072 const i32_words = try emitArithKernelWords(allocator, .i32, unary_test_body(ArithDialect.NegOp).emit);
2073 defer allocator.free(i32_words);
2074 try testing.expect(containsOpcode(i32_words, SpirvOp.SNegate));
2075 }
2076
2077 test "spirv codegen emits arith.neg over unsigned integers as zero subtract" {
2078 const testing = std.testing;
2079 const allocator = testing.allocator;
2080
2081 const words = try emitArithKernelWords(allocator, .u32, unary_test_body(ArithDialect.NegOp).emit);
2082 defer allocator.free(words);
2083 try testing.expect(containsOpcode(words, SpirvOp.ISub));
2084 try testing.expect(!containsOpcode(words, SpirvOp.SNegate));
2085 try testing.expect(!containsOpcode(words, SpirvOp.FNegate));
2086 }
2087
2088 test "spirv codegen emits arith.exp / log / tanh via GLSL.std.450" {
2089 const testing = std.testing;
2090 const allocator = testing.allocator;
2091
2092 const cases: []const UnaryExtInstCase = &.{
2093 .{
2094 .name = "exp",
2095 .opcode = GLSLstd450.Exp,
2096 .emit = unary_test_body(ArithDialect.ExpOp).emit,
2097 },
2098 .{
2099 .name = "log",
2100 .opcode = GLSLstd450.Log,
2101 .emit = unary_test_body(ArithDialect.LogOp).emit,
2102 },
2103 .{
2104 .name = "tanh",
2105 .opcode = GLSLstd450.Tanh,
2106 .emit = unary_test_body(ArithDialect.TanhOp).emit,
2107 },
2108 };
2109
2110 for (cases) |case| {
2111 const words = try emitArithKernelWords(allocator, .f32, case.emit);
2112 defer allocator.free(words);
2113
2114 const set_id = findExtInstImportId(words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2115 try testing.expect(containsExtInst(words, set_id, case.opcode));
2116
2117 const import_idx = firstOpcodeIndex(words, SpirvOp.ExtInstImport) orelse return error.TestExpectedImport;
2118 const mem_model_idx = firstOpcodeIndex(words, SpirvOp.MemoryModel) orelse return error.TestExpectedMemoryModel;
2119 try testing.expect(import_idx < mem_model_idx);
2120 }
2121 }
2122
2123 test "spirv codegen rejects arith.exp on f64 (GLSL.std.450 requires 16/32-bit)" {
2124 const testing = std.testing;
2125 const allocator = testing.allocator;
2126
2127 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2128 defer ctx.deinit(allocator);
2129
2130 const loc = ir.Location.getUnknown();
2131 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2132
2133 const f64_type = try ArithDialect.getScalarType(&ctx, .f64);
2134 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f64_type, .device);
2135
2136 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
2137 try module.getBodyBlock().addOperation(func.op);
2138
2139 const entry = func.getEntryBlock();
2140 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2141 try entry.addOperation(gid.op);
2142 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), f64_type);
2143 try entry.addOperation(load.op);
2144 const exp = try ArithDialect.ExpOp.create(&ctx, loc, load.getResult());
2145 try entry.addOperation(exp.op);
2146 const store = try MemrefDialect.StoreOp.create(&ctx, loc, exp.getResult(), func.getArgument(0), gid.getResult());
2147 try entry.addOperation(store.op);
2148 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2149 try entry.addOperation(ret.op);
2150
2151 var codegen = SpirvCodegen.init(allocator);
2152 defer codegen.deinit();
2153 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2154 }
2155
2156 test "spirv codegen emits arith.max / arith.min via GLSL.std.450" {
2157 const testing = std.testing;
2158 const allocator = testing.allocator;
2159
2160 const f32_max_words = try emitMinMaxKernel(allocator, .f32, .max);
2161 defer allocator.free(f32_max_words);
2162 {
2163 const set_id = findExtInstImportId(f32_max_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2164 try testing.expect(containsExtInst(f32_max_words, set_id, GLSLstd450.NMax));
2165 try testing.expect(!containsExtInst(f32_max_words, set_id, GLSLstd450.FMax));
2166 }
2167
2168 const f32_min_words = try emitMinMaxKernel(allocator, .f32, .min);
2169 defer allocator.free(f32_min_words);
2170 {
2171 const set_id = findExtInstImportId(f32_min_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2172 try testing.expect(containsExtInst(f32_min_words, set_id, GLSLstd450.NMin));
2173 try testing.expect(!containsExtInst(f32_min_words, set_id, GLSLstd450.FMin));
2174 }
2175
2176 const i32_max_words = try emitMinMaxKernel(allocator, .i32, .max);
2177 defer allocator.free(i32_max_words);
2178 {
2179 const set_id = findExtInstImportId(i32_max_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2180 try testing.expect(containsExtInst(i32_max_words, set_id, GLSLstd450.SMax));
2181 }
2182
2183 const idx_min_words = try emitMinMaxKernel(allocator, .index, .min);
2184 defer allocator.free(idx_min_words);
2185 {
2186 const set_id = findExtInstImportId(idx_min_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2187 try testing.expect(containsExtInst(idx_min_words, set_id, GLSLstd450.UMin));
2188 }
2189 }
2190
2191 test "spirv codegen does not import GLSL.std.450 unless an ext-inst op is used" {
2192 const testing = std.testing;
2193 const allocator = testing.allocator;
2194
2195 const words = try emitArithKernelWords(allocator, .f32, emit_test_add);
2196 defer allocator.free(words);
2197
2198 try testing.expectEqual(@as(?usize, null), firstOpcodeIndex(words, SpirvOp.ExtInstImport));
2199 try testing.expectEqual(@as(?u32, null), findExtInstImportId(words, "GLSL.std.450"));
2200 }
2201
2202 test "spirv codegen requires Float64 capability for f64 kernels" {
2203 const testing = std.testing;
2204 const allocator = testing.allocator;
2205
2206 const words = try emitArithKernelWords(allocator, .f64, emit_test_add);
2207 defer allocator.free(words);
2208
2209 try testing.expect(containsCapability(words, SpvCapability.Float64));
2210 try testing.expect(!containsCapability(words, SpvCapability.Float16));
2211 try testing.expect(!containsCapability(words, SpvCapability.Int64));
2212 }
2213
2214 test "spirv codegen requires Float16 capability for f16 kernels" {
2215 const testing = std.testing;
2216 const allocator = testing.allocator;
2217
2218 const words = try emitArithKernelWords(allocator, .f16, emit_test_add);
2219 defer allocator.free(words);
2220
2221 try testing.expect(containsCapability(words, SpvCapability.Float16));
2222 try testing.expect(!containsCapability(words, SpvCapability.Float64));
2223 }
2224
2225 test "spirv codegen requires Int8 / Int16 / Int64 capabilities" {
2226 const testing = std.testing;
2227 const allocator = testing.allocator;
2228
2229 const cases: []const ScalarCapabilityCase = &.{
2230 .{ .elem = .i8, .cap = SpvCapability.Int8 },
2231 .{ .elem = .i16, .cap = SpvCapability.Int16 },
2232 .{ .elem = .i64, .cap = SpvCapability.Int64 },
2233 };
2234
2235 for (cases) |case| {
2236 const words = try emitArithKernelWords(allocator, case.elem, emit_test_add);
2237 defer allocator.free(words);
2238 try testing.expect(containsCapability(words, case.cap));
2239 }
2240 }
2241
2242 test "spirv codegen does not require non-32-bit capabilities for f32 / i32 / index kernels" {
2243 const testing = std.testing;
2244 const allocator = testing.allocator;
2245
2246 const cases = [_]ArithDialect.ScalarTypeKind{ .f32, .i32, .index };
2247 for (cases) |elem| {
2248 const words = try emitArithKernelWords(allocator, elem, emit_test_add);
2249 defer allocator.free(words);
2250 try testing.expect(!containsCapability(words, SpvCapability.Float16));
2251 try testing.expect(!containsCapability(words, SpvCapability.Float64));
2252 try testing.expect(!containsCapability(words, SpvCapability.Int8));
2253 try testing.expect(!containsCapability(words, SpvCapability.Int16));
2254 try testing.expect(!containsCapability(words, SpvCapability.Int64));
2255 }
2256 }
2257
2258 test "spirv codegen emits arith.sqrt over f32 / f64 via GLSL.std.450" {
2259 const testing = std.testing;
2260 const allocator = testing.allocator;
2261
2262 const f32_words = try emitArithKernelWords(allocator, .f32, unary_test_body(ArithDialect.SqrtOp).emit);
2263 defer allocator.free(f32_words);
2264 {
2265 const set_id = findExtInstImportId(f32_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2266 try testing.expect(containsExtInst(f32_words, set_id, GLSLstd450.Sqrt));
2267 }
2268
2269 const f64_words = try emitArithKernelWords(allocator, .f64, unary_test_body(ArithDialect.SqrtOp).emit);
2270 defer allocator.free(f64_words);
2271 {
2272 try testing.expect(containsCapability(f64_words, SpvCapability.Float64));
2273 const set_id = findExtInstImportId(f64_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2274 try testing.expect(containsExtInst(f64_words, set_id, GLSLstd450.Sqrt));
2275 }
2276 }
2277
2278 test "spirv codegen emits arith.abs as FAbs / SAbs and aliases unsigned" {
2279 const testing = std.testing;
2280 const allocator = testing.allocator;
2281
2282 const f32_words = try emitArithKernelWords(allocator, .f32, unary_test_body(ArithDialect.AbsOp).emit);
2283 defer allocator.free(f32_words);
2284 {
2285 const set_id = findExtInstImportId(f32_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2286 try testing.expect(containsExtInst(f32_words, set_id, GLSLstd450.FAbs));
2287 try testing.expect(!containsExtInst(f32_words, set_id, GLSLstd450.SAbs));
2288 }
2289
2290 const i32_words = try emitArithKernelWords(allocator, .i32, unary_test_body(ArithDialect.AbsOp).emit);
2291 defer allocator.free(i32_words);
2292 {
2293 const set_id = findExtInstImportId(i32_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2294 try testing.expect(containsExtInst(i32_words, set_id, GLSLstd450.SAbs));
2295 try testing.expect(!containsExtInst(i32_words, set_id, GLSLstd450.FAbs));
2296 }
2297
2298 const u32_words = try emitArithKernelWords(allocator, .u32, unary_test_body(ArithDialect.AbsOp).emit);
2299 defer allocator.free(u32_words);
2300 try testing.expect(!containsOpcode(u32_words, SpirvOp.ExtInst));
2301 }
2302
2303 test "spirv codegen emits arith.exp over f16 (Float16 capability now wired)" {
2304 const testing = std.testing;
2305 const allocator = testing.allocator;
2306
2307 const words = try emitArithKernelWords(allocator, .f16, unary_test_body(ArithDialect.ExpOp).emit);
2308 defer allocator.free(words);
2309
2310 try testing.expect(containsCapability(words, SpvCapability.Float16));
2311 const set_id = findExtInstImportId(words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2312 try testing.expect(containsExtInst(words, set_id, GLSLstd450.Exp));
2313 }
2314
2315 test "spirv codegen emits arith.sin / cos / tan / floor / trunc via GLSL.std.450" {
2316 const testing = std.testing;
2317 const allocator = testing.allocator;
2318
2319 const cases: []const UnaryExtInstCase = &.{
2320 .{
2321 .opcode = GLSLstd450.Sin,
2322 .name = "sin",
2323 .emit = unary_test_body(ArithDialect.SinOp).emit,
2324 },
2325 .{
2326 .opcode = GLSLstd450.Cos,
2327 .name = "cos",
2328 .emit = unary_test_body(ArithDialect.CosOp).emit,
2329 },
2330 .{
2331 .opcode = GLSLstd450.Tan,
2332 .name = "tan",
2333 .emit = unary_test_body(ArithDialect.TanOp).emit,
2334 },
2335 .{
2336 .opcode = GLSLstd450.Floor,
2337 .name = "floor",
2338 .emit = unary_test_body(ArithDialect.FloorOp).emit,
2339 },
2340 .{
2341 .opcode = GLSLstd450.Trunc,
2342 .name = "trunc",
2343 .emit = unary_test_body(ArithDialect.TruncOp).emit,
2344 },
2345 };
2346
2347 for (cases) |case| {
2348 const f32_words = try emitArithKernelWords(allocator, .f32, case.emit);
2349 defer allocator.free(f32_words);
2350 const set_id = findExtInstImportId(f32_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2351 try testing.expect(containsExtInst(f32_words, set_id, case.opcode));
2352
2353 const f16_words = try emitArithKernelWords(allocator, .f16, case.emit);
2354 defer allocator.free(f16_words);
2355 try testing.expect(containsCapability(f16_words, SpvCapability.Float16));
2356 const f16_set_id = findExtInstImportId(f16_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2357 try testing.expect(containsExtInst(f16_words, f16_set_id, case.opcode));
2358 }
2359 }
2360
2361 test "spirv codegen emits arith.round as sign-preserving floor composition" {
2362 const testing = std.testing;
2363 const allocator = testing.allocator;
2364
2365 const f32_words = try emitArithKernelWords(allocator, .f32, unary_test_body(ArithDialect.RoundOp).emit);
2366 defer allocator.free(f32_words);
2367 const set_id = findExtInstImportId(f32_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2368 try testing.expect(containsExtInst(f32_words, set_id, GLSLstd450.FAbs));
2369 try testing.expect(containsExtInst(f32_words, set_id, GLSLstd450.Floor));
2370 try testing.expect(containsOpcode(f32_words, SpirvOp.BitwiseAnd));
2371 try testing.expect(containsOpcode(f32_words, SpirvOp.BitwiseOr));
2372 try testing.expect(containsOpcode(f32_words, SpirvOp.Bitcast));
2373
2374 const f16_words = try emitArithKernelWords(allocator, .f16, unary_test_body(ArithDialect.RoundOp).emit);
2375 defer allocator.free(f16_words);
2376 try testing.expect(containsCapability(f16_words, SpvCapability.Float16));
2377 try testing.expect(containsCapability(f16_words, SpvCapability.Int16));
2378 const f16_set_id = findExtInstImportId(f16_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2379 try testing.expect(containsExtInst(f16_words, f16_set_id, GLSLstd450.FAbs));
2380 try testing.expect(containsExtInst(f16_words, f16_set_id, GLSLstd450.Floor));
2381 try testing.expect(containsOpcode(f16_words, SpirvOp.BitwiseAnd));
2382 try testing.expect(containsOpcode(f16_words, SpirvOp.BitwiseOr));
2383 try testing.expect(containsOpcode(f16_words, SpirvOp.Bitcast));
2384 }
2385
2386 test "spirv codegen rejects arith.sin on f64" {
2387 const testing = std.testing;
2388 const allocator = testing.allocator;
2389
2390 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2391 defer ctx.deinit(allocator);
2392
2393 const loc = ir.Location.getUnknown();
2394 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2395 const f64_type = try ArithDialect.getScalarType(&ctx, .f64);
2396 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f64_type, .device);
2397 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
2398 try module.getBodyBlock().addOperation(func.op);
2399
2400 const entry = func.getEntryBlock();
2401 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2402 try entry.addOperation(gid.op);
2403 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), f64_type);
2404 try entry.addOperation(load.op);
2405 const sin = try ArithDialect.SinOp.create(&ctx, loc, load.getResult());
2406 try entry.addOperation(sin.op);
2407 const store = try MemrefDialect.StoreOp.create(&ctx, loc, sin.getResult(), func.getArgument(0), gid.getResult());
2408 try entry.addOperation(store.op);
2409 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2410 try entry.addOperation(ret.op);
2411
2412 var codegen = SpirvCodegen.init(allocator);
2413 defer codegen.deinit();
2414 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2415 }
2416
2417 test "spirv codegen emits arith.pow and arith.atan2 as binary OpExtInst" {
2418 const testing = std.testing;
2419 const allocator = testing.allocator;
2420
2421 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2422 defer ctx.deinit(allocator);
2423
2424 const loc = ir.Location.getUnknown();
2425 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2426 const f32_type = try ArithDialect.getScalarType(&ctx, .f32);
2427 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
2428
2429 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{ memref_type, memref_type });
2430 try module.getBodyBlock().addOperation(func.op);
2431
2432 const entry = func.getEntryBlock();
2433 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2434 try entry.addOperation(gid.op);
2435 const base_load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), f32_type);
2436 try entry.addOperation(base_load.op);
2437 const exp_load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(1), gid.getResult(), f32_type);
2438 try entry.addOperation(exp_load.op);
2439 const pow = try ArithDialect.PowOp.create(&ctx, loc, base_load.getResult(), exp_load.getResult());
2440 try entry.addOperation(pow.op);
2441 const atan2 = try ArithDialect.Atan2Op.create(&ctx, loc, pow.getResult(), exp_load.getResult());
2442 try entry.addOperation(atan2.op);
2443 const store = try MemrefDialect.StoreOp.create(&ctx, loc, atan2.getResult(), func.getArgument(0), gid.getResult());
2444 try entry.addOperation(store.op);
2445 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2446 try entry.addOperation(ret.op);
2447
2448 var codegen = SpirvCodegen.init(allocator);
2449 defer codegen.deinit();
2450 const words = try codegen.emitModuleWords(module.op);
2451 defer allocator.free(words);
2452
2453 const set_id = findExtInstImportId(words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2454 try testing.expect(containsExtInst(words, set_id, GLSLstd450.Pow));
2455 try testing.expect(containsExtInst(words, set_id, GLSLstd450.Atan2));
2456
2457 const wc = extInstWordCount(words, set_id, GLSLstd450.Pow) orelse return error.TestExpectedExtInst;
2458 try testing.expectEqual(@as(u32, 7), wc);
2459 const atan2_wc = extInstWordCount(words, set_id, GLSLstd450.Atan2) orelse return error.TestExpectedExtInst;
2460 try testing.expectEqual(@as(u32, 7), atan2_wc);
2461 }
2462
2463 test "spirv codegen rejects arith.pow on f64" {
2464 const testing = std.testing;
2465 const allocator = testing.allocator;
2466
2467 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2468 defer ctx.deinit(allocator);
2469
2470 const loc = ir.Location.getUnknown();
2471 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2472 const f64_type = try ArithDialect.getScalarType(&ctx, .f64);
2473 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f64_type, .device);
2474
2475 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{ memref_type, memref_type });
2476 try module.getBodyBlock().addOperation(func.op);
2477
2478 const entry = func.getEntryBlock();
2479 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2480 try entry.addOperation(gid.op);
2481 const a = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), f64_type);
2482 try entry.addOperation(a.op);
2483 const b = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(1), gid.getResult(), f64_type);
2484 try entry.addOperation(b.op);
2485 const pow = try ArithDialect.PowOp.create(&ctx, loc, a.getResult(), b.getResult());
2486 try entry.addOperation(pow.op);
2487 const store = try MemrefDialect.StoreOp.create(&ctx, loc, pow.getResult(), func.getArgument(0), gid.getResult());
2488 try entry.addOperation(store.op);
2489 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2490 try entry.addOperation(ret.op);
2491
2492 var codegen = SpirvCodegen.init(allocator);
2493 defer codegen.deinit();
2494 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2495 }
2496
2497 test "spirv codegen emits arith.abs over f64 with Float64 capability" {
2498 const testing = std.testing;
2499 const allocator = testing.allocator;
2500
2501 const words = try emitArithKernelWords(allocator, .f64, unary_test_body(ArithDialect.AbsOp).emit);
2502 defer allocator.free(words);
2503
2504 try testing.expect(containsCapability(words, SpvCapability.Float64));
2505 const set_id = findExtInstImportId(words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2506 try testing.expect(containsExtInst(words, set_id, GLSLstd450.FAbs));
2507 }
2508
2509 test "spirv codegen emits arith.and / or / xor via OpBitwise{And,Or,Xor}" {
2510 const testing = std.testing;
2511 const allocator = testing.allocator;
2512
2513 const and_words = try emitArithKernelWords(
2514 allocator,
2515 .i32,
2516 integer_constant_test_body(ArithDialect.AndOp, 1).emit,
2517 );
2518 defer allocator.free(and_words);
2519 try testing.expect(containsOpcode(and_words, SpirvOp.BitwiseAnd));
2520
2521 const or_words = try emitArithKernelWords(
2522 allocator,
2523 .i32,
2524 integer_constant_test_body(ArithDialect.OrOp, 1).emit,
2525 );
2526 defer allocator.free(or_words);
2527 try testing.expect(containsOpcode(or_words, SpirvOp.BitwiseOr));
2528
2529 const xor_words = try emitArithKernelWords(
2530 allocator,
2531 .i32,
2532 integer_constant_test_body(ArithDialect.XorOp, 1).emit,
2533 );
2534 defer allocator.free(xor_words);
2535 try testing.expect(containsOpcode(xor_words, SpirvOp.BitwiseXor));
2536 }
2537
2538 test "spirv codegen emits arith.not as OpNot" {
2539 const testing = std.testing;
2540 const allocator = testing.allocator;
2541
2542 const words = try emitArithKernelWords(allocator, .i32, unary_test_body(ArithDialect.NotOp).emit);
2543 defer allocator.free(words);
2544 try testing.expect(containsOpcode(words, SpirvOp.Not));
2545 }
2546
2547 test "spirv codegen emits arith.not over bool as OpLogicalNot" {
2548 const testing = std.testing;
2549 const allocator = testing.allocator;
2550
2551 const words = try emitArithKernelWords(allocator, .bool, unary_test_body(ArithDialect.NotOp).emit);
2552 defer allocator.free(words);
2553 try testing.expect(containsOpcode(words, SpirvOp.LogicalNot));
2554 }
2555
2556 test "spirv codegen emits arith.shl / shr / ushr via OpShift{LeftLogical,RightArithmetic,RightLogical}" {
2557 const testing = std.testing;
2558 const allocator = testing.allocator;
2559
2560 const shl_words = try emitArithKernelWords(
2561 allocator,
2562 .i32,
2563 integer_constant_test_body(ArithDialect.ShlOp, 2).emit,
2564 );
2565 defer allocator.free(shl_words);
2566 try testing.expect(containsOpcode(shl_words, SpirvOp.ShiftLeftLogical));
2567
2568 const shr_words = try emitArithKernelWords(
2569 allocator,
2570 .i32,
2571 integer_constant_test_body(ArithDialect.ShrOp, 2).emit,
2572 );
2573 defer allocator.free(shr_words);
2574 try testing.expect(containsOpcode(shr_words, SpirvOp.ShiftRightArithmetic));
2575
2576 const ushr_words = try emitArithKernelWords(
2577 allocator,
2578 .i32,
2579 integer_constant_test_body(ArithDialect.UshrOp, 2).emit,
2580 );
2581 defer allocator.free(ushr_words);
2582 try testing.expect(containsOpcode(ushr_words, SpirvOp.ShiftRightLogical));
2583 }
2584
2585 test "spirv codegen rejects arith.shr over arith.index (unsigned + arithmetic shift mismatch)" {
2586 const testing = std.testing;
2587 const allocator = testing.allocator;
2588
2589 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2590 defer ctx.deinit(allocator);
2591
2592 const loc = ir.Location.getUnknown();
2593 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2594 const idx_type = try ArithDialect.getIndexType(&ctx);
2595 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, idx_type, .device);
2596
2597 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
2598 try module.getBodyBlock().addOperation(func.op);
2599
2600 const entry = func.getEntryBlock();
2601 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2602 try entry.addOperation(gid.op);
2603 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), idx_type);
2604 try entry.addOperation(load.op);
2605 const c = try ArithDialect.ConstantOp.createInt(&ctx, loc, idx_type, 2);
2606 try entry.addOperation(c.op);
2607 const shr = try ArithDialect.ShrOp.create(&ctx, loc, load.getResult(), c.getResult());
2608 try entry.addOperation(shr.op);
2609 const store = try MemrefDialect.StoreOp.create(&ctx, loc, shr.getResult(), func.getArgument(0), gid.getResult());
2610 try entry.addOperation(store.op);
2611 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2612 try entry.addOperation(ret.op);
2613
2614 var codegen = SpirvCodegen.init(allocator);
2615 defer codegen.deinit();
2616 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2617 }
2618
2619 test "spirv codegen rejects arith.and over float (bitwise needs arith.bitcast first)" {
2620 const testing = std.testing;
2621 const allocator = testing.allocator;
2622
2623 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2624 defer ctx.deinit(allocator);
2625
2626 const loc = ir.Location.getUnknown();
2627 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2628 const f32_type = try ArithDialect.getScalarType(&ctx, .f32);
2629 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
2630
2631 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
2632 try module.getBodyBlock().addOperation(func.op);
2633
2634 const entry = func.getEntryBlock();
2635 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2636 try entry.addOperation(gid.op);
2637 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), f32_type);
2638 try entry.addOperation(load.op);
2639 const c = try ArithDialect.ConstantOp.createFloat(&ctx, loc, f32_type, 1.0);
2640 try entry.addOperation(c.op);
2641 const op = try ArithDialect.AndOp.create(&ctx, loc, load.getResult(), c.getResult());
2642 try entry.addOperation(op.op);
2643 const store = try MemrefDialect.StoreOp.create(&ctx, loc, op.getResult(), func.getArgument(0), gid.getResult());
2644 try entry.addOperation(store.op);
2645 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2646 try entry.addOperation(ret.op);
2647
2648 var codegen = SpirvCodegen.init(allocator);
2649 defer codegen.deinit();
2650 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2651 }
2652
2653 test "spirv codegen rejects arith.shl with width-mismatched shift count (cross-backend strict-equality policy)" {
2654 const testing = std.testing;
2655 const allocator = testing.allocator;
2656
2657 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2658 defer ctx.deinit(allocator);
2659
2660 const loc = ir.Location.getUnknown();
2661 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2662 const i32_type = try ArithDialect.getScalarType(&ctx, .i32);
2663 const i64_type = try ArithDialect.getScalarType(&ctx, .i64);
2664 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, i32_type, .device);
2665
2666 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
2667 try module.getBodyBlock().addOperation(func.op);
2668
2669 const entry = func.getEntryBlock();
2670 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2671 try entry.addOperation(gid.op);
2672 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), i32_type);
2673 try entry.addOperation(load.op);
2674 const c64 = try ArithDialect.ConstantOp.createInt(&ctx, loc, i64_type, 2);
2675 try entry.addOperation(c64.op);
2676 const shl = try ArithDialect.ShlOp.create(&ctx, loc, load.getResult(), c64.getResult());
2677 try entry.addOperation(shl.op);
2678 const store = try MemrefDialect.StoreOp.create(&ctx, loc, shl.getResult(), func.getArgument(0), gid.getResult());
2679 try entry.addOperation(store.op);
2680 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2681 try entry.addOperation(ret.op);
2682
2683 var codegen = SpirvCodegen.init(allocator);
2684 defer codegen.deinit();
2685 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2686 }
2687
2688 test "spirv codegen capability dedup: requireCapability emits at most once" {
2689 const testing = std.testing;
2690 const allocator = testing.allocator;
2691
2692 var codegen = SpirvCodegen.init(allocator);
2693 defer codegen.deinit();
2694
2695 try codegen.requireCapability(SpvCapability.Float64);
2696 try codegen.requireCapability(SpvCapability.Float64);
2697
2698 const words = try codegen.builder.toWords(allocator);
2699 defer allocator.free(words);
2700
2701 try testing.expectEqual(@as(usize, 1), countCapability(words, SpvCapability.Float64));
2702 try testing.expectEqual(@as(usize, 0), countCapability(words, SpvCapability.Float16));
2703 }
2704
2705 test "spirv codegen emits OpConvertUToF for arith.cast index → f32 (iota lowering shape)" {
2706 const testing = std.testing;
2707 const allocator = testing.allocator;
2708
2709 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2710 defer ctx.deinit(allocator);
2711
2712 const loc = ir.Location.getUnknown();
2713 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2714 const f32_type = try ArithDialect.getScalarType(&ctx, .f32);
2715 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
2716
2717 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
2718 try module.getBodyBlock().addOperation(func.op);
2719
2720 const entry = func.getEntryBlock();
2721 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2722 try entry.addOperation(gid.op);
2723 const cast = try ArithDialect.CastOp.create(&ctx, loc, gid.getResult(), f32_type);
2724 try entry.addOperation(cast.op);
2725 const store = try MemrefDialect.StoreOp.create(&ctx, loc, cast.getResult(), func.getArgument(0), gid.getResult());
2726 try entry.addOperation(store.op);
2727 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2728 try entry.addOperation(ret.op);
2729
2730 var codegen = SpirvCodegen.init(allocator);
2731 defer codegen.deinit();
2732 const words = try codegen.emitModuleWords(module.op);
2733 defer allocator.free(words);
2734
2735 try testing.expect(containsOpcode(words, SpirvOp.ConvertUToF));
2736 }
2737
2738 test "spirv codegen emits OpBitcast for arith.bitcast i32 -> f32 (different kinds, same width)" {
2739 const testing = std.testing;
2740 const allocator = testing.allocator;
2741
2742 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2743 defer ctx.deinit(allocator);
2744
2745 const loc = ir.Location.getUnknown();
2746 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2747 const i32_type = try ArithDialect.getScalarType(&ctx, .i32);
2748 const f32_type = try ArithDialect.getScalarType(&ctx, .f32);
2749 const i32_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, i32_type, .device);
2750 const f32_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
2751
2752 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{ i32_memref, f32_memref });
2753 try module.getBodyBlock().addOperation(func.op);
2754
2755 const entry = func.getEntryBlock();
2756 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2757 try entry.addOperation(gid.op);
2758 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), i32_type);
2759 try entry.addOperation(load.op);
2760 const bc = try ArithDialect.BitcastOp.create(&ctx, loc, load.getResult(), f32_type);
2761 try entry.addOperation(bc.op);
2762 const store = try MemrefDialect.StoreOp.create(&ctx, loc, bc.getResult(), func.getArgument(1), gid.getResult());
2763 try entry.addOperation(store.op);
2764 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2765 try entry.addOperation(ret.op);
2766
2767 var codegen = SpirvCodegen.init(allocator);
2768 defer codegen.deinit();
2769 const words = try codegen.emitModuleWords(module.op);
2770 defer allocator.free(words);
2771
2772 try testing.expect(containsOpcode(words, SpirvOp.Bitcast));
2773 }
2774
2775 test "spirv codegen aliases same-kind arith.bitcast (no OpBitcast emission)" {
2776 const testing = std.testing;
2777 const allocator = testing.allocator;
2778
2779 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2780 defer ctx.deinit(allocator);
2781
2782 const loc = ir.Location.getUnknown();
2783 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2784 const f32_type = try ArithDialect.getScalarType(&ctx, .f32);
2785 const f32_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, f32_type, .device);
2786
2787 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{f32_memref});
2788 try module.getBodyBlock().addOperation(func.op);
2789
2790 const entry = func.getEntryBlock();
2791 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2792 try entry.addOperation(gid.op);
2793 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), f32_type);
2794 try entry.addOperation(load.op);
2795 const bc = try ArithDialect.BitcastOp.create(&ctx, loc, load.getResult(), f32_type);
2796 try entry.addOperation(bc.op);
2797 const store = try MemrefDialect.StoreOp.create(&ctx, loc, bc.getResult(), func.getArgument(0), gid.getResult());
2798 try entry.addOperation(store.op);
2799 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2800 try entry.addOperation(ret.op);
2801
2802 var codegen = SpirvCodegen.init(allocator);
2803 defer codegen.deinit();
2804 const words = try codegen.emitModuleWords(module.op);
2805 defer allocator.free(words);
2806
2807 try testing.expect(!containsOpcode(words, SpirvOp.Bitcast));
2808 }
2809
2810 test "spirv codegen rejects arith.bitcast over arith.bool" {
2811 const testing = std.testing;
2812 const allocator = testing.allocator;
2813
2814 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2815 defer ctx.deinit(allocator);
2816
2817 const loc = ir.Location.getUnknown();
2818 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2819 const bool_type = try ArithDialect.getScalarType(&ctx, .bool);
2820 const i64_type = try ArithDialect.getScalarType(&ctx, .i64);
2821 const bool_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, bool_type, .device);
2822
2823 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{bool_memref});
2824 try module.getBodyBlock().addOperation(func.op);
2825
2826 const entry = func.getEntryBlock();
2827 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2828 try entry.addOperation(gid.op);
2829 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), bool_type);
2830 try entry.addOperation(load.op);
2831 const bc = try ArithDialect.BitcastOp.create(&ctx, loc, load.getResult(), i64_type);
2832 try entry.addOperation(bc.op);
2833 const store = try MemrefDialect.StoreOp.create(&ctx, loc, bc.getResult(), func.getArgument(0), gid.getResult());
2834 try entry.addOperation(store.op);
2835 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2836 try entry.addOperation(ret.op);
2837
2838 var codegen = SpirvCodegen.init(allocator);
2839 defer codegen.deinit();
2840 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2841 }
2842
2843 test "spirv codegen emits OpExtInst Fma for arith.fma over f32" {
2844 const testing = std.testing;
2845 const allocator = testing.allocator;
2846
2847 const f32_words = try emitArithKernelWords(allocator, .f32, emit_test_fma_f32);
2848 defer allocator.free(f32_words);
2849
2850 const set_id = findExtInstImportId(f32_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2851 try testing.expect(containsExtInst(f32_words, set_id, GLSLstd450.Fma));
2852 const wc = extInstWordCount(f32_words, set_id, GLSLstd450.Fma) orelse return error.TestExpectedExtInst;
2853 try testing.expectEqual(@as(u32, 8), wc);
2854 }
2855
2856 test "spirv codegen emits arith.fma over f16 with Float16 capability" {
2857 const testing = std.testing;
2858 const allocator = testing.allocator;
2859
2860 const f16_words = try emitArithKernelWords(allocator, .f16, emit_test_fma_f16);
2861 defer allocator.free(f16_words);
2862
2863 try testing.expect(containsCapability(f16_words, SpvCapability.Float16));
2864 const set_id = findExtInstImportId(f16_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2865 try testing.expect(containsExtInst(f16_words, set_id, GLSLstd450.Fma));
2866 }
2867
2868 test "spirv codegen emits arith.fma over f64 with Float64 capability" {
2869 const testing = std.testing;
2870 const allocator = testing.allocator;
2871
2872 const f64_words = try emitArithKernelWords(allocator, .f64, emit_test_fma_f64);
2873 defer allocator.free(f64_words);
2874
2875 try testing.expect(containsCapability(f64_words, SpvCapability.Float64));
2876 const set_id = findExtInstImportId(f64_words, "GLSL.std.450") orelse return error.TestExpectedGlslImport;
2877 try testing.expect(containsExtInst(f64_words, set_id, GLSLstd450.Fma));
2878 }
2879
2880 test "spirv codegen rejects arith.fma over integer types" {
2881 const testing = std.testing;
2882 const allocator = testing.allocator;
2883
2884 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2885 defer ctx.deinit(allocator);
2886
2887 const loc = ir.Location.getUnknown();
2888 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2889 const i32_type = try ArithDialect.getScalarType(&ctx, .i32);
2890 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, i32_type, .device);
2891
2892 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{memref_type});
2893 try module.getBodyBlock().addOperation(func.op);
2894
2895 const entry = func.getEntryBlock();
2896 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2897 try entry.addOperation(gid.op);
2898 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), i32_type);
2899 try entry.addOperation(load.op);
2900 const a = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 2);
2901 try entry.addOperation(a.op);
2902 const b = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 3);
2903 try entry.addOperation(b.op);
2904 const fma = try ArithDialect.FmaOp.create(&ctx, loc, a.getResult(), b.getResult(), load.getResult());
2905 try entry.addOperation(fma.op);
2906 const store = try MemrefDialect.StoreOp.create(&ctx, loc, fma.getResult(), func.getArgument(0), gid.getResult());
2907 try entry.addOperation(store.op);
2908 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2909 try entry.addOperation(ret.op);
2910
2911 var codegen = SpirvCodegen.init(allocator);
2912 defer codegen.deinit();
2913 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2914 }
2915
2916 test "spirv codegen rejects width-mismatched arith.bitcast (i32 -> f64)" {
2917 const testing = std.testing;
2918 const allocator = testing.allocator;
2919
2920 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2921 defer ctx.deinit(allocator);
2922
2923 const loc = ir.Location.getUnknown();
2924 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2925 const i32_type = try ArithDialect.getScalarType(&ctx, .i32);
2926 const f64_type = try ArithDialect.getScalarType(&ctx, .f64);
2927 const i32_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, i32_type, .device);
2928 const f64_memref = try MemrefDialect.getMemrefTypeDynamic(&ctx, f64_type, .device);
2929
2930 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{ i32_memref, f64_memref });
2931 try module.getBodyBlock().addOperation(func.op);
2932
2933 const entry = func.getEntryBlock();
2934 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2935 try entry.addOperation(gid.op);
2936 const load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), i32_type);
2937 try entry.addOperation(load.op);
2938 const bc = try ArithDialect.BitcastOp.create(&ctx, loc, load.getResult(), f64_type);
2939 try entry.addOperation(bc.op);
2940 const store = try MemrefDialect.StoreOp.create(&ctx, loc, bc.getResult(), func.getArgument(1), gid.getResult());
2941 try entry.addOperation(store.op);
2942 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2943 try entry.addOperation(ret.op);
2944
2945 var codegen = SpirvCodegen.init(allocator);
2946 defer codegen.deinit();
2947 try testing.expectError(error.UnsupportedType, codegen.emitModuleWords(module.op));
2948 }
2949
2950 fn emitMinMaxKernel(
2951 allocator: std.mem.Allocator,
2952 elem_kind: ArithDialect.ScalarTypeKind,
2953 op_kind: enum { max, min },
2954 ) ![]u32 {
2955 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2956 defer ctx.deinit(allocator);
2957
2958 const loc = ir.Location.getUnknown();
2959 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
2960
2961 const elem_type = try ArithDialect.getScalarType(&ctx, elem_kind);
2962 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, elem_type, .device);
2963
2964 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{ memref_type, memref_type });
2965 try module.getBodyBlock().addOperation(func.op);
2966
2967 const entry = func.getEntryBlock();
2968 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
2969 try entry.addOperation(gid.op);
2970 const lhs_load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(0), gid.getResult(), elem_type);
2971 try entry.addOperation(lhs_load.op);
2972 const rhs_load = try MemrefDialect.LoadOp.create(&ctx, loc, func.getArgument(1), gid.getResult(), elem_type);
2973 try entry.addOperation(rhs_load.op);
2974
2975 const stored = switch (op_kind) {
2976 .max => blk: {
2977 const op = try ArithDialect.MaxOp.create(&ctx, loc, lhs_load.getResult(), rhs_load.getResult());
2978 try entry.addOperation(op.op);
2979 break :blk op.getResult();
2980 },
2981 .min => blk: {
2982 const op = try ArithDialect.MinOp.create(&ctx, loc, lhs_load.getResult(), rhs_load.getResult());
2983 try entry.addOperation(op.op);
2984 break :blk op.getResult();
2985 },
2986 };
2987
2988 const store = try MemrefDialect.StoreOp.create(&ctx, loc, stored, func.getArgument(0), gid.getResult());
2989 try entry.addOperation(store.op);
2990 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
2991 try entry.addOperation(ret.op);
2992
2993 var codegen = SpirvCodegen.init(allocator);
2994 defer codegen.deinit();
2995 return codegen.emitModuleWords(module.op);
2996 }
2997
2998 test "spirv codegen lowers scalar arguments to one push constant block" {
2999 const testing = std.testing;
3000 const allocator = testing.allocator;
3001
3002 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
3003 defer ctx.deinit(allocator);
3004
3005 const loc = ir.Location.getUnknown();
3006 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
3007
3008 const memref_elem = try ArithDialect.getScalarType(&ctx, .f32);
3009 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ctx, memref_elem, .device);
3010 const u32_type = try ArithDialect.getScalarType(&ctx, .u32);
3011 const f32_type = try ArithDialect.getScalarType(&ctx, .f32);
3012
3013 var func = try FuncDialect.FuncOp.createKernel(&ctx, loc, "kernel", &.{ memref_type, u32_type, f32_type });
3014 try module.getBodyBlock().addOperation(func.op);
3015
3016 const entry = func.getEntryBlock();
3017 const gid = try GpuDialect.GlobalIdxOp.create(&ctx, loc, .x);
3018 try entry.addOperation(gid.op);
3019
3020 const store = try MemrefDialect.StoreOp.create(&ctx, loc, func.getArgument(2), func.getArgument(0), gid.getResult());
3021 try entry.addOperation(store.op);
3022
3023 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{});
3024 try entry.addOperation(ret.op);
3025
3026 var codegen = SpirvCodegen.init(allocator);
3027 defer codegen.deinit();
3028
3029 const words = try codegen.emitModuleWords(module.op);
3030 defer allocator.free(words);
3031
3032 try testing.expect(containsOpcode(words, SpirvOp.SpecConstant));
3033 try testing.expect(containsOpcode(words, SpirvOp.SpecConstantComposite));
3034 try testing.expect(containsVariableWithStorageClass(words, SpvStorageClass.PushConstant));
3035 try testing.expectEqual(@as(usize, 1), countBindingDecorations(words));
3036 try testing.expect(containsMemberOffset(words, 0));
3037 try testing.expect(containsMemberOffset(words, 4));
3038 }
3039
3040 fn containsVariableWithStorageClass(words: []const u32, storage_class: u32) bool {
3041 if (words.len < 5) return false;
3042 var i: usize = 5;
3043 while (i < words.len) {
3044 const word = words[i];
3045 const word_count = word >> 16;
3046 const op = @as(u16, @truncate(word));
3047 if (op == SpirvOp.Variable and word_count >= 4 and words[i + 3] == storage_class) return true;
3048 if (word_count == 0) break;
3049 i += word_count;
3050 }
3051 return false;
3052 }
3053
3054 fn countBindingDecorations(words: []const u32) usize {
3055 if (words.len < 5) return 0;
3056 var count: usize = 0;
3057 var i: usize = 5;
3058 while (i < words.len) {
3059 const word = words[i];
3060 const word_count = word >> 16;
3061 const op = @as(u16, @truncate(word));
3062 if (op == SpirvOp.Decorate and word_count == 4 and words[i + 2] == SpvDecoration.Binding) count += 1;
3063 if (word_count == 0) break;
3064 i += word_count;
3065 }
3066 return count;
3067 }
3068
3069 fn containsMemberOffset(words: []const u32, offset: u32) bool {
3070 if (words.len < 5) return false;
3071 var i: usize = 5;
3072 while (i < words.len) {
3073 const word = words[i];
3074 const word_count = word >> 16;
3075 const op = @as(u16, @truncate(word));
3076 if (op == SpirvOp.MemberDecorate and word_count == 5 and words[i + 3] == SpvDecoration.Offset and words[i + 4] == offset) return true;
3077 if (word_count == 0) break;
3078 i += word_count;
3079 }
3080 return false;
3081 }
3082
3083 fn containsOpcode(words: []const u32, opcode: u16) bool {
3084 if (words.len < 5) return false;
3085 var i: usize = 5;
3086 while (i < words.len) {
3087 const word = words[i];
3088 const word_count = word >> 16;
3089 const op = @as(u16, @truncate(word));
3090 if (op == opcode) return true;
3091 if (word_count == 0) break;
3092 i += word_count;
3093 }
3094 return false;
3095 }
3096
3097 fn loopMergesImmediatelyPrecedeBranches(words: []const u32) bool {
3098 if (words.len < 5) return false;
3099 var found = false;
3100 var i: usize = 5;
3101 while (i < words.len) {
3102 const word_count = words[i] >> 16;
3103 const op = @as(u16, @truncate(words[i]));
3104 if (op == SpirvOp.LoopMerge) {
3105 found = true;
3106 const next = i + word_count;
3107 if (next >= words.len) return false;
3108 const next_op = @as(u16, @truncate(words[next]));
3109 if (next_op != SpirvOp.Branch and next_op != SpirvOp.BranchConditional) return false;
3110 }
3111 if (word_count == 0) return false;
3112 i += word_count;
3113 }
3114 return found;
3115 }
3116
3117 fn findTypeBool(words: []const u32) ?u32 {
3118 if (words.len < 5) return null;
3119 var i: usize = 5;
3120 while (i < words.len) {
3121 const word = words[i];
3122 const word_count = word >> 16;
3123 const op = @as(u16, @truncate(word));
3124 if (op == SpirvOp.TypeBool and word_count == 2 and i + 1 < words.len) return words[i + 1];
3125 if (word_count == 0) break;
3126 i += word_count;
3127 }
3128 return null;
3129 }
3130
3131 fn findTypeInt(words: []const u32, width: u32, signedness: u32) ?u32 {
3132 if (words.len < 5) return null;
3133 var i: usize = 5;
3134 while (i < words.len) {
3135 const word = words[i];
3136 const word_count = word >> 16;
3137 const op = @as(u16, @truncate(word));
3138 if (op == SpirvOp.TypeInt and word_count == 4 and i + 3 < words.len and words[i + 2] == width and words[i + 3] == signedness) {
3139 return words[i + 1];
3140 }
3141 if (word_count == 0) break;
3142 i += word_count;
3143 }
3144 return null;
3145 }
3146
3147 fn containsRuntimeArrayElement(words: []const u32, elem_type_id: u32) bool {
3148 if (words.len < 5) return false;
3149 var i: usize = 5;
3150 while (i < words.len) {
3151 const word = words[i];
3152 const word_count = word >> 16;
3153 const op = @as(u16, @truncate(word));
3154 if (op == SpirvOp.TypeRuntimeArray and word_count == 3 and i + 2 < words.len and words[i + 2] == elem_type_id) {
3155 return true;
3156 }
3157 if (word_count == 0) break;
3158 i += word_count;
3159 }
3160 return false;
3161 }
3162
3163 fn containsCapability(words: []const u32, capability: u32) bool {
3164 if (words.len < 6) return false;
3165 var i: usize = 5;
3166 while (i < words.len) {
3167 const word = words[i];
3168 const word_count = word >> 16;
3169 const op = @as(u16, @truncate(word));
3170 if (op == SpirvOp.Capability and word_count >= 2 and i + 1 < words.len) {
3171 if (words[i + 1] == capability) return true;
3172 }
3173 if (word_count == 0) break;
3174 i += word_count;
3175 }
3176 return false;
3177 }
3178
3179 fn containsExtension(words: []const u32, extension: []const u8) bool {
3180 if (words.len < 6) return false;
3181 var i: usize = 5;
3182 while (i < words.len) {
3183 const word = words[i];
3184 const word_count = word >> 16;
3185 const op = @as(u16, @truncate(word));
3186 if (op == SpirvOp.Extension and word_count >= 2 and i + word_count <= words.len) {
3187 const bytes = std.mem.sliceAsBytes(words[i + 1 .. i + word_count]);
3188 const end = std.mem.indexOfScalar(u8, bytes, 0) orelse bytes.len;
3189 if (std.mem.eql(u8, bytes[0..end], extension)) return true;
3190 }
3191 if (word_count == 0) break;
3192 i += word_count;
3193 }
3194 return false;
3195 }
3196
3197 fn countCapability(words: []const u32, capability: u32) usize {
3198 if (words.len < 6) return 0;
3199 var i: usize = 5;
3200 var count: usize = 0;
3201 while (i < words.len) {
3202 const word = words[i];
3203 const word_count = word >> 16;
3204 const op = @as(u16, @truncate(word));
3205 if (op == SpirvOp.Capability and word_count >= 2 and i + 1 < words.len) {
3206 if (words[i + 1] == capability) count += 1;
3207 }
3208 if (word_count == 0) break;
3209 i += word_count;
3210 }
3211 return count;
3212 }
3213
3214 fn extInstWordCount(words: []const u32, set_id: u32, ext_opcode: u32) ?u32 {
3215 if (words.len < 6) return null;
3216 var i: usize = 5;
3217 while (i < words.len) {
3218 const word = words[i];
3219 const word_count = word >> 16;
3220 const op = @as(u16, @truncate(word));
3221 if (op == SpirvOp.ExtInst and word_count >= 5 and i + word_count <= words.len) {
3222 if (words[i + 3] == set_id and words[i + 4] == ext_opcode) return word_count;
3223 }
3224 if (word_count == 0) break;
3225 i += word_count;
3226 }
3227 return null;
3228 }
3229
3230 fn countDecorations(words: []const u32, decoration: u32, literal: ?u32) usize {
3231 if (words.len < 6) return 0;
3232 var i: usize = 5;
3233 var count: usize = 0;
3234 while (i < words.len) {
3235 const word = words[i];
3236 const word_count = word >> 16;
3237 const op = @as(u16, @truncate(word));
3238 if (op == SpirvOp.Decorate and word_count >= 3 and i + 2 < words.len) {
3239 if (words[i + 2] == decoration) {
3240 if (literal) |value| {
3241 if (word_count >= 4 and i + 3 < words.len and words[i + 3] == value) {
3242 count += 1;
3243 }
3244 } else {
3245 count += 1;
3246 }
3247 }
3248 }
3249 if (word_count == 0) break;
3250 i += word_count;
3251 }
3252 return count;
3253 }
3254
3255 fn containsDecoration(words: []const u32, decoration: u32, literal: u32) bool {
3256 return countDecorations(words, decoration, literal) > 0;
3257 }
3258
3259 fn wordsToBytes(allocator: std.mem.Allocator, words: []const u32) ![]u8 {
3260 const byte_len = words.len * @sizeOf(u32);
3261 const bytes = try allocator.alloc(u8, byte_len);
3262 var offset: usize = 0;
3263 for (words) |word| {
3264 std.mem.writeInt(u32, bytes[offset..][0..4], word, .little);
3265 offset += 4;
3266 }
3267 return bytes;
3268 }
3269
3270 fn maybeRunSpirvVal(bytes: []const u8) !void {
3271 const testing = std.testing;
3272 const allocator = testing.allocator;
3273
3274 var tmp = testing.tmpDir(.{});
3275 defer tmp.cleanup();
3276
3277 const tmp_dir = try std.fmt.allocPrint(allocator, ".zig-cache/tmp/{s}", .{tmp.sub_path});
3278 defer allocator.free(tmp_dir);
3279 const spirv_path = try std.fmt.allocPrint(allocator, "{s}/spirv-val.spv", .{tmp_dir});
3280 defer allocator.free(spirv_path);
3281
3282 try sys.fs.writeFile(spirv_path, bytes);
3283
3284 const result = sys.process.run(allocator, sys.fs.debugIo(), .{
3285 .argv = &.{ "spirv-val", spirv_path },
3286 }) catch return;
3287 defer allocator.free(result.stdout);
3288 defer allocator.free(result.stderr);
3289
3290 if (sys.process.exitCode(result.term) != 0) return error.SpirvValFailed;
3291 }
3292
3293 fn buildVecAddKernelJob(ir_ctx: *ir.Context) !*ir.Operation {
3294 const loc = ir.Location.getUnknown();
3295 const module = try BuiltinDialect.ModuleOp.create(ir_ctx, loc);
3296 const module_block = module.getBodyBlock();
3297
3298 const elem_type = try ArithDialect.getI32Type(ir_ctx);
3299 const memref_type = try MemrefDialect.getMemrefType1D(ir_ctx, 4, elem_type, .device);
3300
3301 var kernel = try FuncDialect.FuncOp.createKernel(
3302 ir_ctx,
3303 loc,
3304 "vec_add",
3305 &.{ memref_type, memref_type, memref_type },
3306 );
3307 try module_block.addOperation(kernel.op);
3308
3309 const entry = kernel.getEntryBlock();
3310 const a_arg = entry.arguments.items[0];
3311 const b_arg = entry.arguments.items[1];
3312 const c_arg = entry.arguments.items[2];
3313
3314 const gid = try GpuDialect.GlobalIdxOp.create(ir_ctx, loc, .x);
3315 try entry.addOperation(gid.op);
3316
3317 const load_a = try MemrefDialect.LoadOp.create(ir_ctx, loc, a_arg, gid.getResult(), elem_type);
3318 try entry.addOperation(load_a.op);
3319 const load_b = try MemrefDialect.LoadOp.create(ir_ctx, loc, b_arg, gid.getResult(), elem_type);
3320 try entry.addOperation(load_b.op);
3321
3322 const add = try ArithDialect.AddOp.create(ir_ctx, loc, load_a.getResult(), load_b.getResult());
3323 try entry.addOperation(add.op);
3324
3325 const store = try MemrefDialect.StoreOp.create(ir_ctx, loc, add.getResult(), c_arg, gid.getResult());
3326 try entry.addOperation(store.op);
3327
3328 const ret = try FuncDialect.ReturnOp.create(ir_ctx, loc, &.{});
3329 try entry.addOperation(ret.op);
3330
3331 return module.op;
3332 }
3333
3334 fn buildReductionKernelJob(ir_ctx: *ir.Context) !*ir.Operation {
3335 const loc = ir.Location.getUnknown();
3336 const module = try BuiltinDialect.ModuleOp.create(ir_ctx, loc);
3337 const module_block = module.getBodyBlock();
3338
3339 const elem_type = try ArithDialect.getI32Type(ir_ctx);
3340 const index_type = try ArithDialect.getIndexType(ir_ctx);
3341 const memref_in = try MemrefDialect.getMemrefType1D(ir_ctx, 32, elem_type, .device);
3342 const memref_out = try MemrefDialect.getMemrefType1D(ir_ctx, 1, elem_type, .device);
3343
3344 var kernel = try FuncDialect.FuncOp.createKernel(
3345 ir_ctx,
3346 loc,
3347 "reduce_sum",
3348 &.{ memref_in, memref_out },
3349 );
3350 try module_block.addOperation(kernel.op);
3351
3352 const entry = kernel.getEntryBlock();
3353 const input_arg = entry.arguments.items[0];
3354 const output_arg = entry.arguments.items[1];
3355
3356 const tid = try GpuDialect.ThreadIdxOp.create(ir_ctx, loc, .x);
3357 try entry.addOperation(tid.op);
3358
3359 const load_in = try MemrefDialect.LoadOp.create(ir_ctx, loc, input_arg, tid.getResult(), elem_type);
3360 try entry.addOperation(load_in.op);
3361
3362 const mask = try ArithDialect.ConstantOp.createInt(ir_ctx, loc, elem_type, -1);
3363 try entry.addOperation(mask.op);
3364
3365 const reduce = try GpuDialect.WarpReduceOp.create(ir_ctx, loc, .add, mask.getResult(), load_in.getResult());
3366 try entry.addOperation(reduce.op);
3367
3368 const zero_idx = try ArithDialect.ConstantOp.createInt(ir_ctx, loc, index_type, 0);
3369 try entry.addOperation(zero_idx.op);
3370
3371 const store_out = try MemrefDialect.StoreOp.create(ir_ctx, loc, reduce.getResult(), output_arg, zero_idx.getResult());
3372 try entry.addOperation(store_out.op);
3373
3374 const ret = try FuncDialect.ReturnOp.create(ir_ctx, loc, &.{});
3375 try entry.addOperation(ret.op);
3376
3377 return module.op;
3378 }