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 }