lib/choir/src/backends/x64/backend.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const alloc_arena = @import("alloc_arena");
3 const ir = @import("../../core/root.zig");
4 const contract = @import("../root.zig").contract;
5 const artifact = @import("../root.zig").artifact;
6 const interface = @import("../root.zig").interface;
7 const machine = @import("../root.zig").machine_code;
8 const passes = @import("../../passes/root.zig");
9 const invoke = @import("invoke.zig");
10 const object = @import("object.zig");
11 const x86_64 = @import("root.zig");
12 const sys = @import("sys");
13
14 const BackendError = interface.BackendError;
15 const Allocator = std.mem.Allocator;
16
17 const supports_x86_64_backend = sys.capabilities.current.supportsX86_64Execution();
18
19 pub const ModuleObjectCompileOptions = struct {
20 max_threads: usize = 1,
21 worker_allocator: ?Allocator = null,
22 verification_diagnostic: ?*contract.VerificationDiagnostic = null,
23 };
24
25 /// Lowers Choir modules for x86_64 and compiles them into machine-code artifacts, object files, or
26 /// modules of `runtime`.
27 pub const Backend = struct {
28 allocator: Allocator,
29 ctx: *ir.Context,
30 runtime: x86_64.JitRuntime,
31 pub const Limits: type = x86_64.JitRuntime.Limits;
32
33 /// Registers `malloc`, `free`, and the math library with a runtime bounded by `limits`, then
34 /// loads the arith dialect. Limits too small for those names give `error.JitRuntimeFull`, and
35 /// either step failing to allocate gives `error.OutOfMemory`.
36 pub fn init(
37 allocator: Allocator,
38 ctx: *ir.Context,
39 limits: Limits,
40 ) (Allocator.Error || error{JitRuntimeFull})!Backend {
41 var runtime = x86_64.JitRuntime.init(allocator, limits);
42 errdefer runtime.deinit();
43 try runtime.registerExternalSymbol("malloc", @intFromPtr(&sys.heap.malloc));
44 try runtime.registerExternalSymbol("free", @intFromPtr(&sys.heap.free));
45 try x86_64.math.registerSymbols(&runtime);
46 try contract.loadArithDialect(ctx);
47 return .{
48 .allocator = allocator,
49 .ctx = ctx,
50 .runtime = runtime,
51 };
52 }
53
54 pub fn deinit(self: *Backend) void {
55 self.runtime.deinit();
56 }
57
58 pub fn verify(self: *Backend, module: *ir.Operation) BackendError!void {
59 return self.verifyWithDiagnostic(module, null);
60 }
61
62 pub fn verifyWithDiagnostic(
63 self: *Backend,
64 module: *ir.Operation,
65 diagnostic: ?*contract.VerificationDiagnostic,
66 ) BackendError!void {
67 _ = self;
68 return contract.verifyModuleWithDiagnostic("backend/x86_64/verify", module, diagnostic);
69 }
70
71 /// The operations `lower` may add for each operation of the module it is
72 /// given.
73 ///
74 /// A CALLER SIZES A CONTEXT FROM THIS AND NOT FROM A CORPUS. `lower` runs
75 /// transforms that build operations, and the operations they build are
76 /// charged to the context the module lives in, so a caller that reserves
77 /// bytes for that context needs a figure it can read before it lowers
78 /// anything. This is that figure, and it is the sum of what each transform
79 /// `lower` runs declares for itself.
80 ///
81 /// Only promotion adds operations today. `legalizeMemrefViews` rewrites the
82 /// views an operation names and adds none, so it contributes nothing here.
83 /// A transform that starts adding operations adds its own declared figure
84 /// to this sum, and a caller reading this sees the change without reading
85 /// this file.
86 pub const operations_added_per_operation: usize =
87 passes.promotion.operations_created_per_operation;
88
89 pub fn lower(self: *Backend, module: *ir.Operation) BackendError!*ir.Operation {
90 return self.lowerWithDiagnostic(module, null);
91 }
92
93 pub fn lowerWithDiagnostic(
94 self: *Backend,
95 module: *ir.Operation,
96 diagnostic: ?*contract.VerificationDiagnostic,
97 ) BackendError!*ir.Operation {
98 passes.memref_views.legalizeMemrefViews(module, self.allocator) catch |err| switch (err) {
99 error.OutOfMemory => return BackendError.OutOfMemory,
100 else => return BackendError.CodeGenFailed,
101 };
102 _ = passes.promotion.promote(module, self.allocator) catch |err| switch (err) {
103 error.OutOfMemory => return BackendError.OutOfMemory,
104 error.InvalidPromotion => return BackendError.CodeGenFailed,
105 };
106 try self.verifyWithDiagnostic(module, diagnostic);
107 try x86_64.legality.checkModule(module);
108 return module;
109 }
110
111 /// Lowers `module` and compiles it into a new module of `runtime`.
112 /// Its code stays mapped until `runtime.release` receives the returned handle.
113 /// Fails with `error.JitRuntimeFull` when `runtime` already holds as many modules or mapped
114 /// bytes as its limits allow, which says nothing about `module` and reports no diagnostic.
115 /// A `release` admits the next one, so retrying without releasing fails the same way.
116 pub fn compile(self: *Backend, module: *ir.Operation) BackendError!x86_64.jit.ModuleHandle {
117 const lowered = try self.lower(module);
118 return self.runtime.compile(lowered) catch |err| switch (err) {
119 error.OutOfMemory => BackendError.OutOfMemory,
120 error.JitRuntimeFull => BackendError.JitRuntimeFull,
121 else => BackendError.JitCompileFailed,
122 };
123 }
124
125 pub fn compileFunctionToMachineCodeWithRelocations(
126 self: *Backend,
127 module: *ir.Operation,
128 function_name: []const u8,
129 ) BackendError!machine.MachineCode {
130 const lowered = try self.lower(module);
131 const func = ir.inspection.functionDefinitionByName(lowered, function_name) orelse return BackendError.FunctionNotFound;
132 return try object.emitFunctionMachineCode(self.allocator, func);
133 }
134
135 pub fn compileFunctionToArtifact(
136 self: *Backend,
137 module: *ir.Operation,
138 function_name: []const u8,
139 ) BackendError!artifact.Artifact {
140 const lowered = try self.lower(module);
141 const func = ir.inspection.functionDefinitionByName(lowered, function_name) orelse return BackendError.FunctionNotFound;
142 return try object.compileFunctionToArtifact(self.allocator, func, function_name);
143 }
144
145 pub fn compileFunctionToObjectFile(
146 self: *Backend,
147 module: *ir.Operation,
148 function_name: []const u8,
149 ) BackendError!artifact.Artifact {
150 const lowered = try self.lower(module);
151 const func = ir.inspection.functionDefinitionByName(lowered, function_name) orelse return BackendError.FunctionNotFound;
152 return try object.compileFunctionToObjectFile(self.allocator, func, function_name);
153 }
154
155 pub fn compileModuleToObjectFile(
156 self: *Backend,
157 module: *ir.Operation,
158 entry_name: []const u8,
159 ) BackendError!artifact.Artifact {
160 return try self.compileModuleToObjectFileWithOptions(module, entry_name, .{});
161 }
162
163 pub fn compileModuleToObjectFileWithOptions(
164 self: *Backend,
165 module: *ir.Operation,
166 entry_name: []const u8,
167 options: ModuleObjectCompileOptions,
168 ) BackendError!artifact.Artifact {
169 const lowered = try self.lowerWithDiagnostic(module, options.verification_diagnostic);
170 return try object.compileModuleToObjectFile(self.allocator, lowered, entry_name, .{
171 .max_threads = options.max_threads,
172 .worker_allocator = options.worker_allocator,
173 });
174 }
175
176 /// Writes the raw machine code of one function.
177 /// Raw bytes carry no relocations, so a function that calls or references a symbol fails.
178 pub fn emit(
179 self: *Backend,
180 module: *ir.Operation,
181 options: interface.EmitOptions,
182 writer: *std.Io.Writer,
183 ) BackendError!void {
184 const entry = options.entry orelse return BackendError.FunctionNotFound;
185 const lowered = try self.lower(module);
186 const func = ir.inspection.functionDefinitionByName(lowered, entry) orelse
187 return BackendError.FunctionNotFound;
188 var machine_code = try object.emitFunctionMachineCode(self.allocator, func);
189 defer machine_code.deinit(self.allocator);
190 if (machine_code.relocations.len != 0) return BackendError.UnsupportedOperation;
191 if (machine_code.data_relocations.len != 0) return BackendError.UnsupportedOperation;
192 writer.writeAll(machine_code.code) catch return BackendError.CodeGenFailed;
193 }
194 };
195
196 pub fn initHandle(allocator: Allocator, ctx: *ir.Context) BackendError!interface.BackendHandle {
197 return interface.initHandle(
198 Backend,
199 allocator,
200 ctx,
201 interface.BackendTarget.x86_64,
202 "x86_64",
203 .{ .artifact = .{ .machine_code = true, .object_file = true } },
204 );
205 }
206
207 fn choirSum7(a: i64, b: i64, c: i64, d: i64, e: i64, f: i64, g: i64) callconv(.c) i64 {
208 return a + b + c + d + e + f + g;
209 }
210
211 fn choirHandleAdd1(value: i64) callconv(.c) i64 {
212 return value + 1;
213 }
214
215 fn choirRegallocAddFive(value: i64) callconv(.c) i64 {
216 return value + 5;
217 }
218
219 fn choirHandleFalseWithDirtyHighBits() callconv(.c) usize {
220 return 0x100;
221 }
222
223 test "x86_64 JIT executes generic arith.add over vec2xf64 through native packed path" {
224 if (!supports_x86_64_backend) return;
225
226 const testing = std.testing;
227 const dialects = @import("../../dialects/root.zig");
228 const ArithDialect = dialects.ArithDialect;
229 const BuiltinDialect = dialects.BuiltinDialect;
230 const FuncDialect = dialects.FuncDialect;
231
232 var arena = alloc_arena.Arena.init(std.testing.allocator);
233 defer arena.deinit();
234 const allocator = arena.allocator();
235
236 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
237 defer ir_ctx.deinit(allocator);
238 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
239
240 const loc = ir.Location.getUnknown();
241 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
242 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 2, dialects.arith.type_names.float64)).?;
243
244 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
245 const module_block = module.getBodyBlock();
246
247 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "generic_add_vec2f64", &.{}, &.{vec_type});
248 try module_block.addOperation(func.op);
249
250 const entry = func.getEntryBlock();
251 var zero = try ArithDialect.VecConstantOp.createFloat(&ir_ctx, loc, vec_type, 0.0);
252 try entry.addOperation(zero.op);
253 var lhs0 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 1.0);
254 try entry.addOperation(lhs0.op);
255 var lhs_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), lhs0.getResult(), 0);
256 try entry.addOperation(lhs_lane0.op);
257 var lhs1 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 2.0);
258 try entry.addOperation(lhs1.op);
259 var lhs = try ArithDialect.InsertOp.create(&ir_ctx, loc, lhs_lane0.getResult(), lhs1.getResult(), 1);
260 try entry.addOperation(lhs.op);
261
262 var rhs0 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 10.0);
263 try entry.addOperation(rhs0.op);
264 var rhs_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), rhs0.getResult(), 0);
265 try entry.addOperation(rhs_lane0.op);
266 var rhs1 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 20.0);
267 try entry.addOperation(rhs1.op);
268 var rhs = try ArithDialect.InsertOp.create(&ir_ctx, loc, rhs_lane0.getResult(), rhs1.getResult(), 1);
269 try entry.addOperation(rhs.op);
270
271 const add = try ArithDialect.AddOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
272 try entry.addOperation(add.op);
273 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{add.getResult()});
274 try entry.addOperation(ret.op);
275
276 var backend = try Backend.init(allocator, &ir_ctx, .testing);
277 defer backend.deinit();
278 const compiled = try backend.compile(module.op);
279
280 const result = try vectorResult(&backend, compiled, "generic_add_vec2f64");
281
282 try testing.expectEqual(artifact.ScalarType.f64, result.element);
283 try testing.expectEqual(@as(u8, 2), result.lanes);
284 try testing.expectEqual(@as(u64, @bitCast(@as(f64, 11.0))), result.bits[0]);
285 try testing.expectEqual(@as(u64, @bitCast(@as(f64, 22.0))), result.bits[1]);
286 }
287
288 test "x86_64 JIT executes vec2xf64 neg through native packed path" {
289 if (!supports_x86_64_backend) return;
290
291 const testing = std.testing;
292 const dialects = @import("../../dialects/root.zig");
293 const ArithDialect = dialects.ArithDialect;
294 const BuiltinDialect = dialects.BuiltinDialect;
295 const FuncDialect = dialects.FuncDialect;
296
297 var arena = alloc_arena.Arena.init(std.testing.allocator);
298 defer arena.deinit();
299 const allocator = arena.allocator();
300
301 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
302 defer ir_ctx.deinit(allocator);
303 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
304
305 const loc = ir.Location.getUnknown();
306 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
307 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 2, dialects.arith.type_names.float64)).?;
308
309 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
310 const module_block = module.getBodyBlock();
311
312 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "neg_vec2f64", &.{}, &.{vec_type});
313 try module_block.addOperation(func.op);
314
315 const entry = func.getEntryBlock();
316 var zero = try ArithDialect.VecConstantOp.createFloat(&ir_ctx, loc, vec_type, 0.0);
317 try entry.addOperation(zero.op);
318 var lane0 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 1.5);
319 try entry.addOperation(lane0.op);
320 var with_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), lane0.getResult(), 0);
321 try entry.addOperation(with_lane0.op);
322 var lane1 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, -2.25);
323 try entry.addOperation(lane1.op);
324 var input = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane0.getResult(), lane1.getResult(), 1);
325 try entry.addOperation(input.op);
326
327 const neg = try ArithDialect.NegOp.create(&ir_ctx, loc, input.getResult());
328 try entry.addOperation(neg.op);
329 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{neg.getResult()});
330 try entry.addOperation(ret.op);
331
332 var backend = try Backend.init(allocator, &ir_ctx, .testing);
333 defer backend.deinit();
334 const compiled = try backend.compile(module.op);
335
336 const result = try vectorResult(&backend, compiled, "neg_vec2f64");
337
338 try testing.expectEqual(artifact.ScalarType.f64, result.element);
339 try testing.expectEqual(@as(u8, 2), result.lanes);
340 try testing.expectEqual(@as(u64, @bitCast(@as(f64, -1.5))), result.bits[0]);
341 try testing.expectEqual(@as(u64, @bitCast(@as(f64, 2.25))), result.bits[1]);
342 }
343
344 test "x86_64 JIT executes vec2xf64 splat and shuffle through native packed path" {
345 if (!supports_x86_64_backend) return;
346
347 const testing = std.testing;
348 const dialects = @import("../../dialects/root.zig");
349 const ArithDialect = dialects.ArithDialect;
350 const BuiltinDialect = dialects.BuiltinDialect;
351 const FuncDialect = dialects.FuncDialect;
352
353 var arena = alloc_arena.Arena.init(std.testing.allocator);
354 defer arena.deinit();
355 const allocator = arena.allocator();
356
357 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
358 defer ir_ctx.deinit(allocator);
359 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
360
361 const loc = ir.Location.getUnknown();
362 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
363 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 2, dialects.arith.type_names.float64)).?;
364
365 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
366 const module_block = module.getBodyBlock();
367
368 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "splat_shuffle_vec2f64", &.{}, &.{vec_type});
369 try module_block.addOperation(func.op);
370
371 const entry = func.getEntryBlock();
372 var fill = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 4.0);
373 try entry.addOperation(fill.op);
374 var splat = try ArithDialect.SplatOp.create(&ir_ctx, loc, fill.getResult(), vec_type);
375 try entry.addOperation(splat.op);
376 var lane0 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 1.25);
377 try entry.addOperation(lane0.op);
378 var input = try ArithDialect.InsertOp.create(&ir_ctx, loc, splat.getResult(), lane0.getResult(), 0);
379 try entry.addOperation(input.op);
380 const indices = [_]i64{ 1, 0 };
381 const shuffle = try ArithDialect.VecShuffleOp.create(&ir_ctx, loc, input.getResult(), vec_type, indices[0..]);
382 try entry.addOperation(shuffle.op);
383 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{shuffle.getResult()});
384 try entry.addOperation(ret.op);
385
386 var backend = try Backend.init(allocator, &ir_ctx, .testing);
387 defer backend.deinit();
388 const compiled = try backend.compile(module.op);
389
390 const result = try vectorResult(&backend, compiled, "splat_shuffle_vec2f64");
391
392 try testing.expectEqual(artifact.ScalarType.f64, result.element);
393 try testing.expectEqual(@as(u8, 2), result.lanes);
394 try testing.expectEqual(@as(u64, @bitCast(@as(f64, 4.0))), result.bits[0]);
395 try testing.expectEqual(@as(u64, @bitCast(@as(f64, 1.25))), result.bits[1]);
396 }
397
398 test "x86_64 JIT executes vec4xf32 through native packed path" {
399 if (!supports_x86_64_backend) return;
400
401 const testing = std.testing;
402 const dialects = @import("../../dialects/root.zig");
403 const ArithDialect = dialects.ArithDialect;
404 const BuiltinDialect = dialects.BuiltinDialect;
405 const FuncDialect = dialects.FuncDialect;
406
407 var arena = alloc_arena.Arena.init(std.testing.allocator);
408 defer arena.deinit();
409 const allocator = arena.allocator();
410
411 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
412 defer ir_ctx.deinit(allocator);
413 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
414
415 const loc = ir.Location.getUnknown();
416 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
417 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.float32)).?;
418
419 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
420 const module_block = module.getBodyBlock();
421
422 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "packed_vec4f32", &.{}, &.{vec_type});
423 try module_block.addOperation(func.op);
424
425 const entry = func.getEntryBlock();
426 var zero = try ArithDialect.VecConstantOp.createFloat(&ir_ctx, loc, vec_type, 0.0);
427 try entry.addOperation(zero.op);
428 var lane0 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 1.0);
429 try entry.addOperation(lane0.op);
430 var with_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), lane0.getResult(), 0);
431 try entry.addOperation(with_lane0.op);
432 var lane1 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 2.0);
433 try entry.addOperation(lane1.op);
434 var with_lane1 = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane0.getResult(), lane1.getResult(), 1);
435 try entry.addOperation(with_lane1.op);
436 var lane2 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 3.0);
437 try entry.addOperation(lane2.op);
438 var with_lane2 = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane1.getResult(), lane2.getResult(), 2);
439 try entry.addOperation(with_lane2.op);
440 var lane3 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 4.0);
441 try entry.addOperation(lane3.op);
442 var input = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane2.getResult(), lane3.getResult(), 3);
443 try entry.addOperation(input.op);
444
445 var fill = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 10.0);
446 try entry.addOperation(fill.op);
447 var splat = try ArithDialect.SplatOp.create(&ir_ctx, loc, fill.getResult(), vec_type);
448 try entry.addOperation(splat.op);
449 const add = try ArithDialect.AddOp.create(&ir_ctx, loc, input.getResult(), splat.getResult());
450 try entry.addOperation(add.op);
451 const indices = [_]i64{ 3, 2, 1, 0 };
452 const shuffle = try ArithDialect.VecShuffleOp.create(&ir_ctx, loc, add.getResult(), vec_type, indices[0..]);
453 try entry.addOperation(shuffle.op);
454 const neg = try ArithDialect.NegOp.create(&ir_ctx, loc, shuffle.getResult());
455 try entry.addOperation(neg.op);
456 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{neg.getResult()});
457 try entry.addOperation(ret.op);
458
459 var backend = try Backend.init(allocator, &ir_ctx, .testing);
460 defer backend.deinit();
461 const compiled = try backend.compile(module.op);
462
463 const result = try vectorResult(&backend, compiled, "packed_vec4f32");
464
465 try testing.expectEqual(artifact.ScalarType.f32, result.element);
466 try testing.expectEqual(@as(u8, 4), result.lanes);
467 try testing.expectEqual(@as(u64, @as(u32, @bitCast(@as(f32, -14.0)))), result.bits[0]);
468 try testing.expectEqual(@as(u64, @as(u32, @bitCast(@as(f32, -13.0)))), result.bits[1]);
469 try testing.expectEqual(@as(u64, @as(u32, @bitCast(@as(f32, -12.0)))), result.bits[2]);
470 try testing.expectEqual(@as(u64, @as(u32, @bitCast(@as(f32, -11.0)))), result.bits[3]);
471 }
472
473 test "x86_64 JIT executes arith.select over vec4xu32 with scalar bool condition" {
474 if (!supports_x86_64_backend) return;
475
476 const testing = std.testing;
477 const dialects = @import("../../dialects/root.zig");
478 const ArithDialect = dialects.ArithDialect;
479 const BuiltinDialect = dialects.BuiltinDialect;
480 const FuncDialect = dialects.FuncDialect;
481
482 const Case = struct {
483 name: []const u8,
484 condition: bool,
485 expected: u64,
486 };
487 const cases = [_]Case{
488 .{ .name = "select_vec4u32_true", .condition = true, .expected = 7 },
489 .{ .name = "select_vec4u32_false", .condition = false, .expected = 3 },
490 };
491
492 for (cases) |case| {
493 var arena = alloc_arena.Arena.init(std.testing.allocator);
494 defer arena.deinit();
495 const allocator = arena.allocator();
496
497 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
498 defer ir_ctx.deinit(allocator);
499 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
500
501 const loc = ir.Location.getUnknown();
502 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
503
504 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
505 const module_block = module.getBodyBlock();
506
507 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, case.name, &.{}, &.{vec_type});
508 try module_block.addOperation(func.op);
509
510 const entry = func.getEntryBlock();
511 var cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, case.condition);
512 try entry.addOperation(cond.op);
513 var true_value = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 7);
514 try entry.addOperation(true_value.op);
515 var false_value = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 3);
516 try entry.addOperation(false_value.op);
517 const selected = try ArithDialect.SelectOp.create(&ir_ctx, loc, cond.getResult(), true_value.getResult(), false_value.getResult());
518 try entry.addOperation(selected.op);
519 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{selected.getResult()});
520 try entry.addOperation(ret.op);
521
522 var backend = try Backend.init(allocator, &ir_ctx, .testing);
523 defer backend.deinit();
524 const compiled = try backend.compile(module.op);
525
526 const result = try vectorResult(&backend, compiled, case.name);
527 try testing.expectEqual(artifact.ScalarType.u32, result.element);
528 try testing.expectEqual(@as(u8, 4), result.lanes);
529 try testing.expectEqual(case.expected, result.bits[0]);
530 try testing.expectEqual(case.expected, result.bits[1]);
531 try testing.expectEqual(case.expected, result.bits[2]);
532 try testing.expectEqual(case.expected, result.bits[3]);
533 }
534 }
535
536 test "x86_64 JIT executes arith.umulhi over vec4xu32 through native packed path" {
537 if (!supports_x86_64_backend) return;
538
539 const testing = std.testing;
540 const dialects = @import("../../dialects/root.zig");
541 const ArithDialect = dialects.ArithDialect;
542 const BuiltinDialect = dialects.BuiltinDialect;
543 const FuncDialect = dialects.FuncDialect;
544
545 var arena = alloc_arena.Arena.init(std.testing.allocator);
546 defer arena.deinit();
547 const allocator = arena.allocator();
548
549 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
550 defer ir_ctx.deinit(allocator);
551 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
552
553 const loc = ir.Location.getUnknown();
554 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
555
556 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
557 const module_block = module.getBodyBlock();
558
559 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "umulhi_vec4u32", &.{}, &.{vec_type});
560 try module_block.addOperation(func.op);
561
562 const entry = func.getEntryBlock();
563 var lhs = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 0xffff_ffff);
564 try entry.addOperation(lhs.op);
565 var rhs = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 2);
566 try entry.addOperation(rhs.op);
567 const high = try ArithDialect.UmulhiOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
568 try entry.addOperation(high.op);
569 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{high.getResult()});
570 try entry.addOperation(ret.op);
571
572 var backend = try Backend.init(allocator, &ir_ctx, .testing);
573 defer backend.deinit();
574 const compiled = try backend.compile(module.op);
575
576 const result = try vectorResult(&backend, compiled, "umulhi_vec4u32");
577 try testing.expectEqual(artifact.ScalarType.u32, result.element);
578 try testing.expectEqual(@as(u8, 4), result.lanes);
579 try testing.expectEqual(@as(u64, 1), result.bits[0]);
580 try testing.expectEqual(@as(u64, 1), result.bits[1]);
581 try testing.expectEqual(@as(u64, 1), result.bits[2]);
582 try testing.expectEqual(@as(u64, 1), result.bits[3]);
583 }
584
585 test "x86_64 JIT executes arith.popcount over vec4xu32 through native packed path" {
586 if (!supports_x86_64_backend) return;
587
588 const testing = std.testing;
589 const dialects = @import("../../dialects/root.zig");
590 const ArithDialect = dialects.ArithDialect;
591 const BuiltinDialect = dialects.BuiltinDialect;
592 const FuncDialect = dialects.FuncDialect;
593
594 var arena = alloc_arena.Arena.init(std.testing.allocator);
595 defer arena.deinit();
596 const allocator = arena.allocator();
597
598 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
599 defer ir_ctx.deinit(allocator);
600 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
601
602 const loc = ir.Location.getUnknown();
603 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
604
605 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
606 const module_block = module.getBodyBlock();
607
608 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "popcount_vec4u32", &.{}, &.{vec_type});
609 try module_block.addOperation(func.op);
610
611 const entry = func.getEntryBlock();
612 var input = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 0xf0f0_00ff);
613 try entry.addOperation(input.op);
614 const count = try ArithDialect.PopCountOp.create(&ir_ctx, loc, input.getResult());
615 try entry.addOperation(count.op);
616 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{count.getResult()});
617 try entry.addOperation(ret.op);
618
619 var backend = try Backend.init(allocator, &ir_ctx, .testing);
620 defer backend.deinit();
621 const compiled = try backend.compile(module.op);
622
623 const result = try vectorResult(&backend, compiled, "popcount_vec4u32");
624 try testing.expectEqual(artifact.ScalarType.u32, result.element);
625 try testing.expectEqual(@as(u8, 4), result.lanes);
626 try testing.expectEqual(@as(u64, 16), result.bits[0]);
627 try testing.expectEqual(@as(u64, 16), result.bits[1]);
628 try testing.expectEqual(@as(u64, 16), result.bits[2]);
629 try testing.expectEqual(@as(u64, 16), result.bits[3]);
630 }
631
632 test "x86_64 JIT executes arith.popcount over vec2xu64 through native packed path" {
633 if (!supports_x86_64_backend) return;
634
635 const testing = std.testing;
636 const dialects = @import("../../dialects/root.zig");
637 const ArithDialect = dialects.ArithDialect;
638 const BuiltinDialect = dialects.BuiltinDialect;
639 const FuncDialect = dialects.FuncDialect;
640
641 var arena = alloc_arena.Arena.init(std.testing.allocator);
642 defer arena.deinit();
643 const allocator = arena.allocator();
644
645 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
646 defer ir_ctx.deinit(allocator);
647 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
648
649 const loc = ir.Location.getUnknown();
650 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 2, dialects.arith.type_names.uint64)).?;
651
652 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
653 const module_block = module.getBodyBlock();
654
655 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "popcount_vec2u64", &.{}, &.{vec_type});
656 try module_block.addOperation(func.op);
657
658 const entry = func.getEntryBlock();
659 var input = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, -1);
660 try entry.addOperation(input.op);
661 const count = try ArithDialect.PopCountOp.create(&ir_ctx, loc, input.getResult());
662 try entry.addOperation(count.op);
663 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{count.getResult()});
664 try entry.addOperation(ret.op);
665
666 var backend = try Backend.init(allocator, &ir_ctx, .testing);
667 defer backend.deinit();
668 const compiled = try backend.compile(module.op);
669
670 const result = try vectorResult(&backend, compiled, "popcount_vec2u64");
671 try testing.expectEqual(artifact.ScalarType.u64, result.element);
672 try testing.expectEqual(@as(u8, 2), result.lanes);
673 try testing.expectEqual(@as(u64, 64), result.bits[0]);
674 try testing.expectEqual(@as(u64, 64), result.bits[1]);
675 }
676
677 test "x86_64 JIT executes generic arith.max over vec4xu32" {
678 if (!supports_x86_64_backend) return;
679
680 const testing = std.testing;
681 const dialects = @import("../../dialects/root.zig");
682 const ArithDialect = dialects.ArithDialect;
683 const BuiltinDialect = dialects.BuiltinDialect;
684 const FuncDialect = dialects.FuncDialect;
685
686 var arena = alloc_arena.Arena.init(std.testing.allocator);
687 defer arena.deinit();
688 const allocator = arena.allocator();
689
690 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
691 defer ir_ctx.deinit(allocator);
692 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
693
694 const loc = ir.Location.getUnknown();
695 const u32_type = try ArithDialect.getScalarType(&ir_ctx, .u32);
696 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
697
698 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
699 const module_block = module.getBodyBlock();
700
701 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "generic_max_vec4u32", &.{}, &.{vec_type});
702 try module_block.addOperation(func.op);
703
704 const entry = func.getEntryBlock();
705 var zero = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 0);
706 try entry.addOperation(zero.op);
707 var lhs0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x8000_0000);
708 try entry.addOperation(lhs0.op);
709 var lhs_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), lhs0.getResult(), 0);
710 try entry.addOperation(lhs_lane0.op);
711 var lhs1 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 1);
712 try entry.addOperation(lhs1.op);
713 var lhs_lane1 = try ArithDialect.InsertOp.create(&ir_ctx, loc, lhs_lane0.getResult(), lhs1.getResult(), 1);
714 try entry.addOperation(lhs_lane1.op);
715 var lhs2 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0xffff_ffff);
716 try entry.addOperation(lhs2.op);
717 var lhs_lane2 = try ArithDialect.InsertOp.create(&ir_ctx, loc, lhs_lane1.getResult(), lhs2.getResult(), 2);
718 try entry.addOperation(lhs_lane2.op);
719 var lhs3 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 7);
720 try entry.addOperation(lhs3.op);
721 var lhs = try ArithDialect.InsertOp.create(&ir_ctx, loc, lhs_lane2.getResult(), lhs3.getResult(), 3);
722 try entry.addOperation(lhs.op);
723
724 var rhs0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 1);
725 try entry.addOperation(rhs0.op);
726 var rhs_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), rhs0.getResult(), 0);
727 try entry.addOperation(rhs_lane0.op);
728 var rhs1 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x8000_0000);
729 try entry.addOperation(rhs1.op);
730 var rhs_lane1 = try ArithDialect.InsertOp.create(&ir_ctx, loc, rhs_lane0.getResult(), rhs1.getResult(), 1);
731 try entry.addOperation(rhs_lane1.op);
732 var rhs2 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 2);
733 try entry.addOperation(rhs2.op);
734 var rhs_lane2 = try ArithDialect.InsertOp.create(&ir_ctx, loc, rhs_lane1.getResult(), rhs2.getResult(), 2);
735 try entry.addOperation(rhs_lane2.op);
736 var rhs3 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0xffff_fffe);
737 try entry.addOperation(rhs3.op);
738 var rhs = try ArithDialect.InsertOp.create(&ir_ctx, loc, rhs_lane2.getResult(), rhs3.getResult(), 3);
739 try entry.addOperation(rhs.op);
740
741 const maximum = try ArithDialect.MaxOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
742 try entry.addOperation(maximum.op);
743 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{maximum.getResult()});
744 try entry.addOperation(ret.op);
745
746 var backend = try Backend.init(allocator, &ir_ctx, .testing);
747 defer backend.deinit();
748 const compiled = try backend.compile(module.op);
749
750 const result = try vectorResult(&backend, compiled, "generic_max_vec4u32");
751
752 try testing.expectEqual(artifact.ScalarType.u32, result.element);
753 try testing.expectEqual(@as(u8, 4), result.lanes);
754 try testing.expectEqual(@as(u64, 0x8000_0000), result.bits[0]);
755 try testing.expectEqual(@as(u64, 0x8000_0000), result.bits[1]);
756 try testing.expectEqual(@as(u64, 0xffff_ffff), result.bits[2]);
757 try testing.expectEqual(@as(u64, 0xffff_fffe), result.bits[3]);
758 }
759
760 test "x86_64 JIT executes generic arith.min over vec4xf32 with NaN lanes" {
761 if (!supports_x86_64_backend) return;
762
763 const testing = std.testing;
764 const dialects = @import("../../dialects/root.zig");
765 const ArithDialect = dialects.ArithDialect;
766 const BuiltinDialect = dialects.BuiltinDialect;
767 const FuncDialect = dialects.FuncDialect;
768
769 var arena = alloc_arena.Arena.init(std.testing.allocator);
770 defer arena.deinit();
771 const allocator = arena.allocator();
772
773 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
774 defer ir_ctx.deinit(allocator);
775 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
776
777 const loc = ir.Location.getUnknown();
778 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
779 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.float32)).?;
780
781 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
782 const module_block = module.getBodyBlock();
783
784 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "generic_min_vec4f32_nan", &.{}, &.{vec_type});
785 try module_block.addOperation(func.op);
786
787 const entry = func.getEntryBlock();
788 var zero = try ArithDialect.VecConstantOp.createFloat(&ir_ctx, loc, vec_type, 0.0);
789 try entry.addOperation(zero.op);
790 var lhs0 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, std.math.nan(f32));
791 try entry.addOperation(lhs0.op);
792 var lhs_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), lhs0.getResult(), 0);
793 try entry.addOperation(lhs_lane0.op);
794 var lhs1 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 5.0);
795 try entry.addOperation(lhs1.op);
796 var lhs_lane1 = try ArithDialect.InsertOp.create(&ir_ctx, loc, lhs_lane0.getResult(), lhs1.getResult(), 1);
797 try entry.addOperation(lhs_lane1.op);
798 var lhs2 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, -8.0);
799 try entry.addOperation(lhs2.op);
800 var lhs_lane2 = try ArithDialect.InsertOp.create(&ir_ctx, loc, lhs_lane1.getResult(), lhs2.getResult(), 2);
801 try entry.addOperation(lhs_lane2.op);
802 var lhs3 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, -5.0);
803 try entry.addOperation(lhs3.op);
804 var lhs = try ArithDialect.InsertOp.create(&ir_ctx, loc, lhs_lane2.getResult(), lhs3.getResult(), 3);
805 try entry.addOperation(lhs.op);
806
807 var rhs0 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 5.0);
808 try entry.addOperation(rhs0.op);
809 var rhs_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), rhs0.getResult(), 0);
810 try entry.addOperation(rhs_lane0.op);
811 var rhs1 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, std.math.nan(f32));
812 try entry.addOperation(rhs1.op);
813 var rhs_lane1 = try ArithDialect.InsertOp.create(&ir_ctx, loc, rhs_lane0.getResult(), rhs1.getResult(), 1);
814 try entry.addOperation(rhs_lane1.op);
815 var rhs2 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, -1.0);
816 try entry.addOperation(rhs2.op);
817 var rhs_lane2 = try ArithDialect.InsertOp.create(&ir_ctx, loc, rhs_lane1.getResult(), rhs2.getResult(), 2);
818 try entry.addOperation(rhs_lane2.op);
819 var rhs3 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, -3.0);
820 try entry.addOperation(rhs3.op);
821 var rhs = try ArithDialect.InsertOp.create(&ir_ctx, loc, rhs_lane2.getResult(), rhs3.getResult(), 3);
822 try entry.addOperation(rhs.op);
823
824 const minimum = try ArithDialect.MinOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
825 try entry.addOperation(minimum.op);
826 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{minimum.getResult()});
827 try entry.addOperation(ret.op);
828
829 var backend = try Backend.init(allocator, &ir_ctx, .testing);
830 defer backend.deinit();
831 const compiled = try backend.compile(module.op);
832
833 const result = try vectorResult(&backend, compiled, "generic_min_vec4f32_nan");
834
835 try testing.expectEqual(artifact.ScalarType.f32, result.element);
836 try testing.expectEqual(@as(u8, 4), result.lanes);
837 try testing.expectEqual(@as(u64, @as(u32, @bitCast(@as(f32, 5.0)))), result.bits[0]);
838 try testing.expectEqual(@as(u64, @as(u32, @bitCast(@as(f32, 5.0)))), result.bits[1]);
839 try testing.expectEqual(@as(u64, @as(u32, @bitCast(@as(f32, -8.0)))), result.bits[2]);
840 try testing.expectEqual(@as(u64, @as(u32, @bitCast(@as(f32, -5.0)))), result.bits[3]);
841 }
842
843 test "x86_64 JIT executes vec4xu32 splat through native packed path" {
844 if (!supports_x86_64_backend) return;
845
846 const testing = std.testing;
847 const dialects = @import("../../dialects/root.zig");
848 const ArithDialect = dialects.ArithDialect;
849 const BuiltinDialect = dialects.BuiltinDialect;
850 const FuncDialect = dialects.FuncDialect;
851
852 var arena = alloc_arena.Arena.init(std.testing.allocator);
853 defer arena.deinit();
854 const allocator = arena.allocator();
855
856 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
857 defer ir_ctx.deinit(allocator);
858 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
859
860 const loc = ir.Location.getUnknown();
861 const u32_type = try ArithDialect.getScalarType(&ir_ctx, .u32);
862 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
863
864 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
865 const module_block = module.getBodyBlock();
866
867 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "splat_vec4u32", &.{}, &.{vec_type});
868 try module_block.addOperation(func.op);
869
870 const entry = func.getEntryBlock();
871 var fill = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x8000_0003);
872 try entry.addOperation(fill.op);
873 var splat = try ArithDialect.SplatOp.create(&ir_ctx, loc, fill.getResult(), vec_type);
874 try entry.addOperation(splat.op);
875 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{splat.getResult()});
876 try entry.addOperation(ret.op);
877
878 var backend = try Backend.init(allocator, &ir_ctx, .testing);
879 defer backend.deinit();
880 const compiled = try backend.compile(module.op);
881
882 const result = try vectorResult(&backend, compiled, "splat_vec4u32");
883
884 try testing.expectEqual(artifact.ScalarType.u32, result.element);
885 try testing.expectEqual(@as(u8, 4), result.lanes);
886 try testing.expectEqual(@as(u64, 0x8000_0003), result.bits[0]);
887 try testing.expectEqual(@as(u64, 0x8000_0003), result.bits[1]);
888 try testing.expectEqual(@as(u64, 0x8000_0003), result.bits[2]);
889 try testing.expectEqual(@as(u64, 0x8000_0003), result.bits[3]);
890 }
891
892 test "x86_64 JIT executes vec4xu32 shuffle through native packed path" {
893 if (!supports_x86_64_backend) return;
894
895 const testing = std.testing;
896 const dialects = @import("../../dialects/root.zig");
897 const ArithDialect = dialects.ArithDialect;
898 const BuiltinDialect = dialects.BuiltinDialect;
899 const FuncDialect = dialects.FuncDialect;
900
901 var arena = alloc_arena.Arena.init(std.testing.allocator);
902 defer arena.deinit();
903 const allocator = arena.allocator();
904
905 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
906 defer ir_ctx.deinit(allocator);
907 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
908
909 const loc = ir.Location.getUnknown();
910 const u32_type = try ArithDialect.getScalarType(&ir_ctx, .u32);
911 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
912
913 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
914 const module_block = module.getBodyBlock();
915
916 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "shuffle_vec4u32", &.{}, &.{vec_type});
917 try module_block.addOperation(func.op);
918
919 const entry = func.getEntryBlock();
920 var zero = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 0);
921 try entry.addOperation(zero.op);
922 var lane0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x8000_0000);
923 try entry.addOperation(lane0.op);
924 var with_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), lane0.getResult(), 0);
925 try entry.addOperation(with_lane0.op);
926 var lane1 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 1);
927 try entry.addOperation(lane1.op);
928 var with_lane1 = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane0.getResult(), lane1.getResult(), 1);
929 try entry.addOperation(with_lane1.op);
930 var lane2 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0xffff_ffff);
931 try entry.addOperation(lane2.op);
932 var with_lane2 = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane1.getResult(), lane2.getResult(), 2);
933 try entry.addOperation(with_lane2.op);
934 var lane3 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 7);
935 try entry.addOperation(lane3.op);
936 var input = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane2.getResult(), lane3.getResult(), 3);
937 try entry.addOperation(input.op);
938
939 const indices = [_]i64{ 3, 2, 1, 0 };
940 const shuffle = try ArithDialect.VecShuffleOp.create(&ir_ctx, loc, input.getResult(), vec_type, indices[0..]);
941 try entry.addOperation(shuffle.op);
942 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{shuffle.getResult()});
943 try entry.addOperation(ret.op);
944
945 var backend = try Backend.init(allocator, &ir_ctx, .testing);
946 defer backend.deinit();
947 const compiled = try backend.compile(module.op);
948
949 const result = try vectorResult(&backend, compiled, "shuffle_vec4u32");
950
951 try testing.expectEqual(artifact.ScalarType.u32, result.element);
952 try testing.expectEqual(@as(u8, 4), result.lanes);
953 try testing.expectEqual(@as(u64, 7), result.bits[0]);
954 try testing.expectEqual(@as(u64, 0xffff_ffff), result.bits[1]);
955 try testing.expectEqual(@as(u64, 1), result.bits[2]);
956 try testing.expectEqual(@as(u64, 0x8000_0000), result.bits[3]);
957 }
958
959 test "x86_64 JIT executes vec4xu32 extract through native packed path" {
960 if (!supports_x86_64_backend) return;
961
962 const testing = std.testing;
963 const dialects = @import("../../dialects/root.zig");
964 const ArithDialect = dialects.ArithDialect;
965 const BuiltinDialect = dialects.BuiltinDialect;
966 const FuncDialect = dialects.FuncDialect;
967
968 var arena = alloc_arena.Arena.init(std.testing.allocator);
969 defer arena.deinit();
970 const allocator = arena.allocator();
971
972 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
973 defer ir_ctx.deinit(allocator);
974 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
975
976 const loc = ir.Location.getUnknown();
977 const u32_type = try ArithDialect.getScalarType(&ir_ctx, .u32);
978 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
979
980 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
981 const module_block = module.getBodyBlock();
982
983 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "extract_vec4u32", &.{}, &.{u32_type});
984 try module_block.addOperation(func.op);
985
986 const entry = func.getEntryBlock();
987 var zero = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 0);
988 try entry.addOperation(zero.op);
989 var lane0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x10);
990 try entry.addOperation(lane0.op);
991 var with_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), lane0.getResult(), 0);
992 try entry.addOperation(with_lane0.op);
993 var lane1 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x20);
994 try entry.addOperation(lane1.op);
995 var with_lane1 = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane0.getResult(), lane1.getResult(), 1);
996 try entry.addOperation(with_lane1.op);
997 var lane2 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x1234_5678);
998 try entry.addOperation(lane2.op);
999 var input = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane1.getResult(), lane2.getResult(), 2);
1000 try entry.addOperation(input.op);
1001 const extract = try ArithDialect.ExtractOp.create(&ir_ctx, loc, input.getResult(), 2, u32_type);
1002 try entry.addOperation(extract.op);
1003
1004 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{extract.getResult()});
1005 try entry.addOperation(ret.op);
1006
1007 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1008 defer backend.deinit();
1009 const compiled = try backend.compile(module.op);
1010
1011 const exec_result = try onlyResult(&backend, compiled, "extract_vec4u32", &.{});
1012 try testing.expectEqual(@as(i64, 0x1234_5678), try integerOf(exec_result));
1013 }
1014
1015 test "x86_64 JIT executes generic arith.neg over vec4xu32 through native packed path" {
1016 if (!supports_x86_64_backend) return;
1017
1018 const testing = std.testing;
1019 const dialects = @import("../../dialects/root.zig");
1020 const ArithDialect = dialects.ArithDialect;
1021 const BuiltinDialect = dialects.BuiltinDialect;
1022 const FuncDialect = dialects.FuncDialect;
1023
1024 var arena = alloc_arena.Arena.init(std.testing.allocator);
1025 defer arena.deinit();
1026 const allocator = arena.allocator();
1027
1028 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1029 defer ir_ctx.deinit(allocator);
1030 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1031
1032 const loc = ir.Location.getUnknown();
1033 const u32_type = try ArithDialect.getScalarType(&ir_ctx, .u32);
1034 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
1035
1036 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1037 const module_block = module.getBodyBlock();
1038
1039 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "generic_neg_vec4u32", &.{}, &.{vec_type});
1040 try module_block.addOperation(func.op);
1041
1042 const entry = func.getEntryBlock();
1043 var zero = try ArithDialect.VecConstantOp.createInt(&ir_ctx, loc, vec_type, 0);
1044 try entry.addOperation(zero.op);
1045 var lane0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x8000_0000);
1046 try entry.addOperation(lane0.op);
1047 var with_lane0 = try ArithDialect.InsertOp.create(&ir_ctx, loc, zero.getResult(), lane0.getResult(), 0);
1048 try entry.addOperation(with_lane0.op);
1049 var lane1 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 1);
1050 try entry.addOperation(lane1.op);
1051 var with_lane1 = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane0.getResult(), lane1.getResult(), 1);
1052 try entry.addOperation(with_lane1.op);
1053 var lane2 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0xffff_ffff);
1054 try entry.addOperation(lane2.op);
1055 var with_lane2 = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane1.getResult(), lane2.getResult(), 2);
1056 try entry.addOperation(with_lane2.op);
1057 var lane3 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 7);
1058 try entry.addOperation(lane3.op);
1059 var input = try ArithDialect.InsertOp.create(&ir_ctx, loc, with_lane2.getResult(), lane3.getResult(), 3);
1060 try entry.addOperation(input.op);
1061
1062 const neg = try ArithDialect.NegOp.create(&ir_ctx, loc, input.getResult());
1063 try entry.addOperation(neg.op);
1064 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{neg.getResult()});
1065 try entry.addOperation(ret.op);
1066
1067 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1068 defer backend.deinit();
1069 const compiled = try backend.compile(module.op);
1070
1071 const result = try vectorResult(&backend, compiled, "generic_neg_vec4u32");
1072
1073 try testing.expectEqual(artifact.ScalarType.u32, result.element);
1074 try testing.expectEqual(@as(u8, 4), result.lanes);
1075 try testing.expectEqual(@as(u64, 0x8000_0000), result.bits[0]);
1076 try testing.expectEqual(@as(u64, 0xffff_ffff), result.bits[1]);
1077 try testing.expectEqual(@as(u64, 1), result.bits[2]);
1078 try testing.expectEqual(@as(u64, 0xffff_fff9), result.bits[3]);
1079 }
1080
1081 test "x86_64 JIT executes vec4xf32 memref add through native packed path" {
1082 if (!supports_x86_64_backend) return;
1083
1084 const testing = std.testing;
1085 const dialects = @import("../../dialects/root.zig");
1086 const ArithDialect = dialects.ArithDialect;
1087 const BuiltinDialect = dialects.BuiltinDialect;
1088 const FuncDialect = dialects.FuncDialect;
1089 const MemrefDialect = dialects.MemrefDialect;
1090
1091 var arena = alloc_arena.Arena.init(std.testing.allocator);
1092 defer arena.deinit();
1093 const allocator = arena.allocator();
1094
1095 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1096 defer ir_ctx.deinit(allocator);
1097 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1098
1099 const loc = ir.Location.getUnknown();
1100 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
1101 const index_type = try ArithDialect.getIndexType(&ir_ctx);
1102 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.float32)).?;
1103 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, f32_type, .host);
1104
1105 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1106 const module_block = module.getBodyBlock();
1107
1108 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "memref_vec4f32_add", &.{ memref_type, memref_type, memref_type }, &.{});
1109 try module_block.addOperation(func.op);
1110
1111 const entry = func.getEntryBlock();
1112 var idx0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
1113 try entry.addOperation(idx0.op);
1114
1115 var lhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(1), idx0.getResult(), vec_type);
1116 try entry.addOperation(lhs.op);
1117 var rhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(2), idx0.getResult(), vec_type);
1118 try entry.addOperation(rhs.op);
1119 const add = try ArithDialect.AddOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
1120 try entry.addOperation(add.op);
1121 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, add.getResult(), func.getArgument(0), idx0.getResult());
1122 try entry.addOperation(store.op);
1123 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{});
1124 try entry.addOperation(ret.op);
1125
1126 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1127 defer backend.deinit();
1128 const compiled = try backend.compile(module.op);
1129
1130 var dst = [_]f32{ 0.0, 0.0, 0.0, 0.0 };
1131 var lhs_data = [_]f32{ 1.0, -2.5, 3.25, 10.0 };
1132 var rhs_data = [_]f32{ 4.0, 2.0, -1.25, -11.5 };
1133
1134 try backend.runtime.call(compiled, "memref_vec4f32_add", &.{
1135 .{ .memref = @intFromPtr(&dst[0]) },
1136 .{ .memref = @intFromPtr(&lhs_data[0]) },
1137 .{ .memref = @intFromPtr(&rhs_data[0]) },
1138 }, &.{});
1139
1140 try testing.expectApproxEqAbs(@as(f32, 5.0), dst[0], 1e-6);
1141 try testing.expectApproxEqAbs(@as(f32, -0.5), dst[1], 1e-6);
1142 try testing.expectApproxEqAbs(@as(f32, 2.0), dst[2], 1e-6);
1143 try testing.expectApproxEqAbs(@as(f32, -1.5), dst[3], 1e-6);
1144 }
1145
1146 test "x86_64 JIT executes vec4xu32 memref add through native packed path" {
1147 if (!supports_x86_64_backend) return;
1148
1149 const testing = std.testing;
1150 const dialects = @import("../../dialects/root.zig");
1151 const ArithDialect = dialects.ArithDialect;
1152 const BuiltinDialect = dialects.BuiltinDialect;
1153 const FuncDialect = dialects.FuncDialect;
1154 const MemrefDialect = dialects.MemrefDialect;
1155
1156 var arena = alloc_arena.Arena.init(std.testing.allocator);
1157 defer arena.deinit();
1158 const allocator = arena.allocator();
1159
1160 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1161 defer ir_ctx.deinit(allocator);
1162 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1163
1164 const loc = ir.Location.getUnknown();
1165 const u32_type = try ArithDialect.getScalarType(&ir_ctx, .u32);
1166 const index_type = try ArithDialect.getIndexType(&ir_ctx);
1167 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 4, dialects.arith.type_names.uint32)).?;
1168 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, u32_type, .host);
1169
1170 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1171 const module_block = module.getBodyBlock();
1172
1173 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "memref_vec4u32_add", &.{ memref_type, memref_type, memref_type }, &.{});
1174 try module_block.addOperation(func.op);
1175
1176 const entry = func.getEntryBlock();
1177 var idx0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
1178 try entry.addOperation(idx0.op);
1179
1180 var lhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(1), idx0.getResult(), vec_type);
1181 try entry.addOperation(lhs.op);
1182 var rhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(2), idx0.getResult(), vec_type);
1183 try entry.addOperation(rhs.op);
1184 const add = try ArithDialect.AddOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
1185 try entry.addOperation(add.op);
1186 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, add.getResult(), func.getArgument(0), idx0.getResult());
1187 try entry.addOperation(store.op);
1188 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{});
1189 try entry.addOperation(ret.op);
1190
1191 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1192 defer backend.deinit();
1193 const compiled = try backend.compile(module.op);
1194
1195 var dst = [_]u32{ 0, 0, 0, 0 };
1196 var lhs_data = [_]u32{ 1, 0x8000_0000, 0xffff_ffff, 7 };
1197 var rhs_data = [_]u32{ 4, 0x8000_0000, 2, 0xffff_fffe };
1198
1199 try backend.runtime.call(compiled, "memref_vec4u32_add", &.{
1200 .{ .memref = @intFromPtr(&dst[0]) },
1201 .{ .memref = @intFromPtr(&lhs_data[0]) },
1202 .{ .memref = @intFromPtr(&rhs_data[0]) },
1203 }, &.{});
1204
1205 try testing.expectEqualSlices(u32, &.{ 5, 0, 1, 5 }, dst[0..]);
1206 }
1207
1208 test "x86_64 JIT executes vec2xf64 memref add through native packed path" {
1209 if (!supports_x86_64_backend) return;
1210
1211 const testing = std.testing;
1212 const dialects = @import("../../dialects/root.zig");
1213 const ArithDialect = dialects.ArithDialect;
1214 const BuiltinDialect = dialects.BuiltinDialect;
1215 const FuncDialect = dialects.FuncDialect;
1216 const MemrefDialect = dialects.MemrefDialect;
1217
1218 var arena = alloc_arena.Arena.init(std.testing.allocator);
1219 defer arena.deinit();
1220 const allocator = arena.allocator();
1221
1222 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1223 defer ir_ctx.deinit(allocator);
1224 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1225
1226 const loc = ir.Location.getUnknown();
1227 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
1228 const index_type = try ArithDialect.getIndexType(&ir_ctx);
1229 const vec_type = (try ArithDialect.getVecType(&ir_ctx, 2, dialects.arith.type_names.float64)).?;
1230 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, f64_type, .host);
1231
1232 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1233 const module_block = module.getBodyBlock();
1234
1235 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "memref_vec2f64_add", &.{ memref_type, memref_type, memref_type }, &.{});
1236 try module_block.addOperation(func.op);
1237
1238 const entry = func.getEntryBlock();
1239 var idx0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
1240 try entry.addOperation(idx0.op);
1241
1242 var lhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(1), idx0.getResult(), vec_type);
1243 try entry.addOperation(lhs.op);
1244 var rhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(2), idx0.getResult(), vec_type);
1245 try entry.addOperation(rhs.op);
1246 const add = try ArithDialect.AddOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
1247 try entry.addOperation(add.op);
1248 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, add.getResult(), func.getArgument(0), idx0.getResult());
1249 try entry.addOperation(store.op);
1250 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{});
1251 try entry.addOperation(ret.op);
1252
1253 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1254 defer backend.deinit();
1255 const compiled = try backend.compile(module.op);
1256
1257 var dst = [_]f64{ 0.0, 0.0 };
1258 var lhs_data = [_]f64{ 1.25, -3.5 };
1259 var rhs_data = [_]f64{ 2.75, 1.0 };
1260
1261 try backend.runtime.call(compiled, "memref_vec2f64_add", &.{
1262 .{ .memref = @intFromPtr(&dst[0]) },
1263 .{ .memref = @intFromPtr(&lhs_data[0]) },
1264 .{ .memref = @intFromPtr(&rhs_data[0]) },
1265 }, &.{});
1266
1267 try testing.expectApproxEqAbs(@as(f64, 4.0), dst[0], 1e-12);
1268 try testing.expectApproxEqAbs(@as(f64, -2.5), dst[1], 1e-12);
1269 }
1270
1271 test "x86_64 external call passes stack args with aligned call frame" {
1272 if (!supports_x86_64_backend) return;
1273
1274 const testing = std.testing;
1275 const dialects = @import("../../dialects/root.zig");
1276 const ArithDialect = dialects.ArithDialect;
1277 const BuiltinDialect = dialects.BuiltinDialect;
1278 const FuncDialect = dialects.FuncDialect;
1279
1280 var arena = alloc_arena.Arena.init(std.testing.allocator);
1281 defer arena.deinit();
1282 const allocator = arena.allocator();
1283
1284 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1285 defer ir_ctx.deinit(allocator);
1286 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1287
1288 const loc = ir.Location.getUnknown();
1289 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1290
1291 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1292 const module_block = module.getBodyBlock();
1293
1294 const sum7_decl = try FuncDialect.FuncOp.createDeclaration(
1295 &ir_ctx,
1296 loc,
1297 "choir_sum7",
1298 &.{ i64_type, i64_type, i64_type, i64_type, i64_type, i64_type, i64_type },
1299 &.{i64_type},
1300 );
1301 try module_block.addOperation(sum7_decl.op);
1302
1303 var func = try FuncDialect.FuncOp.create(
1304 &ir_ctx,
1305 loc,
1306 "main",
1307 &.{},
1308 &.{i64_type},
1309 );
1310 try module_block.addOperation(func.op);
1311
1312 const entry = func.getEntryBlock();
1313 var c1 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
1314 try entry.addOperation(c1.op);
1315 var c2 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 2);
1316 try entry.addOperation(c2.op);
1317 var c3 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 3);
1318 try entry.addOperation(c3.op);
1319 var c4 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 4);
1320 try entry.addOperation(c4.op);
1321 var c5 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 5);
1322 try entry.addOperation(c5.op);
1323 var c6 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 6);
1324 try entry.addOperation(c6.op);
1325 var c7 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 7);
1326 try entry.addOperation(c7.op);
1327
1328 var call_op = try FuncDialect.CallOp.create(
1329 &ir_ctx,
1330 loc,
1331 "choir_sum7",
1332 &.{
1333 c1.getResult(),
1334 c2.getResult(),
1335 c3.getResult(),
1336 c4.getResult(),
1337 c5.getResult(),
1338 c6.getResult(),
1339 c7.getResult(),
1340 },
1341 &.{i64_type},
1342 );
1343 try entry.addOperation(call_op.op);
1344
1345 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{call_op.getResult(0).?});
1346 try entry.addOperation(ret.op);
1347
1348 var emitter = x86_64.Emitter.init(allocator);
1349 defer emitter.deinit();
1350 try emitter.emitFunction(func.op);
1351 try testing.expectEqual(x86_64.GPR.rdi, emitter.value_locations.get(c4.getResult()).?);
1352
1353 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1354 defer backend.deinit();
1355
1356 try backend.runtime.registerExternalSymbol("choir_sum7", @intFromPtr(&choirSum7));
1357 const compiled = try backend.compile(module.op);
1358
1359 const exec_result = try onlyResult(&backend, compiled, "main", &.{});
1360 try testing.expectEqual(@as(i64, 28), try integerOf(exec_result));
1361 }
1362
1363 test "x86_64 external call uses register homes for scalar args and results" {
1364 if (!supports_x86_64_backend) return;
1365
1366 const testing = std.testing;
1367 const dialects = @import("../../dialects/root.zig");
1368 const ArithDialect = dialects.ArithDialect;
1369 const BuiltinDialect = dialects.BuiltinDialect;
1370 const FuncDialect = dialects.FuncDialect;
1371
1372 var arena = alloc_arena.Arena.init(std.testing.allocator);
1373 defer arena.deinit();
1374 const allocator = arena.allocator();
1375
1376 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1377 defer ir_ctx.deinit(allocator);
1378 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1379
1380 const loc = ir.Location.getUnknown();
1381 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1382
1383 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1384 const module_block = module.getBodyBlock();
1385
1386 const decl_name = "choir_regalloc_add_five";
1387 const declaration = try FuncDialect.FuncOp.createDeclaration(
1388 &ir_ctx,
1389 loc,
1390 decl_name,
1391 &.{i64_type},
1392 &.{i64_type},
1393 );
1394 try module_block.addOperation(declaration.op);
1395
1396 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "call_regalloc_homes", &.{i64_type}, &.{i64_type});
1397 try module_block.addOperation(func.op);
1398
1399 const entry = func.getEntryBlock();
1400 var one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
1401 try entry.addOperation(one.op);
1402 var shifted = try ArithDialect.AddOp.create(&ir_ctx, loc, func.getArgument(0), one.getResult());
1403 try entry.addOperation(shifted.op);
1404 var call = try FuncDialect.CallOp.create(&ir_ctx, loc, decl_name, &.{shifted.getResult()}, &.{i64_type});
1405 try entry.addOperation(call.op);
1406 const call_result = call.getResult(0) orelse return error.TestFailure;
1407 var two = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 2);
1408 try entry.addOperation(two.op);
1409 var result = try ArithDialect.AddOp.create(&ir_ctx, loc, call_result, two.getResult());
1410 try entry.addOperation(result.op);
1411 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result.getResult()});
1412 try entry.addOperation(ret.op);
1413
1414 var emitter = x86_64.Emitter.init(allocator);
1415 defer emitter.deinit();
1416 try emitter.emitFunction(func.op);
1417 try testing.expect(emitter.value_locations.contains(shifted.getResult()));
1418 try testing.expect(emitter.value_locations.contains(call_result));
1419
1420 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1421 defer backend.deinit();
1422
1423 try backend.runtime.registerExternalSymbol(decl_name, @intFromPtr(&choirRegallocAddFive));
1424 const compiled = try backend.compile(module.op);
1425
1426 const exec_result = try onlyResult(
1427 &backend,
1428 compiled,
1429 "call_regalloc_homes",
1430 &.{.{ .i64 = 10 }},
1431 );
1432 try testing.expectEqual(@as(i64, 18), try integerOf(exec_result));
1433 }
1434
1435 test "x86_64 backend links externs and exposes compiled symbols" {
1436 if (!supports_x86_64_backend) return;
1437
1438 const testing = std.testing;
1439 const dialects = @import("../../dialects/root.zig");
1440 const ArithDialect = dialects.ArithDialect;
1441 const BuiltinDialect = dialects.BuiltinDialect;
1442 const FuncDialect = dialects.FuncDialect;
1443
1444 var arena = alloc_arena.Arena.init(std.testing.allocator);
1445 defer arena.deinit();
1446 const allocator = arena.allocator();
1447
1448 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1449 defer ir_ctx.deinit(allocator);
1450 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1451
1452 const loc = ir.Location.getUnknown();
1453 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1454
1455 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1456 const module_block = module.getBodyBlock();
1457
1458 const decl_name = "choir_handle_add1";
1459 const declaration = try FuncDialect.FuncOp.createDeclaration(
1460 &ir_ctx,
1461 loc,
1462 decl_name,
1463 &.{i64_type},
1464 &.{i64_type},
1465 );
1466 try module_block.addOperation(declaration.op);
1467
1468 var caller = try FuncDialect.FuncOp.create(&ir_ctx, loc, "call_handle_decl", &.{}, &.{i64_type});
1469 try module_block.addOperation(caller.op);
1470
1471 const entry = caller.getEntryBlock();
1472 var forty_one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 41);
1473 try entry.addOperation(forty_one.op);
1474
1475 var call = try FuncDialect.CallOp.create(&ir_ctx, loc, decl_name, &.{forty_one.getResult()}, &.{i64_type});
1476 try entry.addOperation(call.op);
1477
1478 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{call.getResult(0).?});
1479 try entry.addOperation(ret.op);
1480
1481 var handle = try initHandle(allocator, &ir_ctx);
1482 defer handle.deinit();
1483 try testing.expect(interface.BackendTarget.x86_64.eql(handle.target));
1484 try testing.expect(handle.capabilities.artifact.object_file);
1485
1486 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1487 defer backend.deinit();
1488 try backend.runtime.registerExternalSymbol(decl_name, @intFromPtr(&choirHandleAdd1));
1489 const compiled = try backend.compile(module.op);
1490
1491 const address = try backend.runtime.functionAddress(compiled, "call_handle_decl");
1492 try testing.expect(address != 0);
1493
1494 const exec_result = try onlyResult(&backend, compiled, "call_handle_decl", &.{});
1495 try testing.expectEqual(@as(i64, 42), try integerOf(exec_result));
1496 }
1497
1498 test "x86_64 backend materializes serialized machine code artifacts" {
1499 if (!supports_x86_64_backend) return;
1500
1501 const testing = std.testing;
1502 const dialects = @import("../../dialects/root.zig");
1503 const ArithDialect = dialects.ArithDialect;
1504 const BuiltinDialect = dialects.BuiltinDialect;
1505 const FuncDialect = dialects.FuncDialect;
1506
1507 var arena = alloc_arena.Arena.init(testing.allocator);
1508 defer arena.deinit();
1509 const ir_allocator = arena.allocator();
1510 var ir_ctx = try ir.Context.init(ir_allocator, ir.Context.Limits.testing);
1511 defer ir_ctx.deinit(ir_allocator);
1512 try dialects.registerAllDialects(&ir_ctx);
1513
1514 const location = ir.Location.getUnknown();
1515 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1516 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, location);
1517 const module_block = module.getBodyBlock();
1518 const external_symbol = "choir_artifact_add_five";
1519 const declaration = try FuncDialect.FuncOp.createDeclaration(
1520 &ir_ctx,
1521 location,
1522 external_symbol,
1523 &.{i64_type},
1524 &.{i64_type},
1525 );
1526 try module_block.addOperation(declaration.op);
1527 const export_symbol = "artifact_call_external";
1528 var function = try FuncDialect.FuncOp.create(
1529 &ir_ctx,
1530 location,
1531 export_symbol,
1532 &.{i64_type},
1533 &.{i64_type},
1534 );
1535 try module_block.addOperation(function.op);
1536 const entry = function.getEntryBlock();
1537 var call = try FuncDialect.CallOp.create(
1538 &ir_ctx,
1539 location,
1540 external_symbol,
1541 &.{function.getArgument(0)},
1542 &.{i64_type},
1543 );
1544 try entry.addOperation(call.op);
1545 const return_op = try FuncDialect.ReturnOp.create(&ir_ctx, location, &.{call.getResult(0).?});
1546 try entry.addOperation(return_op.op);
1547
1548 var producer = try Backend.init(testing.allocator, &ir_ctx, .testing);
1549 defer producer.deinit();
1550 var serialized = try producer.compileFunctionToArtifact(module.op, export_symbol);
1551 var serialized_owned = true;
1552 defer if (serialized_owned) serialized.deinit();
1553 try testing.expectEqual(artifact.ArtifactKind.machine_code, serialized.metadata.kind);
1554 try testing.expect(serialized.linkage.hasProvided(export_symbol));
1555 try testing.expectEqual(@as(usize, 1), serialized.linkage.relocations.items.len);
1556
1557 var consumer = x86_64.JitRuntime.init(testing.allocator, .testing);
1558 defer consumer.deinit();
1559 try consumer.registerExternalSymbol(external_symbol, @intFromPtr(&choirRegallocAddFive));
1560 const loaded = try consumer.loadMachineCodeArtifact(&serialized);
1561 const reloaded = try consumer.loadMachineCodeArtifact(&serialized);
1562 const address = try consumer.functionAddress(loaded, export_symbol);
1563 try testing.expect(address != try consumer.functionAddress(reloaded, export_symbol));
1564
1565 serialized.deinit();
1566 serialized_owned = false;
1567 try consumer.release(reloaded);
1568 const loaded_entry: *const fn (i64) callconv(.c) i64 = @ptrFromInt(address);
1569 try testing.expectEqual(@as(i64, 42), loaded_entry(37));
1570 }
1571
1572 test "x86_64 backend canonicalizes extern bool results" {
1573 if (!supports_x86_64_backend) return;
1574
1575 const testing = std.testing;
1576 const dialects = @import("../../dialects/root.zig");
1577 const ArithDialect = dialects.ArithDialect;
1578 const BuiltinDialect = dialects.BuiltinDialect;
1579 const FuncDialect = dialects.FuncDialect;
1580 const ScfDialect = dialects.ScfDialect;
1581
1582 var arena = alloc_arena.Arena.init(std.testing.allocator);
1583 defer arena.deinit();
1584 const allocator = arena.allocator();
1585
1586 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1587 defer ir_ctx.deinit(allocator);
1588 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1589
1590 const loc = ir.Location.getUnknown();
1591 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
1592 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1593
1594 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1595 const module_block = module.getBodyBlock();
1596
1597 const decl_name = "choir_handle_false_dirty";
1598 const declaration = try FuncDialect.FuncOp.createDeclaration(&ir_ctx, loc, decl_name, &.{}, &.{bool_type});
1599 try module_block.addOperation(declaration.op);
1600
1601 var caller = try FuncDialect.FuncOp.create(&ir_ctx, loc, "call_dirty_bool_decl", &.{}, &.{i64_type});
1602 try module_block.addOperation(caller.op);
1603
1604 const entry = caller.getEntryBlock();
1605 var call = try FuncDialect.CallOp.create(&ir_ctx, loc, decl_name, &.{}, &.{bool_type});
1606 try entry.addOperation(call.op);
1607
1608 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, call.getResult(0).?, &.{i64_type});
1609 try entry.addOperation(if_op.op);
1610
1611 var one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
1612 try if_op.getThenBlock().addOperation(one.op);
1613 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{one.getResult()});
1614 try if_op.getThenBlock().addOperation(then_yield.op);
1615
1616 var two = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 2);
1617 try if_op.getElseBlock().?.addOperation(two.op);
1618 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{two.getResult()});
1619 try if_op.getElseBlock().?.addOperation(else_yield.op);
1620
1621 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{if_op.getResult(0).?});
1622 try entry.addOperation(ret.op);
1623
1624 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1625 defer backend.deinit();
1626 try backend.runtime.registerExternalSymbol(decl_name, @intFromPtr(&choirHandleFalseWithDirtyHighBits));
1627 const compiled = try backend.compile(module.op);
1628
1629 const exec_result = try onlyResult(&backend, compiled, "call_dirty_bool_decl", &.{});
1630 try testing.expectEqual(@as(i64, 2), try integerOf(exec_result));
1631 }
1632
1633 test "x86_64 register allocation preserves scf if branch homes" {
1634 if (!supports_x86_64_backend) return;
1635
1636 const testing = std.testing;
1637 const dialects = @import("../../dialects/root.zig");
1638 const ArithDialect = dialects.ArithDialect;
1639 const BuiltinDialect = dialects.BuiltinDialect;
1640 const FuncDialect = dialects.FuncDialect;
1641 const ScfDialect = dialects.ScfDialect;
1642
1643 var arena = alloc_arena.Arena.init(std.testing.allocator);
1644 defer arena.deinit();
1645 const allocator = arena.allocator();
1646
1647 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1648 defer ir_ctx.deinit(allocator);
1649 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1650
1651 const loc = ir.Location.getUnknown();
1652 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
1653 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1654
1655 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1656 const module_block = module.getBodyBlock();
1657
1658 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "branch_regalloc", &.{bool_type}, &.{i64_type});
1659 try module_block.addOperation(func.op);
1660
1661 const entry = func.getEntryBlock();
1662 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, func.getArgument(0), &.{i64_type});
1663 try entry.addOperation(if_op.op);
1664
1665 const then_block = if_op.getThenBlock();
1666 var then_a = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 10);
1667 try then_block.addOperation(then_a.op);
1668 var then_b = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 32);
1669 try then_block.addOperation(then_b.op);
1670 var then_sum = try ArithDialect.AddOp.create(&ir_ctx, loc, then_a.getResult(), then_b.getResult());
1671 try then_block.addOperation(then_sum.op);
1672 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{then_sum.getResult()});
1673 try then_block.addOperation(then_yield.op);
1674
1675 const else_block = if_op.getElseBlock().?;
1676 var else_a = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 100);
1677 try else_block.addOperation(else_a.op);
1678 var else_b = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 23);
1679 try else_block.addOperation(else_b.op);
1680 var else_sum = try ArithDialect.AddOp.create(&ir_ctx, loc, else_a.getResult(), else_b.getResult());
1681 try else_block.addOperation(else_sum.op);
1682 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{else_sum.getResult()});
1683 try else_block.addOperation(else_yield.op);
1684
1685 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{if_op.getResult(0).?});
1686 try entry.addOperation(ret.op);
1687
1688 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1689 defer backend.deinit();
1690
1691 const compiled = try backend.compile(module.op);
1692
1693 const true_result = try onlyResult(
1694 &backend,
1695 compiled,
1696 "branch_regalloc",
1697 &.{.{ .bool = true }},
1698 );
1699 try testing.expectEqual(@as(i64, 42), try integerOf(true_result));
1700
1701 const false_result = try onlyResult(
1702 &backend,
1703 compiled,
1704 "branch_regalloc",
1705 &.{.{ .bool = false }},
1706 );
1707 try testing.expectEqual(@as(i64, 123), try integerOf(false_result));
1708 }
1709
1710 test "x86_64 scf.if yields resolve register-home swaps" {
1711 if (!supports_x86_64_backend) return;
1712
1713 const testing = std.testing;
1714 const dialects = @import("../../dialects/root.zig");
1715 const ArithDialect = dialects.ArithDialect;
1716 const BuiltinDialect = dialects.BuiltinDialect;
1717 const FuncDialect = dialects.FuncDialect;
1718 const ScfDialect = dialects.ScfDialect;
1719
1720 var arena = alloc_arena.Arena.init(std.testing.allocator);
1721 defer arena.deinit();
1722 const allocator = arena.allocator();
1723
1724 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1725 defer ir_ctx.deinit(allocator);
1726 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1727
1728 const loc = ir.Location.getUnknown();
1729 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
1730 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1731
1732 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1733 const module_block = module.getBodyBlock();
1734
1735 var func = try FuncDialect.FuncOp.create(
1736 &ir_ctx,
1737 loc,
1738 "branch_swap_regalloc",
1739 &.{ bool_type, i64_type, i64_type, i64_type },
1740 &.{i64_type},
1741 );
1742 try module_block.addOperation(func.op);
1743
1744 const entry = func.getEntryBlock();
1745 const cond = func.getArgument(0);
1746 const keep = func.getArgument(1);
1747 const lhs = func.getArgument(2);
1748 const rhs = func.getArgument(3);
1749
1750 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, cond, &.{ i64_type, i64_type });
1751 try entry.addOperation(if_op.op);
1752
1753 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{ rhs, lhs });
1754 try if_op.getThenBlock().addOperation(then_yield.op);
1755
1756 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{ lhs, rhs });
1757 try if_op.getElseBlock().?.addOperation(else_yield.op);
1758
1759 const first = if_op.getResult(0) orelse return error.TestFailure;
1760 const second = if_op.getResult(1) orelse return error.TestFailure;
1761 var diff = try ArithDialect.SubOp.create(&ir_ctx, loc, first, second);
1762 try entry.addOperation(diff.op);
1763 var result = try ArithDialect.AddOp.create(&ir_ctx, loc, diff.getResult(), keep);
1764 try entry.addOperation(result.op);
1765 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result.getResult()});
1766 try entry.addOperation(ret.op);
1767
1768 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1769 defer backend.deinit();
1770
1771 const compiled = try backend.compile(module.op);
1772
1773 const true_result = try onlyResult(&backend, compiled, "branch_swap_regalloc", &.{
1774 .{ .bool = true }, .{ .i64 = 100 }, .{ .i64 = 10 }, .{ .i64 = 3 },
1775 });
1776 try testing.expectEqual(@as(i64, 93), try integerOf(true_result));
1777
1778 const false_result = try onlyResult(&backend, compiled, "branch_swap_regalloc", &.{
1779 .{ .bool = false }, .{ .i64 = 100 }, .{ .i64 = 10 }, .{ .i64 = 3 },
1780 });
1781 try testing.expectEqual(@as(i64, 107), try integerOf(false_result));
1782 }
1783
1784 test "x86_64 scf.while zero exit spills live merge values" {
1785 if (!supports_x86_64_backend) return;
1786
1787 const testing = std.testing;
1788 const dialects = @import("../../dialects/root.zig");
1789 const ArithDialect = dialects.ArithDialect;
1790 const BuiltinDialect = dialects.BuiltinDialect;
1791 const FuncDialect = dialects.FuncDialect;
1792 const ScfDialect = dialects.ScfDialect;
1793
1794 const Helpers = struct {
1795 fn countdownWhile(
1796 ctx: *ir.Context,
1797 loc: ir.Location,
1798 i64_type: ir.Type,
1799 entry: *ir.Block,
1800 initial_count: *ir.Value,
1801 initial_acc: *ir.Value,
1802 ) !ScfDialect.WhileOp {
1803 var while_op = try ScfDialect.WhileOp.create(ctx, loc, &.{ initial_count, initial_acc }, &.{ i64_type, i64_type });
1804 try entry.addOperation(while_op.op);
1805
1806 const before = while_op.getBeforeBlock();
1807 const before_count = before.arguments.items[0];
1808 const before_acc = before.arguments.items[1];
1809 var before_zero = try ArithDialect.ConstantOp.createInt(ctx, loc, i64_type, 0);
1810 try before.addOperation(before_zero.op);
1811 var keep_going = try ArithDialect.CmpOp.create(ctx, loc, .gt, before_count, before_zero.getResult());
1812 try before.addOperation(keep_going.op);
1813 const condition = try ScfDialect.ConditionOp.create(ctx, loc, keep_going.getResult(), &.{ before_count, before_acc });
1814 try before.addOperation(condition.op);
1815
1816 const after = while_op.getAfterBlock();
1817 const after_count = after.arguments.items[0];
1818 const after_acc = after.arguments.items[1];
1819 var one = try ArithDialect.ConstantOp.createInt(ctx, loc, i64_type, 1);
1820 try after.addOperation(one.op);
1821 var next_count = try ArithDialect.SubOp.create(ctx, loc, after_count, one.getResult());
1822 try after.addOperation(next_count.op);
1823 var next_acc = try ArithDialect.AddOp.create(ctx, loc, after_acc, after_count);
1824 try after.addOperation(next_acc.op);
1825 const yield = try ScfDialect.YieldOp.create(ctx, loc, &.{ next_count.getResult(), next_acc.getResult() });
1826 try after.addOperation(yield.op);
1827
1828 return while_op;
1829 }
1830 };
1831
1832 var arena = alloc_arena.Arena.init(std.testing.allocator);
1833 defer arena.deinit();
1834 const allocator = arena.allocator();
1835
1836 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1837 defer ir_ctx.deinit(allocator);
1838 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1839
1840 const loc = ir.Location.getUnknown();
1841 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1842
1843 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1844 const module_block = module.getBodyBlock();
1845
1846 var helper = try FuncDialect.FuncOp.create(&ir_ctx, loc, "while_merge_helper", &.{ i64_type, i64_type, i64_type, i64_type }, &.{i64_type});
1847 try module_block.addOperation(helper.op);
1848 const helper_entry = helper.getEntryBlock();
1849 var helper_sum = try ArithDialect.AddOp.create(&ir_ctx, loc, helper.getArgument(0), helper.getArgument(1));
1850 try helper_entry.addOperation(helper_sum.op);
1851 var helper_xor = try ArithDialect.XorOp.create(&ir_ctx, loc, helper.getArgument(2), helper.getArgument(3));
1852 try helper_entry.addOperation(helper_xor.op);
1853 var helper_three = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 3);
1854 try helper_entry.addOperation(helper_three.op);
1855 var helper_scaled = try ArithDialect.MulOp.create(&ir_ctx, loc, helper_sum.getResult(), helper_three.getResult());
1856 try helper_entry.addOperation(helper_scaled.op);
1857 var helper_result = try ArithDialect.SubOp.create(&ir_ctx, loc, helper_scaled.getResult(), helper_xor.getResult());
1858 try helper_entry.addOperation(helper_result.op);
1859 const helper_ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{helper_result.getResult()});
1860 try helper_entry.addOperation(helper_ret.op);
1861
1862 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "while_merge_spill", &.{}, &.{i64_type});
1863 try module_block.addOperation(func.op);
1864 const entry = func.getEntryBlock();
1865
1866 var seed = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, -32768);
1867 try entry.addOperation(seed.op);
1868 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
1869 try entry.addOperation(zero.op);
1870 var first_while = try Helpers.countdownWhile(&ir_ctx, loc, i64_type, entry, zero.getResult(), seed.getResult());
1871 var first_call = try FuncDialect.CallOp.create(&ir_ctx, loc, "while_merge_helper", &.{ seed.getResult(), seed.getResult(), first_while.op.getResult(0).?, seed.getResult() }, &.{i64_type});
1872 try entry.addOperation(first_call.op);
1873 var neg_two = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, -2);
1874 try entry.addOperation(neg_two.op);
1875 var div = try ArithDialect.DivOp.create(&ir_ctx, loc, seed.getResult(), neg_two.getResult());
1876 try entry.addOperation(div.op);
1877 var second_zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
1878 try entry.addOperation(second_zero.op);
1879 var second_while = try Helpers.countdownWhile(&ir_ctx, loc, i64_type, entry, second_zero.getResult(), neg_two.getResult());
1880 var second_call = try FuncDialect.CallOp.create(&ir_ctx, loc, "while_merge_helper", &.{ first_while.op.getResult(0).?, second_while.op.getResult(0).?, first_call.getResult(0).?, first_while.op.getResult(1).? }, &.{i64_type});
1881 try entry.addOperation(second_call.op);
1882 var salt = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, -21051);
1883 try entry.addOperation(salt.op);
1884 const min = try ArithDialect.MinOp.create(&ir_ctx, loc, seed.getResult(), second_while.op.getResult(1).?);
1885 try entry.addOperation(min.op);
1886 const cmp = try ArithDialect.CmpOp.create(&ir_ctx, loc, .ne, first_while.op.getResult(0).?, first_call.getResult(0).?);
1887 try entry.addOperation(cmp.op);
1888 var thousand = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, -1000);
1889 try entry.addOperation(thousand.op);
1890 var rem = try ArithDialect.RemOp.create(&ir_ctx, loc, second_while.op.getResult(1).?, thousand.getResult());
1891 try entry.addOperation(rem.op);
1892 var third_call = try FuncDialect.CallOp.create(&ir_ctx, loc, "while_merge_helper", &.{ second_while.op.getResult(1).?, second_call.getResult(0).?, first_call.getResult(0).?, first_while.op.getResult(1).? }, &.{i64_type});
1893 try entry.addOperation(third_call.op);
1894 var shift_count = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 42);
1895 try entry.addOperation(shift_count.op);
1896 const shifted = try ArithDialect.ShlOp.create(&ir_ctx, loc, rem.getResult(), shift_count.getResult());
1897 try entry.addOperation(shifted.op);
1898 var final_call = try FuncDialect.CallOp.create(&ir_ctx, loc, "while_merge_helper", &.{ div.getResult(), third_call.getResult(0).?, salt.getResult(), second_call.getResult(0).? }, &.{i64_type});
1899 try entry.addOperation(final_call.op);
1900 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{final_call.getResult(0).?});
1901 try entry.addOperation(ret.op);
1902
1903 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1904 defer backend.deinit();
1905
1906 const compiled = try backend.compile(module.op);
1907 const result = try onlyResult(&backend, compiled, "while_merge_spill", &.{});
1908 try testing.expectEqual(@as(i64, -1633751), try integerOf(result));
1909 }
1910
1911 test "x86_64 memref alloc returns sane pointer" {
1912 if (!supports_x86_64_backend) return;
1913
1914 const testing = std.testing;
1915 const dialects = @import("../../dialects/root.zig");
1916 const ArithDialect = dialects.ArithDialect;
1917 const BuiltinDialect = dialects.BuiltinDialect;
1918 const FuncDialect = dialects.FuncDialect;
1919 const MemrefDialect = dialects.MemrefDialect;
1920
1921 var arena = alloc_arena.Arena.init(std.testing.allocator);
1922 defer arena.deinit();
1923 const allocator = arena.allocator();
1924
1925 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1926 defer ir_ctx.deinit(allocator);
1927 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1928
1929 const loc = ir.Location.getUnknown();
1930 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1931 const index_type = try ArithDialect.getIndexType(&ir_ctx);
1932 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, i64_type, .host);
1933
1934 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1935 const module_block = module.getBodyBlock();
1936
1937 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "alloc_only", &.{}, &.{memref_type});
1938 try module_block.addOperation(func.op);
1939
1940 const entry = func.getEntryBlock();
1941 var size = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 4);
1942 try entry.addOperation(size.op);
1943
1944 var alloc = try MemrefDialect.AllocOp.createDynamic(&ir_ctx, loc, size.getResult(), memref_type);
1945 try entry.addOperation(alloc.op);
1946
1947 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{alloc.getResult()});
1948 try entry.addOperation(ret.op);
1949
1950 var backend = try Backend.init(allocator, &ir_ctx, .testing);
1951 defer backend.deinit();
1952 const compiled = try backend.compile(module.op);
1953
1954 const exec_result = try onlyResult(&backend, compiled, "alloc_only", &.{});
1955 const ptr_int = exec_result.memref;
1956 try testing.expect(ptr_int != 0);
1957 try testing.expect(@as(u64, @intCast(ptr_int)) > 0x10000);
1958 try testing.expect(@as(u64, @intCast(ptr_int)) % 8 == 0);
1959
1960 sys.heap.freeAddress(@as(usize, @intCast(ptr_int)));
1961 }
1962
1963 test "x86_64 memref alloc/free roundtrip" {
1964 if (!supports_x86_64_backend) return;
1965
1966 const testing = std.testing;
1967 const dialects = @import("../../dialects/root.zig");
1968 const ArithDialect = dialects.ArithDialect;
1969 const BuiltinDialect = dialects.BuiltinDialect;
1970 const FuncDialect = dialects.FuncDialect;
1971 const MemrefDialect = dialects.MemrefDialect;
1972
1973 var arena = alloc_arena.Arena.init(std.testing.allocator);
1974 defer arena.deinit();
1975 const allocator = arena.allocator();
1976
1977 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
1978 defer ir_ctx.deinit(allocator);
1979 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
1980
1981 const loc = ir.Location.getUnknown();
1982 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
1983 const index_type = try ArithDialect.getIndexType(&ir_ctx);
1984 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, i64_type, .host);
1985
1986 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
1987 const module_block = module.getBodyBlock();
1988
1989 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "alloc_free", &.{}, &.{i64_type});
1990 try module_block.addOperation(func.op);
1991
1992 const entry = func.getEntryBlock();
1993 var size = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
1994 try entry.addOperation(size.op);
1995 var alloc = try MemrefDialect.AllocOp.createDynamic(&ir_ctx, loc, size.getResult(), memref_type);
1996 try entry.addOperation(alloc.op);
1997
1998 var idx0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
1999 try entry.addOperation(idx0.op);
2000 var value = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 77);
2001 try entry.addOperation(value.op);
2002
2003 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, value.getResult(), alloc.getResult(), idx0.getResult());
2004 try entry.addOperation(store.op);
2005 var load = try MemrefDialect.LoadOp.create(&ir_ctx, loc, alloc.getResult(), idx0.getResult(), i64_type);
2006 try entry.addOperation(load.op);
2007
2008 const dealloc = try MemrefDialect.DeallocOp.create(&ir_ctx, loc, alloc.getResult());
2009 try entry.addOperation(dealloc.op);
2010
2011 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{load.getResult()});
2012 try entry.addOperation(ret.op);
2013
2014 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2015 defer backend.deinit();
2016 const compiled = try backend.compile(module.op);
2017
2018 const exec_result = try onlyResult(&backend, compiled, "alloc_free", &.{});
2019 try testing.expectEqual(@as(i64, 77), try integerOf(exec_result));
2020 }
2021
2022 test "x86_64 memref f32 add load" {
2023 if (!supports_x86_64_backend) return;
2024
2025 const testing = std.testing;
2026 const dialects = @import("../../dialects/root.zig");
2027 const ArithDialect = dialects.ArithDialect;
2028 const BuiltinDialect = dialects.BuiltinDialect;
2029 const FuncDialect = dialects.FuncDialect;
2030 const MemrefDialect = dialects.MemrefDialect;
2031
2032 var arena = alloc_arena.Arena.init(std.testing.allocator);
2033 defer arena.deinit();
2034 const allocator = arena.allocator();
2035
2036 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2037 defer ir_ctx.deinit(allocator);
2038 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2039
2040 const loc = ir.Location.getUnknown();
2041 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
2042 const index_type = try ArithDialect.getIndexType(&ir_ctx);
2043 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, f32_type, .host);
2044
2045 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2046 const module_block = module.getBodyBlock();
2047
2048 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "f32_memref_add", &.{}, &.{f32_type});
2049 try module_block.addOperation(func.op);
2050
2051 const entry = func.getEntryBlock();
2052 var size = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 2);
2053 try entry.addOperation(size.op);
2054 var alloc = try MemrefDialect.AllocOp.createDynamic(&ir_ctx, loc, size.getResult(), memref_type);
2055 try entry.addOperation(alloc.op);
2056
2057 var idx0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
2058 try entry.addOperation(idx0.op);
2059 var idx1 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
2060 try entry.addOperation(idx1.op);
2061
2062 var v0 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 1.25);
2063 try entry.addOperation(v0.op);
2064 var v1 = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 2.75);
2065 try entry.addOperation(v1.op);
2066
2067 const store0 = try MemrefDialect.StoreOp.create(&ir_ctx, loc, v0.getResult(), alloc.getResult(), idx0.getResult());
2068 try entry.addOperation(store0.op);
2069 const store1 = try MemrefDialect.StoreOp.create(&ir_ctx, loc, v1.getResult(), alloc.getResult(), idx1.getResult());
2070 try entry.addOperation(store1.op);
2071
2072 var load0 = try MemrefDialect.LoadOp.create(&ir_ctx, loc, alloc.getResult(), idx0.getResult(), f32_type);
2073 try entry.addOperation(load0.op);
2074 var load1 = try MemrefDialect.LoadOp.create(&ir_ctx, loc, alloc.getResult(), idx1.getResult(), f32_type);
2075 try entry.addOperation(load1.op);
2076
2077 var add = try ArithDialect.AddOp.create(&ir_ctx, loc, load0.getResult(), load1.getResult());
2078 try entry.addOperation(add.op);
2079
2080 const dealloc = try MemrefDialect.DeallocOp.create(&ir_ctx, loc, alloc.getResult());
2081 try entry.addOperation(dealloc.op);
2082
2083 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{add.getResult()});
2084 try entry.addOperation(ret.op);
2085
2086 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2087 defer backend.deinit();
2088 const compiled = try backend.compile(module.op);
2089
2090 const exec_result = try onlyResult(&backend, compiled, "f32_memref_add", &.{});
2091 try testing.expectApproxEqAbs(@as(f64, 4.0), try floatOf(exec_result), 1e-6);
2092 }
2093
2094 test "x86_64 f32 dot product loop" {
2095 if (!supports_x86_64_backend) return;
2096
2097 const testing = std.testing;
2098 const dialects = @import("../../dialects/root.zig");
2099 const ArithDialect = dialects.ArithDialect;
2100 const BuiltinDialect = dialects.BuiltinDialect;
2101 const FuncDialect = dialects.FuncDialect;
2102 const MemrefDialect = dialects.MemrefDialect;
2103 const ScfDialect = dialects.ScfDialect;
2104
2105 var arena = alloc_arena.Arena.init(std.testing.allocator);
2106 defer arena.deinit();
2107 const allocator = arena.allocator();
2108
2109 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2110 defer ir_ctx.deinit(allocator);
2111 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2112
2113 const loc = ir.Location.getUnknown();
2114 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
2115 const index_type = try ArithDialect.getIndexType(&ir_ctx);
2116 const memref_type = try MemrefDialect.getMemrefType1D(&ir_ctx, 4, f32_type, .host);
2117
2118 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2119 const module_block = module.getBodyBlock();
2120
2121 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "dot_f32_loop", &.{}, &.{f32_type});
2122 try module_block.addOperation(func.op);
2123
2124 const entry = func.getEntryBlock();
2125 var alloc_a = try MemrefDialect.AllocOp.createStatic(&ir_ctx, loc, memref_type);
2126 try entry.addOperation(alloc_a.op);
2127 var alloc_b = try MemrefDialect.AllocOp.createStatic(&ir_ctx, loc, memref_type);
2128 try entry.addOperation(alloc_b.op);
2129
2130 const values = [_]f64{ 1.0, 2.0, 3.0, 4.0 };
2131 for (values, 0..) |val, i| {
2132 var idx = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, @intCast(i));
2133 try entry.addOperation(idx.op);
2134 var fval = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, val);
2135 try entry.addOperation(fval.op);
2136
2137 const store_a = try MemrefDialect.StoreOp.create(&ir_ctx, loc, fval.getResult(), alloc_a.getResult(), idx.getResult());
2138 try entry.addOperation(store_a.op);
2139 const store_b = try MemrefDialect.StoreOp.create(&ir_ctx, loc, fval.getResult(), alloc_b.getResult(), idx.getResult());
2140 try entry.addOperation(store_b.op);
2141 }
2142
2143 var lo = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
2144 try entry.addOperation(lo.op);
2145 var hi = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 4);
2146 try entry.addOperation(hi.op);
2147 var step = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
2148 try entry.addOperation(step.op);
2149 var acc_init = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 0.0);
2150 try entry.addOperation(acc_init.op);
2151
2152 var for_op = try ScfDialect.ForOp.create(
2153 &ir_ctx,
2154 loc,
2155 lo.getResult(),
2156 hi.getResult(),
2157 step.getResult(),
2158 &.{acc_init.getResult()},
2159 &.{f32_type},
2160 );
2161 try entry.addOperation(for_op.op);
2162
2163 const body = for_op.getBodyBlock();
2164 const iv = body.arguments.items[0];
2165 const acc = body.arguments.items[1];
2166
2167 var load_a = try MemrefDialect.LoadOp.create(&ir_ctx, loc, alloc_a.getResult(), iv, f32_type);
2168 try body.addOperation(load_a.op);
2169 var load_b = try MemrefDialect.LoadOp.create(&ir_ctx, loc, alloc_b.getResult(), iv, f32_type);
2170 try body.addOperation(load_b.op);
2171
2172 var mul = try ArithDialect.MulOp.create(&ir_ctx, loc, load_a.getResult(), load_b.getResult());
2173 try body.addOperation(mul.op);
2174 var add = try ArithDialect.AddOp.create(&ir_ctx, loc, acc, mul.getResult());
2175 try body.addOperation(add.op);
2176
2177 const yield_op = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{add.getResult()});
2178 try body.addOperation(yield_op.op);
2179
2180 const dealloc_a = try MemrefDialect.DeallocOp.create(&ir_ctx, loc, alloc_a.getResult());
2181 try entry.addOperation(dealloc_a.op);
2182 const dealloc_b = try MemrefDialect.DeallocOp.create(&ir_ctx, loc, alloc_b.getResult());
2183 try entry.addOperation(dealloc_b.op);
2184
2185 const result = for_op.getResult(0) orelse return error.TestFailure;
2186 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result});
2187 try entry.addOperation(ret.op);
2188
2189 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2190 defer backend.deinit();
2191 const compiled = try backend.compile(module.op);
2192
2193 const exec_result = try onlyResult(&backend, compiled, "dot_f32_loop", &.{});
2194 try testing.expectApproxEqAbs(@as(f64, 30.0), try floatOf(exec_result), 1e-6);
2195 }
2196
2197 test "x86_64 scf while carries integer values" {
2198 if (!supports_x86_64_backend) return;
2199
2200 const testing = std.testing;
2201 const dialects = @import("../../dialects/root.zig");
2202 const ArithDialect = dialects.ArithDialect;
2203 const BuiltinDialect = dialects.BuiltinDialect;
2204 const FuncDialect = dialects.FuncDialect;
2205 const ScfDialect = dialects.ScfDialect;
2206
2207 var arena = alloc_arena.Arena.init(std.testing.allocator);
2208 defer arena.deinit();
2209 const allocator = arena.allocator();
2210
2211 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2212 defer ir_ctx.deinit(allocator);
2213 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2214
2215 const loc = ir.Location.getUnknown();
2216 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
2217
2218 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2219 const module_block = module.getBodyBlock();
2220
2221 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "while_sum", &.{}, &.{i64_type});
2222 try module_block.addOperation(func.op);
2223
2224 const entry = func.getEntryBlock();
2225 var n_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 5);
2226 try entry.addOperation(n_init.op);
2227 var acc_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
2228 try entry.addOperation(acc_init.op);
2229
2230 var while_op = try ScfDialect.WhileOp.create(
2231 &ir_ctx,
2232 loc,
2233 &.{ n_init.getResult(), acc_init.getResult() },
2234 &.{ i64_type, i64_type },
2235 );
2236 try entry.addOperation(while_op.op);
2237
2238 const before = while_op.getBeforeBlock();
2239 const before_n = before.arguments.items[0];
2240 const before_acc = before.arguments.items[1];
2241 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
2242 try before.addOperation(zero.op);
2243 var keep_going = try ArithDialect.CmpOp.create(&ir_ctx, loc, .gt, before_n, zero.getResult());
2244 try before.addOperation(keep_going.op);
2245 const condition = try ScfDialect.ConditionOp.create(&ir_ctx, loc, keep_going.getResult(), &.{ before_n, before_acc });
2246 try before.addOperation(condition.op);
2247
2248 const after = while_op.getAfterBlock();
2249 const after_n = after.arguments.items[0];
2250 const after_acc = after.arguments.items[1];
2251 var one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
2252 try after.addOperation(one.op);
2253 var next_n = try ArithDialect.SubOp.create(&ir_ctx, loc, after_n, one.getResult());
2254 try after.addOperation(next_n.op);
2255 var next_acc = try ArithDialect.AddOp.create(&ir_ctx, loc, after_acc, after_n);
2256 try after.addOperation(next_acc.op);
2257 const yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{ next_n.getResult(), next_acc.getResult() });
2258 try after.addOperation(yield.op);
2259
2260 const result = while_op.op.getResult(1) orelse return error.TestFailure;
2261 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result});
2262 try entry.addOperation(ret.op);
2263
2264 var emitter = x86_64.Emitter.init(allocator);
2265 defer emitter.deinit();
2266 try emitter.emitFunction(func.op);
2267 try testing.expect((emitter.value_locations.get(before_n) orelse return error.TestFailure).isCallerSaved());
2268 try testing.expect((emitter.value_locations.get(before_acc) orelse return error.TestFailure).isCallerSaved());
2269 try testing.expect((emitter.value_locations.get(result) orelse return error.TestFailure).isCallerSaved());
2270 try testing.expectEqual(@as(usize, 0), emitter.reserved_callee_saved);
2271
2272 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2273 defer backend.deinit();
2274 const compiled = try backend.compile(module.op);
2275
2276 const exec_result = try onlyResult(&backend, compiled, "while_sum", &.{});
2277 try testing.expectEqual(@as(i64, 15), try integerOf(exec_result));
2278 }
2279
2280 test "x86_64 scf.while result feeds post-loop select" {
2281 if (!supports_x86_64_backend) return;
2282
2283 const testing = std.testing;
2284 const dialects = @import("../../dialects/root.zig");
2285 const ArithDialect = dialects.ArithDialect;
2286 const BuiltinDialect = dialects.BuiltinDialect;
2287 const FuncDialect = dialects.FuncDialect;
2288 const ScfDialect = dialects.ScfDialect;
2289
2290 var arena = alloc_arena.Arena.init(std.testing.allocator);
2291 defer arena.deinit();
2292 const allocator = arena.allocator();
2293
2294 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2295 defer ir_ctx.deinit(allocator);
2296 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2297
2298 const loc = ir.Location.getUnknown();
2299 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
2300
2301 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2302 const module_block = module.getBodyBlock();
2303
2304 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "while_select", &.{i64_type}, &.{i64_type});
2305 try module_block.addOperation(func.op);
2306
2307 const entry = func.getEntryBlock();
2308 const pressure_seed = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, -18823);
2309 try entry.addOperation(pressure_seed.op);
2310 var true_const = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
2311 try entry.addOperation(true_const.op);
2312 const pressure_false = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false);
2313 try entry.addOperation(pressure_false.op);
2314 var count_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 2);
2315 try entry.addOperation(count_init.op);
2316
2317 var while_op = try ScfDialect.WhileOp.create(
2318 &ir_ctx,
2319 loc,
2320 &.{ count_init.getResult(), func.getArgument(0) },
2321 &.{ i64_type, i64_type },
2322 );
2323 try entry.addOperation(while_op.op);
2324
2325 const before = while_op.getBeforeBlock();
2326 const before_count = before.arguments.items[0];
2327 const before_acc = before.arguments.items[1];
2328 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
2329 try before.addOperation(zero.op);
2330 var keep_going = try ArithDialect.CmpOp.create(&ir_ctx, loc, .gt, before_count, zero.getResult());
2331 try before.addOperation(keep_going.op);
2332 const condition = try ScfDialect.ConditionOp.create(&ir_ctx, loc, keep_going.getResult(), &.{ before_count, before_acc });
2333 try before.addOperation(condition.op);
2334
2335 const after = while_op.getAfterBlock();
2336 const after_count = after.arguments.items[0];
2337 const after_acc = after.arguments.items[1];
2338 var one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
2339 try after.addOperation(one.op);
2340 var next_count = try ArithDialect.SubOp.create(&ir_ctx, loc, after_count, one.getResult());
2341 try after.addOperation(next_count.op);
2342 var next_acc = try ArithDialect.AddOp.create(&ir_ctx, loc, after_acc, after_count);
2343 try after.addOperation(next_acc.op);
2344 const yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{ next_count.getResult(), next_acc.getResult() });
2345 try after.addOperation(yield.op);
2346
2347 const pressure_select = try ArithDialect.SelectOp.create(&ir_ctx, loc, true_const.getResult(), func.getArgument(0), while_op.op.getResult(1).?);
2348 try entry.addOperation(pressure_select.op);
2349 const pressure_extra = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, -20440);
2350 try entry.addOperation(pressure_extra.op);
2351
2352 var result_select = try ArithDialect.SelectOp.create(&ir_ctx, loc, true_const.getResult(), while_op.op.getResult(0).?, func.getArgument(0));
2353 try entry.addOperation(result_select.op);
2354 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_select.getResult()});
2355 try entry.addOperation(ret.op);
2356
2357 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2358 defer backend.deinit();
2359 const compiled = try backend.compile(module.op);
2360
2361 const exec_result = try onlyResult(&backend, compiled, "while_select", &.{.{ .i64 = -32767 }});
2362 try testing.expectEqual(@as(i64, 0), try integerOf(exec_result));
2363 }
2364
2365 test "x86_64 register homes initialize stack-passed arguments" {
2366 if (!supports_x86_64_backend) return;
2367
2368 const testing = std.testing;
2369 const dialects = @import("../../dialects/root.zig");
2370 const ArithDialect = dialects.ArithDialect;
2371 const BuiltinDialect = dialects.BuiltinDialect;
2372 const FuncDialect = dialects.FuncDialect;
2373
2374 var arena = alloc_arena.Arena.init(std.testing.allocator);
2375 defer arena.deinit();
2376 const allocator = arena.allocator();
2377
2378 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2379 defer ir_ctx.deinit(allocator);
2380 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2381
2382 const loc = ir.Location.getUnknown();
2383 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
2384 const arg_types = [_]ir.Type{ i64_type, i64_type, i64_type, i64_type, i64_type, i64_type, i64_type };
2385
2386 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2387 const module_block = module.getBodyBlock();
2388
2389 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "stack_arg_home", &arg_types, &.{i64_type});
2390 try module_block.addOperation(func.op);
2391
2392 const entry = func.getEntryBlock();
2393 const first_arg = func.getArgument(0);
2394 const stack_arg = func.getArgument(6);
2395
2396 var acc = try ArithDialect.AddOp.create(&ir_ctx, loc, stack_arg, stack_arg);
2397 try entry.addOperation(acc.op);
2398 inline for (0..5) |_| {
2399 const next = try ArithDialect.AddOp.create(&ir_ctx, loc, acc.getResult(), stack_arg);
2400 try entry.addOperation(next.op);
2401 acc = next;
2402 }
2403 var result = try ArithDialect.AddOp.create(&ir_ctx, loc, acc.getResult(), first_arg);
2404 try entry.addOperation(result.op);
2405 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result.getResult()});
2406 try entry.addOperation(ret.op);
2407
2408 var emitter = x86_64.Emitter.init(allocator);
2409 defer emitter.deinit();
2410 try emitter.emitFunction(func.op);
2411 try testing.expect(emitter.value_locations.contains(stack_arg));
2412
2413 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2414 defer backend.deinit();
2415 const compiled = try backend.compile(module.op);
2416
2417 const exec_result = try onlyResult(&backend, compiled, "stack_arg_home", &.{
2418 .{ .i64 = 1 }, .{ .i64 = 2 }, .{ .i64 = 3 }, .{ .i64 = 4 },
2419 .{ .i64 = 5 }, .{ .i64 = 6 }, .{ .i64 = 11 },
2420 });
2421 try testing.expectEqual(@as(i64, 78), try integerOf(exec_result));
2422 }
2423
2424 test "x86_64 scf.while preserves overlapping carried values" {
2425 if (!supports_x86_64_backend) return;
2426
2427 const testing = std.testing;
2428 const dialects = @import("../../dialects/root.zig");
2429 const ArithDialect = dialects.ArithDialect;
2430 const BuiltinDialect = dialects.BuiltinDialect;
2431 const FuncDialect = dialects.FuncDialect;
2432 const ScfDialect = dialects.ScfDialect;
2433
2434 var arena = alloc_arena.Arena.init(std.testing.allocator);
2435 defer arena.deinit();
2436 const allocator = arena.allocator();
2437
2438 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2439 defer ir_ctx.deinit(allocator);
2440 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2441
2442 const loc = ir.Location.getUnknown();
2443 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
2444
2445 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2446 const module_block = module.getBodyBlock();
2447
2448 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "while_swap_loop", &.{}, &.{i64_type});
2449 try module_block.addOperation(func.op);
2450
2451 const entry = func.getEntryBlock();
2452 var n_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 5);
2453 try entry.addOperation(n_init.op);
2454 var x_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 10);
2455 try entry.addOperation(x_init.op);
2456 var y_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
2457 try entry.addOperation(y_init.op);
2458 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
2459 try entry.addOperation(zero.op);
2460 var one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
2461 try entry.addOperation(one.op);
2462
2463 var while_op = try ScfDialect.WhileOp.create(
2464 &ir_ctx,
2465 loc,
2466 &.{ n_init.getResult(), x_init.getResult(), y_init.getResult() },
2467 &.{ i64_type, i64_type, i64_type },
2468 );
2469 try entry.addOperation(while_op.op);
2470
2471 const before = while_op.getBeforeBlock();
2472 const before_n = before.arguments.items[0];
2473 const before_x = before.arguments.items[1];
2474 const before_y = before.arguments.items[2];
2475 var keep_going = try ArithDialect.CmpOp.create(&ir_ctx, loc, .gt, before_n, zero.getResult());
2476 try before.addOperation(keep_going.op);
2477 const condition = try ScfDialect.ConditionOp.create(&ir_ctx, loc, keep_going.getResult(), &.{ before_n, before_x, before_y });
2478 try before.addOperation(condition.op);
2479
2480 const after = while_op.getAfterBlock();
2481 const after_n = after.arguments.items[0];
2482 const after_x = after.arguments.items[1];
2483 const after_y = after.arguments.items[2];
2484 var next_n = try ArithDialect.SubOp.create(&ir_ctx, loc, after_n, one.getResult());
2485 try after.addOperation(next_n.op);
2486 const loop_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{ next_n.getResult(), after_y, after_x });
2487 try after.addOperation(loop_yield.op);
2488
2489 const result_x = while_op.op.getResult(1) orelse return error.TestFailure;
2490 const result_y = while_op.op.getResult(2) orelse return error.TestFailure;
2491 var difference = try ArithDialect.SubOp.create(&ir_ctx, loc, result_x, result_y);
2492 try entry.addOperation(difference.op);
2493 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{difference.getResult()});
2494 try entry.addOperation(ret.op);
2495
2496 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2497 defer backend.deinit();
2498 const compiled = try backend.compile(module.op);
2499
2500 const exec_result = try onlyResult(&backend, compiled, "while_swap_loop", &.{});
2501 try testing.expectEqual(@as(i64, -9), try integerOf(exec_result));
2502 }
2503
2504 test "x86_64 scf.while preserves high arity overlapping carried values" {
2505 if (!supports_x86_64_backend) return;
2506
2507 const testing = std.testing;
2508 const dialects = @import("../../dialects/root.zig");
2509 const ArithDialect = dialects.ArithDialect;
2510 const BuiltinDialect = dialects.BuiltinDialect;
2511 const FuncDialect = dialects.FuncDialect;
2512 const ScfDialect = dialects.ScfDialect;
2513
2514 var arena = alloc_arena.Arena.init(std.testing.allocator);
2515 defer arena.deinit();
2516 const allocator = arena.allocator();
2517
2518 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2519 defer ir_ctx.deinit(allocator);
2520 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2521
2522 const loc = ir.Location.getUnknown();
2523 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
2524
2525 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2526 const module_block = module.getBodyBlock();
2527
2528 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "while_rotate_loop", &.{}, &.{i64_type});
2529 try module_block.addOperation(func.op);
2530
2531 const entry = func.getEntryBlock();
2532 var n_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
2533 try entry.addOperation(n_init.op);
2534 var v0_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
2535 try entry.addOperation(v0_init.op);
2536 var v1_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 2);
2537 try entry.addOperation(v1_init.op);
2538 var v2_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 3);
2539 try entry.addOperation(v2_init.op);
2540 var v3_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 4);
2541 try entry.addOperation(v3_init.op);
2542 var v4_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 5);
2543 try entry.addOperation(v4_init.op);
2544 var v5_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 6);
2545 try entry.addOperation(v5_init.op);
2546 var v6_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 7);
2547 try entry.addOperation(v6_init.op);
2548 var v7_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 8);
2549 try entry.addOperation(v7_init.op);
2550 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
2551 try entry.addOperation(zero.op);
2552 var one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
2553 try entry.addOperation(one.op);
2554
2555 var while_op = try ScfDialect.WhileOp.create(
2556 &ir_ctx,
2557 loc,
2558 &.{
2559 n_init.getResult(),
2560 v0_init.getResult(),
2561 v1_init.getResult(),
2562 v2_init.getResult(),
2563 v3_init.getResult(),
2564 v4_init.getResult(),
2565 v5_init.getResult(),
2566 v6_init.getResult(),
2567 v7_init.getResult(),
2568 },
2569 &.{
2570 i64_type,
2571 i64_type,
2572 i64_type,
2573 i64_type,
2574 i64_type,
2575 i64_type,
2576 i64_type,
2577 i64_type,
2578 i64_type,
2579 },
2580 );
2581 try entry.addOperation(while_op.op);
2582
2583 const before = while_op.getBeforeBlock();
2584 const before_n = before.arguments.items[0];
2585 var keep_going = try ArithDialect.CmpOp.create(&ir_ctx, loc, .gt, before_n, zero.getResult());
2586 try before.addOperation(keep_going.op);
2587 const condition = try ScfDialect.ConditionOp.create(&ir_ctx, loc, keep_going.getResult(), &.{
2588 before_n,
2589 before.arguments.items[1],
2590 before.arguments.items[2],
2591 before.arguments.items[3],
2592 before.arguments.items[4],
2593 before.arguments.items[5],
2594 before.arguments.items[6],
2595 before.arguments.items[7],
2596 before.arguments.items[8],
2597 });
2598 try before.addOperation(condition.op);
2599
2600 const after = while_op.getAfterBlock();
2601 const after_n = after.arguments.items[0];
2602 var next_n = try ArithDialect.SubOp.create(&ir_ctx, loc, after_n, one.getResult());
2603 try after.addOperation(next_n.op);
2604 const loop_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{
2605 next_n.getResult(),
2606 after.arguments.items[2],
2607 after.arguments.items[3],
2608 after.arguments.items[4],
2609 after.arguments.items[5],
2610 after.arguments.items[6],
2611 after.arguments.items[7],
2612 after.arguments.items[8],
2613 after.arguments.items[1],
2614 });
2615 try after.addOperation(loop_yield.op);
2616
2617 const result_v0 = while_op.op.getResult(1) orelse return error.TestFailure;
2618 const result_v7 = while_op.op.getResult(8) orelse return error.TestFailure;
2619 var difference = try ArithDialect.SubOp.create(&ir_ctx, loc, result_v0, result_v7);
2620 try entry.addOperation(difference.op);
2621 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{difference.getResult()});
2622 try entry.addOperation(ret.op);
2623
2624 var emitter = x86_64.Emitter.init(allocator);
2625 defer emitter.deinit();
2626 try emitter.emitFunction(func.op);
2627 var uses_r11 = false;
2628 var homes = emitter.value_locations.iterator();
2629 while (homes.next()) |home_entry| {
2630 if (home_entry.value_ptr.* == .r11) {
2631 uses_r11 = true;
2632 break;
2633 }
2634 }
2635 try testing.expect(uses_r11);
2636
2637 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2638 defer backend.deinit();
2639 const compiled = try backend.compile(module.op);
2640
2641 const exec_result = try onlyResult(&backend, compiled, "while_rotate_loop", &.{});
2642 try testing.expectEqual(@as(i64, 1), try integerOf(exec_result));
2643 }
2644
2645 test "x86_64 memref i8 load/store" {
2646 if (!supports_x86_64_backend) return;
2647
2648 const testing = std.testing;
2649 const dialects = @import("../../dialects/root.zig");
2650 const ArithDialect = dialects.ArithDialect;
2651 const BuiltinDialect = dialects.BuiltinDialect;
2652 const FuncDialect = dialects.FuncDialect;
2653 const MemrefDialect = dialects.MemrefDialect;
2654
2655 var arena = alloc_arena.Arena.init(std.testing.allocator);
2656 defer arena.deinit();
2657 const allocator = arena.allocator();
2658
2659 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2660 defer ir_ctx.deinit(allocator);
2661 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2662
2663 const loc = ir.Location.getUnknown();
2664 const i8_type = try ArithDialect.getScalarType(&ir_ctx, .i8);
2665 const index_type = try ArithDialect.getIndexType(&ir_ctx);
2666 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, i8_type, .host);
2667
2668 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2669 const module_block = module.getBodyBlock();
2670
2671 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "i8_memref_roundtrip", &.{}, &.{i8_type});
2672 try module_block.addOperation(func.op);
2673
2674 const entry = func.getEntryBlock();
2675 var size = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
2676 try entry.addOperation(size.op);
2677 var alloc = try MemrefDialect.AllocOp.createDynamic(&ir_ctx, loc, size.getResult(), memref_type);
2678 try entry.addOperation(alloc.op);
2679
2680 var idx0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
2681 try entry.addOperation(idx0.op);
2682 var value = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i8_type, 42);
2683 try entry.addOperation(value.op);
2684
2685 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, value.getResult(), alloc.getResult(), idx0.getResult());
2686 try entry.addOperation(store.op);
2687 var load = try MemrefDialect.LoadOp.create(&ir_ctx, loc, alloc.getResult(), idx0.getResult(), i8_type);
2688 try entry.addOperation(load.op);
2689
2690 const dealloc = try MemrefDialect.DeallocOp.create(&ir_ctx, loc, alloc.getResult());
2691 try entry.addOperation(dealloc.op);
2692
2693 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{load.getResult()});
2694 try entry.addOperation(ret.op);
2695
2696 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2697 defer backend.deinit();
2698 const compiled = try backend.compile(module.op);
2699
2700 const exec_result = try onlyResult(&backend, compiled, "i8_memref_roundtrip", &.{});
2701 try testing.expectEqual(@as(i64, 42), try integerOf(exec_result));
2702 }
2703
2704 test "x86_64 memref i16 load/store" {
2705 if (!supports_x86_64_backend) return;
2706
2707 const testing = std.testing;
2708 const dialects = @import("../../dialects/root.zig");
2709 const ArithDialect = dialects.ArithDialect;
2710 const BuiltinDialect = dialects.BuiltinDialect;
2711 const FuncDialect = dialects.FuncDialect;
2712 const MemrefDialect = dialects.MemrefDialect;
2713
2714 var arena = alloc_arena.Arena.init(std.testing.allocator);
2715 defer arena.deinit();
2716 const allocator = arena.allocator();
2717
2718 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2719 defer ir_ctx.deinit(allocator);
2720 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2721
2722 const loc = ir.Location.getUnknown();
2723 const i16_type = try ArithDialect.getScalarType(&ir_ctx, .i16);
2724 const index_type = try ArithDialect.getIndexType(&ir_ctx);
2725 const memref_type = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, i16_type, .host);
2726
2727 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2728 const module_block = module.getBodyBlock();
2729
2730 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "i16_memref_roundtrip", &.{}, &.{i16_type});
2731 try module_block.addOperation(func.op);
2732
2733 const entry = func.getEntryBlock();
2734 var size = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
2735 try entry.addOperation(size.op);
2736 var alloc = try MemrefDialect.AllocOp.createDynamic(&ir_ctx, loc, size.getResult(), memref_type);
2737 try entry.addOperation(alloc.op);
2738
2739 var idx0 = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
2740 try entry.addOperation(idx0.op);
2741 var value = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i16_type, 32000);
2742 try entry.addOperation(value.op);
2743
2744 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, value.getResult(), alloc.getResult(), idx0.getResult());
2745 try entry.addOperation(store.op);
2746 var load = try MemrefDialect.LoadOp.create(&ir_ctx, loc, alloc.getResult(), idx0.getResult(), i16_type);
2747 try entry.addOperation(load.op);
2748
2749 const dealloc = try MemrefDialect.DeallocOp.create(&ir_ctx, loc, alloc.getResult());
2750 try entry.addOperation(dealloc.op);
2751
2752 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{load.getResult()});
2753 try entry.addOperation(ret.op);
2754
2755 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2756 defer backend.deinit();
2757 const compiled = try backend.compile(module.op);
2758
2759 const exec_result = try onlyResult(&backend, compiled, "i16_memref_roundtrip", &.{});
2760 try testing.expectEqual(@as(i64, 32000), try integerOf(exec_result));
2761 }
2762
2763 const ScalarOpKind = enum {
2764 neg,
2765 max,
2766 min,
2767 sqrt,
2768 abs,
2769 sin,
2770 cos,
2771 tan,
2772 exp,
2773 log,
2774 tanh,
2775 floor,
2776 pow,
2777 div,
2778 rem,
2779 umulhi,
2780 band,
2781 bor,
2782 bxor,
2783 bnot,
2784 popcount,
2785 shl,
2786 shr,
2787 ushr,
2788 };
2789
2790 fn runIntOp(
2791 comptime kind: ScalarOpKind,
2792 int_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
2793 name: []const u8,
2794 c_lhs: i64,
2795 c_rhs: i64,
2796 ) !i64 {
2797 const dialects = @import("../../dialects/root.zig");
2798 const ArithDialect = dialects.ArithDialect;
2799 const BuiltinDialect = dialects.BuiltinDialect;
2800 const FuncDialect = dialects.FuncDialect;
2801
2802 var arena = alloc_arena.Arena.init(std.testing.allocator);
2803 defer arena.deinit();
2804 const allocator = arena.allocator();
2805
2806 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2807 defer ir_ctx.deinit(allocator);
2808 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2809
2810 const loc = ir.Location.getUnknown();
2811 const int_type = try ArithDialect.getScalarType(&ir_ctx, int_kind);
2812
2813 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2814 const module_block = module.getBodyBlock();
2815
2816 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{int_type});
2817 try module_block.addOperation(func.op);
2818
2819 const entry = func.getEntryBlock();
2820 var lhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_lhs);
2821 try entry.addOperation(lhs.op);
2822
2823 const result_value = switch (kind) {
2824 .neg => blk: {
2825 const op = try ArithDialect.NegOp.create(&ir_ctx, loc, lhs.getResult());
2826 try entry.addOperation(op.op);
2827 break :blk op.getResult();
2828 },
2829 .abs => blk: {
2830 const op = try ArithDialect.AbsOp.create(&ir_ctx, loc, lhs.getResult());
2831 try entry.addOperation(op.op);
2832 break :blk op.getResult();
2833 },
2834 .max => blk: {
2835 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2836 try entry.addOperation(rhs.op);
2837 const op = try ArithDialect.MaxOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2838 try entry.addOperation(op.op);
2839 break :blk op.getResult();
2840 },
2841 .min => blk: {
2842 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2843 try entry.addOperation(rhs.op);
2844 const op = try ArithDialect.MinOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2845 try entry.addOperation(op.op);
2846 break :blk op.getResult();
2847 },
2848 .div => blk: {
2849 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2850 try entry.addOperation(rhs.op);
2851 const op = try ArithDialect.DivOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2852 try entry.addOperation(op.op);
2853 break :blk op.getResult();
2854 },
2855 .rem => blk: {
2856 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2857 try entry.addOperation(rhs.op);
2858 const op = try ArithDialect.RemOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2859 try entry.addOperation(op.op);
2860 break :blk op.getResult();
2861 },
2862 .umulhi => blk: {
2863 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2864 try entry.addOperation(rhs.op);
2865 const op = try ArithDialect.UmulhiOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2866 try entry.addOperation(op.op);
2867 break :blk op.getResult();
2868 },
2869 .band => blk: {
2870 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2871 try entry.addOperation(rhs.op);
2872 const op = try ArithDialect.AndOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2873 try entry.addOperation(op.op);
2874 break :blk op.getResult();
2875 },
2876 .bor => blk: {
2877 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2878 try entry.addOperation(rhs.op);
2879 const op = try ArithDialect.OrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2880 try entry.addOperation(op.op);
2881 break :blk op.getResult();
2882 },
2883 .bxor => blk: {
2884 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2885 try entry.addOperation(rhs.op);
2886 const op = try ArithDialect.XorOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2887 try entry.addOperation(op.op);
2888 break :blk op.getResult();
2889 },
2890 .bnot => blk: {
2891 const op = try ArithDialect.NotOp.create(&ir_ctx, loc, lhs.getResult());
2892 try entry.addOperation(op.op);
2893 break :blk op.getResult();
2894 },
2895 .popcount => blk: {
2896 const op = try ArithDialect.PopCountOp.create(&ir_ctx, loc, lhs.getResult());
2897 try entry.addOperation(op.op);
2898 break :blk op.getResult();
2899 },
2900 .shl => blk: {
2901 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2902 try entry.addOperation(rhs.op);
2903 const op = try ArithDialect.ShlOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2904 try entry.addOperation(op.op);
2905 break :blk op.getResult();
2906 },
2907 .shr => blk: {
2908 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2909 try entry.addOperation(rhs.op);
2910 const op = try ArithDialect.ShrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2911 try entry.addOperation(op.op);
2912 break :blk op.getResult();
2913 },
2914 .ushr => blk: {
2915 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, int_type, c_rhs);
2916 try entry.addOperation(rhs.op);
2917 const op = try ArithDialect.UshrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
2918 try entry.addOperation(op.op);
2919 break :blk op.getResult();
2920 },
2921 .sqrt, .sin, .cos, .tan, .exp, .log, .tanh, .floor, .pow => @compileError(
2922 "runIntOp does not handle float-only ops — use runFloatUnaryOp / runFloatBinaryOpRaw",
2923 ),
2924 };
2925
2926 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_value});
2927 try entry.addOperation(ret.op);
2928
2929 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2930 defer backend.deinit();
2931 const compiled = try backend.compile(module.op);
2932
2933 const exec_result = try onlyResult(&backend, compiled, name, &.{});
2934 return integerOf(exec_result);
2935 }
2936
2937 fn runFloatUnaryOp(
2938 comptime kind: ScalarOpKind,
2939 float_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
2940 name: []const u8,
2941 c_lhs: f64,
2942 ) !f64 {
2943 const dialects = @import("../../dialects/root.zig");
2944 const ArithDialect = dialects.ArithDialect;
2945 const BuiltinDialect = dialects.BuiltinDialect;
2946 const FuncDialect = dialects.FuncDialect;
2947
2948 var arena = alloc_arena.Arena.init(std.testing.allocator);
2949 defer arena.deinit();
2950 const allocator = arena.allocator();
2951
2952 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
2953 defer ir_ctx.deinit(allocator);
2954 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
2955
2956 const loc = ir.Location.getUnknown();
2957 const float_type = try ArithDialect.getScalarType(&ir_ctx, float_kind);
2958
2959 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
2960 const module_block = module.getBodyBlock();
2961
2962 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{float_type});
2963 try module_block.addOperation(func.op);
2964
2965 const entry = func.getEntryBlock();
2966 var lhs = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, float_type, c_lhs);
2967 try entry.addOperation(lhs.op);
2968
2969 const result_value = switch (kind) {
2970 .sqrt => blk: {
2971 const op = try ArithDialect.SqrtOp.create(&ir_ctx, loc, lhs.getResult());
2972 try entry.addOperation(op.op);
2973 break :blk op.getResult();
2974 },
2975 .floor => blk: {
2976 const op = try ArithDialect.FloorOp.create(&ir_ctx, loc, lhs.getResult());
2977 try entry.addOperation(op.op);
2978 break :blk op.getResult();
2979 },
2980 .neg => blk: {
2981 const op = try ArithDialect.NegOp.create(&ir_ctx, loc, lhs.getResult());
2982 try entry.addOperation(op.op);
2983 break :blk op.getResult();
2984 },
2985 .abs => blk: {
2986 const op = try ArithDialect.AbsOp.create(&ir_ctx, loc, lhs.getResult());
2987 try entry.addOperation(op.op);
2988 break :blk op.getResult();
2989 },
2990 else => @compileError("runFloatUnaryOp only handles .sqrt/.floor/.neg/.abs"),
2991 };
2992
2993 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_value});
2994 try entry.addOperation(ret.op);
2995
2996 var backend = try Backend.init(allocator, &ir_ctx, .testing);
2997 defer backend.deinit();
2998 const compiled = try backend.compile(module.op);
2999
3000 const exec_result = try onlyResult(&backend, compiled, name, &.{});
3001 return floatOf(exec_result);
3002 }
3003
3004 fn runFloatUnaryOpRaw(
3005 comptime kind: ScalarOpKind,
3006 comptime FT: type,
3007 name: []const u8,
3008 c_lhs: FT,
3009 ) !FT {
3010 const dialects = @import("../../dialects/root.zig");
3011 const ArithDialect = dialects.ArithDialect;
3012 const BuiltinDialect = dialects.BuiltinDialect;
3013 const FuncDialect = dialects.FuncDialect;
3014
3015 const float_kind: ArithDialect.ScalarTypeKind = comptime blk: {
3016 if (FT == f32) break :blk .f32;
3017 if (FT == f64) break :blk .f64;
3018 @compileError("runFloatUnaryOpRaw: FT must be f32 or f64");
3019 };
3020
3021 var arena = alloc_arena.Arena.init(std.testing.allocator);
3022 defer arena.deinit();
3023 const allocator = arena.allocator();
3024
3025 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
3026 defer ir_ctx.deinit(allocator);
3027 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
3028
3029 const loc = ir.Location.getUnknown();
3030 const float_type = try ArithDialect.getScalarType(&ir_ctx, float_kind);
3031
3032 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
3033 const module_block = module.getBodyBlock();
3034
3035 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{float_type});
3036 try module_block.addOperation(func.op);
3037
3038 const entry = func.getEntryBlock();
3039 var lhs = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, float_type, @floatCast(c_lhs));
3040 try entry.addOperation(lhs.op);
3041
3042 const result_value = switch (kind) {
3043 .neg => blk: {
3044 const op = try ArithDialect.NegOp.create(&ir_ctx, loc, lhs.getResult());
3045 try entry.addOperation(op.op);
3046 break :blk op.getResult();
3047 },
3048 .abs => blk: {
3049 const op = try ArithDialect.AbsOp.create(&ir_ctx, loc, lhs.getResult());
3050 try entry.addOperation(op.op);
3051 break :blk op.getResult();
3052 },
3053 .sqrt => blk: {
3054 const op = try ArithDialect.SqrtOp.create(&ir_ctx, loc, lhs.getResult());
3055 try entry.addOperation(op.op);
3056 break :blk op.getResult();
3057 },
3058 .floor => blk: {
3059 const op = try ArithDialect.FloorOp.create(&ir_ctx, loc, lhs.getResult());
3060 try entry.addOperation(op.op);
3061 break :blk op.getResult();
3062 },
3063 .sin => blk: {
3064 const op = try ArithDialect.SinOp.create(&ir_ctx, loc, lhs.getResult());
3065 try entry.addOperation(op.op);
3066 break :blk op.getResult();
3067 },
3068 .cos => blk: {
3069 const op = try ArithDialect.CosOp.create(&ir_ctx, loc, lhs.getResult());
3070 try entry.addOperation(op.op);
3071 break :blk op.getResult();
3072 },
3073 .tan => blk: {
3074 const op = try ArithDialect.TanOp.create(&ir_ctx, loc, lhs.getResult());
3075 try entry.addOperation(op.op);
3076 break :blk op.getResult();
3077 },
3078 .exp => blk: {
3079 const op = try ArithDialect.ExpOp.create(&ir_ctx, loc, lhs.getResult());
3080 try entry.addOperation(op.op);
3081 break :blk op.getResult();
3082 },
3083 .log => blk: {
3084 const op = try ArithDialect.LogOp.create(&ir_ctx, loc, lhs.getResult());
3085 try entry.addOperation(op.op);
3086 break :blk op.getResult();
3087 },
3088 .tanh => blk: {
3089 const op = try ArithDialect.TanhOp.create(&ir_ctx, loc, lhs.getResult());
3090 try entry.addOperation(op.op);
3091 break :blk op.getResult();
3092 },
3093 else => @compileError("runFloatUnaryOpRaw only handles unary float ops"),
3094 };
3095
3096 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_value});
3097 try entry.addOperation(ret.op);
3098
3099 var backend = try Backend.init(allocator, &ir_ctx, .testing);
3100 defer backend.deinit();
3101 const compiled = try backend.compile(module.op);
3102
3103 const FnPtr = *const fn () callconv(.c) FT;
3104 const fn_ptr = try backend.runtime.getFunction(compiled, name, FnPtr);
3105 return fn_ptr();
3106 }
3107
3108 fn runFloatBinaryOpRaw(
3109 comptime kind: ScalarOpKind,
3110 comptime FT: type,
3111 name: []const u8,
3112 c_lhs: FT,
3113 c_rhs: FT,
3114 ) !FT {
3115 const dialects = @import("../../dialects/root.zig");
3116 const ArithDialect = dialects.ArithDialect;
3117 const BuiltinDialect = dialects.BuiltinDialect;
3118 const FuncDialect = dialects.FuncDialect;
3119
3120 const float_kind: ArithDialect.ScalarTypeKind = comptime blk: {
3121 if (FT == f32) break :blk .f32;
3122 if (FT == f64) break :blk .f64;
3123 @compileError("runFloatBinaryOpRaw: FT must be f32 or f64");
3124 };
3125
3126 var arena = alloc_arena.Arena.init(std.testing.allocator);
3127 defer arena.deinit();
3128 const allocator = arena.allocator();
3129
3130 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
3131 defer ir_ctx.deinit(allocator);
3132 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
3133
3134 const loc = ir.Location.getUnknown();
3135 const float_type = try ArithDialect.getScalarType(&ir_ctx, float_kind);
3136
3137 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
3138 const module_block = module.getBodyBlock();
3139
3140 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{float_type});
3141 try module_block.addOperation(func.op);
3142
3143 const entry = func.getEntryBlock();
3144 var lhs = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, float_type, @floatCast(c_lhs));
3145 try entry.addOperation(lhs.op);
3146 var rhs = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, float_type, @floatCast(c_rhs));
3147 try entry.addOperation(rhs.op);
3148
3149 const result_value = switch (kind) {
3150 .pow => blk: {
3151 const op = try ArithDialect.PowOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
3152 try entry.addOperation(op.op);
3153 break :blk op.getResult();
3154 },
3155 .max => blk: {
3156 const op = try ArithDialect.MaxOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
3157 try entry.addOperation(op.op);
3158 break :blk op.getResult();
3159 },
3160 .min => blk: {
3161 const op = try ArithDialect.MinOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
3162 try entry.addOperation(op.op);
3163 break :blk op.getResult();
3164 },
3165 else => @compileError("runFloatBinaryOpRaw only handles .pow/.max/.min"),
3166 };
3167
3168 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_value});
3169 try entry.addOperation(ret.op);
3170
3171 var backend = try Backend.init(allocator, &ir_ctx, .testing);
3172 defer backend.deinit();
3173 const compiled = try backend.compile(module.op);
3174
3175 const FnPtr = *const fn () callconv(.c) FT;
3176 const fn_ptr = try backend.runtime.getFunction(compiled, name, FnPtr);
3177 return fn_ptr();
3178 }
3179
3180 const FloatTernaryOpKind = enum { fma };
3181
3182 fn runFloatTernaryOpRaw(
3183 comptime kind: FloatTernaryOpKind,
3184 comptime FT: type,
3185 name: []const u8,
3186 c_a: FT,
3187 c_b: FT,
3188 c_c: FT,
3189 ) !FT {
3190 const dialects = @import("../../dialects/root.zig");
3191 const ArithDialect = dialects.ArithDialect;
3192 const BuiltinDialect = dialects.BuiltinDialect;
3193 const FuncDialect = dialects.FuncDialect;
3194
3195 const float_kind: ArithDialect.ScalarTypeKind = comptime blk: {
3196 if (FT == f32) break :blk .f32;
3197 if (FT == f64) break :blk .f64;
3198 @compileError("runFloatTernaryOpRaw: FT must be f32 or f64");
3199 };
3200
3201 var arena = alloc_arena.Arena.init(std.testing.allocator);
3202 defer arena.deinit();
3203 const allocator = arena.allocator();
3204
3205 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
3206 defer ir_ctx.deinit(allocator);
3207 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
3208
3209 const loc = ir.Location.getUnknown();
3210 const float_type = try ArithDialect.getScalarType(&ir_ctx, float_kind);
3211
3212 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
3213 const module_block = module.getBodyBlock();
3214
3215 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{float_type});
3216 try module_block.addOperation(func.op);
3217
3218 const entry = func.getEntryBlock();
3219 var a = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, float_type, @floatCast(c_a));
3220 try entry.addOperation(a.op);
3221 var b = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, float_type, @floatCast(c_b));
3222 try entry.addOperation(b.op);
3223 var c = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, float_type, @floatCast(c_c));
3224 try entry.addOperation(c.op);
3225
3226 const result_value = switch (kind) {
3227 .fma => blk: {
3228 const op = try ArithDialect.FmaOp.create(&ir_ctx, loc, a.getResult(), b.getResult(), c.getResult());
3229 try entry.addOperation(op.op);
3230 break :blk op.getResult();
3231 },
3232 };
3233
3234 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_value});
3235 try entry.addOperation(ret.op);
3236
3237 var backend = try Backend.init(allocator, &ir_ctx, .testing);
3238 defer backend.deinit();
3239 const compiled = try backend.compile(module.op);
3240
3241 const FnPtr = *const fn () callconv(.c) FT;
3242 const fn_ptr = try backend.runtime.getFunction(compiled, name, FnPtr);
3243 return fn_ptr();
3244 }
3245
3246 test "x86_64 arith.neg over i32 negates value (signed return marshalling)" {
3247 if (!supports_x86_64_backend) return;
3248 const testing = std.testing;
3249
3250 try testing.expectEqual(@as(i64, -5), try runIntOp(.neg, .i32, "neg_i32_pos5", 5, 0));
3251 try testing.expectEqual(@as(i64, 7), try runIntOp(.neg, .i32, "neg_i32_neg7", -7, 0));
3252 try testing.expectEqual(@as(i64, 0), try runIntOp(.neg, .i32, "neg_i32_zero", 0, 0));
3253 try testing.expectEqual(
3254 @as(i64, std.math.minInt(i32)),
3255 try runIntOp(.neg, .i32, "neg_i32_min", std.math.minInt(i32), 0),
3256 );
3257 }
3258
3259 test "x86_64 arith.neg over i64 negates value (and wraps INT_MIN)" {
3260 if (!supports_x86_64_backend) return;
3261 const testing = std.testing;
3262
3263 try testing.expectEqual(@as(i64, -5), try runIntOp(.neg, .i64, "neg_i64_pos", 5, 0));
3264 try testing.expectEqual(@as(i64, 7), try runIntOp(.neg, .i64, "neg_i64_neg", -7, 0));
3265 try testing.expectEqual(@as(i64, 0), try runIntOp(.neg, .i64, "neg_i64_zero", 0, 0));
3266 try testing.expectEqual(
3267 @as(i64, std.math.minInt(i64)),
3268 try runIntOp(.neg, .i64, "neg_i64_min", std.math.minInt(i64), 0),
3269 );
3270 }
3271
3272 test "x86_64 arith.max over i32 picks the larger operand" {
3273 if (!supports_x86_64_backend) return;
3274 const testing = std.testing;
3275
3276 try testing.expectEqual(@as(i64, 7), try runIntOp(.max, .i32, "max_i32_a", 3, 7));
3277 try testing.expectEqual(@as(i64, 7), try runIntOp(.max, .i32, "max_i32_b", 7, 3));
3278 try testing.expectEqual(@as(i64, 0), try runIntOp(.max, .i32, "max_i32_eq", 0, 0));
3279 try testing.expectEqual(@as(i64, -3), try runIntOp(.max, .i32, "max_i32_neg", -5, -3));
3280 }
3281
3282 test "x86_64 arith.max over i64 picks the larger operand" {
3283 if (!supports_x86_64_backend) return;
3284 const testing = std.testing;
3285
3286 try testing.expectEqual(@as(i64, std.math.maxInt(i64)), try runIntOp(
3287 .max,
3288 .i64,
3289 "max_i64",
3290 std.math.minInt(i64),
3291 std.math.maxInt(i64),
3292 ));
3293 try testing.expectEqual(@as(i64, -3), try runIntOp(.max, .i64, "max_i64_neg", -5, -3));
3294 }
3295
3296 test "x86_64 arith.min over i32 picks the smaller operand" {
3297 if (!supports_x86_64_backend) return;
3298 const testing = std.testing;
3299
3300 try testing.expectEqual(@as(i64, 3), try runIntOp(.min, .i32, "min_i32_a", 3, 7));
3301 try testing.expectEqual(@as(i64, 3), try runIntOp(.min, .i32, "min_i32_b", 7, 3));
3302 try testing.expectEqual(@as(i64, 0), try runIntOp(.min, .i32, "min_i32_zero", 0, 5));
3303 try testing.expectEqual(@as(i64, -5), try runIntOp(.min, .i32, "min_i32_neg", -5, -3));
3304 }
3305
3306 test "x86_64 arith.min over i64 picks the smaller operand" {
3307 if (!supports_x86_64_backend) return;
3308 const testing = std.testing;
3309
3310 try testing.expectEqual(@as(i64, std.math.minInt(i64)), try runIntOp(
3311 .min,
3312 .i64,
3313 "min_i64",
3314 std.math.minInt(i64),
3315 std.math.maxInt(i64),
3316 ));
3317 try testing.expectEqual(@as(i64, -5), try runIntOp(.min, .i64, "min_i64_neg", -5, -3));
3318 }
3319
3320 test "x86_64 arith.div over integers returns quotient" {
3321 if (!supports_x86_64_backend) return;
3322 const testing = std.testing;
3323
3324 try testing.expectEqual(@as(i64, 3), try runIntOp(.div, .i32, "div_i32_pos", 17, 5));
3325 try testing.expectEqual(@as(i64, -4), try runIntOp(.div, .i64, "div_i64_neg", -21, 5));
3326 }
3327
3328 test "x86_64 arith.rem over integers returns remainder" {
3329 if (!supports_x86_64_backend) return;
3330 const testing = std.testing;
3331
3332 try testing.expectEqual(@as(i64, 2), try runIntOp(.rem, .i32, "rem_i32_pos", 17, 5));
3333 try testing.expectEqual(@as(i64, -1), try runIntOp(.rem, .i64, "rem_i64_neg", -21, 5));
3334 }
3335
3336 test "x86_64 arith.div and rem over u64 use unsigned high-bit semantics" {
3337 if (!supports_x86_64_backend) return;
3338 const testing = std.testing;
3339
3340 const high_bit: i64 = @bitCast(@as(u64, 0x8000_0000_0000_0000));
3341 try testing.expectEqual(
3342 @as(i64, @bitCast(@as(u64, 0x4000_0000_0000_0000))),
3343 try runIntOp(.div, .u64, "div_u64_high_bit_2", high_bit, 2),
3344 );
3345 try testing.expectEqual(
3346 @as(i64, @bitCast(@as(u64, 0x7fff_ffff_ffff_ffff))),
3347 try runIntOp(.div, .u64, "div_u64_max_2", -1, 2),
3348 );
3349 try testing.expectEqual(@as(i64, 1), try runIntOp(.rem, .u64, "rem_u64_max_2", -1, 2));
3350 }
3351
3352 test "x86_64 arith.div and rem over u32 use unsigned high-bit semantics" {
3353 if (!supports_x86_64_backend) return;
3354 const testing = std.testing;
3355
3356 try testing.expectEqual(@as(i64, 0x4000_0000), try runIntOp(.div, .u32, "div_u32_high_bit_2", 0x8000_0000, 2));
3357 try testing.expectEqual(@as(i64, 0x7fff_ffff), try runIntOp(.div, .u32, "div_u32_max_2", 0xffff_ffff, 2));
3358 try testing.expectEqual(@as(i64, 1), try runIntOp(.rem, .u32, "rem_u32_max_2", 0xffff_ffff, 2));
3359 }
3360
3361 test "x86_64 arith.umulhi over integers returns unsigned high half" {
3362 if (!supports_x86_64_backend) return;
3363 const testing = std.testing;
3364
3365 try testing.expectEqual(@as(i64, 1), try runIntOp(.umulhi, .u32, "umulhi_u32_max_2", 0xffff_ffff, 2));
3366 try testing.expectEqual(@as(i64, 2), try runIntOp(.umulhi, .i32, "umulhi_i32_high_4", std.math.minInt(i32), 4));
3367 try testing.expectEqual(@as(i64, 1), try runIntOp(.umulhi, .u64, "umulhi_u64_max_2", -1, 2));
3368 try testing.expectEqual(@as(i64, 1), try runIntOp(.umulhi, .i64, "umulhi_i64_high_2", std.math.minInt(i64), 2));
3369 try testing.expectEqual(@as(i64, 1), try runIntOp(.umulhi, .u8, "umulhi_u8_max_2", 0xff, 2));
3370 try testing.expectEqual(@as(i64, -2), try runIntOp(.umulhi, .i16, "umulhi_i16_max_max", -1, -1));
3371 }
3372
3373 test "x86_64 arith.popcount over integers counts bits" {
3374 if (!supports_x86_64_backend) return;
3375 const testing = std.testing;
3376
3377 try testing.expectEqual(@as(i64, 0), try runIntOp(.popcount, .i32, "popcount_i32_zero", 0, 0));
3378 try testing.expectEqual(@as(i64, 16), try runIntOp(.popcount, .u32, "popcount_u32_half", 0x5555_5555, 0));
3379 try testing.expectEqual(@as(i64, 32), try runIntOp(.popcount, .i32, "popcount_i32_full", -1, 0));
3380 try testing.expectEqual(@as(i64, 64), try runIntOp(.popcount, .u64, "popcount_u64_full", -1, 0));
3381 try testing.expectEqual(@as(i64, 8), try runIntOp(.popcount, .u8, "popcount_u8_full", 0xff, 0));
3382 try testing.expectEqual(@as(i64, 64), try runIntOp(.popcount, .index, "popcount_index_full", -1, 0));
3383 }
3384
3385 test "x86_64 arith.neg over i8 / i16 wraps INT_MIN (movsx load path)" {
3386 if (!supports_x86_64_backend) return;
3387 const testing = std.testing;
3388
3389 try testing.expectEqual(@as(i64, -7), try runIntOp(.neg, .i8, "neg_i8_pos", 7, 0));
3390 try testing.expectEqual(@as(i64, 50), try runIntOp(.neg, .i8, "neg_i8_neg", -50, 0));
3391 try testing.expectEqual(
3392 @as(i64, std.math.minInt(i8)),
3393 try runIntOp(.neg, .i8, "neg_i8_min", std.math.minInt(i8), 0),
3394 );
3395
3396 try testing.expectEqual(@as(i64, -1234), try runIntOp(.neg, .i16, "neg_i16_pos", 1234, 0));
3397 try testing.expectEqual(@as(i64, 9999), try runIntOp(.neg, .i16, "neg_i16_neg", -9999, 0));
3398 try testing.expectEqual(
3399 @as(i64, std.math.minInt(i16)),
3400 try runIntOp(.neg, .i16, "neg_i16_min", std.math.minInt(i16), 0),
3401 );
3402 }
3403
3404 test "x86_64 arith.max / min over i8 / i16 with negative operands" {
3405 if (!supports_x86_64_backend) return;
3406 const testing = std.testing;
3407
3408 try testing.expectEqual(@as(i64, -3), try runIntOp(.max, .i8, "max_i8_neg", -7, -3));
3409 try testing.expectEqual(@as(i64, -7), try runIntOp(.min, .i8, "min_i8_neg", -7, -3));
3410 try testing.expectEqual(@as(i64, std.math.maxInt(i8)), try runIntOp(
3411 .max,
3412 .i8,
3413 "max_i8_extremes",
3414 std.math.minInt(i8),
3415 std.math.maxInt(i8),
3416 ));
3417
3418 try testing.expectEqual(@as(i64, -3), try runIntOp(.max, .i16, "max_i16_neg", -7, -3));
3419 try testing.expectEqual(@as(i64, -7), try runIntOp(.min, .i16, "min_i16_neg", -7, -3));
3420 }
3421
3422 test "x86_64 arith.max over arith.index uses unsigned cmov (cmovb path)" {
3423 if (!supports_x86_64_backend) return;
3424 const testing = std.testing;
3425 const dialects = @import("../../dialects/root.zig");
3426 const ArithDialect = dialects.ArithDialect;
3427 const BuiltinDialect = dialects.BuiltinDialect;
3428 const FuncDialect = dialects.FuncDialect;
3429
3430 var arena = alloc_arena.Arena.init(std.testing.allocator);
3431 defer arena.deinit();
3432 const allocator = arena.allocator();
3433
3434 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
3435 defer ir_ctx.deinit(allocator);
3436 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
3437
3438 const loc = ir.Location.getUnknown();
3439 const idx_type = try ArithDialect.getIndexType(&ir_ctx);
3440
3441 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
3442 const module_block = module.getBodyBlock();
3443
3444 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "max_index", &.{}, &.{idx_type});
3445 try module_block.addOperation(func.op);
3446
3447 const entry = func.getEntryBlock();
3448 var lhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, -1);
3449 try entry.addOperation(lhs.op);
3450 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, 1);
3451 try entry.addOperation(rhs.op);
3452 var max_op = try ArithDialect.MaxOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
3453 try entry.addOperation(max_op.op);
3454 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{max_op.getResult()});
3455 try entry.addOperation(ret.op);
3456
3457 var backend = try Backend.init(allocator, &ir_ctx, .testing);
3458 defer backend.deinit();
3459 const compiled = try backend.compile(module.op);
3460
3461 const exec_result = try onlyResult(&backend, compiled, "max_index", &.{});
3462 try testing.expectEqual(@as(i64, -1), try integerOf(exec_result));
3463 }
3464
3465 test "x86_64 arith.min over arith.index uses unsigned cmov (cmova path)" {
3466 if (!supports_x86_64_backend) return;
3467 const testing = std.testing;
3468 const dialects = @import("../../dialects/root.zig");
3469 const ArithDialect = dialects.ArithDialect;
3470 const BuiltinDialect = dialects.BuiltinDialect;
3471 const FuncDialect = dialects.FuncDialect;
3472
3473 var arena = alloc_arena.Arena.init(std.testing.allocator);
3474 defer arena.deinit();
3475 const allocator = arena.allocator();
3476
3477 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
3478 defer ir_ctx.deinit(allocator);
3479 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
3480
3481 const loc = ir.Location.getUnknown();
3482 const idx_type = try ArithDialect.getIndexType(&ir_ctx);
3483
3484 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
3485 const module_block = module.getBodyBlock();
3486
3487 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "min_index", &.{}, &.{idx_type});
3488 try module_block.addOperation(func.op);
3489
3490 const entry = func.getEntryBlock();
3491 var lhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, 1);
3492 try entry.addOperation(lhs.op);
3493 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, -1);
3494 try entry.addOperation(rhs.op);
3495 var min_op = try ArithDialect.MinOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
3496 try entry.addOperation(min_op.op);
3497 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{min_op.getResult()});
3498 try entry.addOperation(ret.op);
3499
3500 var backend = try Backend.init(allocator, &ir_ctx, .testing);
3501 defer backend.deinit();
3502 const compiled = try backend.compile(module.op);
3503
3504 const exec_result = try onlyResult(&backend, compiled, "min_index", &.{});
3505 try testing.expectEqual(@as(i64, 1), try integerOf(exec_result));
3506 }
3507
3508 test "x86_64 arith.max / min over arith.index — taken cmov path (reversed operands)" {
3509 if (!supports_x86_64_backend) return;
3510 const testing = std.testing;
3511 const dialects = @import("../../dialects/root.zig");
3512 const ArithDialect = dialects.ArithDialect;
3513 const BuiltinDialect = dialects.BuiltinDialect;
3514 const FuncDialect = dialects.FuncDialect;
3515
3516 inline for (.{
3517 .{ .op = ScalarOpKind.max, .lhs = @as(i64, 1), .rhs = @as(i64, -1), .name = "max_idx_taken", .expected = @as(i64, -1) },
3518 .{ .op = ScalarOpKind.min, .lhs = @as(i64, -1), .rhs = @as(i64, 1), .name = "min_idx_taken", .expected = @as(i64, 1) },
3519 }) |case| {
3520 var arena = alloc_arena.Arena.init(std.testing.allocator);
3521 defer arena.deinit();
3522 const allocator = arena.allocator();
3523
3524 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
3525 defer ir_ctx.deinit(allocator);
3526 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
3527
3528 const loc = ir.Location.getUnknown();
3529 const idx_type = try ArithDialect.getIndexType(&ir_ctx);
3530
3531 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
3532 const module_block = module.getBodyBlock();
3533
3534 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, case.name, &.{}, &.{idx_type});
3535 try module_block.addOperation(func.op);
3536
3537 const entry = func.getEntryBlock();
3538 var lhs_op = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, case.lhs);
3539 try entry.addOperation(lhs_op.op);
3540 var rhs_op = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, case.rhs);
3541 try entry.addOperation(rhs_op.op);
3542 const result_value = switch (case.op) {
3543 .max => blk: {
3544 const op = try ArithDialect.MaxOp.create(&ir_ctx, loc, lhs_op.getResult(), rhs_op.getResult());
3545 try entry.addOperation(op.op);
3546 break :blk op.getResult();
3547 },
3548 .min => blk: {
3549 const op = try ArithDialect.MinOp.create(&ir_ctx, loc, lhs_op.getResult(), rhs_op.getResult());
3550 try entry.addOperation(op.op);
3551 break :blk op.getResult();
3552 },
3553 else => unreachable,
3554 };
3555 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_value});
3556 try entry.addOperation(ret.op);
3557
3558 var backend = try Backend.init(allocator, &ir_ctx, .testing);
3559 defer backend.deinit();
3560 const compiled = try backend.compile(module.op);
3561
3562 const exec_result = try onlyResult(&backend, compiled, case.name, &.{});
3563 try testing.expectEqual(case.expected, try integerOf(exec_result));
3564 }
3565 }
3566
3567 test "x86_64 arith.max and min over u64 use unsigned order" {
3568 if (!supports_x86_64_backend) return;
3569 const testing = std.testing;
3570
3571 try testing.expectEqual(@as(i64, -1), try runIntOp(.max, .u64, "max_u64_high", -1, 1));
3572 try testing.expectEqual(@as(i64, 1), try runIntOp(.min, .u64, "min_u64_high", -1, 1));
3573 }
3574
3575 test "x86_64 arith.max and min over u32 use unsigned order" {
3576 if (!supports_x86_64_backend) return;
3577 const testing = std.testing;
3578
3579 try testing.expectEqual(@as(i64, 0x8000_0000), try runIntOp(.max, .u32, "max_u32_high", 0x8000_0000, 1));
3580 try testing.expectEqual(@as(i64, 1), try runIntOp(.min, .u32, "min_u32_high", 0x8000_0000, 1));
3581 }
3582
3583 test "x86_64 arith.max and min over u8 and u16 use unsigned order" {
3584 if (!supports_x86_64_backend) return;
3585 const testing = std.testing;
3586
3587 try testing.expectEqual(@as(i64, 0xff), try runIntOp(.max, .u8, "max_u8_high", 0xff, 1));
3588 try testing.expectEqual(@as(i64, 1), try runIntOp(.min, .u8, "min_u8_high", 0xff, 1));
3589 try testing.expectEqual(@as(i64, 0x8000), try runIntOp(.max, .u16, "max_u16_high", 0x8000, 1));
3590 try testing.expectEqual(@as(i64, 1), try runIntOp(.min, .u16, "min_u16_high", 0x8000, 1));
3591 }
3592
3593 test "x86_64 arith.min over i8 INT_MIN" {
3594 if (!supports_x86_64_backend) return;
3595 const testing = std.testing;
3596
3597 try testing.expectEqual(@as(i64, std.math.minInt(i8)), try runIntOp(
3598 .min,
3599 .i8,
3600 "min_i8_extremes",
3601 std.math.minInt(i8),
3602 std.math.maxInt(i8),
3603 ));
3604 }
3605
3606 test "x86_64 arith.sqrt over f32 / f64 (sqrtss / sqrtsd)" {
3607 if (!supports_x86_64_backend) return;
3608 const testing = std.testing;
3609
3610 try testing.expectApproxEqAbs(@as(f64, 2.0), try runFloatUnaryOp(.sqrt, .f32, "sqrt_f32_4", 4.0), 1e-6);
3611 try testing.expectApproxEqAbs(@as(f64, 0.0), try runFloatUnaryOp(.sqrt, .f32, "sqrt_f32_0", 0.0), 1e-6);
3612 try testing.expectApproxEqAbs(@as(f64, 3.0), try runFloatUnaryOp(.sqrt, .f64, "sqrt_f64_9", 9.0), 1e-12);
3613 try testing.expectApproxEqAbs(@as(f64, 4.0), try runFloatUnaryOp(.sqrt, .f64, "sqrt_f64_16", 16.0), 1e-12);
3614 }
3615
3616 test "x86_64 arith.abs over signed integers (mov+neg+cmovns)" {
3617 if (!supports_x86_64_backend) return;
3618 const testing = std.testing;
3619
3620 try testing.expectEqual(@as(i64, 5), try runIntOp(.abs, .i32, "abs_i32_pos", 5, 0));
3621 try testing.expectEqual(@as(i64, 5), try runIntOp(.abs, .i32, "abs_i32_neg", -5, 0));
3622 try testing.expectEqual(@as(i64, 0), try runIntOp(.abs, .i32, "abs_i32_zero", 0, 0));
3623 try testing.expectEqual(
3624 @as(i64, std.math.minInt(i32)),
3625 try runIntOp(.abs, .i32, "abs_i32_min", std.math.minInt(i32), 0),
3626 );
3627
3628 try testing.expectEqual(@as(i64, 1234567890), try runIntOp(.abs, .i64, "abs_i64_pos", 1234567890, 0));
3629 try testing.expectEqual(@as(i64, 1234567890), try runIntOp(.abs, .i64, "abs_i64_neg", -1234567890, 0));
3630 try testing.expectEqual(
3631 @as(i64, std.math.minInt(i64)),
3632 try runIntOp(.abs, .i64, "abs_i64_min", std.math.minInt(i64), 0),
3633 );
3634
3635 try testing.expectEqual(@as(i64, 7), try runIntOp(.abs, .i8, "abs_i8_neg", -7, 0));
3636 try testing.expectEqual(
3637 @as(i64, std.math.minInt(i8)),
3638 try runIntOp(.abs, .i8, "abs_i8_min", std.math.minInt(i8), 0),
3639 );
3640 try testing.expectEqual(@as(i64, 1234), try runIntOp(.abs, .i16, "abs_i16_neg", -1234, 0));
3641 try testing.expectEqual(
3642 @as(i64, std.math.minInt(i16)),
3643 try runIntOp(.abs, .i16, "abs_i16_min", std.math.minInt(i16), 0),
3644 );
3645 }
3646
3647 test "x86_64 arith.abs over u64 is identity" {
3648 if (!supports_x86_64_backend) return;
3649 const testing = std.testing;
3650
3651 const high_bit: i64 = @bitCast(@as(u64, 0x8000_0000_0000_0000));
3652 try testing.expectEqual(high_bit, try runIntOp(.abs, .u64, "abs_u64_high_bit", high_bit, 0));
3653 try testing.expectEqual(@as(i64, -1), try runIntOp(.abs, .u64, "abs_u64_max", -1, 0));
3654 }
3655
3656 test "x86_64 arith.neg over f32 (pcmpeqd + pslld + xorps sign-flip)" {
3657 if (!supports_x86_64_backend) return;
3658 const testing = std.testing;
3659
3660 try testing.expectEqual(@as(f32, -5.0), try runFloatUnaryOpRaw(.neg, f32, "neg_f32_pos", 5.0));
3661 try testing.expectEqual(@as(f32, 3.5), try runFloatUnaryOpRaw(.neg, f32, "neg_f32_neg", -3.5));
3662
3663 const neg_pos_zero = try runFloatUnaryOpRaw(.neg, f32, "neg_f32_pzero", 0.0);
3664 try testing.expectEqual(@as(u32, 0x80000000), @as(u32, @bitCast(neg_pos_zero)));
3665 const neg_neg_zero = try runFloatUnaryOpRaw(.neg, f32, "neg_f32_nzero", -0.0);
3666 try testing.expectEqual(@as(u32, 0x00000000), @as(u32, @bitCast(neg_neg_zero)));
3667
3668 const neg_nan = try runFloatUnaryOpRaw(.neg, f32, "neg_f32_nan", std.math.nan(f32));
3669 try testing.expect(std.math.isNan(neg_nan));
3670 }
3671
3672 test "x86_64 arith.neg over f64 (pcmpeqd + psllq + xorpd sign-flip)" {
3673 if (!supports_x86_64_backend) return;
3674 const testing = std.testing;
3675
3676 try testing.expectEqual(@as(f64, -5.0), try runFloatUnaryOpRaw(.neg, f64, "neg_f64_pos", 5.0));
3677 try testing.expectEqual(@as(f64, 3.5), try runFloatUnaryOpRaw(.neg, f64, "neg_f64_neg", -3.5));
3678
3679 const neg_pos_zero = try runFloatUnaryOpRaw(.neg, f64, "neg_f64_pzero", 0.0);
3680 try testing.expectEqual(@as(u64, 0x8000000000000000), @as(u64, @bitCast(neg_pos_zero)));
3681 const neg_neg_zero = try runFloatUnaryOpRaw(.neg, f64, "neg_f64_nzero", -0.0);
3682 try testing.expectEqual(@as(u64, 0x0000000000000000), @as(u64, @bitCast(neg_neg_zero)));
3683
3684 const neg_nan = try runFloatUnaryOpRaw(.neg, f64, "neg_f64_nan", std.math.nan(f64));
3685 try testing.expect(std.math.isNan(neg_nan));
3686 }
3687
3688 test "x86_64 arith.abs over f32 (pcmpeqd + psrld + andps sign-clear)" {
3689 if (!supports_x86_64_backend) return;
3690 const testing = std.testing;
3691
3692 try testing.expectEqual(@as(f32, 5.0), try runFloatUnaryOpRaw(.abs, f32, "abs_f32_pos", 5.0));
3693 try testing.expectEqual(@as(f32, 3.5), try runFloatUnaryOpRaw(.abs, f32, "abs_f32_neg", -3.5));
3694
3695 const abs_neg_zero = try runFloatUnaryOpRaw(.abs, f32, "abs_f32_nzero", -0.0);
3696 try testing.expectEqual(@as(u32, 0x00000000), @as(u32, @bitCast(abs_neg_zero)));
3697 const abs_pos_zero = try runFloatUnaryOpRaw(.abs, f32, "abs_f32_pzero", 0.0);
3698 try testing.expectEqual(@as(u32, 0x00000000), @as(u32, @bitCast(abs_pos_zero)));
3699
3700 const abs_nan = try runFloatUnaryOpRaw(.abs, f32, "abs_f32_nan", std.math.nan(f32));
3701 try testing.expect(std.math.isNan(abs_nan));
3702 }
3703
3704 test "x86_64 arith.abs over f64 (pcmpeqd + psrlq + andpd sign-clear)" {
3705 if (!supports_x86_64_backend) return;
3706 const testing = std.testing;
3707
3708 try testing.expectEqual(@as(f64, 5.0), try runFloatUnaryOpRaw(.abs, f64, "abs_f64_pos", 5.0));
3709 try testing.expectEqual(@as(f64, 3.5), try runFloatUnaryOpRaw(.abs, f64, "abs_f64_neg", -3.5));
3710
3711 const abs_neg_zero = try runFloatUnaryOpRaw(.abs, f64, "abs_f64_nzero", -0.0);
3712 try testing.expectEqual(@as(u64, 0x0000000000000000), @as(u64, @bitCast(abs_neg_zero)));
3713 const abs_pos_zero = try runFloatUnaryOpRaw(.abs, f64, "abs_f64_pzero", 0.0);
3714 try testing.expectEqual(@as(u64, 0x0000000000000000), @as(u64, @bitCast(abs_pos_zero)));
3715
3716 const abs_nan = try runFloatUnaryOpRaw(.abs, f64, "abs_f64_nan", std.math.nan(f64));
3717 try testing.expect(std.math.isNan(abs_nan));
3718 }
3719
3720 test "x86_64 arith.floor over f32 / f64 via libm floorf / floor" {
3721 if (!supports_x86_64_backend) return;
3722 const testing = std.testing;
3723
3724 try testing.expectEqual(@as(f32, 1.0), try runFloatUnaryOpRaw(.floor, f32, "floor_f32_pos", 1.75));
3725 try testing.expectEqual(@as(f32, -2.0), try runFloatUnaryOpRaw(.floor, f32, "floor_f32_neg", -1.25));
3726 try testing.expectEqual(@as(f64, 1.0), try runFloatUnaryOpRaw(.floor, f64, "floor_f64_pos", 1.75));
3727 try testing.expectEqual(@as(f64, -2.0), try runFloatUnaryOpRaw(.floor, f64, "floor_f64_neg", -1.25));
3728 const floor_nan = try runFloatUnaryOpRaw(.floor, f64, "floor_f64_nan", std.math.nan(f64));
3729 try testing.expect(std.math.isNan(floor_nan));
3730 }
3731
3732 test "x86_64 arith.sin over f32 / f64 via libm sinf / sin" {
3733 if (!supports_x86_64_backend) return;
3734 const testing = std.testing;
3735
3736 try testing.expectApproxEqAbs(
3737 @as(f32, 0.0),
3738 try runFloatUnaryOpRaw(.sin, f32, "sin_f32_zero", 0.0),
3739 1e-5,
3740 );
3741 try testing.expectApproxEqAbs(
3742 @as(f32, 1.0),
3743 try runFloatUnaryOpRaw(.sin, f32, "sin_f32_pi_over_2", std.math.pi / 2.0),
3744 1e-5,
3745 );
3746 try testing.expectApproxEqAbs(
3747 @as(f64, 0.0),
3748 try runFloatUnaryOpRaw(.sin, f64, "sin_f64_zero", 0.0),
3749 1e-12,
3750 );
3751 try testing.expectApproxEqAbs(
3752 @as(f64, 1.0),
3753 try runFloatUnaryOpRaw(.sin, f64, "sin_f64_pi_over_2", std.math.pi / 2.0),
3754 1e-12,
3755 );
3756 const sin_nan = try runFloatUnaryOpRaw(.sin, f64, "sin_f64_nan", std.math.nan(f64));
3757 try testing.expect(std.math.isNan(sin_nan));
3758 }
3759
3760 test "x86_64 arith.cos over f32 / f64 via libm cosf / cos" {
3761 if (!supports_x86_64_backend) return;
3762 const testing = std.testing;
3763
3764 try testing.expectApproxEqAbs(
3765 @as(f32, 1.0),
3766 try runFloatUnaryOpRaw(.cos, f32, "cos_f32_zero", 0.0),
3767 1e-5,
3768 );
3769 try testing.expectApproxEqAbs(
3770 @as(f32, 0.0),
3771 try runFloatUnaryOpRaw(.cos, f32, "cos_f32_pi_over_2", std.math.pi / 2.0),
3772 1e-5,
3773 );
3774 try testing.expectApproxEqAbs(
3775 @as(f64, 1.0),
3776 try runFloatUnaryOpRaw(.cos, f64, "cos_f64_zero", 0.0),
3777 1e-12,
3778 );
3779 try testing.expectApproxEqAbs(
3780 @as(f64, -1.0),
3781 try runFloatUnaryOpRaw(.cos, f64, "cos_f64_pi", std.math.pi),
3782 1e-12,
3783 );
3784 }
3785
3786 test "x86_64 arith.tan over f32 / f64 via libm tanf / tan" {
3787 if (!supports_x86_64_backend) return;
3788 const testing = std.testing;
3789
3790 try testing.expectApproxEqAbs(
3791 @as(f32, 0.0),
3792 try runFloatUnaryOpRaw(.tan, f32, "tan_f32_zero", 0.0),
3793 1e-5,
3794 );
3795 try testing.expectApproxEqAbs(
3796 @as(f64, 1.0),
3797 try runFloatUnaryOpRaw(.tan, f64, "tan_f64_pi_over_4", std.math.pi / 4.0),
3798 1e-12,
3799 );
3800 }
3801
3802 test "x86_64 arith.exp over f32 / f64 via libm expf / exp" {
3803 if (!supports_x86_64_backend) return;
3804 const testing = std.testing;
3805
3806 try testing.expectApproxEqAbs(
3807 @as(f32, 1.0),
3808 try runFloatUnaryOpRaw(.exp, f32, "exp_f32_zero", 0.0),
3809 1e-5,
3810 );
3811 try testing.expectApproxEqAbs(
3812 @as(f32, std.math.e),
3813 try runFloatUnaryOpRaw(.exp, f32, "exp_f32_one", 1.0),
3814 1e-5,
3815 );
3816 try testing.expectApproxEqAbs(
3817 @as(f64, 1.0),
3818 try runFloatUnaryOpRaw(.exp, f64, "exp_f64_zero", 0.0),
3819 1e-12,
3820 );
3821 try testing.expectApproxEqAbs(
3822 @as(f64, std.math.e),
3823 try runFloatUnaryOpRaw(.exp, f64, "exp_f64_one", 1.0),
3824 1e-12,
3825 );
3826 }
3827
3828 test "x86_64 arith.log over f32 / f64 via libm logf / log" {
3829 if (!supports_x86_64_backend) return;
3830 const testing = std.testing;
3831
3832 try testing.expectApproxEqAbs(
3833 @as(f32, 0.0),
3834 try runFloatUnaryOpRaw(.log, f32, "log_f32_one", 1.0),
3835 1e-5,
3836 );
3837 try testing.expectApproxEqAbs(
3838 @as(f32, 1.0),
3839 try runFloatUnaryOpRaw(.log, f32, "log_f32_e", std.math.e),
3840 1e-5,
3841 );
3842 try testing.expectApproxEqAbs(
3843 @as(f64, 0.0),
3844 try runFloatUnaryOpRaw(.log, f64, "log_f64_one", 1.0),
3845 1e-12,
3846 );
3847 try testing.expectApproxEqAbs(
3848 @as(f64, 1.0),
3849 try runFloatUnaryOpRaw(.log, f64, "log_f64_e", std.math.e),
3850 1e-12,
3851 );
3852 }
3853
3854 test "x86_64 arith.tanh over f32 / f64 via libm tanhf / tanh" {
3855 if (!supports_x86_64_backend) return;
3856 const testing = std.testing;
3857
3858 try testing.expectApproxEqAbs(
3859 @as(f32, 0.0),
3860 try runFloatUnaryOpRaw(.tanh, f32, "tanh_f32_zero", 0.0),
3861 1e-5,
3862 );
3863 try testing.expectApproxEqAbs(
3864 @as(f32, 1.0),
3865 try runFloatUnaryOpRaw(.tanh, f32, "tanh_f32_large", 20.0),
3866 1e-5,
3867 );
3868 try testing.expectApproxEqAbs(
3869 @as(f64, 0.0),
3870 try runFloatUnaryOpRaw(.tanh, f64, "tanh_f64_zero", 0.0),
3871 1e-12,
3872 );
3873 try testing.expectApproxEqAbs(
3874 @as(f64, -1.0),
3875 try runFloatUnaryOpRaw(.tanh, f64, "tanh_f64_neg_large", -20.0),
3876 1e-12,
3877 );
3878 }
3879
3880 test "x86_64 arith.pow over f32 / f64 via libm powf / pow" {
3881 if (!supports_x86_64_backend) return;
3882 const testing = std.testing;
3883
3884 try testing.expectApproxEqAbs(
3885 @as(f32, 1024.0),
3886 try runFloatBinaryOpRaw(.pow, f32, "pow_f32_2_10", 2.0, 10.0),
3887 1e-3,
3888 );
3889 try testing.expectApproxEqAbs(
3890 @as(f32, std.math.sqrt2),
3891 try runFloatBinaryOpRaw(.pow, f32, "pow_f32_2_half", 2.0, 0.5),
3892 1e-5,
3893 );
3894 try testing.expectApproxEqAbs(
3895 @as(f64, 1024.0),
3896 try runFloatBinaryOpRaw(.pow, f64, "pow_f64_2_10", 2.0, 10.0),
3897 1e-9,
3898 );
3899 try testing.expectApproxEqAbs(
3900 @as(f64, std.math.sqrt2),
3901 try runFloatBinaryOpRaw(.pow, f64, "pow_f64_2_half", 2.0, 0.5),
3902 1e-12,
3903 );
3904 }
3905
3906 test "x86_64 arith.pow rejects mismatched base / exponent types" {
3907 if (!supports_x86_64_backend) return;
3908 const testing = std.testing;
3909 const dialects = @import("../../dialects/root.zig");
3910 const ArithDialect = dialects.ArithDialect;
3911 const BuiltinDialect = dialects.BuiltinDialect;
3912 const FuncDialect = dialects.FuncDialect;
3913
3914 var arena = alloc_arena.Arena.init(std.testing.allocator);
3915 defer arena.deinit();
3916 const allocator = arena.allocator();
3917
3918 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
3919 defer ir_ctx.deinit(allocator);
3920 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
3921
3922 const loc = ir.Location.getUnknown();
3923 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
3924 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
3925
3926 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
3927 const module_block = module.getBodyBlock();
3928
3929 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "pow_mismatched", &.{}, &.{f32_type});
3930 try module_block.addOperation(func.op);
3931
3932 const entry = func.getEntryBlock();
3933 var base = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 2.0);
3934 try entry.addOperation(base.op);
3935 var exponent = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 3.0);
3936 try entry.addOperation(exponent.op);
3937 var op = try ArithDialect.PowOp.create(&ir_ctx, loc, base.getResult(), exponent.getResult());
3938 try entry.addOperation(op.op);
3939
3940 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{op.getResult()});
3941 try entry.addOperation(ret.op);
3942
3943 var backend = try Backend.init(allocator, &ir_ctx, .testing);
3944 defer backend.deinit();
3945 try testing.expectError(error.VerificationFailed, backend.compile(module.op));
3946 }
3947
3948 test "x86_64 arith.max over f32: ordered + NaN-aware contract" {
3949 if (!supports_x86_64_backend) return;
3950 const testing = std.testing;
3951
3952 try testing.expectEqual(@as(f32, 7.0), try runFloatBinaryOpRaw(.max, f32, "max_f32_ord_a", 3.0, 7.0));
3953 try testing.expectEqual(@as(f32, 7.0), try runFloatBinaryOpRaw(.max, f32, "max_f32_ord_b", 7.0, 3.0));
3954 try testing.expectEqual(@as(f32, -3.0), try runFloatBinaryOpRaw(.max, f32, "max_f32_negs", -5.0, -3.0));
3955 try testing.expectEqual(
3956 @as(f32, 5.0),
3957 try runFloatBinaryOpRaw(.max, f32, "max_f32_nan_lhs", std.math.nan(f32), 5.0),
3958 );
3959 try testing.expectEqual(
3960 @as(f32, 5.0),
3961 try runFloatBinaryOpRaw(.max, f32, "max_f32_nan_rhs", 5.0, std.math.nan(f32)),
3962 );
3963 const both_nan = try runFloatBinaryOpRaw(.max, f32, "max_f32_both_nan", std.math.nan(f32), std.math.nan(f32));
3964 try testing.expect(std.math.isNan(both_nan));
3965 }
3966
3967 test "x86_64 arith.max over f64: ordered + NaN-aware contract" {
3968 if (!supports_x86_64_backend) return;
3969 const testing = std.testing;
3970
3971 try testing.expectEqual(@as(f64, 7.0), try runFloatBinaryOpRaw(.max, f64, "max_f64_ord_a", 3.0, 7.0));
3972 try testing.expectEqual(@as(f64, 7.0), try runFloatBinaryOpRaw(.max, f64, "max_f64_ord_b", 7.0, 3.0));
3973 try testing.expectEqual(@as(f64, -3.0), try runFloatBinaryOpRaw(.max, f64, "max_f64_negs", -5.0, -3.0));
3974 try testing.expectEqual(
3975 @as(f64, 5.0),
3976 try runFloatBinaryOpRaw(.max, f64, "max_f64_nan_lhs", std.math.nan(f64), 5.0),
3977 );
3978 try testing.expectEqual(
3979 @as(f64, 5.0),
3980 try runFloatBinaryOpRaw(.max, f64, "max_f64_nan_rhs", 5.0, std.math.nan(f64)),
3981 );
3982 const both_nan = try runFloatBinaryOpRaw(.max, f64, "max_f64_both_nan", std.math.nan(f64), std.math.nan(f64));
3983 try testing.expect(std.math.isNan(both_nan));
3984 }
3985
3986 test "x86_64 arith.min over f32: ordered + NaN-aware contract" {
3987 if (!supports_x86_64_backend) return;
3988 const testing = std.testing;
3989
3990 try testing.expectEqual(@as(f32, 3.0), try runFloatBinaryOpRaw(.min, f32, "min_f32_ord_a", 3.0, 7.0));
3991 try testing.expectEqual(@as(f32, 3.0), try runFloatBinaryOpRaw(.min, f32, "min_f32_ord_b", 7.0, 3.0));
3992 try testing.expectEqual(@as(f32, -5.0), try runFloatBinaryOpRaw(.min, f32, "min_f32_negs", -5.0, -3.0));
3993 try testing.expectEqual(
3994 @as(f32, 5.0),
3995 try runFloatBinaryOpRaw(.min, f32, "min_f32_nan_lhs", std.math.nan(f32), 5.0),
3996 );
3997 try testing.expectEqual(
3998 @as(f32, 5.0),
3999 try runFloatBinaryOpRaw(.min, f32, "min_f32_nan_rhs", 5.0, std.math.nan(f32)),
4000 );
4001 const both_nan = try runFloatBinaryOpRaw(.min, f32, "min_f32_both_nan", std.math.nan(f32), std.math.nan(f32));
4002 try testing.expect(std.math.isNan(both_nan));
4003 }
4004
4005 test "x86_64 arith.min over f64: ordered + NaN-aware contract" {
4006 if (!supports_x86_64_backend) return;
4007 const testing = std.testing;
4008
4009 try testing.expectEqual(@as(f64, 3.0), try runFloatBinaryOpRaw(.min, f64, "min_f64_ord_a", 3.0, 7.0));
4010 try testing.expectEqual(@as(f64, 3.0), try runFloatBinaryOpRaw(.min, f64, "min_f64_ord_b", 7.0, 3.0));
4011 try testing.expectEqual(@as(f64, -5.0), try runFloatBinaryOpRaw(.min, f64, "min_f64_negs", -5.0, -3.0));
4012 try testing.expectEqual(
4013 @as(f64, 5.0),
4014 try runFloatBinaryOpRaw(.min, f64, "min_f64_nan_lhs", std.math.nan(f64), 5.0),
4015 );
4016 try testing.expectEqual(
4017 @as(f64, 5.0),
4018 try runFloatBinaryOpRaw(.min, f64, "min_f64_nan_rhs", 5.0, std.math.nan(f64)),
4019 );
4020 const both_nan = try runFloatBinaryOpRaw(.min, f64, "min_f64_both_nan", std.math.nan(f64), std.math.nan(f64));
4021 try testing.expect(std.math.isNan(both_nan));
4022 }
4023
4024 test "x86_64 Backend.init wires the arith verifier traits even on a bare context (tick-33 [P2] closure)" {
4025 if (!supports_x86_64_backend) return;
4026 const testing = std.testing;
4027 const dialects = @import("../../dialects/root.zig");
4028 const ArithDialect = dialects.ArithDialect;
4029 const BuiltinDialect = dialects.BuiltinDialect;
4030 const FuncDialect = dialects.FuncDialect;
4031
4032 var arena = alloc_arena.Arena.init(std.testing.allocator);
4033 defer arena.deinit();
4034 const allocator = arena.allocator();
4035
4036 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4037 defer ir_ctx.deinit(allocator);
4038 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4039
4040 const loc = ir.Location.getUnknown();
4041 const i32_type = try ArithDialect.getI32Type(&ir_ctx);
4042 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
4043
4044 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4045 const module_block = module.getBodyBlock();
4046 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "add_mismatched", &.{}, &.{i32_type});
4047 try module_block.addOperation(func.op);
4048
4049 const entry = func.getEntryBlock();
4050 var lhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 1);
4051 try entry.addOperation(lhs.op);
4052 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 2);
4053 try entry.addOperation(rhs.op);
4054 var add = try ArithDialect.AddOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
4055 try entry.addOperation(add.op);
4056 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{add.getResult()});
4057 try entry.addOperation(ret.op);
4058
4059 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4060 defer backend.deinit();
4061 try testing.expectError(error.VerificationFailed, backend.compile(module.op));
4062 }
4063
4064 test "x86_64 arith.max rejects mismatched lhs / rhs types" {
4065 if (!supports_x86_64_backend) return;
4066 const testing = std.testing;
4067 const dialects = @import("../../dialects/root.zig");
4068 const ArithDialect = dialects.ArithDialect;
4069 const BuiltinDialect = dialects.BuiltinDialect;
4070 const FuncDialect = dialects.FuncDialect;
4071
4072 var arena = alloc_arena.Arena.init(std.testing.allocator);
4073 defer arena.deinit();
4074 const allocator = arena.allocator();
4075
4076 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4077 defer ir_ctx.deinit(allocator);
4078 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4079
4080 const loc = ir.Location.getUnknown();
4081 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
4082 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
4083
4084 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4085 const module_block = module.getBodyBlock();
4086
4087 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "max_mismatched", &.{}, &.{f32_type});
4088 try module_block.addOperation(func.op);
4089
4090 const entry = func.getEntryBlock();
4091 var lhs = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 2.0);
4092 try entry.addOperation(lhs.op);
4093 var rhs = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 3.0);
4094 try entry.addOperation(rhs.op);
4095 var op = try ArithDialect.MaxOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
4096 try entry.addOperation(op.op);
4097
4098 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{op.getResult()});
4099 try entry.addOperation(ret.op);
4100
4101 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4102 defer backend.deinit();
4103 try testing.expectError(error.VerificationFailed, backend.compile(module.op));
4104 }
4105
4106 test "x86_64 arith.max/min handle signed zero ordered inputs" {
4107 if (!supports_x86_64_backend) return;
4108 const testing = std.testing;
4109
4110 try testing.expectEqual(
4111 @as(f32, 0.0),
4112 try runFloatBinaryOpRaw(.max, f32, "max_f32_signed_zeros", 0.0, -0.0),
4113 );
4114 try testing.expectEqual(
4115 @as(f32, 0.0),
4116 try runFloatBinaryOpRaw(.min, f32, "min_f32_signed_zeros", 0.0, -0.0),
4117 );
4118 try testing.expectEqual(
4119 @as(f64, 0.0),
4120 try runFloatBinaryOpRaw(.max, f64, "max_f64_signed_zeros", 0.0, -0.0),
4121 );
4122 try testing.expectEqual(
4123 @as(f64, 0.0),
4124 try runFloatBinaryOpRaw(.min, f64, "min_f64_signed_zeros", 0.0, -0.0),
4125 );
4126 }
4127
4128 test "x86_64 arith.and over signed integers" {
4129 if (!supports_x86_64_backend) return;
4130 const testing = std.testing;
4131
4132 try testing.expectEqual(@as(i64, 0x0F00), try runIntOp(.band, .i32, "and_i32_a", 0xFF00, 0x0F0F));
4133 try testing.expectEqual(@as(i64, 0), try runIntOp(.band, .i32, "and_i32_zero", 0xFFFF, 0));
4134 try testing.expectEqual(@as(i64, 0x0F0F), try runIntOp(.band, .i32, "and_i32_id", 0xFFFF, 0x0F0F));
4135 try testing.expectEqual(@as(i64, 0x0F00), try runIntOp(.band, .i64, "and_i64_a", 0xFF00, 0x0F0F));
4136 try testing.expectEqual(
4137 @as(i64, std.math.maxInt(i64)),
4138 try runIntOp(.band, .i64, "and_i64_mask_high", -1, std.math.maxInt(i64)),
4139 );
4140 }
4141
4142 test "x86_64 arith.or over signed integers" {
4143 if (!supports_x86_64_backend) return;
4144 const testing = std.testing;
4145
4146 try testing.expectEqual(@as(i64, 0xFF0F), try runIntOp(.bor, .i32, "or_i32_a", 0xFF00, 0x0F0F));
4147 try testing.expectEqual(@as(i64, 0xFFFF), try runIntOp(.bor, .i32, "or_i32_full", 0xFF00, 0x00FF));
4148 try testing.expectEqual(@as(i64, 0xFF0F), try runIntOp(.bor, .i64, "or_i64_a", 0xFF00, 0x0F0F));
4149 }
4150
4151 test "x86_64 arith.xor over signed integers" {
4152 if (!supports_x86_64_backend) return;
4153 const testing = std.testing;
4154
4155 try testing.expectEqual(@as(i64, 0xF00F), try runIntOp(.bxor, .i32, "xor_i32_a", 0xFF00, 0x0F0F));
4156 try testing.expectEqual(@as(i64, 0), try runIntOp(.bxor, .i32, "xor_i32_self", 0x12345678, 0x12345678));
4157 try testing.expectEqual(@as(i64, 0x12345678), try runIntOp(.bxor, .i32, "xor_i32_id", 0x12345678, 0));
4158 try testing.expectEqual(@as(i64, 0xF00F), try runIntOp(.bxor, .i64, "xor_i64_a", 0xFF00, 0x0F0F));
4159 }
4160
4161 test "x86_64 arith.not over signed integers" {
4162 if (!supports_x86_64_backend) return;
4163 const testing = std.testing;
4164
4165 try testing.expectEqual(@as(i64, -1), try runIntOp(.bnot, .i32, "not_i32_zero", 0, 0));
4166 try testing.expectEqual(@as(i64, 0), try runIntOp(.bnot, .i32, "not_i32_neg1", -1, 0));
4167 try testing.expectEqual(@as(i64, -6), try runIntOp(.bnot, .i32, "not_i32_5", 5, 0));
4168 try testing.expectEqual(@as(i64, -1), try runIntOp(.bnot, .i64, "not_i64_zero", 0, 0));
4169 try testing.expectEqual(@as(i64, 0), try runIntOp(.bnot, .i64, "not_i64_neg1", -1, 0));
4170 }
4171
4172 test "x86_64 arith.shl over signed integers (shift left)" {
4173 if (!supports_x86_64_backend) return;
4174 const testing = std.testing;
4175
4176 try testing.expectEqual(@as(i64, 16), try runIntOp(.shl, .i32, "shl_i32_1_4", 1, 4));
4177 try testing.expectEqual(@as(i64, 1024), try runIntOp(.shl, .i32, "shl_i32_1_10", 1, 10));
4178 try testing.expectEqual(@as(i64, 5), try runIntOp(.shl, .i32, "shl_i32_5_0", 5, 0));
4179 try testing.expectEqual(@as(i64, 16), try runIntOp(.shl, .i64, "shl_i64_1_4", 1, 4));
4180 try testing.expectEqual(
4181 @as(i64, @as(i64, 1) << 32),
4182 try runIntOp(.shl, .i64, "shl_i64_1_32", 1, 32),
4183 );
4184 }
4185
4186 test "x86_64 arith.shr over signed integers (arithmetic / sign-extending)" {
4187 if (!supports_x86_64_backend) return;
4188 const testing = std.testing;
4189
4190 try testing.expectEqual(@as(i64, 4), try runIntOp(.shr, .i32, "shr_i32_16_2", 16, 2));
4191 try testing.expectEqual(@as(i64, -4), try runIntOp(.shr, .i32, "shr_i32_neg16_2", -16, 2));
4192 try testing.expectEqual(@as(i64, -16), try runIntOp(.shr, .i32, "shr_i32_neg16_0", -16, 0));
4193 try testing.expectEqual(@as(i64, 4), try runIntOp(.shr, .i64, "shr_i64_16_2", 16, 2));
4194 try testing.expectEqual(@as(i64, -1), try runIntOp(.shr, .i64, "shr_i64_neg1_5", -1, 5));
4195 }
4196
4197 test "x86_64 arith.ushr over signed integers (logical / zero-extending)" {
4198 if (!supports_x86_64_backend) return;
4199 const testing = std.testing;
4200
4201 try testing.expectEqual(@as(i64, 4), try runIntOp(.ushr, .i32, "ushr_i32_16_2", 16, 2));
4202 try testing.expectEqual(
4203 @as(i64, 0x3FFFFFFC),
4204 try runIntOp(.ushr, .i32, "ushr_i32_neg16_2", -16, 2),
4205 );
4206 try testing.expectEqual(
4207 @as(i64, std.math.maxInt(i64)),
4208 try runIntOp(.ushr, .i64, "ushr_i64_neg1_1", -1, 1),
4209 );
4210 try testing.expectEqual(@as(i64, 127), try runIntOp(.ushr, .i8, "ushr_i8_neg1_1", -1, 1));
4211 }
4212
4213 test "x86_64 arith.ushr over u64 is logical" {
4214 if (!supports_x86_64_backend) return;
4215 const testing = std.testing;
4216
4217 try testing.expectEqual(@as(i64, std.math.maxInt(i64)), try runIntOp(.ushr, .u64, "ushr_u64_max_1", -1, 1));
4218 }
4219
4220 test "x86_64 arith.ushr over u32 is logical" {
4221 if (!supports_x86_64_backend) return;
4222 const testing = std.testing;
4223
4224 try testing.expectEqual(@as(i64, 0x4000_0000), try runIntOp(.ushr, .u32, "ushr_u32_high_1", 0x8000_0000, 1));
4225 }
4226
4227 test "x86_64 arith.shr over arith.index is rejected (unsigned-but-arithmetic mismatch)" {
4228 if (!supports_x86_64_backend) return;
4229 const testing = std.testing;
4230 const dialects = @import("../../dialects/root.zig");
4231 const ArithDialect = dialects.ArithDialect;
4232 const BuiltinDialect = dialects.BuiltinDialect;
4233 const FuncDialect = dialects.FuncDialect;
4234
4235 var arena = alloc_arena.Arena.init(std.testing.allocator);
4236 defer arena.deinit();
4237 const allocator = arena.allocator();
4238
4239 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4240 defer ir_ctx.deinit(allocator);
4241 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4242
4243 const loc = ir.Location.getUnknown();
4244 const idx_type = try ArithDialect.getScalarType(&ir_ctx, .index);
4245
4246 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4247 const module_block = module.getBodyBlock();
4248
4249 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "shr_index", &.{}, &.{idx_type});
4250 try module_block.addOperation(func.op);
4251
4252 const entry = func.getEntryBlock();
4253 var lhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, 16);
4254 try entry.addOperation(lhs.op);
4255 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, 2);
4256 try entry.addOperation(rhs.op);
4257 var op = try ArithDialect.ShrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
4258 try entry.addOperation(op.op);
4259 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{op.getResult()});
4260 try entry.addOperation(ret.op);
4261
4262 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4263 defer backend.deinit();
4264 try testing.expectError(error.JitCompileFailed, backend.compile(module.op));
4265 }
4266
4267 test "x86_64 arith.shr over u64 is rejected" {
4268 if (!supports_x86_64_backend) return;
4269 const testing = std.testing;
4270 const dialects = @import("../../dialects/root.zig");
4271 const ArithDialect = dialects.ArithDialect;
4272 const BuiltinDialect = dialects.BuiltinDialect;
4273 const FuncDialect = dialects.FuncDialect;
4274
4275 var arena = alloc_arena.Arena.init(std.testing.allocator);
4276 defer arena.deinit();
4277 const allocator = arena.allocator();
4278
4279 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4280 defer ir_ctx.deinit(allocator);
4281 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4282
4283 const loc = ir.Location.getUnknown();
4284 const u64_type = try ArithDialect.getScalarType(&ir_ctx, .u64);
4285
4286 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4287 const module_block = module.getBodyBlock();
4288
4289 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "shr_u64", &.{}, &.{u64_type});
4290 try module_block.addOperation(func.op);
4291
4292 const entry = func.getEntryBlock();
4293 var lhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u64_type, -1);
4294 try entry.addOperation(lhs.op);
4295 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u64_type, 1);
4296 try entry.addOperation(rhs.op);
4297 var op = try ArithDialect.ShrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
4298 try entry.addOperation(op.op);
4299 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{op.getResult()});
4300 try entry.addOperation(ret.op);
4301
4302 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4303 defer backend.deinit();
4304 try testing.expectError(error.JitCompileFailed, backend.compile(module.op));
4305 }
4306
4307 test "x86_64 arith.shr over u32 is rejected" {
4308 if (!supports_x86_64_backend) return;
4309 const testing = std.testing;
4310 const dialects = @import("../../dialects/root.zig");
4311 const ArithDialect = dialects.ArithDialect;
4312 const BuiltinDialect = dialects.BuiltinDialect;
4313 const FuncDialect = dialects.FuncDialect;
4314
4315 var arena = alloc_arena.Arena.init(std.testing.allocator);
4316 defer arena.deinit();
4317 const allocator = arena.allocator();
4318
4319 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4320 defer ir_ctx.deinit(allocator);
4321 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4322
4323 const loc = ir.Location.getUnknown();
4324 const u32_type = try ArithDialect.getScalarType(&ir_ctx, .u32);
4325
4326 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4327 const module_block = module.getBodyBlock();
4328
4329 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "shr_u32", &.{}, &.{u32_type});
4330 try module_block.addOperation(func.op);
4331
4332 const entry = func.getEntryBlock();
4333 var lhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0x8000_0000);
4334 try entry.addOperation(lhs.op);
4335 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 1);
4336 try entry.addOperation(rhs.op);
4337 var op = try ArithDialect.ShrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
4338 try entry.addOperation(op.op);
4339 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{op.getResult()});
4340 try entry.addOperation(ret.op);
4341
4342 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4343 defer backend.deinit();
4344 try testing.expectError(error.JitCompileFailed, backend.compile(module.op));
4345 }
4346
4347 fn runCastInt(
4348 input_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
4349 output_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
4350 name: []const u8,
4351 c_in: i64,
4352 ) !i64 {
4353 const dialects = @import("../../dialects/root.zig");
4354 const ArithDialect = dialects.ArithDialect;
4355 const BuiltinDialect = dialects.BuiltinDialect;
4356 const FuncDialect = dialects.FuncDialect;
4357
4358 var arena = alloc_arena.Arena.init(std.testing.allocator);
4359 defer arena.deinit();
4360 const allocator = arena.allocator();
4361
4362 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4363 defer ir_ctx.deinit(allocator);
4364 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4365
4366 const loc = ir.Location.getUnknown();
4367 const in_type = try ArithDialect.getScalarType(&ir_ctx, input_kind);
4368 const out_type = try ArithDialect.getScalarType(&ir_ctx, output_kind);
4369
4370 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4371 const module_block = module.getBodyBlock();
4372
4373 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{out_type});
4374 try module_block.addOperation(func.op);
4375
4376 const entry = func.getEntryBlock();
4377 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, in_type, c_in);
4378 try entry.addOperation(c.op);
4379 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), out_type);
4380 try entry.addOperation(cast.op);
4381 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
4382 try entry.addOperation(ret.op);
4383
4384 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4385 defer backend.deinit();
4386 const compiled = try backend.compile(module.op);
4387 const exec_result = try onlyResult(&backend, compiled, name, &.{});
4388 return integerOf(exec_result);
4389 }
4390
4391 fn runCastFloatToFloat(
4392 comptime InFT: type,
4393 comptime OutFT: type,
4394 name: []const u8,
4395 c_in: InFT,
4396 ) !OutFT {
4397 const dialects = @import("../../dialects/root.zig");
4398 const ArithDialect = dialects.ArithDialect;
4399 const BuiltinDialect = dialects.BuiltinDialect;
4400 const FuncDialect = dialects.FuncDialect;
4401
4402 const in_kind: ArithDialect.ScalarTypeKind = comptime if (InFT == f32) .f32 else if (InFT == f64) .f64 else @compileError("InFT");
4403 const out_kind: ArithDialect.ScalarTypeKind = comptime if (OutFT == f32) .f32 else if (OutFT == f64) .f64 else @compileError("OutFT");
4404
4405 var arena = alloc_arena.Arena.init(std.testing.allocator);
4406 defer arena.deinit();
4407 const allocator = arena.allocator();
4408
4409 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4410 defer ir_ctx.deinit(allocator);
4411 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4412
4413 const loc = ir.Location.getUnknown();
4414 const in_type = try ArithDialect.getScalarType(&ir_ctx, in_kind);
4415 const out_type = try ArithDialect.getScalarType(&ir_ctx, out_kind);
4416
4417 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4418 const module_block = module.getBodyBlock();
4419
4420 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{out_type});
4421 try module_block.addOperation(func.op);
4422
4423 const entry = func.getEntryBlock();
4424 var c = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, in_type, @floatCast(c_in));
4425 try entry.addOperation(c.op);
4426 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), out_type);
4427 try entry.addOperation(cast.op);
4428 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
4429 try entry.addOperation(ret.op);
4430
4431 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4432 defer backend.deinit();
4433 const compiled = try backend.compile(module.op);
4434
4435 const FnPtr = *const fn () callconv(.c) OutFT;
4436 const fn_ptr = try backend.runtime.getFunction(compiled, name, FnPtr);
4437 return fn_ptr();
4438 }
4439
4440 fn runCastIntToFloat(
4441 input_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
4442 comptime OutFT: type,
4443 name: []const u8,
4444 c_in: i64,
4445 ) !OutFT {
4446 const dialects = @import("../../dialects/root.zig");
4447 const ArithDialect = dialects.ArithDialect;
4448 const BuiltinDialect = dialects.BuiltinDialect;
4449 const FuncDialect = dialects.FuncDialect;
4450
4451 const out_kind: ArithDialect.ScalarTypeKind = comptime if (OutFT == f32) .f32 else if (OutFT == f64) .f64 else @compileError("OutFT");
4452
4453 var arena = alloc_arena.Arena.init(std.testing.allocator);
4454 defer arena.deinit();
4455 const allocator = arena.allocator();
4456
4457 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4458 defer ir_ctx.deinit(allocator);
4459 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4460
4461 const loc = ir.Location.getUnknown();
4462 const in_type = try ArithDialect.getScalarType(&ir_ctx, input_kind);
4463 const out_type = try ArithDialect.getScalarType(&ir_ctx, out_kind);
4464
4465 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4466 const module_block = module.getBodyBlock();
4467
4468 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{out_type});
4469 try module_block.addOperation(func.op);
4470
4471 const entry = func.getEntryBlock();
4472 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, in_type, c_in);
4473 try entry.addOperation(c.op);
4474 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), out_type);
4475 try entry.addOperation(cast.op);
4476 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
4477 try entry.addOperation(ret.op);
4478
4479 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4480 defer backend.deinit();
4481 const compiled = try backend.compile(module.op);
4482
4483 const FnPtr = *const fn () callconv(.c) OutFT;
4484 const fn_ptr = try backend.runtime.getFunction(compiled, name, FnPtr);
4485 return fn_ptr();
4486 }
4487
4488 fn runCastFloatToInt(
4489 comptime InFT: type,
4490 output_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
4491 name: []const u8,
4492 c_in: InFT,
4493 ) !i64 {
4494 const dialects = @import("../../dialects/root.zig");
4495 const ArithDialect = dialects.ArithDialect;
4496 const BuiltinDialect = dialects.BuiltinDialect;
4497 const FuncDialect = dialects.FuncDialect;
4498
4499 const in_kind: ArithDialect.ScalarTypeKind = comptime if (InFT == f32) .f32 else if (InFT == f64) .f64 else @compileError("InFT");
4500
4501 var arena = alloc_arena.Arena.init(std.testing.allocator);
4502 defer arena.deinit();
4503 const allocator = arena.allocator();
4504
4505 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4506 defer ir_ctx.deinit(allocator);
4507 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4508
4509 const loc = ir.Location.getUnknown();
4510 const in_type = try ArithDialect.getScalarType(&ir_ctx, in_kind);
4511 const out_type = try ArithDialect.getScalarType(&ir_ctx, output_kind);
4512
4513 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4514 const module_block = module.getBodyBlock();
4515
4516 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{out_type});
4517 try module_block.addOperation(func.op);
4518
4519 const entry = func.getEntryBlock();
4520 var c = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, in_type, @floatCast(c_in));
4521 try entry.addOperation(c.op);
4522 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), out_type);
4523 try entry.addOperation(cast.op);
4524 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
4525 try entry.addOperation(ret.op);
4526
4527 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4528 defer backend.deinit();
4529 const compiled = try backend.compile(module.op);
4530 const exec_result = try onlyResult(&backend, compiled, name, &.{});
4531 return integerOf(exec_result);
4532 }
4533
4534 test "x86_64 arith.cast: int -> int (sign-extend, truncate, same-kind)" {
4535 if (!supports_x86_64_backend) return;
4536 const testing = std.testing;
4537
4538 try testing.expectEqual(@as(i64, -5), try runCastInt(.i32, .i64, "cast_i32_i64_neg", -5));
4539 try testing.expectEqual(@as(i64, 1234567), try runCastInt(.i32, .i64, "cast_i32_i64_pos", 1234567));
4540 try testing.expectEqual(@as(i64, 5), try runCastInt(.i64, .i32, "cast_i64_i32_pos", 5));
4541 try testing.expectEqual(@as(i64, -5), try runCastInt(.i64, .i32, "cast_i64_i32_neg", -5));
4542 try testing.expectEqual(@as(i64, 42), try runCastInt(.i32, .i32, "cast_i32_i32_same", 42));
4543 try testing.expectEqual(@as(i64, 100), try runCastInt(.index, .i64, "cast_index_i64", 100));
4544 try testing.expectEqual(@as(i64, -1), try runCastInt(.i32, .i8, "cast_i32_i8_neg1", -1));
4545 }
4546
4547 test "x86_64 arith.cast: u32 int casts preserve unsigned high bits" {
4548 if (!supports_x86_64_backend) return;
4549 const testing = std.testing;
4550
4551 try testing.expectEqual(@as(i64, 0x8000_0000), try runCastInt(.u32, .i64, "cast_u32_i64_high", 0x8000_0000));
4552 try testing.expectEqual(@as(i64, 0xffff_ffff), try runCastInt(.u32, .index, "cast_u32_index_max", 0xffff_ffff));
4553 try testing.expectEqual(@as(i64, 0xffff_ffff), try runCastInt(.i64, .u32, "cast_i64_u32_neg1", -1));
4554 try testing.expectEqual(@as(i64, -1), try runCastInt(.u32, .i32, "cast_u32_i32_max", 0xffff_ffff));
4555 }
4556
4557 test "x86_64 arith.cast: float -> float (cvtss2sd / cvtsd2ss)" {
4558 if (!supports_x86_64_backend) return;
4559 const testing = std.testing;
4560
4561 try testing.expectEqual(@as(f64, 3.5), try runCastFloatToFloat(f32, f64, "cast_f32_f64", 3.5));
4562 try testing.expectEqual(@as(f64, -2.0), try runCastFloatToFloat(f32, f64, "cast_f32_f64_neg", -2.0));
4563 try testing.expectEqual(@as(f32, 3.5), try runCastFloatToFloat(f64, f32, "cast_f64_f32", 3.5));
4564 try testing.expectEqual(@as(f32, 1.25), try runCastFloatToFloat(f32, f32, "cast_f32_f32_same", 1.25));
4565 }
4566
4567 test "x86_64 arith.cast: int -> float (cvtsi2ss / cvtsi2sd)" {
4568 if (!supports_x86_64_backend) return;
4569 const testing = std.testing;
4570
4571 try testing.expectEqual(@as(f32, 5.0), try runCastIntToFloat(.i32, f32, "cast_i32_f32", 5));
4572 try testing.expectEqual(@as(f32, -7.0), try runCastIntToFloat(.i32, f32, "cast_i32_f32_neg", -7));
4573 try testing.expectEqual(@as(f64, 1234567890.0), try runCastIntToFloat(.i32, f64, "cast_i32_f64", 1234567890));
4574 try testing.expectEqual(@as(f64, 100.0), try runCastIntToFloat(.i64, f64, "cast_i64_f64", 100));
4575 try testing.expectEqual(@as(f64, 42.0), try runCastIntToFloat(.index, f64, "cast_index_f64", 42));
4576 }
4577
4578 test "x86_64 arith.cast: u32 -> float uses unsigned high-bit values" {
4579 if (!supports_x86_64_backend) return;
4580 const testing = std.testing;
4581
4582 try testing.expectEqual(@as(f64, 2147483648.0), try runCastIntToFloat(.u32, f64, "cast_u32_f64_2pow31", 0x8000_0000));
4583 try testing.expectEqual(@as(f64, 4294967295.0), try runCastIntToFloat(.u32, f64, "cast_u32_f64_max", 0xffff_ffff));
4584 try testing.expectEqual(@as(f32, @floatFromInt(@as(u32, 0xffff_ffff))), try runCastIntToFloat(.u32, f32, "cast_u32_f32_max", 0xffff_ffff));
4585 }
4586
4587 test "x86_64 arith.cast: arith.index -> f64 high-bit values (Hacker's Delight)" {
4588 if (!supports_x86_64_backend) return;
4589 const testing = std.testing;
4590
4591 try testing.expectEqual(
4592 @as(f64, 0.0),
4593 try runCastIntToFloat(.index, f64, "cast_index_f64_zero", 0),
4594 );
4595 try testing.expectEqual(
4596 @as(f64, 1.0),
4597 try runCastIntToFloat(.index, f64, "cast_index_f64_one", 1),
4598 );
4599 try testing.expectEqual(
4600 @as(f64, @floatFromInt(@as(i64, std.math.maxInt(i64)))),
4601 try runCastIntToFloat(.index, f64, "cast_index_f64_i64max", std.math.maxInt(i64)),
4602 );
4603
4604 const two_to_63: f64 = @bitCast(@as(u64, 0x43E0000000000000));
4605 try testing.expectEqual(
4606 two_to_63,
4607 try runCastIntToFloat(.index, f64, "cast_index_f64_2pow63", @as(i64, @bitCast(@as(u64, 0x8000000000000000)))),
4608 );
4609
4610 const expected_sticky: f64 = @bitCast(@as(u64, 0x43E0000000000001));
4611 try testing.expectEqual(
4612 expected_sticky,
4613 try runCastIntToFloat(.index, f64, "cast_index_f64_sticky", @as(i64, @bitCast(@as(u64, 0x8000000000000401)))),
4614 );
4615
4616 const two_to_64: f64 = @bitCast(@as(u64, 0x43F0000000000000));
4617 try testing.expectEqual(
4618 two_to_64,
4619 try runCastIntToFloat(.index, f64, "cast_index_f64_u64max", @as(i64, @bitCast(@as(u64, 0xFFFFFFFFFFFFFFFF)))),
4620 );
4621 }
4622
4623 test "x86_64 arith.cast: u64 -> f64 high-bit values" {
4624 if (!supports_x86_64_backend) return;
4625 const testing = std.testing;
4626
4627 const two_to_64: f64 = @bitCast(@as(u64, 0x43F0000000000000));
4628 try testing.expectEqual(
4629 two_to_64,
4630 try runCastIntToFloat(.u64, f64, "cast_u64_f64_u64max", -1),
4631 );
4632 }
4633
4634 test "x86_64 arith.cast: arith.index -> f32 high-bit values (Hacker's Delight)" {
4635 if (!supports_x86_64_backend) return;
4636 const testing = std.testing;
4637
4638 const two_to_63_f32: f32 = @bitCast(@as(u32, 0x5F000000));
4639 try testing.expectEqual(
4640 two_to_63_f32,
4641 try runCastIntToFloat(.index, f32, "cast_index_f32_2pow63", @as(i64, @bitCast(@as(u64, 0x8000000000000000)))),
4642 );
4643
4644 const two_to_64_f32: f32 = @bitCast(@as(u32, 0x5F800000));
4645 try testing.expectEqual(
4646 two_to_64_f32,
4647 try runCastIntToFloat(.index, f32, "cast_index_f32_u64max", @as(i64, @bitCast(@as(u64, 0xFFFFFFFFFFFFFFFF)))),
4648 );
4649 }
4650
4651 test "x86_64 arith.cast: float -> int (cvttss2si / cvttsd2si truncating)" {
4652 if (!supports_x86_64_backend) return;
4653 const testing = std.testing;
4654
4655 try testing.expectEqual(@as(i64, 3), try runCastFloatToInt(f32, .i32, "cast_f32_i32_pos", 3.7));
4656 try testing.expectEqual(@as(i64, -3), try runCastFloatToInt(f32, .i32, "cast_f32_i32_neg", -3.7));
4657 try testing.expectEqual(@as(i64, 7), try runCastFloatToInt(f64, .i32, "cast_f64_i32", 7.99));
4658 try testing.expectEqual(@as(i64, 100), try runCastFloatToInt(f64, .i64, "cast_f64_i64", 100.5));
4659 try testing.expectEqual(@as(i64, -7), try runCastFloatToInt(f64, .i64, "cast_f64_i64_neg", -7.99));
4660 }
4661
4662 test "x86_64 arith.cast: float -> u32 uses unsigned high-bit values" {
4663 if (!supports_x86_64_backend) return;
4664 const testing = std.testing;
4665
4666 try testing.expectEqual(@as(i64, 0x8000_0000), try runCastFloatToInt(f64, .u32, "cast_f64_u32_2pow31", 2147483648.0));
4667 try testing.expectEqual(@as(i64, 0xffff_ffff), try runCastFloatToInt(f64, .u32, "cast_f64_u32_max", 4294967295.0));
4668 try testing.expectEqual(@as(i64, 0x8000_0000), try runCastFloatToInt(f32, .u32, "cast_f32_u32_2pow31", @as(f32, 2147483648.0)));
4669 }
4670
4671 test "x86_64 arith.cast: f64 -> arith.index for high-bit values (2^63-split)" {
4672 if (!supports_x86_64_backend) return;
4673 const testing = std.testing;
4674
4675 try testing.expectEqual(
4676 @as(i64, 0),
4677 try runCastFloatToInt(f64, .index, "cast_f64_idx_zero", 0.0),
4678 );
4679 try testing.expectEqual(
4680 @as(i64, 1),
4681 try runCastFloatToInt(f64, .index, "cast_f64_idx_one", 1.0),
4682 );
4683
4684 const low_boundary_f64: f64 = @bitCast(@as(u64, 0x43DFFFFFFFFFFFFF));
4685 const low_boundary_idx: i64 = @bitCast(@as(u64, 0x7FFFFFFFFFFFFC00));
4686 try testing.expectEqual(
4687 low_boundary_idx,
4688 try runCastFloatToInt(f64, .index, "cast_f64_idx_low_boundary", low_boundary_f64),
4689 );
4690
4691 const two_to_63_f64: f64 = @bitCast(@as(u64, 0x43E0000000000000));
4692 const two_to_63_idx: i64 = @bitCast(@as(u64, 0x8000000000000000));
4693 try testing.expectEqual(
4694 two_to_63_idx,
4695 try runCastFloatToInt(f64, .index, "cast_f64_idx_2pow63", two_to_63_f64),
4696 );
4697
4698 const residual_f64: f64 = @bitCast(@as(u64, 0x43E0000000000001));
4699 const residual_idx: i64 = @bitCast(@as(u64, 0x8000000000000800));
4700 try testing.expectEqual(
4701 residual_idx,
4702 try runCastFloatToInt(f64, .index, "cast_f64_idx_residual", residual_f64),
4703 );
4704
4705 const near_max_f64: f64 = @bitCast(@as(u64, 0x43EFFFFFFFFFFFFF));
4706 const near_max_idx: i64 = @bitCast(@as(u64, 0xFFFFFFFFFFFFF800));
4707 try testing.expectEqual(
4708 near_max_idx,
4709 try runCastFloatToInt(f64, .index, "cast_f64_idx_near_max", near_max_f64),
4710 );
4711 }
4712
4713 test "x86_64 arith.cast: f64 -> u64 high-bit values" {
4714 if (!supports_x86_64_backend) return;
4715 const testing = std.testing;
4716
4717 const two_to_63_f64: f64 = @bitCast(@as(u64, 0x43E0000000000000));
4718 const two_to_63_u64: i64 = @bitCast(@as(u64, 0x8000000000000000));
4719 try testing.expectEqual(
4720 two_to_63_u64,
4721 try runCastFloatToInt(f64, .u64, "cast_f64_u64_2pow63", two_to_63_f64),
4722 );
4723 }
4724
4725 test "x86_64 arith.cast: f32 -> arith.index for high-bit values (2^63-split)" {
4726 if (!supports_x86_64_backend) return;
4727 const testing = std.testing;
4728
4729 const two_to_63_f32: f32 = @bitCast(@as(u32, 0x5F000000));
4730 const two_to_63_idx: i64 = @bitCast(@as(u64, 0x8000000000000000));
4731 try testing.expectEqual(
4732 two_to_63_idx,
4733 try runCastFloatToInt(f32, .index, "cast_f32_idx_2pow63", two_to_63_f32),
4734 );
4735
4736 try testing.expectEqual(
4737 @as(i64, 1),
4738 try runCastFloatToInt(f32, .index, "cast_f32_idx_one", 1.0),
4739 );
4740 }
4741
4742 test "x86_64 arith.cast: i32 -> index sign-extends (movsxd, post-review fix)" {
4743 if (!supports_x86_64_backend) return;
4744 const testing = std.testing;
4745
4746 try testing.expectEqual(@as(i64, -1), try runCastInt(.i32, .index, "cast_i32_index_neg1", -1));
4747 try testing.expectEqual(@as(i64, 100), try runCastInt(.i32, .index, "cast_i32_index_pos", 100));
4748 }
4749
4750 test "x86_64 arith.cast: bool -> bool identity passes (post-review fix)" {
4751 if (!supports_x86_64_backend) return;
4752 const testing = std.testing;
4753 const dialects = @import("../../dialects/root.zig");
4754 const ArithDialect = dialects.ArithDialect;
4755 const BuiltinDialect = dialects.BuiltinDialect;
4756 const FuncDialect = dialects.FuncDialect;
4757
4758 var arena = alloc_arena.Arena.init(std.testing.allocator);
4759 defer arena.deinit();
4760 const allocator = arena.allocator();
4761
4762 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4763 defer ir_ctx.deinit(allocator);
4764 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4765
4766 const loc = ir.Location.getUnknown();
4767 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
4768
4769 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4770 const module_block = module.getBodyBlock();
4771 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "cast_bool_bool_id", &.{}, &.{bool_type});
4772 try module_block.addOperation(func.op);
4773
4774 const entry = func.getEntryBlock();
4775 var c = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
4776 try entry.addOperation(c.op);
4777 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), bool_type);
4778 try entry.addOperation(cast.op);
4779 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
4780 try entry.addOperation(ret.op);
4781
4782 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4783 defer backend.deinit();
4784 const compiled = try backend.compile(module.op);
4785 const result = try onlyResult(&backend, compiled, "cast_bool_bool_id", &.{});
4786 try testing.expectEqual(@as(i64, 1), try integerOf(result));
4787 }
4788
4789 fn runBoolBitwise(
4790 comptime kind: ScalarOpKind,
4791 name: []const u8,
4792 lhs_value: bool,
4793 rhs_value: bool,
4794 ) !i64 {
4795 const dialects = @import("../../dialects/root.zig");
4796 const ArithDialect = dialects.ArithDialect;
4797 const BuiltinDialect = dialects.BuiltinDialect;
4798 const FuncDialect = dialects.FuncDialect;
4799
4800 var arena = alloc_arena.Arena.init(std.testing.allocator);
4801 defer arena.deinit();
4802 const allocator = arena.allocator();
4803
4804 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4805 defer ir_ctx.deinit(allocator);
4806 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4807
4808 const loc = ir.Location.getUnknown();
4809 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4810 const module_block = module.getBodyBlock();
4811 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
4812 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{bool_type});
4813 try module_block.addOperation(func.op);
4814
4815 const entry = func.getEntryBlock();
4816 var lhs = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, lhs_value);
4817 try entry.addOperation(lhs.op);
4818 var rhs = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, rhs_value);
4819 try entry.addOperation(rhs.op);
4820
4821 const result_value = switch (kind) {
4822 .band => blk: {
4823 const op = try ArithDialect.AndOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
4824 try entry.addOperation(op.op);
4825 break :blk op.getResult();
4826 },
4827 .bor => blk: {
4828 const op = try ArithDialect.OrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
4829 try entry.addOperation(op.op);
4830 break :blk op.getResult();
4831 },
4832 .bxor => blk: {
4833 const op = try ArithDialect.XorOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
4834 try entry.addOperation(op.op);
4835 break :blk op.getResult();
4836 },
4837 else => return error.UnsupportedType,
4838 };
4839 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_value});
4840 try entry.addOperation(ret.op);
4841
4842 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4843 defer backend.deinit();
4844 const compiled = try backend.compile(module.op);
4845 const result = try onlyResult(&backend, compiled, name, &.{});
4846 return integerOf(result);
4847 }
4848
4849 test "x86_64 arith.and/or/xor over bool slots" {
4850 if (!supports_x86_64_backend) return;
4851 const testing = std.testing;
4852
4853 try testing.expectEqual(@as(i64, 0), try runBoolBitwise(.band, "and_bool_false", true, false));
4854 try testing.expectEqual(@as(i64, 1), try runBoolBitwise(.band, "and_bool_true", true, true));
4855 try testing.expectEqual(@as(i64, 1), try runBoolBitwise(.bor, "or_bool_true", true, false));
4856 try testing.expectEqual(@as(i64, 0), try runBoolBitwise(.bor, "or_bool_false", false, false));
4857 try testing.expectEqual(@as(i64, 1), try runBoolBitwise(.bxor, "xor_bool_true", true, false));
4858 try testing.expectEqual(@as(i64, 0), try runBoolBitwise(.bxor, "xor_bool_false", true, true));
4859 }
4860
4861 fn runBoolNot(name: []const u8, bool_val: bool) !i64 {
4862 const dialects = @import("../../dialects/root.zig");
4863 const ArithDialect = dialects.ArithDialect;
4864 const BuiltinDialect = dialects.BuiltinDialect;
4865 const FuncDialect = dialects.FuncDialect;
4866
4867 var arena = alloc_arena.Arena.init(std.testing.allocator);
4868 defer arena.deinit();
4869 const allocator = arena.allocator();
4870
4871 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4872 defer ir_ctx.deinit(allocator);
4873 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4874
4875 const loc = ir.Location.getUnknown();
4876 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4877 const module_block = module.getBodyBlock();
4878 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
4879 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{bool_type});
4880 try module_block.addOperation(func.op);
4881
4882 const entry = func.getEntryBlock();
4883 var value = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, bool_val);
4884 try entry.addOperation(value.op);
4885 var not_op = try ArithDialect.NotOp.create(&ir_ctx, loc, value.getResult());
4886 try entry.addOperation(not_op.op);
4887 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{not_op.getResult()});
4888 try entry.addOperation(ret.op);
4889
4890 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4891 defer backend.deinit();
4892 const compiled = try backend.compile(module.op);
4893 const result = try onlyResult(&backend, compiled, name, &.{});
4894 return integerOf(result);
4895 }
4896
4897 test "x86_64 arith.not over bool slots" {
4898 if (!supports_x86_64_backend) return;
4899 const testing = std.testing;
4900
4901 try testing.expectEqual(@as(i64, 0), try runBoolNot("not_bool_true", true));
4902 try testing.expectEqual(@as(i64, 1), try runBoolNot("not_bool_false", false));
4903 }
4904
4905 fn runCastBoolToInt(
4906 output_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
4907 name: []const u8,
4908 bool_val: bool,
4909 ) !i64 {
4910 const dialects = @import("../../dialects/root.zig");
4911 const ArithDialect = dialects.ArithDialect;
4912 const BuiltinDialect = dialects.BuiltinDialect;
4913 const FuncDialect = dialects.FuncDialect;
4914
4915 var arena = alloc_arena.Arena.init(std.testing.allocator);
4916 defer arena.deinit();
4917 const allocator = arena.allocator();
4918
4919 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4920 defer ir_ctx.deinit(allocator);
4921 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4922
4923 const loc = ir.Location.getUnknown();
4924 const out_type = try ArithDialect.getScalarType(&ir_ctx, output_kind);
4925
4926 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4927 const module_block = module.getBodyBlock();
4928
4929 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{out_type});
4930 try module_block.addOperation(func.op);
4931
4932 const entry = func.getEntryBlock();
4933 var c = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, bool_val);
4934 try entry.addOperation(c.op);
4935 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), out_type);
4936 try entry.addOperation(cast.op);
4937 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
4938 try entry.addOperation(ret.op);
4939
4940 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4941 defer backend.deinit();
4942 const compiled = try backend.compile(module.op);
4943 const exec = try onlyResult(&backend, compiled, name, &.{});
4944 return integerOf(exec);
4945 }
4946
4947 fn runCastBoolToFloat(
4948 comptime OutFT: type,
4949 name: []const u8,
4950 bool_val: bool,
4951 ) !OutFT {
4952 const dialects = @import("../../dialects/root.zig");
4953 const ArithDialect = dialects.ArithDialect;
4954 const BuiltinDialect = dialects.BuiltinDialect;
4955 const FuncDialect = dialects.FuncDialect;
4956
4957 const out_kind: ArithDialect.ScalarTypeKind = comptime if (OutFT == f32) .f32 else if (OutFT == f64) .f64 else @compileError("OutFT");
4958
4959 var arena = alloc_arena.Arena.init(std.testing.allocator);
4960 defer arena.deinit();
4961 const allocator = arena.allocator();
4962
4963 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
4964 defer ir_ctx.deinit(allocator);
4965 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
4966
4967 const loc = ir.Location.getUnknown();
4968 const out_type = try ArithDialect.getScalarType(&ir_ctx, out_kind);
4969
4970 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
4971 const module_block = module.getBodyBlock();
4972
4973 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{out_type});
4974 try module_block.addOperation(func.op);
4975
4976 const entry = func.getEntryBlock();
4977 var c = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, bool_val);
4978 try entry.addOperation(c.op);
4979 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), out_type);
4980 try entry.addOperation(cast.op);
4981 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
4982 try entry.addOperation(ret.op);
4983
4984 var backend = try Backend.init(allocator, &ir_ctx, .testing);
4985 defer backend.deinit();
4986 const compiled = try backend.compile(module.op);
4987
4988 const FnPtr = *const fn () callconv(.c) OutFT;
4989 const fn_ptr = try backend.runtime.getFunction(compiled, name, FnPtr);
4990 return fn_ptr();
4991 }
4992
4993 fn runCastIntToBool(
4994 input_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
4995 name: []const u8,
4996 c_in: i64,
4997 ) !i64 {
4998 const dialects = @import("../../dialects/root.zig");
4999 const ArithDialect = dialects.ArithDialect;
5000 const BuiltinDialect = dialects.BuiltinDialect;
5001 const FuncDialect = dialects.FuncDialect;
5002
5003 var arena = alloc_arena.Arena.init(std.testing.allocator);
5004 defer arena.deinit();
5005 const allocator = arena.allocator();
5006
5007 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5008 defer ir_ctx.deinit(allocator);
5009 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5010
5011 const loc = ir.Location.getUnknown();
5012 const in_type = try ArithDialect.getScalarType(&ir_ctx, input_kind);
5013 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
5014
5015 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5016 const module_block = module.getBodyBlock();
5017
5018 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{bool_type});
5019 try module_block.addOperation(func.op);
5020
5021 const entry = func.getEntryBlock();
5022 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, in_type, c_in);
5023 try entry.addOperation(c.op);
5024 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), bool_type);
5025 try entry.addOperation(cast.op);
5026 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
5027 try entry.addOperation(ret.op);
5028
5029 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5030 defer backend.deinit();
5031 const compiled = try backend.compile(module.op);
5032 const exec = try onlyResult(&backend, compiled, name, &.{});
5033 return integerOf(exec);
5034 }
5035
5036 fn runCastFloatToBool(
5037 comptime InFT: type,
5038 name: []const u8,
5039 c_in: InFT,
5040 ) !i64 {
5041 const dialects = @import("../../dialects/root.zig");
5042 const ArithDialect = dialects.ArithDialect;
5043 const BuiltinDialect = dialects.BuiltinDialect;
5044 const FuncDialect = dialects.FuncDialect;
5045
5046 const in_kind: ArithDialect.ScalarTypeKind = comptime if (InFT == f32) .f32 else if (InFT == f64) .f64 else @compileError("InFT");
5047
5048 var arena = alloc_arena.Arena.init(std.testing.allocator);
5049 defer arena.deinit();
5050 const allocator = arena.allocator();
5051
5052 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5053 defer ir_ctx.deinit(allocator);
5054 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5055
5056 const loc = ir.Location.getUnknown();
5057 const in_type = try ArithDialect.getScalarType(&ir_ctx, in_kind);
5058 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
5059
5060 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5061 const module_block = module.getBodyBlock();
5062
5063 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{bool_type});
5064 try module_block.addOperation(func.op);
5065
5066 const entry = func.getEntryBlock();
5067 var c = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, in_type, @floatCast(c_in));
5068 try entry.addOperation(c.op);
5069 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), bool_type);
5070 try entry.addOperation(cast.op);
5071 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{cast.getResult()});
5072 try entry.addOperation(ret.op);
5073
5074 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5075 defer backend.deinit();
5076 const compiled = try backend.compile(module.op);
5077 const exec = try onlyResult(&backend, compiled, name, &.{});
5078 return integerOf(exec);
5079 }
5080
5081 test "x86_64 arith.cast: bool -> integer (slot copy with width adjust)" {
5082 if (!supports_x86_64_backend) return;
5083 const testing = std.testing;
5084
5085 try testing.expectEqual(@as(i64, 1), try runCastBoolToInt(.i32, "cast_bool_i32_t", true));
5086 try testing.expectEqual(@as(i64, 0), try runCastBoolToInt(.i32, "cast_bool_i32_f", false));
5087 try testing.expectEqual(@as(i64, 1), try runCastBoolToInt(.i8, "cast_bool_i8_t", true));
5088 try testing.expectEqual(@as(i64, 1), try runCastBoolToInt(.i16, "cast_bool_i16_t", true));
5089 try testing.expectEqual(@as(i64, 1), try runCastBoolToInt(.i64, "cast_bool_i64_t", true));
5090 try testing.expectEqual(@as(i64, 0), try runCastBoolToInt(.i64, "cast_bool_i64_f", false));
5091 }
5092
5093 test "x86_64 arith.cast: bool -> float (cvtsi2ss / cvtsi2sd on 0 or 1)" {
5094 if (!supports_x86_64_backend) return;
5095 const testing = std.testing;
5096
5097 try testing.expectEqual(@as(f32, 1.0), try runCastBoolToFloat(f32, "cast_bool_f32_t", true));
5098 try testing.expectEqual(@as(f32, 0.0), try runCastBoolToFloat(f32, "cast_bool_f32_f", false));
5099 try testing.expectEqual(@as(f64, 1.0), try runCastBoolToFloat(f64, "cast_bool_f64_t", true));
5100 try testing.expectEqual(@as(f64, 0.0), try runCastBoolToFloat(f64, "cast_bool_f64_f", false));
5101 }
5102
5103 test "x86_64 arith.cast: integer -> bool (cmp + setne)" {
5104 if (!supports_x86_64_backend) return;
5105 const testing = std.testing;
5106
5107 try testing.expectEqual(@as(i64, 1), try runCastIntToBool(.i32, "cast_i32_bool_5", 5));
5108 try testing.expectEqual(@as(i64, 0), try runCastIntToBool(.i32, "cast_i32_bool_0", 0));
5109 try testing.expectEqual(@as(i64, 1), try runCastIntToBool(.i32, "cast_i32_bool_neg", -1));
5110 try testing.expectEqual(@as(i64, 1), try runCastIntToBool(.i64, "cast_i64_bool_max", std.math.maxInt(i64)));
5111 try testing.expectEqual(@as(i64, 0), try runCastIntToBool(.i64, "cast_i64_bool_zero", 0));
5112 }
5113
5114 test "x86_64 arith.cast: float -> bool (ordered-NE; NaN and ±0 are false)" {
5115 if (!supports_x86_64_backend) return;
5116 const testing = std.testing;
5117
5118 try testing.expectEqual(@as(i64, 1), try runCastFloatToBool(f32, "cast_f32_bool_pos", 3.5));
5119 try testing.expectEqual(@as(i64, 1), try runCastFloatToBool(f32, "cast_f32_bool_neg", -3.5));
5120 try testing.expectEqual(@as(i64, 0), try runCastFloatToBool(f32, "cast_f32_bool_pzero", 0.0));
5121 try testing.expectEqual(@as(i64, 0), try runCastFloatToBool(f32, "cast_f32_bool_nzero", -0.0));
5122 try testing.expectEqual(@as(i64, 0), try runCastFloatToBool(f32, "cast_f32_bool_nan", std.math.nan(f32)));
5123
5124 try testing.expectEqual(@as(i64, 1), try runCastFloatToBool(f64, "cast_f64_bool_pos", 1e10));
5125 try testing.expectEqual(@as(i64, 0), try runCastFloatToBool(f64, "cast_f64_bool_zero", 0.0));
5126 try testing.expectEqual(@as(i64, 0), try runCastFloatToBool(f64, "cast_f64_bool_nan", std.math.nan(f64)));
5127 }
5128
5129 test "x86_64 arith.bitcast: i32 ↔ f32 round-trips bit pattern" {
5130 if (!supports_x86_64_backend) return;
5131 const testing = std.testing;
5132 const dialects = @import("../../dialects/root.zig");
5133 const ArithDialect = dialects.ArithDialect;
5134 const BuiltinDialect = dialects.BuiltinDialect;
5135 const FuncDialect = dialects.FuncDialect;
5136
5137 var arena = alloc_arena.Arena.init(std.testing.allocator);
5138 defer arena.deinit();
5139 const allocator = arena.allocator();
5140
5141 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5142 defer ir_ctx.deinit(allocator);
5143 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5144
5145 const loc = ir.Location.getUnknown();
5146 const i32_type = try ArithDialect.getI32Type(&ir_ctx);
5147 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
5148
5149 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5150 const module_block = module.getBodyBlock();
5151 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "bitcast_round_trip", &.{}, &.{i32_type});
5152 try module_block.addOperation(func.op);
5153
5154 const entry = func.getEntryBlock();
5155 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 0x40490FDB);
5156 try entry.addOperation(c.op);
5157 var as_f32 = try ArithDialect.BitcastOp.create(&ir_ctx, loc, c.getResult(), f32_type);
5158 try entry.addOperation(as_f32.op);
5159 var back_to_i32 = try ArithDialect.BitcastOp.create(&ir_ctx, loc, as_f32.getResult(), i32_type);
5160 try entry.addOperation(back_to_i32.op);
5161 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{back_to_i32.getResult()});
5162 try entry.addOperation(ret.op);
5163
5164 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5165 defer backend.deinit();
5166 const compiled = try backend.compile(module.op);
5167 const exec = try onlyResult(&backend, compiled, "bitcast_round_trip", &.{});
5168 try testing.expectEqual(@as(i64, 0x40490FDB), try integerOf(exec));
5169 }
5170
5171 test "x86_64 arith.bitcast: u32 ↔ f32 round-trips high bit pattern" {
5172 if (!supports_x86_64_backend) return;
5173 const testing = std.testing;
5174 const dialects = @import("../../dialects/root.zig");
5175 const ArithDialect = dialects.ArithDialect;
5176 const BuiltinDialect = dialects.BuiltinDialect;
5177 const FuncDialect = dialects.FuncDialect;
5178
5179 var arena = alloc_arena.Arena.init(std.testing.allocator);
5180 defer arena.deinit();
5181 const allocator = arena.allocator();
5182
5183 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5184 defer ir_ctx.deinit(allocator);
5185 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5186
5187 const loc = ir.Location.getUnknown();
5188 const u32_type = try ArithDialect.getScalarType(&ir_ctx, .u32);
5189 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
5190
5191 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5192 const module_block = module.getBodyBlock();
5193 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "bitcast_u32_round_trip", &.{}, &.{u32_type});
5194 try module_block.addOperation(func.op);
5195
5196 const entry = func.getEntryBlock();
5197 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, u32_type, 0xFFC00001);
5198 try entry.addOperation(c.op);
5199 var as_f32 = try ArithDialect.BitcastOp.create(&ir_ctx, loc, c.getResult(), f32_type);
5200 try entry.addOperation(as_f32.op);
5201 var back_to_u32 = try ArithDialect.BitcastOp.create(&ir_ctx, loc, as_f32.getResult(), u32_type);
5202 try entry.addOperation(back_to_u32.op);
5203 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{back_to_u32.getResult()});
5204 try entry.addOperation(ret.op);
5205
5206 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5207 defer backend.deinit();
5208 const compiled = try backend.compile(module.op);
5209 const exec = try onlyResult(&backend, compiled, "bitcast_u32_round_trip", &.{});
5210 try testing.expectEqual(@as(i64, 0xFFC00001), try integerOf(exec));
5211 }
5212
5213 test "x86_64 arith.bitcast: i64 ↔ f64 round-trips bit pattern" {
5214 if (!supports_x86_64_backend) return;
5215 const testing = std.testing;
5216 const dialects = @import("../../dialects/root.zig");
5217 const ArithDialect = dialects.ArithDialect;
5218 const BuiltinDialect = dialects.BuiltinDialect;
5219 const FuncDialect = dialects.FuncDialect;
5220
5221 var arena = alloc_arena.Arena.init(std.testing.allocator);
5222 defer arena.deinit();
5223 const allocator = arena.allocator();
5224
5225 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5226 defer ir_ctx.deinit(allocator);
5227 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5228
5229 const loc = ir.Location.getUnknown();
5230 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
5231 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
5232
5233 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5234 const module_block = module.getBodyBlock();
5235 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "bitcast_64_round_trip", &.{}, &.{i64_type});
5236 try module_block.addOperation(func.op);
5237
5238 const entry = func.getEntryBlock();
5239 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, @bitCast(@as(u64, 0x400921FB54442D18)));
5240 try entry.addOperation(c.op);
5241 var as_f64 = try ArithDialect.BitcastOp.create(&ir_ctx, loc, c.getResult(), f64_type);
5242 try entry.addOperation(as_f64.op);
5243 var back_to_i64 = try ArithDialect.BitcastOp.create(&ir_ctx, loc, as_f64.getResult(), i64_type);
5244 try entry.addOperation(back_to_i64.op);
5245 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{back_to_i64.getResult()});
5246 try entry.addOperation(ret.op);
5247
5248 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5249 defer backend.deinit();
5250 const compiled = try backend.compile(module.op);
5251 const exec = try onlyResult(&backend, compiled, "bitcast_64_round_trip", &.{});
5252 try testing.expectEqual(
5253 @as(i64, @bitCast(@as(u64, 0x400921FB54442D18))),
5254 try integerOf(exec),
5255 );
5256 }
5257
5258 test "x86_64 arith.bitcast: width mismatch + bool reject cleanly" {
5259 if (!supports_x86_64_backend) return;
5260 const testing = std.testing;
5261 const dialects = @import("../../dialects/root.zig");
5262 const ArithDialect = dialects.ArithDialect;
5263 const BuiltinDialect = dialects.BuiltinDialect;
5264 const FuncDialect = dialects.FuncDialect;
5265
5266 {
5267 var arena = alloc_arena.Arena.init(std.testing.allocator);
5268 defer arena.deinit();
5269 const allocator = arena.allocator();
5270
5271 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5272 defer ir_ctx.deinit(allocator);
5273 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5274
5275 const loc = ir.Location.getUnknown();
5276 const i32_type = try ArithDialect.getI32Type(&ir_ctx);
5277 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
5278
5279 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5280 const module_block = module.getBodyBlock();
5281 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "bitcast_width_mismatch", &.{}, &.{f64_type});
5282 try module_block.addOperation(func.op);
5283
5284 const entry = func.getEntryBlock();
5285 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 1);
5286 try entry.addOperation(c.op);
5287 var bc = try ArithDialect.BitcastOp.create(&ir_ctx, loc, c.getResult(), f64_type);
5288 try entry.addOperation(bc.op);
5289 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{bc.getResult()});
5290 try entry.addOperation(ret.op);
5291
5292 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5293 defer backend.deinit();
5294 try testing.expectError(error.JitCompileFailed, backend.compile(module.op));
5295 }
5296
5297 {
5298 var arena = alloc_arena.Arena.init(std.testing.allocator);
5299 defer arena.deinit();
5300 const allocator = arena.allocator();
5301
5302 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5303 defer ir_ctx.deinit(allocator);
5304 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5305
5306 const loc = ir.Location.getUnknown();
5307 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
5308
5309 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5310 const module_block = module.getBodyBlock();
5311 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "bitcast_bool", &.{}, &.{i64_type});
5312 try module_block.addOperation(func.op);
5313
5314 const entry = func.getEntryBlock();
5315 var c = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
5316 try entry.addOperation(c.op);
5317 var bc = try ArithDialect.BitcastOp.create(&ir_ctx, loc, c.getResult(), i64_type);
5318 try entry.addOperation(bc.op);
5319 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{bc.getResult()});
5320 try entry.addOperation(ret.op);
5321
5322 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5323 defer backend.deinit();
5324 try testing.expectError(error.JitCompileFailed, backend.compile(module.op));
5325 }
5326 }
5327
5328 test "x86_64 arith.bitcast: index reject (post-review whitelist fix)" {
5329 if (!supports_x86_64_backend) return;
5330 const testing = std.testing;
5331 const dialects = @import("../../dialects/root.zig");
5332 const ArithDialect = dialects.ArithDialect;
5333 const BuiltinDialect = dialects.BuiltinDialect;
5334 const FuncDialect = dialects.FuncDialect;
5335
5336 var arena = alloc_arena.Arena.init(std.testing.allocator);
5337 defer arena.deinit();
5338 const allocator = arena.allocator();
5339
5340 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5341 defer ir_ctx.deinit(allocator);
5342 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5343
5344 const loc = ir.Location.getUnknown();
5345 const idx_type = try ArithDialect.getIndexType(&ir_ctx);
5346 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
5347
5348 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5349 const module_block = module.getBodyBlock();
5350 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "bitcast_index", &.{}, &.{i64_type});
5351 try module_block.addOperation(func.op);
5352
5353 const entry = func.getEntryBlock();
5354 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, idx_type, 7);
5355 try entry.addOperation(c.op);
5356 var bc = try ArithDialect.BitcastOp.create(&ir_ctx, loc, c.getResult(), i64_type);
5357 try entry.addOperation(bc.op);
5358 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{bc.getResult()});
5359 try entry.addOperation(ret.op);
5360
5361 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5362 defer backend.deinit();
5363 try testing.expectError(error.JitCompileFailed, backend.compile(module.op));
5364 }
5365
5366 fn runSelectInt(
5367 output_kind: @import("../../dialects/root.zig").ArithDialect.ScalarTypeKind,
5368 name: []const u8,
5369 cond: bool,
5370 true_val: i64,
5371 false_val: i64,
5372 ) !i64 {
5373 const dialects = @import("../../dialects/root.zig");
5374 const ArithDialect = dialects.ArithDialect;
5375 const BuiltinDialect = dialects.BuiltinDialect;
5376 const FuncDialect = dialects.FuncDialect;
5377
5378 var arena = alloc_arena.Arena.init(std.testing.allocator);
5379 defer arena.deinit();
5380 const allocator = arena.allocator();
5381
5382 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5383 defer ir_ctx.deinit(allocator);
5384 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5385
5386 const loc = ir.Location.getUnknown();
5387 const out_type = try ArithDialect.getScalarType(&ir_ctx, output_kind);
5388
5389 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5390 const module_block = module.getBodyBlock();
5391
5392 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{out_type});
5393 try module_block.addOperation(func.op);
5394
5395 const entry = func.getEntryBlock();
5396 var c_cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, cond);
5397 try entry.addOperation(c_cond.op);
5398 const c_true_op = if (output_kind == .bool)
5399 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true_val != 0)
5400 else
5401 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, out_type, true_val);
5402 try entry.addOperation(c_true_op.op);
5403 const c_false_op = if (output_kind == .bool)
5404 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false_val != 0)
5405 else
5406 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, out_type, false_val);
5407 try entry.addOperation(c_false_op.op);
5408 var sel = try ArithDialect.SelectOp.create(&ir_ctx, loc, c_cond.getResult(), c_true_op.getResult(), c_false_op.getResult());
5409 try entry.addOperation(sel.op);
5410 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{sel.getResult()});
5411 try entry.addOperation(ret.op);
5412
5413 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5414 defer backend.deinit();
5415 const compiled = try backend.compile(module.op);
5416 const exec = try onlyResult(&backend, compiled, name, &.{});
5417 return integerOf(exec);
5418 }
5419
5420 fn runSelectFloat(
5421 comptime FT: type,
5422 name: []const u8,
5423 cond: bool,
5424 true_val: FT,
5425 false_val: FT,
5426 ) !FT {
5427 const dialects = @import("../../dialects/root.zig");
5428 const ArithDialect = dialects.ArithDialect;
5429 const BuiltinDialect = dialects.BuiltinDialect;
5430 const FuncDialect = dialects.FuncDialect;
5431
5432 const out_kind: ArithDialect.ScalarTypeKind = comptime if (FT == f32) .f32 else if (FT == f64) .f64 else @compileError("FT");
5433
5434 var arena = alloc_arena.Arena.init(std.testing.allocator);
5435 defer arena.deinit();
5436 const allocator = arena.allocator();
5437
5438 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5439 defer ir_ctx.deinit(allocator);
5440 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5441
5442 const loc = ir.Location.getUnknown();
5443 const out_type = try ArithDialect.getScalarType(&ir_ctx, out_kind);
5444
5445 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5446 const module_block = module.getBodyBlock();
5447
5448 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{out_type});
5449 try module_block.addOperation(func.op);
5450
5451 const entry = func.getEntryBlock();
5452 var c_cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, cond);
5453 try entry.addOperation(c_cond.op);
5454 var c_true = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, out_type, @floatCast(true_val));
5455 try entry.addOperation(c_true.op);
5456 var c_false = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, out_type, @floatCast(false_val));
5457 try entry.addOperation(c_false.op);
5458 var sel = try ArithDialect.SelectOp.create(&ir_ctx, loc, c_cond.getResult(), c_true.getResult(), c_false.getResult());
5459 try entry.addOperation(sel.op);
5460 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{sel.getResult()});
5461 try entry.addOperation(ret.op);
5462
5463 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5464 defer backend.deinit();
5465 const compiled = try backend.compile(module.op);
5466
5467 const FnPtr = *const fn () callconv(.c) FT;
5468 const fn_ptr = try backend.runtime.getFunction(compiled, name, FnPtr);
5469 return fn_ptr();
5470 }
5471
5472 test "x86_64 arith.select: integer cmov path (i32, i64, index, bool)" {
5473 if (!supports_x86_64_backend) return;
5474 const testing = std.testing;
5475
5476 try testing.expectEqual(@as(i64, 7), try runSelectInt(.i32, "select_i32_t", true, 7, -3));
5477 try testing.expectEqual(@as(i64, -3), try runSelectInt(.i32, "select_i32_f", false, 7, -3));
5478
5479 try testing.expectEqual(
5480 @as(i64, std.math.maxInt(i64)),
5481 try runSelectInt(.i64, "select_i64_max", true, std.math.maxInt(i64), -1),
5482 );
5483 try testing.expectEqual(
5484 @as(i64, std.math.minInt(i64)),
5485 try runSelectInt(.i64, "select_i64_min", false, 0, std.math.minInt(i64)),
5486 );
5487
5488 try testing.expectEqual(@as(i64, 100), try runSelectInt(.index, "select_idx_t", true, 100, 200));
5489
5490 try testing.expectEqual(@as(i64, 1), try runSelectInt(.bool, "select_bool_tt", true, 1, 0));
5491 try testing.expectEqual(@as(i64, 0), try runSelectInt(.bool, "select_bool_ff", false, 1, 0));
5492 }
5493
5494 test "x86_64 arith.select: float branched path (f32, f64)" {
5495 if (!supports_x86_64_backend) return;
5496 const testing = std.testing;
5497
5498 try testing.expectEqual(@as(f32, 3.5), try runSelectFloat(f32, "select_f32_t", true, 3.5, -2.0));
5499 try testing.expectEqual(@as(f32, -2.0), try runSelectFloat(f32, "select_f32_f", false, 3.5, -2.0));
5500 try testing.expectEqual(@as(f64, 3.5), try runSelectFloat(f64, "select_f64_t", true, 3.5, -2.0));
5501 try testing.expectEqual(@as(f64, -2.0), try runSelectFloat(f64, "select_f64_f", false, 3.5, -2.0));
5502 }
5503
5504 test "x86_64 arith.select inside scf.if body (dispatchOp parity)" {
5505 if (!supports_x86_64_backend) return;
5506 const testing = std.testing;
5507 const dialects = @import("../../dialects/root.zig");
5508 const ArithDialect = dialects.ArithDialect;
5509 const BuiltinDialect = dialects.BuiltinDialect;
5510 const FuncDialect = dialects.FuncDialect;
5511 const ScfDialect = dialects.ScfDialect;
5512
5513 var arena = alloc_arena.Arena.init(std.testing.allocator);
5514 defer arena.deinit();
5515 const allocator = arena.allocator();
5516
5517 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5518 defer ir_ctx.deinit(allocator);
5519 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5520
5521 const loc = ir.Location.getUnknown();
5522 const i32_type = try ArithDialect.getI32Type(&ir_ctx);
5523
5524 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5525 const module_block = module.getBodyBlock();
5526 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "select_in_if", &.{}, &.{i32_type});
5527 try module_block.addOperation(func.op);
5528
5529 const entry = func.getEntryBlock();
5530 var outer_cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
5531 try entry.addOperation(outer_cond.op);
5532 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, outer_cond.getResult(), &.{i32_type});
5533 try entry.addOperation(if_op.op);
5534
5535 const then_block = if_op.getThenBlock();
5536 var inner_cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false);
5537 try then_block.addOperation(inner_cond.op);
5538 var t_val = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 42);
5539 try then_block.addOperation(t_val.op);
5540 var f_val = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, -7);
5541 try then_block.addOperation(f_val.op);
5542 var sel = try ArithDialect.SelectOp.create(&ir_ctx, loc, inner_cond.getResult(), t_val.getResult(), f_val.getResult());
5543 try then_block.addOperation(sel.op);
5544 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{sel.getResult()});
5545 try then_block.addOperation(then_yield.op);
5546
5547 const else_block = if_op.getElseBlock().?;
5548 var else_const = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 0);
5549 try else_block.addOperation(else_const.op);
5550 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{else_const.getResult()});
5551 try else_block.addOperation(else_yield.op);
5552
5553 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{if_op.op.getResult(0).?});
5554 try entry.addOperation(ret.op);
5555
5556 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5557 defer backend.deinit();
5558 const compiled = try backend.compile(module.op);
5559 const exec = try onlyResult(&backend, compiled, "select_in_if", &.{});
5560 try testing.expectEqual(@as(i64, -7), try integerOf(exec));
5561 }
5562
5563 test "x86_64 arith.select rejects mismatched true/false types (trait verifier fires)" {
5564 if (!supports_x86_64_backend) return;
5565 const testing = std.testing;
5566 const dialects = @import("../../dialects/root.zig");
5567 const ArithDialect = dialects.ArithDialect;
5568 const BuiltinDialect = dialects.BuiltinDialect;
5569 const FuncDialect = dialects.FuncDialect;
5570
5571 var arena = alloc_arena.Arena.init(std.testing.allocator);
5572 defer arena.deinit();
5573 const allocator = arena.allocator();
5574
5575 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5576 defer ir_ctx.deinit(allocator);
5577 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5578
5579 const loc = ir.Location.getUnknown();
5580 const i32_type = try ArithDialect.getI32Type(&ir_ctx);
5581 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
5582
5583 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5584 const module_block = module.getBodyBlock();
5585 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "select_mismatched", &.{}, &.{i32_type});
5586 try module_block.addOperation(func.op);
5587
5588 const entry = func.getEntryBlock();
5589 var c_cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
5590 try entry.addOperation(c_cond.op);
5591 var c_true = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 1);
5592 try entry.addOperation(c_true.op);
5593 var c_false = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 2);
5594 try entry.addOperation(c_false.op);
5595 var sel = try ArithDialect.SelectOp.create(&ir_ctx, loc, c_cond.getResult(), c_true.getResult(), c_false.getResult());
5596 try entry.addOperation(sel.op);
5597 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{sel.getResult()});
5598 try entry.addOperation(ret.op);
5599
5600 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5601 defer backend.deinit();
5602 try testing.expectError(error.VerificationFailed, backend.compile(module.op));
5603 }
5604
5605 test "x86_64 arith.fma over f32 computes a*b+c with libm fmaf" {
5606 if (!supports_x86_64_backend) return;
5607 const testing = std.testing;
5608 try testing.expectEqual(
5609 @as(f32, 10.0),
5610 try runFloatTernaryOpRaw(.fma, f32, "fma_f32_basic", 2.0, 3.0, 4.0),
5611 );
5612 }
5613
5614 test "x86_64 arith.fma over f64 computes a*b+c with libm fma" {
5615 if (!supports_x86_64_backend) return;
5616 const testing = std.testing;
5617 try testing.expectEqual(
5618 @as(f64, 10.0),
5619 try runFloatTernaryOpRaw(.fma, f64, "fma_f64_basic", 2.0, 3.0, 4.0),
5620 );
5621 }
5622
5623 test "x86_64 arith.fma proves single-rounding via representable sentinel (f64)" {
5624 if (!supports_x86_64_backend) return;
5625 const testing = std.testing;
5626 const a: f64 = @bitCast(@as(u64, 0x3FF0000000000001));
5627 const b: f64 = a;
5628 const c: f64 = @bitCast(@as(u64, 0xBFF0000000000002));
5629 const expected: f64 = @bitCast(@as(u64, 0x3970000000000000));
5630 try testing.expectEqual(
5631 expected,
5632 try runFloatTernaryOpRaw(.fma, f64, "fma_f64_single_round", a, b, c),
5633 );
5634 }
5635
5636 test "x86_64 arith.fma proves single-rounding via representable sentinel (f32)" {
5637 if (!supports_x86_64_backend) return;
5638 const testing = std.testing;
5639 const a: f32 = @bitCast(@as(u32, 0x3F800001));
5640 const b: f32 = a;
5641 const c: f32 = @bitCast(@as(u32, 0xBF800002));
5642 const expected: f32 = @bitCast(@as(u32, 0x28800000));
5643 try testing.expectEqual(
5644 expected,
5645 try runFloatTernaryOpRaw(.fma, f32, "fma_f32_single_round", a, b, c),
5646 );
5647 }
5648
5649 test "x86_64 arith.fma rejects f16, integers, and mismatched operand types" {
5650 if (!supports_x86_64_backend) return;
5651 const testing = std.testing;
5652 const dialects = @import("../../dialects/root.zig");
5653 const ArithDialect = dialects.ArithDialect;
5654 const BuiltinDialect = dialects.BuiltinDialect;
5655 const FuncDialect = dialects.FuncDialect;
5656
5657 {
5658 var arena = alloc_arena.Arena.init(std.testing.allocator);
5659 defer arena.deinit();
5660 const allocator = arena.allocator();
5661
5662 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5663 defer ir_ctx.deinit(allocator);
5664 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5665 const loc = ir.Location.getUnknown();
5666 const i32_type = try ArithDialect.getScalarType(&ir_ctx, .i32);
5667
5668 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5669 const module_block = module.getBodyBlock();
5670 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "fma_int_reject", &.{}, &.{i32_type});
5671 try module_block.addOperation(func.op);
5672
5673 const entry = func.getEntryBlock();
5674 var ca = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 2);
5675 try entry.addOperation(ca.op);
5676 var cb = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 3);
5677 try entry.addOperation(cb.op);
5678 var cc = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 4);
5679 try entry.addOperation(cc.op);
5680
5681 var state = ir.Operation.State.init(ArithDialect.FmaOp.operation_name, loc);
5682 state.addOperands(&.{ ca.getResult(), cb.getResult(), cc.getResult() });
5683 state.addTypes(&.{i32_type});
5684 const fma_op = try ir_ctx.createOperation(state);
5685 try entry.addOperation(fma_op);
5686 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{fma_op.getResult(0).?});
5687 try entry.addOperation(ret.op);
5688
5689 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5690 defer backend.deinit();
5691 try testing.expectError(error.JitCompileFailed, backend.compile(module.op));
5692 }
5693
5694 {
5695 var arena = alloc_arena.Arena.init(std.testing.allocator);
5696 defer arena.deinit();
5697 const allocator = arena.allocator();
5698
5699 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5700 defer ir_ctx.deinit(allocator);
5701 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5702 const loc = ir.Location.getUnknown();
5703 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
5704 const f64_type = try ArithDialect.getF64Type(&ir_ctx);
5705
5706 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5707 const module_block = module.getBodyBlock();
5708 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "fma_mismatch", &.{}, &.{f32_type});
5709 try module_block.addOperation(func.op);
5710
5711 const entry = func.getEntryBlock();
5712 var ca = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 1.0);
5713 try entry.addOperation(ca.op);
5714 var cb = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f32_type, 2.0);
5715 try entry.addOperation(cb.op);
5716 var cc = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 3.0);
5717 try entry.addOperation(cc.op);
5718 var fma = try ArithDialect.FmaOp.create(&ir_ctx, loc, ca.getResult(), cb.getResult(), cc.getResult());
5719 try entry.addOperation(fma.op);
5720 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{fma.getResult()});
5721 try entry.addOperation(ret.op);
5722
5723 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5724 defer backend.deinit();
5725 try testing.expectError(error.VerificationFailed, backend.compile(module.op));
5726 }
5727 }
5728
5729 test "x86_64 arith.cast: f16 still rejected (no Float16 capability)" {
5730 if (!supports_x86_64_backend) return;
5731 const testing = std.testing;
5732 const dialects = @import("../../dialects/root.zig");
5733 const ArithDialect = dialects.ArithDialect;
5734 const BuiltinDialect = dialects.BuiltinDialect;
5735 const FuncDialect = dialects.FuncDialect;
5736
5737 {
5738 var arena = alloc_arena.Arena.init(std.testing.allocator);
5739 defer arena.deinit();
5740 const allocator = arena.allocator();
5741
5742 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5743 defer ir_ctx.deinit(allocator);
5744 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5745
5746 const loc = ir.Location.getUnknown();
5747 const i32_type = try ArithDialect.getI32Type(&ir_ctx);
5748 const f16_type = try ArithDialect.getF16Type(&ir_ctx);
5749
5750 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5751 const module_block = module.getBodyBlock();
5752 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "cast_i32_f16", &.{}, &.{i32_type});
5753 try module_block.addOperation(func.op);
5754
5755 const entry = func.getEntryBlock();
5756 var c = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 1);
5757 try entry.addOperation(c.op);
5758 var cast = try ArithDialect.CastOp.create(&ir_ctx, loc, c.getResult(), f16_type);
5759 try entry.addOperation(cast.op);
5760 var back = try ArithDialect.CastOp.create(&ir_ctx, loc, cast.getResult(), i32_type);
5761 try entry.addOperation(back.op);
5762 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{back.getResult()});
5763 try entry.addOperation(ret.op);
5764
5765 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5766 defer backend.deinit();
5767 try testing.expectError(error.JitCompileFailed, backend.compile(module.op));
5768 }
5769 }
5770
5771 fn runUnaryFloatInScfIf(
5772 name: []const u8,
5773 op_kind: enum { sqrt, exp },
5774 input: f64,
5775 else_val: f64,
5776 expected: f64,
5777 tol: f64,
5778 ) !void {
5779 const dialects = @import("../../dialects/root.zig");
5780 const ArithDialect = dialects.ArithDialect;
5781 const BuiltinDialect = dialects.BuiltinDialect;
5782 const FuncDialect = dialects.FuncDialect;
5783 const ScfDialect = dialects.ScfDialect;
5784
5785 var arena = alloc_arena.Arena.init(std.testing.allocator);
5786 defer arena.deinit();
5787 const allocator = arena.allocator();
5788
5789 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5790 defer ir_ctx.deinit(allocator);
5791 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5792
5793 const loc = ir.Location.getUnknown();
5794 const f64_type = try ArithDialect.getF64Type(&ir_ctx);
5795 const bool_type = try ArithDialect.getScalarType(&ir_ctx, .bool);
5796
5797 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5798 const module_block = module.getBodyBlock();
5799 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, name, &.{}, &.{f64_type});
5800 try module_block.addOperation(func.op);
5801
5802 const entry = func.getEntryBlock();
5803 var cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
5804 try entry.addOperation(cond.op);
5805
5806 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, cond.getResult(), &.{f64_type});
5807 try entry.addOperation(if_op.op);
5808
5809 const then_block = if_op.getThenBlock();
5810 var then_const = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, input);
5811 try then_block.addOperation(then_const.op);
5812 const then_result = blk: {
5813 switch (op_kind) {
5814 .sqrt => {
5815 var op = try ArithDialect.SqrtOp.create(&ir_ctx, loc, then_const.getResult());
5816 try then_block.addOperation(op.op);
5817 break :blk op.getResult();
5818 },
5819 .exp => {
5820 var op = try ArithDialect.ExpOp.create(&ir_ctx, loc, then_const.getResult());
5821 try then_block.addOperation(op.op);
5822 break :blk op.getResult();
5823 },
5824 }
5825 };
5826 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{then_result});
5827 try then_block.addOperation(then_yield.op);
5828
5829 const else_block = if_op.getElseBlock().?;
5830 var else_const = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, else_val);
5831 try else_block.addOperation(else_const.op);
5832 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{else_const.getResult()});
5833 try else_block.addOperation(else_yield.op);
5834
5835 const result = if_op.op.getResult(0).?;
5836 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result});
5837 try entry.addOperation(ret.op);
5838 _ = bool_type;
5839
5840 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5841 defer backend.deinit();
5842 const compiled = try backend.compile(module.op);
5843 const exec_result = try onlyResult(&backend, compiled, name, &.{});
5844 const got = try floatOf(exec_result);
5845 try std.testing.expectApproxEqAbs(expected, got, tol);
5846 }
5847
5848 test "x86_64 arith.sqrt inside scf.if body (emitArithSqrt via dispatchOp)" {
5849 if (!supports_x86_64_backend) return;
5850 try runUnaryFloatInScfIf("sqrt_in_if", .sqrt, 16.0, 0.0, 4.0, 1e-12);
5851 }
5852
5853 test "x86_64 arith.exp inside scf.if body (emitArithLibmUnary via dispatchOp)" {
5854 if (!supports_x86_64_backend) return;
5855 try runUnaryFloatInScfIf("exp_in_if", .exp, 0.0, 0.0, 1.0, 1e-12);
5856 }
5857
5858 test "x86_64 arith.pow inside scf.if body (emitArithLibmBinary via dispatchOp)" {
5859 if (!supports_x86_64_backend) return;
5860 const testing = std.testing;
5861 const dialects = @import("../../dialects/root.zig");
5862 const ArithDialect = dialects.ArithDialect;
5863 const BuiltinDialect = dialects.BuiltinDialect;
5864 const FuncDialect = dialects.FuncDialect;
5865 const ScfDialect = dialects.ScfDialect;
5866
5867 var arena = alloc_arena.Arena.init(std.testing.allocator);
5868 defer arena.deinit();
5869 const allocator = arena.allocator();
5870
5871 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5872 defer ir_ctx.deinit(allocator);
5873 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5874
5875 const loc = ir.Location.getUnknown();
5876 const f64_type = try ArithDialect.getF64Type(&ir_ctx);
5877
5878 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5879 const module_block = module.getBodyBlock();
5880 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "pow_in_if", &.{}, &.{f64_type});
5881 try module_block.addOperation(func.op);
5882
5883 const entry = func.getEntryBlock();
5884 var cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
5885 try entry.addOperation(cond.op);
5886
5887 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, cond.getResult(), &.{f64_type});
5888 try entry.addOperation(if_op.op);
5889
5890 const then_block = if_op.getThenBlock();
5891 var base = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 2.0);
5892 try then_block.addOperation(base.op);
5893 var exp = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 10.0);
5894 try then_block.addOperation(exp.op);
5895 var pow = try ArithDialect.PowOp.create(&ir_ctx, loc, base.getResult(), exp.getResult());
5896 try then_block.addOperation(pow.op);
5897 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{pow.getResult()});
5898 try then_block.addOperation(then_yield.op);
5899
5900 const else_block = if_op.getElseBlock().?;
5901 var else_const = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 0.0);
5902 try else_block.addOperation(else_const.op);
5903 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{else_const.getResult()});
5904 try else_block.addOperation(else_yield.op);
5905
5906 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{if_op.op.getResult(0).?});
5907 try entry.addOperation(ret.op);
5908
5909 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5910 defer backend.deinit();
5911 const compiled = try backend.compile(module.op);
5912 const exec = try onlyResult(&backend, compiled, "pow_in_if", &.{});
5913 try testing.expectApproxEqAbs(@as(f64, 1024.0), try floatOf(exec), 1e-9);
5914 }
5915
5916 test "x86_64 arith.and / arith.not inside scf.if body (emitArithBitwiseBinary + emitArithNot via dispatchOp)" {
5917 if (!supports_x86_64_backend) return;
5918 const testing = std.testing;
5919 const dialects = @import("../../dialects/root.zig");
5920 const ArithDialect = dialects.ArithDialect;
5921 const BuiltinDialect = dialects.BuiltinDialect;
5922 const FuncDialect = dialects.FuncDialect;
5923 const ScfDialect = dialects.ScfDialect;
5924
5925 var arena = alloc_arena.Arena.init(std.testing.allocator);
5926 defer arena.deinit();
5927 const allocator = arena.allocator();
5928
5929 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5930 defer ir_ctx.deinit(allocator);
5931 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5932
5933 const loc = ir.Location.getUnknown();
5934 const i32_type = try ArithDialect.getI32Type(&ir_ctx);
5935
5936 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5937 const module_block = module.getBodyBlock();
5938 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "and_not_in_if", &.{}, &.{i32_type});
5939 try module_block.addOperation(func.op);
5940
5941 const entry = func.getEntryBlock();
5942 var cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
5943 try entry.addOperation(cond.op);
5944
5945 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, cond.getResult(), &.{i32_type});
5946 try entry.addOperation(if_op.op);
5947
5948 const then_block = if_op.getThenBlock();
5949 var lhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 0xFF00);
5950 try then_block.addOperation(lhs.op);
5951 var rhs = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 0x0F0F);
5952 try then_block.addOperation(rhs.op);
5953 var and_op = try ArithDialect.AndOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
5954 try then_block.addOperation(and_op.op);
5955 var not_op = try ArithDialect.NotOp.create(&ir_ctx, loc, and_op.getResult());
5956 try then_block.addOperation(not_op.op);
5957 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{not_op.getResult()});
5958 try then_block.addOperation(then_yield.op);
5959
5960 const else_block = if_op.getElseBlock().?;
5961 var else_const = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 0);
5962 try else_block.addOperation(else_const.op);
5963 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{else_const.getResult()});
5964 try else_block.addOperation(else_yield.op);
5965
5966 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{if_op.op.getResult(0).?});
5967 try entry.addOperation(ret.op);
5968
5969 var backend = try Backend.init(allocator, &ir_ctx, .testing);
5970 defer backend.deinit();
5971 const compiled = try backend.compile(module.op);
5972 const exec = try onlyResult(&backend, compiled, "and_not_in_if", &.{});
5973 try testing.expectEqual(@as(i64, -3841), try integerOf(exec));
5974 }
5975
5976 test "x86_64 arith.shl inside scf.if body (emitArithShift via dispatchOp)" {
5977 if (!supports_x86_64_backend) return;
5978 const testing = std.testing;
5979 const dialects = @import("../../dialects/root.zig");
5980 const ArithDialect = dialects.ArithDialect;
5981 const BuiltinDialect = dialects.BuiltinDialect;
5982 const FuncDialect = dialects.FuncDialect;
5983 const ScfDialect = dialects.ScfDialect;
5984
5985 var arena = alloc_arena.Arena.init(std.testing.allocator);
5986 defer arena.deinit();
5987 const allocator = arena.allocator();
5988
5989 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
5990 defer ir_ctx.deinit(allocator);
5991 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
5992
5993 const loc = ir.Location.getUnknown();
5994 const i32_type = try ArithDialect.getI32Type(&ir_ctx);
5995
5996 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
5997 const module_block = module.getBodyBlock();
5998 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "shl_in_if", &.{}, &.{i32_type});
5999 try module_block.addOperation(func.op);
6000
6001 const entry = func.getEntryBlock();
6002 var cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
6003 try entry.addOperation(cond.op);
6004
6005 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, cond.getResult(), &.{i32_type});
6006 try entry.addOperation(if_op.op);
6007
6008 const then_block = if_op.getThenBlock();
6009 var v = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 1);
6010 try then_block.addOperation(v.op);
6011 var n = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 4);
6012 try then_block.addOperation(n.op);
6013 var shl = try ArithDialect.ShlOp.create(&ir_ctx, loc, v.getResult(), n.getResult());
6014 try then_block.addOperation(shl.op);
6015 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{shl.getResult()});
6016 try then_block.addOperation(then_yield.op);
6017
6018 const else_block = if_op.getElseBlock().?;
6019 var else_const = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i32_type, 0);
6020 try else_block.addOperation(else_const.op);
6021 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{else_const.getResult()});
6022 try else_block.addOperation(else_yield.op);
6023
6024 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{if_op.op.getResult(0).?});
6025 try entry.addOperation(ret.op);
6026
6027 var backend = try Backend.init(allocator, &ir_ctx, .testing);
6028 defer backend.deinit();
6029 const compiled = try backend.compile(module.op);
6030 const exec = try onlyResult(&backend, compiled, "shl_in_if", &.{});
6031 try testing.expectEqual(@as(i64, 16), try integerOf(exec));
6032 }
6033
6034 test "x86_64 arith.max (float branch) inside scf.if body (emitArithMinMax via dispatchOp)" {
6035 if (!supports_x86_64_backend) return;
6036 const testing = std.testing;
6037 const dialects = @import("../../dialects/root.zig");
6038 const ArithDialect = dialects.ArithDialect;
6039 const BuiltinDialect = dialects.BuiltinDialect;
6040 const FuncDialect = dialects.FuncDialect;
6041 const ScfDialect = dialects.ScfDialect;
6042
6043 var arena = alloc_arena.Arena.init(std.testing.allocator);
6044 defer arena.deinit();
6045 const allocator = arena.allocator();
6046
6047 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
6048 defer ir_ctx.deinit(allocator);
6049 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
6050
6051 const loc = ir.Location.getUnknown();
6052 const f64_type = try ArithDialect.getF64Type(&ir_ctx);
6053
6054 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
6055 const module_block = module.getBodyBlock();
6056 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "max_in_if", &.{}, &.{f64_type});
6057 try module_block.addOperation(func.op);
6058
6059 const entry = func.getEntryBlock();
6060 var cond = try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true);
6061 try entry.addOperation(cond.op);
6062
6063 var if_op = try ScfDialect.IfOp.create(&ir_ctx, loc, cond.getResult(), &.{f64_type});
6064 try entry.addOperation(if_op.op);
6065
6066 const then_block = if_op.getThenBlock();
6067 var a = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 3.0);
6068 try then_block.addOperation(a.op);
6069 var b = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 7.0);
6070 try then_block.addOperation(b.op);
6071 var max_op = try ArithDialect.MaxOp.create(&ir_ctx, loc, a.getResult(), b.getResult());
6072 try then_block.addOperation(max_op.op);
6073 const then_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{max_op.getResult()});
6074 try then_block.addOperation(then_yield.op);
6075
6076 const else_block = if_op.getElseBlock().?;
6077 var else_const = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 0.0);
6078 try else_block.addOperation(else_const.op);
6079 const else_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{else_const.getResult()});
6080 try else_block.addOperation(else_yield.op);
6081
6082 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{if_op.op.getResult(0).?});
6083 try entry.addOperation(ret.op);
6084
6085 var backend = try Backend.init(allocator, &ir_ctx, .testing);
6086 defer backend.deinit();
6087 const compiled = try backend.compile(module.op);
6088 const exec = try onlyResult(&backend, compiled, "max_in_if", &.{});
6089 try testing.expectEqual(@as(f64, 7.0), try floatOf(exec));
6090 }
6091
6092 test "x86_64 unsupported scalar arith dtype/op pairs are rejected" {
6093 if (!supports_x86_64_backend) return;
6094 const testing = std.testing;
6095 const dialects = @import("../../dialects/root.zig");
6096 const ArithDialect = dialects.ArithDialect;
6097 const BuiltinDialect = dialects.BuiltinDialect;
6098 const FuncDialect = dialects.FuncDialect;
6099
6100 const RejectCase = struct {
6101 dtype: ArithDialect.ScalarTypeKind,
6102 op_kind: ScalarOpKind,
6103 };
6104 const cases: []const RejectCase = &.{
6105 .{ .dtype = .index, .op_kind = .neg },
6106 .{ .dtype = .i32, .op_kind = .sqrt },
6107 .{ .dtype = .i64, .op_kind = .sqrt },
6108 .{ .dtype = .index, .op_kind = .sqrt },
6109 .{ .dtype = .bool, .op_kind = .sqrt },
6110 .{ .dtype = .i32, .op_kind = .floor },
6111 .{ .dtype = .bool, .op_kind = .floor },
6112 .{ .dtype = .index, .op_kind = .abs },
6113 .{ .dtype = .f16, .op_kind = .neg },
6114 .{ .dtype = .f16, .op_kind = .abs },
6115 .{ .dtype = .f16, .op_kind = .sqrt },
6116 .{ .dtype = .f16, .op_kind = .floor },
6117 .{ .dtype = .f16, .op_kind = .max },
6118 .{ .dtype = .f16, .op_kind = .min },
6119 .{ .dtype = .f16, .op_kind = .sin },
6120 .{ .dtype = .f16, .op_kind = .cos },
6121 .{ .dtype = .f16, .op_kind = .tan },
6122 .{ .dtype = .f16, .op_kind = .exp },
6123 .{ .dtype = .f16, .op_kind = .log },
6124 .{ .dtype = .f16, .op_kind = .tanh },
6125 .{ .dtype = .f16, .op_kind = .pow },
6126 .{ .dtype = .bool, .op_kind = .neg },
6127 .{ .dtype = .bool, .op_kind = .abs },
6128 .{ .dtype = .bool, .op_kind = .max },
6129 .{ .dtype = .bool, .op_kind = .min },
6130 .{ .dtype = .bool, .op_kind = .div },
6131 .{ .dtype = .bool, .op_kind = .rem },
6132 .{ .dtype = .bool, .op_kind = .umulhi },
6133 .{ .dtype = .bool, .op_kind = .popcount },
6134 .{ .dtype = .bool, .op_kind = .sin },
6135 .{ .dtype = .bool, .op_kind = .pow },
6136 .{ .dtype = .i32, .op_kind = .sin },
6137 .{ .dtype = .i32, .op_kind = .cos },
6138 .{ .dtype = .i32, .op_kind = .tan },
6139 .{ .dtype = .i32, .op_kind = .exp },
6140 .{ .dtype = .i32, .op_kind = .log },
6141 .{ .dtype = .i32, .op_kind = .tanh },
6142 .{ .dtype = .i32, .op_kind = .pow },
6143 .{ .dtype = .index, .op_kind = .pow },
6144 .{ .dtype = .f32, .op_kind = .band },
6145 .{ .dtype = .f32, .op_kind = .bor },
6146 .{ .dtype = .f32, .op_kind = .bxor },
6147 .{ .dtype = .f32, .op_kind = .bnot },
6148 .{ .dtype = .f32, .op_kind = .shl },
6149 .{ .dtype = .f32, .op_kind = .shr },
6150 .{ .dtype = .f32, .op_kind = .ushr },
6151 .{ .dtype = .f32, .op_kind = .rem },
6152 .{ .dtype = .f32, .op_kind = .umulhi },
6153 .{ .dtype = .f32, .op_kind = .popcount },
6154 .{ .dtype = .f64, .op_kind = .band },
6155 .{ .dtype = .f64, .op_kind = .shl },
6156 .{ .dtype = .f64, .op_kind = .rem },
6157 .{ .dtype = .f64, .op_kind = .umulhi },
6158 .{ .dtype = .f64, .op_kind = .popcount },
6159 .{ .dtype = .index, .op_kind = .umulhi },
6160 .{ .dtype = .bool, .op_kind = .shl },
6161 .{ .dtype = .f16, .op_kind = .band },
6162 .{ .dtype = .f16, .op_kind = .bor },
6163 .{ .dtype = .f16, .op_kind = .bxor },
6164 .{ .dtype = .f16, .op_kind = .bnot },
6165 .{ .dtype = .f16, .op_kind = .shl },
6166 .{ .dtype = .f16, .op_kind = .shr },
6167 .{ .dtype = .f16, .op_kind = .ushr },
6168 .{ .dtype = .f16, .op_kind = .umulhi },
6169 .{ .dtype = .f16, .op_kind = .popcount },
6170 };
6171
6172 for (cases) |case| {
6173 var arena = alloc_arena.Arena.init(std.testing.allocator);
6174 defer arena.deinit();
6175 const allocator = arena.allocator();
6176
6177 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
6178 defer ir_ctx.deinit(allocator);
6179 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
6180
6181 const loc = ir.Location.getUnknown();
6182 const t = try ArithDialect.getScalarType(&ir_ctx, case.dtype);
6183
6184 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
6185 const module_block = module.getBodyBlock();
6186
6187 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "rejected", &.{}, &.{t});
6188 try module_block.addOperation(func.op);
6189
6190 const entry = func.getEntryBlock();
6191 var lhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6192 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 1.0)
6193 else if (case.dtype == .bool)
6194 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, true)
6195 else
6196 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 1);
6197 try entry.addOperation(lhs.op);
6198
6199 const result_value = switch (case.op_kind) {
6200 .neg => blk: {
6201 const op = try ArithDialect.NegOp.create(&ir_ctx, loc, lhs.getResult());
6202 try entry.addOperation(op.op);
6203 break :blk op.getResult();
6204 },
6205 .abs => blk: {
6206 const op = try ArithDialect.AbsOp.create(&ir_ctx, loc, lhs.getResult());
6207 try entry.addOperation(op.op);
6208 break :blk op.getResult();
6209 },
6210 .sqrt => blk: {
6211 const op = try ArithDialect.SqrtOp.create(&ir_ctx, loc, lhs.getResult());
6212 try entry.addOperation(op.op);
6213 break :blk op.getResult();
6214 },
6215 .floor => blk: {
6216 const op = try ArithDialect.FloorOp.create(&ir_ctx, loc, lhs.getResult());
6217 try entry.addOperation(op.op);
6218 break :blk op.getResult();
6219 },
6220 .max => blk: {
6221 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6222 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6223 else if (case.dtype == .bool)
6224 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6225 else
6226 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6227 try entry.addOperation(rhs.op);
6228 const op = try ArithDialect.MaxOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6229 try entry.addOperation(op.op);
6230 break :blk op.getResult();
6231 },
6232 .min => blk: {
6233 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6234 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6235 else if (case.dtype == .bool)
6236 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6237 else
6238 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6239 try entry.addOperation(rhs.op);
6240 const op = try ArithDialect.MinOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6241 try entry.addOperation(op.op);
6242 break :blk op.getResult();
6243 },
6244 .div => blk: {
6245 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6246 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6247 else if (case.dtype == .bool)
6248 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6249 else
6250 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6251 try entry.addOperation(rhs.op);
6252 const op = try ArithDialect.DivOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6253 try entry.addOperation(op.op);
6254 break :blk op.getResult();
6255 },
6256 .rem => blk: {
6257 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6258 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6259 else if (case.dtype == .bool)
6260 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6261 else
6262 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6263 try entry.addOperation(rhs.op);
6264 const op = try ArithDialect.RemOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6265 try entry.addOperation(op.op);
6266 break :blk op.getResult();
6267 },
6268 .umulhi => blk: {
6269 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6270 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6271 else if (case.dtype == .bool)
6272 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6273 else
6274 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6275 try entry.addOperation(rhs.op);
6276 const op = try ArithDialect.UmulhiOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6277 try entry.addOperation(op.op);
6278 break :blk op.getResult();
6279 },
6280 .popcount => blk: {
6281 const op = try ArithDialect.PopCountOp.create(&ir_ctx, loc, lhs.getResult());
6282 try entry.addOperation(op.op);
6283 break :blk op.getResult();
6284 },
6285 .sin => blk: {
6286 const op = try ArithDialect.SinOp.create(&ir_ctx, loc, lhs.getResult());
6287 try entry.addOperation(op.op);
6288 break :blk op.getResult();
6289 },
6290 .cos => blk: {
6291 const op = try ArithDialect.CosOp.create(&ir_ctx, loc, lhs.getResult());
6292 try entry.addOperation(op.op);
6293 break :blk op.getResult();
6294 },
6295 .tan => blk: {
6296 const op = try ArithDialect.TanOp.create(&ir_ctx, loc, lhs.getResult());
6297 try entry.addOperation(op.op);
6298 break :blk op.getResult();
6299 },
6300 .exp => blk: {
6301 const op = try ArithDialect.ExpOp.create(&ir_ctx, loc, lhs.getResult());
6302 try entry.addOperation(op.op);
6303 break :blk op.getResult();
6304 },
6305 .log => blk: {
6306 const op = try ArithDialect.LogOp.create(&ir_ctx, loc, lhs.getResult());
6307 try entry.addOperation(op.op);
6308 break :blk op.getResult();
6309 },
6310 .tanh => blk: {
6311 const op = try ArithDialect.TanhOp.create(&ir_ctx, loc, lhs.getResult());
6312 try entry.addOperation(op.op);
6313 break :blk op.getResult();
6314 },
6315 .pow => blk: {
6316 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6317 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6318 else if (case.dtype == .bool)
6319 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6320 else
6321 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6322 try entry.addOperation(rhs.op);
6323 const op = try ArithDialect.PowOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6324 try entry.addOperation(op.op);
6325 break :blk op.getResult();
6326 },
6327 .band => blk: {
6328 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6329 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6330 else if (case.dtype == .bool)
6331 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6332 else
6333 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6334 try entry.addOperation(rhs.op);
6335 const op = try ArithDialect.AndOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6336 try entry.addOperation(op.op);
6337 break :blk op.getResult();
6338 },
6339 .bor => blk: {
6340 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6341 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6342 else if (case.dtype == .bool)
6343 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6344 else
6345 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6346 try entry.addOperation(rhs.op);
6347 const op = try ArithDialect.OrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6348 try entry.addOperation(op.op);
6349 break :blk op.getResult();
6350 },
6351 .bxor => blk: {
6352 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6353 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6354 else if (case.dtype == .bool)
6355 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6356 else
6357 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6358 try entry.addOperation(rhs.op);
6359 const op = try ArithDialect.XorOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6360 try entry.addOperation(op.op);
6361 break :blk op.getResult();
6362 },
6363 .bnot => blk: {
6364 const op = try ArithDialect.NotOp.create(&ir_ctx, loc, lhs.getResult());
6365 try entry.addOperation(op.op);
6366 break :blk op.getResult();
6367 },
6368 .shl => blk: {
6369 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6370 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6371 else if (case.dtype == .bool)
6372 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6373 else
6374 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6375 try entry.addOperation(rhs.op);
6376 const op = try ArithDialect.ShlOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6377 try entry.addOperation(op.op);
6378 break :blk op.getResult();
6379 },
6380 .shr => blk: {
6381 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6382 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6383 else if (case.dtype == .bool)
6384 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6385 else
6386 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6387 try entry.addOperation(rhs.op);
6388 const op = try ArithDialect.ShrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6389 try entry.addOperation(op.op);
6390 break :blk op.getResult();
6391 },
6392 .ushr => blk: {
6393 var rhs = if (case.dtype == .f16 or case.dtype == .f32 or case.dtype == .f64)
6394 try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, t, 2.0)
6395 else if (case.dtype == .bool)
6396 try ArithDialect.ConstantOp.createBool(&ir_ctx, loc, false)
6397 else
6398 try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, t, 2);
6399 try entry.addOperation(rhs.op);
6400 const op = try ArithDialect.UshrOp.create(&ir_ctx, loc, lhs.getResult(), rhs.getResult());
6401 try entry.addOperation(op.op);
6402 break :blk op.getResult();
6403 },
6404 };
6405
6406 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result_value});
6407 try entry.addOperation(ret.op);
6408
6409 var backend = try Backend.init(allocator, &ir_ctx, .testing);
6410 defer backend.deinit();
6411 const expected = if (case.dtype == .f16)
6412 error.UnsupportedOperation
6413 else
6414 error.JitCompileFailed;
6415 try testing.expectError(expected, backend.compile(module.op));
6416 }
6417 }
6418
6419 test "x86_64 scf.while loop-invariant home survives body register reuse" {
6420 if (!supports_x86_64_backend) return;
6421
6422 const testing = std.testing;
6423 const dialects = @import("../../dialects/root.zig");
6424 const ArithDialect = dialects.ArithDialect;
6425 const BuiltinDialect = dialects.BuiltinDialect;
6426 const FuncDialect = dialects.FuncDialect;
6427 const ScfDialect = dialects.ScfDialect;
6428
6429 var arena = alloc_arena.Arena.init(std.testing.allocator);
6430 defer arena.deinit();
6431 const allocator = arena.allocator();
6432
6433 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
6434 defer ir_ctx.deinit(allocator);
6435 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
6436
6437 const loc = ir.Location.getUnknown();
6438 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
6439
6440 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
6441 const module_block = module.getBodyBlock();
6442
6443 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "while_invariant", &.{i64_type}, &.{i64_type});
6444 try module_block.addOperation(func.op);
6445 const limit = func.getArgument(0);
6446
6447 const entry = func.getEntryBlock();
6448 var fuel_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 100);
6449 try entry.addOperation(fuel_init.op);
6450 var acc_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
6451 try entry.addOperation(acc_init.op);
6452
6453 var while_op = try ScfDialect.WhileOp.create(
6454 &ir_ctx,
6455 loc,
6456 &.{ fuel_init.getResult(), acc_init.getResult() },
6457 &.{ i64_type, i64_type },
6458 );
6459 try entry.addOperation(while_op.op);
6460
6461 const before = while_op.getBeforeBlock();
6462 const before_fuel = before.arguments.items[0];
6463 const before_acc = before.arguments.items[1];
6464 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
6465 try before.addOperation(zero.op);
6466 var has_fuel = try ArithDialect.CmpOp.create(&ir_ctx, loc, .gt, before_fuel, zero.getResult());
6467 try before.addOperation(has_fuel.op);
6468 var below = try ArithDialect.CmpOp.create(&ir_ctx, loc, .lt, before_acc, limit);
6469 try before.addOperation(below.op);
6470 var both = try ArithDialect.AndOp.create(&ir_ctx, loc, has_fuel.getResult(), below.getResult());
6471 try before.addOperation(both.op);
6472 const condition = try ScfDialect.ConditionOp.create(&ir_ctx, loc, both.getResult(), &.{ before_fuel, before_acc });
6473 try before.addOperation(condition.op);
6474
6475 const after = while_op.getAfterBlock();
6476 const after_fuel = after.arguments.items[0];
6477 const after_acc = after.arguments.items[1];
6478 var one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 1);
6479 try after.addOperation(one.op);
6480 var fuel_next = try ArithDialect.SubOp.create(&ir_ctx, loc, after_fuel, one.getResult());
6481 try after.addOperation(fuel_next.op);
6482
6483 var chain_zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
6484 try after.addOperation(chain_zero.op);
6485 var chain: *ir.Value = chain_zero.getResult();
6486 var chain_index: usize = 0;
6487 while (chain_index < 10) : (chain_index += 1) {
6488 var link = try ArithDialect.SubOp.create(&ir_ctx, loc, chain, one.getResult());
6489 try after.addOperation(link.op);
6490 chain = link.getResult();
6491 }
6492 var zero_term = try ArithDialect.SubOp.create(&ir_ctx, loc, chain, chain);
6493 try after.addOperation(zero_term.op);
6494 var acc_bump = try ArithDialect.AddOp.create(&ir_ctx, loc, after_acc, one.getResult());
6495 try after.addOperation(acc_bump.op);
6496 var acc_next = try ArithDialect.AddOp.create(&ir_ctx, loc, acc_bump.getResult(), zero_term.getResult());
6497 try after.addOperation(acc_next.op);
6498 const yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{ fuel_next.getResult(), acc_next.getResult() });
6499 try after.addOperation(yield.op);
6500
6501 const result = while_op.op.getResult(1) orelse return error.TestFailure;
6502 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{result});
6503 try entry.addOperation(ret.op);
6504
6505 var backend = try Backend.init(allocator, &ir_ctx, .testing);
6506 defer backend.deinit();
6507 const compiled = try backend.compile(module.op);
6508
6509 const exec_result =
6510 try onlyResult(&backend, compiled, "while_invariant", &.{.{ .i64 = 10 }});
6511 try testing.expectEqual(@as(i64, 10), try integerOf(exec_result));
6512 }
6513
6514 test "x86_64 scf.for memref kernel is register allocated and correct" {
6515 if (!supports_x86_64_backend) return;
6516
6517 const testing = std.testing;
6518 const dialects = @import("../../dialects/root.zig");
6519 const ArithDialect = dialects.ArithDialect;
6520 const BuiltinDialect = dialects.BuiltinDialect;
6521 const FuncDialect = dialects.FuncDialect;
6522 const MemrefDialect = dialects.MemrefDialect;
6523 const ScfDialect = dialects.ScfDialect;
6524
6525 var arena = alloc_arena.Arena.init(std.testing.allocator);
6526 defer arena.deinit();
6527 const allocator = arena.allocator();
6528
6529 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
6530 defer ir_ctx.deinit(allocator);
6531 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
6532
6533 const loc = ir.Location.getUnknown();
6534 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
6535 const index_type = try ArithDialect.getIndexType(&ir_ctx);
6536 const memref_i64 = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, i64_type, .host);
6537
6538 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
6539 const module_block = module.getBodyBlock();
6540
6541 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "for_sum", &.{ i64_type, memref_i64, memref_i64 }, &.{});
6542 try module_block.addOperation(func.op);
6543 const entry = func.getEntryBlock();
6544
6545 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
6546 try entry.addOperation(zero.op);
6547 var bound = try ArithDialect.CastOp.create(&ir_ctx, loc, func.getArgument(0), index_type);
6548 try entry.addOperation(bound.op);
6549 var step = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
6550 try entry.addOperation(step.op);
6551 var acc_init = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, i64_type, 0);
6552 try entry.addOperation(acc_init.op);
6553
6554 var for_op = try ScfDialect.ForOp.create(&ir_ctx, loc, zero.getResult(), bound.getResult(), step.getResult(), &.{acc_init.getResult()}, &.{i64_type});
6555 try entry.addOperation(for_op.op);
6556 const body = for_op.getBodyBlock();
6557 const iv = body.arguments.items[0];
6558 const acc = body.arguments.items[1];
6559
6560 var x = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(1), iv, i64_type);
6561 try body.addOperation(x.op);
6562 var next = try ArithDialect.AddOp.create(&ir_ctx, loc, acc, x.getResult());
6563 try body.addOperation(next.op);
6564 const yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{next.getResult()});
6565 try body.addOperation(yield.op);
6566
6567 const total = for_op.op.getResult(0) orelse return error.TestFailure;
6568 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, total, func.getArgument(2), zero.getResult());
6569 try entry.addOperation(store.op);
6570 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{});
6571 try entry.addOperation(ret.op);
6572
6573 var emitter = x86_64.Emitter.init(allocator);
6574 defer emitter.deinit();
6575 try emitter.emitFunction(func.op);
6576 try testing.expect(emitter.value_location_ranges.items.len > 0);
6577
6578 var backend = try Backend.init(allocator, &ir_ctx, .testing);
6579 defer backend.deinit();
6580 const compiled = try backend.compile(module.op);
6581
6582 const SumFn = *const fn (i64, [*]const i64, [*]i64) callconv(.c) void;
6583 const kernel = try backend.runtime.getFunction(compiled, "for_sum", SumFn);
6584
6585 var input: [100]i64 = undefined;
6586 for (&input, 0..) |*slot, i| slot.* = @intCast(i);
6587 var output: [1]i64 = .{-1};
6588 kernel(input.len, &input, &output);
6589 try testing.expectEqual(@as(i64, 4950), output[0]);
6590 }
6591
6592 test "x86_64 nested scf.for float accumulator stays slot carried" {
6593 if (!supports_x86_64_backend) return;
6594
6595 const testing = std.testing;
6596 const dialects = @import("../../dialects/root.zig");
6597 const ArithDialect = dialects.ArithDialect;
6598 const BuiltinDialect = dialects.BuiltinDialect;
6599 const FuncDialect = dialects.FuncDialect;
6600 const MemrefDialect = dialects.MemrefDialect;
6601 const ScfDialect = dialects.ScfDialect;
6602
6603 var arena = alloc_arena.Arena.init(std.testing.allocator);
6604 defer arena.deinit();
6605 const allocator = arena.allocator();
6606
6607 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
6608 defer ir_ctx.deinit(allocator);
6609 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
6610
6611 const loc = ir.Location.getUnknown();
6612 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
6613 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
6614 const index_type = try ArithDialect.getIndexType(&ir_ctx);
6615 const memref_f64 = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, f64_type, .host);
6616
6617 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
6618 const module_block = module.getBodyBlock();
6619
6620 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "for_dot", &.{ i64_type, memref_f64, memref_f64, memref_f64 }, &.{});
6621 try module_block.addOperation(func.op);
6622 const entry = func.getEntryBlock();
6623
6624 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
6625 try entry.addOperation(zero.op);
6626 var bound = try ArithDialect.CastOp.create(&ir_ctx, loc, func.getArgument(0), index_type);
6627 try entry.addOperation(bound.op);
6628 var step = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
6629 try entry.addOperation(step.op);
6630 var acc_init = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 0.0);
6631 try entry.addOperation(acc_init.op);
6632
6633 var for_op = try ScfDialect.ForOp.create(&ir_ctx, loc, zero.getResult(), bound.getResult(), step.getResult(), &.{acc_init.getResult()}, &.{f64_type});
6634 try entry.addOperation(for_op.op);
6635 const body = for_op.getBodyBlock();
6636 const iv = body.arguments.items[0];
6637 const acc = body.arguments.items[1];
6638
6639 var x = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(1), iv, f64_type);
6640 try body.addOperation(x.op);
6641 var y = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(2), iv, f64_type);
6642 try body.addOperation(y.op);
6643 var product = try ArithDialect.MulOp.create(&ir_ctx, loc, x.getResult(), y.getResult());
6644 try body.addOperation(product.op);
6645 var next = try ArithDialect.AddOp.create(&ir_ctx, loc, acc, product.getResult());
6646 try body.addOperation(next.op);
6647 const yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{next.getResult()});
6648 try body.addOperation(yield.op);
6649
6650 const total = for_op.op.getResult(0) orelse return error.TestFailure;
6651 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, total, func.getArgument(3), zero.getResult());
6652 try entry.addOperation(store.op);
6653 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{});
6654 try entry.addOperation(ret.op);
6655
6656 var backend = try Backend.init(allocator, &ir_ctx, .testing);
6657 defer backend.deinit();
6658 const compiled = try backend.compile(module.op);
6659
6660 const DotFn = *const fn (i64, [*]const f64, [*]const f64, [*]f64) callconv(.c) void;
6661 const kernel = try backend.runtime.getFunction(compiled, "for_dot", DotFn);
6662
6663 var xs: [64]f64 = undefined;
6664 var ys: [64]f64 = undefined;
6665 var expected: f64 = 0.0;
6666 for (&xs, &ys, 0..) |*xv, *yv, i| {
6667 xv.* = @floatFromInt(i);
6668 yv.* = 0.5;
6669 expected += @as(f64, @floatFromInt(i)) * 0.5;
6670 }
6671 var out: [1]f64 = .{-1.0};
6672 kernel(xs.len, &xs, &ys, &out);
6673 try testing.expectEqual(expected, out[0]);
6674 }
6675
6676 test "x86_64 float loop result survives trailing body ops" {
6677 if (!supports_x86_64_backend) return;
6678
6679 const testing = std.testing;
6680 const dialects = @import("../../dialects/root.zig");
6681 const ArithDialect = dialects.ArithDialect;
6682 const BuiltinDialect = dialects.BuiltinDialect;
6683 const FuncDialect = dialects.FuncDialect;
6684 const MemrefDialect = dialects.MemrefDialect;
6685 const ScfDialect = dialects.ScfDialect;
6686
6687 var arena = alloc_arena.Arena.init(std.testing.allocator);
6688 defer arena.deinit();
6689 const allocator = arena.allocator();
6690
6691 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
6692 defer ir_ctx.deinit(allocator);
6693 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
6694
6695 const loc = ir.Location.getUnknown();
6696 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
6697 const f64_type = try ArithDialect.getScalarType(&ir_ctx, .f64);
6698 const index_type = try ArithDialect.getIndexType(&ir_ctx);
6699 const memref_f64 = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, f64_type, .host);
6700
6701 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
6702 const module_block = module.getBodyBlock();
6703
6704 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "sum_and_squares", &.{ i64_type, memref_f64, memref_f64, memref_f64 }, &.{});
6705 try module_block.addOperation(func.op);
6706 const entry = func.getEntryBlock();
6707
6708 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
6709 try entry.addOperation(zero.op);
6710 var one_index = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
6711 try entry.addOperation(one_index.op);
6712 var bound = try ArithDialect.CastOp.create(&ir_ctx, loc, func.getArgument(0), index_type);
6713 try entry.addOperation(bound.op);
6714 var acc_init = try ArithDialect.ConstantOp.createFloat(&ir_ctx, loc, f64_type, 0.0);
6715 try entry.addOperation(acc_init.op);
6716
6717 var for_op = try ScfDialect.ForOp.create(&ir_ctx, loc, zero.getResult(), bound.getResult(), one_index.getResult(), &.{acc_init.getResult()}, &.{f64_type});
6718 try entry.addOperation(for_op.op);
6719 {
6720 const body = for_op.getBodyBlock();
6721 const iv = body.arguments.items[0];
6722 const acc = body.arguments.items[1];
6723 var x = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(1), iv, f64_type);
6724 try body.addOperation(x.op);
6725 var next = try ArithDialect.AddOp.create(&ir_ctx, loc, acc, x.getResult());
6726 try body.addOperation(next.op);
6727 var square = try ArithDialect.MulOp.create(&ir_ctx, loc, x.getResult(), x.getResult());
6728 try body.addOperation(square.op);
6729 const square_store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, square.getResult(), func.getArgument(2), iv);
6730 try body.addOperation(square_store.op);
6731 const yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{next.getResult()});
6732 try body.addOperation(yield.op);
6733 }
6734 const total = for_op.op.getResult(0) orelse return error.TestFailure;
6735
6736 const total_store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, total, func.getArgument(3), zero.getResult());
6737 try entry.addOperation(total_store.op);
6738 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{});
6739 try entry.addOperation(ret.op);
6740
6741 var backend = try Backend.init(allocator, &ir_ctx, .testing);
6742 defer backend.deinit();
6743 const compiled = try backend.compile(module.op);
6744
6745 const KernelFn = *const fn (i64, [*]const f64, [*]f64, [*]f64) callconv(.c) void;
6746 const kernel = try backend.runtime.getFunction(compiled, "sum_and_squares", KernelFn);
6747
6748 var xs: [16]f64 = undefined;
6749 var expected_sum: f64 = 0.0;
6750 for (&xs, 0..) |*xv, i| {
6751 xv.* = @floatFromInt(i + 1);
6752 expected_sum += xv.*;
6753 }
6754 var squares: [16]f64 = undefined;
6755 var out: [1]f64 = .{-1.0};
6756 kernel(xs.len, &xs, &squares, &out);
6757 try testing.expectEqual(expected_sum, out[0]);
6758 try testing.expectEqual(@as(f64, 256.0), squares[15]);
6759 }
6760
6761 test "x86_64 packed vec4xf32 loop keeps values in xmm registers" {
6762 if (!supports_x86_64_backend) return;
6763
6764 const dialects = @import("../../dialects/root.zig");
6765 const ArithDialect = dialects.ArithDialect;
6766 const BuiltinDialect = dialects.BuiltinDialect;
6767 const FuncDialect = dialects.FuncDialect;
6768 const MemrefDialect = dialects.MemrefDialect;
6769 const ScfDialect = dialects.ScfDialect;
6770
6771 var arena = alloc_arena.Arena.init(std.testing.allocator);
6772 defer arena.deinit();
6773 const allocator = arena.allocator();
6774
6775 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
6776 defer ir_ctx.deinit(allocator);
6777 try @import("../../dialects/root.zig").registerAllDialects(&ir_ctx);
6778
6779 const loc = ir.Location.getUnknown();
6780 const i64_type = try ArithDialect.getScalarType(&ir_ctx, .i64);
6781 const f32_type = try ArithDialect.getScalarType(&ir_ctx, .f32);
6782 const index_type = try ArithDialect.getIndexType(&ir_ctx);
6783 const vec_type = try ir_ctx.getDialectTypeFromName("arith.vec4xf32");
6784 const memref_f32 = try MemrefDialect.getMemrefTypeDynamic(&ir_ctx, f32_type, .host);
6785
6786 const module = try BuiltinDialect.ModuleOp.create(&ir_ctx, loc);
6787 const module_block = module.getBodyBlock();
6788
6789 var func = try FuncDialect.FuncOp.create(&ir_ctx, loc, "vaffine", &.{ i64_type, f32_type, memref_f32, memref_f32, memref_f32 }, &.{});
6790 try module_block.addOperation(func.op);
6791 const entry = func.getEntryBlock();
6792
6793 var zero = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 0);
6794 try entry.addOperation(zero.op);
6795 var four = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 4);
6796 try entry.addOperation(four.op);
6797 var bound = try ArithDialect.CastOp.create(&ir_ctx, loc, func.getArgument(0), index_type);
6798 try entry.addOperation(bound.op);
6799
6800 var remainder = try ArithDialect.RemOp.create(&ir_ctx, loc, bound.getResult(), four.getResult());
6801 try entry.addOperation(remainder.op);
6802 var vector_upper = try ArithDialect.SubOp.create(&ir_ctx, loc, bound.getResult(), remainder.getResult());
6803 try entry.addOperation(vector_upper.op);
6804 var one = try ArithDialect.ConstantOp.createInt(&ir_ctx, loc, index_type, 1);
6805 try entry.addOperation(one.op);
6806
6807 var for_op = try ScfDialect.ForOp.create(&ir_ctx, loc, zero.getResult(), vector_upper.getResult(), four.getResult(), &.{}, &.{});
6808 try entry.addOperation(for_op.op);
6809 const body = for_op.getBodyBlock();
6810 const iv = body.arguments.items[0];
6811
6812 var scale_vec = try ArithDialect.SplatOp.create(&ir_ctx, loc, func.getArgument(1), vec_type);
6813 try body.addOperation(scale_vec.op);
6814 var lhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(2), iv, vec_type);
6815 try body.addOperation(lhs.op);
6816 var rhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(3), iv, vec_type);
6817 try body.addOperation(rhs.op);
6818 var scaled = try ArithDialect.MulOp.create(&ir_ctx, loc, scale_vec.getResult(), lhs.getResult());
6819 try body.addOperation(scaled.op);
6820 var sum = try ArithDialect.AddOp.create(&ir_ctx, loc, scaled.getResult(), rhs.getResult());
6821 try body.addOperation(sum.op);
6822 const store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, sum.getResult(), func.getArgument(4), iv);
6823 try body.addOperation(store.op);
6824 const yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{});
6825 try body.addOperation(yield.op);
6826
6827 var tail_op = try ScfDialect.ForOp.create(&ir_ctx, loc, vector_upper.getResult(), bound.getResult(), one.getResult(), &.{}, &.{});
6828 try entry.addOperation(tail_op.op);
6829 const tail = tail_op.getBodyBlock();
6830 const tail_iv = tail.arguments.items[0];
6831 var tail_lhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(2), tail_iv, f32_type);
6832 try tail.addOperation(tail_lhs.op);
6833 var tail_rhs = try MemrefDialect.LoadOp.create(&ir_ctx, loc, func.getArgument(3), tail_iv, f32_type);
6834 try tail.addOperation(tail_rhs.op);
6835 var tail_scaled = try ArithDialect.MulOp.create(&ir_ctx, loc, func.getArgument(1), tail_lhs.getResult());
6836 try tail.addOperation(tail_scaled.op);
6837 var tail_sum = try ArithDialect.AddOp.create(&ir_ctx, loc, tail_scaled.getResult(), tail_rhs.getResult());
6838 try tail.addOperation(tail_sum.op);
6839 const tail_store = try MemrefDialect.StoreOp.create(&ir_ctx, loc, tail_sum.getResult(), func.getArgument(4), tail_iv);
6840 try tail.addOperation(tail_store.op);
6841 const tail_yield = try ScfDialect.YieldOp.create(&ir_ctx, loc, &.{});
6842 try tail.addOperation(tail_yield.op);
6843
6844 const ret = try FuncDialect.ReturnOp.create(&ir_ctx, loc, &.{});
6845 try entry.addOperation(ret.op);
6846
6847 var emitter = x86_64.Emitter.init(allocator);
6848 defer emitter.deinit();
6849 try emitter.emitFunction(func.op);
6850 var packed_ranges: usize = 0;
6851 for (emitter.xmm_location_ranges.items) |range| {
6852 if (emitter.vector_slot_map.get(range.value) != null) packed_ranges += 1;
6853 }
6854 try std.testing.expect(packed_ranges > 0);
6855
6856 var backend = try Backend.init(allocator, &ir_ctx, .testing);
6857 defer backend.deinit();
6858 const compiled = try backend.compile(module.op);
6859
6860 const AffineFn = *const fn (i64, f32, [*]const f32, [*]const f32, [*]f32) callconv(.c) void;
6861 const kernel = try backend.runtime.getFunction(compiled, "vaffine", AffineFn);
6862
6863 var lhs_buf: [19]f32 = undefined;
6864 var rhs_buf: [19]f32 = undefined;
6865 for (&lhs_buf, &rhs_buf, 0..) |*a, *b, i| {
6866 a.* = @floatFromInt(i + 1);
6867 b.* = @floatFromInt(100 + i);
6868 }
6869 var out: [19]f32 = @splat(-1.0);
6870 kernel(out.len, 2.0, &lhs_buf, &rhs_buf, &out);
6871 for (out, 0..) |value, i| {
6872 const expected = @as(f32, @floatFromInt(i + 1)) * 2.0 + @as(f32, @floatFromInt(100 + i));
6873 try std.testing.expectEqual(expected, value);
6874 }
6875 }
6876
6877 test "x86_64 byte arguments widen through a stored unsigned result" {
6878 if (!supports_x86_64_backend) return;
6879 const dialects = @import("../../dialects/root.zig");
6880 const arith = dialects.ArithDialect;
6881 const memref = dialects.MemrefDialect;
6882 const allocator = std.testing.allocator;
6883 var context = try ir.Context.init(allocator, ir.Context.Limits.testing);
6884 defer context.deinit(allocator);
6885 try dialects.registerAllDialects(&context);
6886 const location = ir.Location.getUnknown();
6887 const byte = try arith.getScalarType(&context, .u8);
6888 const word = try arith.getScalarType(&context, .i64);
6889 const index_type = try arith.getScalarType(&context, .index);
6890 const buffer = try memref.getMemrefType1D(&context, 1, byte, .host);
6891 const module = try dialects.BuiltinDialect.ModuleOp.create(&context, location);
6892 var function = try dialects.FuncDialect.FuncOp.create(
6893 &context,
6894 location,
6895 "byte_sum",
6896 &.{ byte, byte },
6897 &.{word},
6898 );
6899 try module.getBodyBlock().addOperation(function.op);
6900 const block = function.getEntryBlock();
6901 var cells = try memref.AllocaOp.createStatic(&context, location, buffer);
6902 try block.addOperation(cells.op);
6903 var zero = try arith.ConstantOp.createInt(&context, location, index_type, 0);
6904 try block.addOperation(zero.op);
6905 var sum = try arith.AddOp.create(
6906 &context,
6907 location,
6908 function.getArgument(0),
6909 function.getArgument(1),
6910 );
6911 try block.addOperation(sum.op);
6912 const store = try memref.StoreOp.create(
6913 &context,
6914 location,
6915 sum.getResult(),
6916 cells.getResult(),
6917 zero.getResult(),
6918 );
6919 try block.addOperation(store.op);
6920 var loaded = try memref.LoadOp.create(
6921 &context,
6922 location,
6923 cells.getResult(),
6924 zero.getResult(),
6925 byte,
6926 );
6927 try block.addOperation(loaded.op);
6928 var widened = try arith.CastOp.create(&context, location, loaded.getResult(), word);
6929 try block.addOperation(widened.op);
6930 const ret = try dialects.FuncDialect.ReturnOp.create(
6931 &context,
6932 location,
6933 &.{widened.getResult()},
6934 );
6935 try block.addOperation(ret.op);
6936 var backend = try Backend.init(allocator, &context, .testing);
6937 defer backend.deinit();
6938 const compiled = try backend.compile(module.op);
6939 const ByteSum = *const fn (u8, u8) callconv(.c) i64;
6940 const call = try backend.runtime.getFunction(compiled, "byte_sum", ByteSum);
6941 for (0..256) |left| for (0..256) |right| {
6942 const expected: i64 = @intCast((left + right) % 256);
6943 try std.testing.expectEqual(expected, call(@intCast(left), @intCast(right)));
6944 };
6945 }
6946
6947 test "x86_64 compare and swap loop increments a word one hundred times" {
6948 if (!supports_x86_64_backend) return;
6949 const dialects = @import("../../dialects/root.zig");
6950 const arith = dialects.ArithDialect;
6951 const memref = dialects.MemrefDialect;
6952 const scf = dialects.ScfDialect;
6953 const allocator = std.testing.allocator;
6954 var context = try ir.Context.init(allocator, ir.Context.Limits.testing);
6955 defer context.deinit(allocator);
6956 try dialects.registerAllDialects(&context);
6957 const location = ir.Location.getUnknown();
6958 const word = try arith.getScalarType(&context, .i64);
6959 const index_type = try arith.getScalarType(&context, .index);
6960 const buffer = try memref.getMemrefType1D(&context, 1, word, .host);
6961 const module = try dialects.BuiltinDialect.ModuleOp.create(&context, location);
6962 var function = try dialects.FuncDialect.FuncOp.create(
6963 &context,
6964 location,
6965 "cas_count",
6966 &.{},
6967 &.{word},
6968 );
6969 try module.getBodyBlock().addOperation(function.op);
6970 const entry = function.getEntryBlock();
6971 const cells = try memref.AllocaOp.createStatic(&context, location, buffer);
6972 try entry.addOperation(cells.op);
6973 const first = try arith.ConstantOp.createInt(&context, location, index_type, 0);
6974 try entry.addOperation(first.op);
6975 const start = try arith.ConstantOp.createInt(&context, location, word, 0);
6976 try entry.addOperation(start.op);
6977 const clear = try memref.AtomicStoreOp.create(
6978 &context,
6979 location,
6980 start.getResult(),
6981 cells.getResult(),
6982 first.getResult(),
6983 .release,
6984 );
6985 try entry.addOperation(clear.op);
6986
6987 const loop = try scf.WhileOp.create(&context, location, &.{start.getResult()}, &.{word});
6988 try entry.addOperation(loop.op);
6989 const before = loop.getBeforeBlock();
6990 const limit = try arith.ConstantOp.createInt(&context, location, word, 100);
6991 try before.addOperation(limit.op);
6992 const below = try arith.CmpOp.create(
6993 &context,
6994 location,
6995 .lt,
6996 before.arguments.items[0],
6997 limit.getResult(),
6998 );
6999 try before.addOperation(below.op);
7000 const condition = try scf.ConditionOp.create(
7001 &context,
7002 location,
7003 below.getResult(),
7004 &.{before.arguments.items[0]},
7005 );
7006 try before.addOperation(condition.op);
7007
7008 const after = loop.getAfterBlock();
7009 const one = try arith.ConstantOp.createInt(&context, location, word, 1);
7010 try after.addOperation(one.op);
7011 const none = try arith.ConstantOp.createInt(&context, location, word, 0);
7012 try after.addOperation(none.op);
7013 const seen = try memref.AtomicLoadOp.create(
7014 &context,
7015 location,
7016 cells.getResult(),
7017 first.getResult(),
7018 word,
7019 .acquire,
7020 );
7021 try after.addOperation(seen.op);
7022 const bumped = try arith.AddOp.create(&context, location, seen.getResult(), one.getResult());
7023 try after.addOperation(bumped.op);
7024 const exchanged = try memref.AtomicCasOp.createOrdered(
7025 &context,
7026 location,
7027 seen.getResult(),
7028 bumped.getResult(),
7029 cells.getResult(),
7030 first.getResult(),
7031 word,
7032 .acq_rel,
7033 );
7034 try after.addOperation(exchanged.op);
7035 const won = try arith.CmpOp.create(
7036 &context,
7037 location,
7038 .eq,
7039 exchanged.getResult(),
7040 seen.getResult(),
7041 );
7042 try after.addOperation(won.op);
7043 const step = try arith.SelectOp.create(
7044 &context,
7045 location,
7046 won.getResult(),
7047 one.getResult(),
7048 none.getResult(),
7049 );
7050 try after.addOperation(step.op);
7051 const advanced = try arith.AddOp.create(
7052 &context,
7053 location,
7054 after.arguments.items[0],
7055 step.getResult(),
7056 );
7057 try after.addOperation(advanced.op);
7058 const yield = try scf.YieldOp.create(&context, location, &.{advanced.getResult()});
7059 try after.addOperation(yield.op);
7060
7061 const total = try memref.AtomicLoadOp.create(
7062 &context,
7063 location,
7064 cells.getResult(),
7065 first.getResult(),
7066 word,
7067 .seq_cst,
7068 );
7069 try entry.addOperation(total.op);
7070 const ret = try dialects.FuncDialect.ReturnOp.create(&context, location, &.{total.getResult()});
7071 try entry.addOperation(ret.op);
7072
7073 var emitter = x86_64.Emitter.init(allocator);
7074 defer emitter.deinit();
7075 try emitter.emitFunction(function.op);
7076 try std.testing.expectEqual(@as(u32, 0), emitter.value_locations.count());
7077 try std.testing.expectEqual(@as(usize, 0), emitter.value_location_ranges.items.len);
7078
7079 var backend = try Backend.init(allocator, &context, .testing);
7080 defer backend.deinit();
7081 const compiled = try backend.compile(module.op);
7082 const answer = try onlyResult(&backend, compiled, "cas_count", &.{});
7083 try std.testing.expectEqual(@as(i64, 100), try integerOf(answer));
7084 }
7085
7086 test "x86_64 half word compare and swap answers what it saw and writes on a match" {
7087 if (!supports_x86_64_backend) return;
7088 const dialects = @import("../../dialects/root.zig");
7089 const arith = dialects.ArithDialect;
7090 const memref = dialects.MemrefDialect;
7091 const allocator = std.testing.allocator;
7092 var context = try ir.Context.init(allocator, ir.Context.Limits.testing);
7093 defer context.deinit(allocator);
7094 try dialects.registerAllDialects(&context);
7095 const location = ir.Location.getUnknown();
7096 const half = try arith.getScalarType(&context, .i32);
7097 const index_type = try arith.getScalarType(&context, .index);
7098 const buffer = try memref.getMemrefType1D(&context, 2, half, .host);
7099 const module = try dialects.BuiltinDialect.ModuleOp.create(&context, location);
7100 var function = try dialects.FuncDialect.FuncOp.create(
7101 &context,
7102 location,
7103 "cas_half",
7104 &.{},
7105 &.{half},
7106 );
7107 try module.getBodyBlock().addOperation(function.op);
7108 const block = function.getEntryBlock();
7109 const cells = try memref.AllocaOp.createStatic(&context, location, buffer);
7110 try block.addOperation(cells.op);
7111 const second = try arith.ConstantOp.createInt(&context, location, index_type, 1);
7112 try block.addOperation(second.op);
7113 var ints: [7]*ir.Value = undefined;
7114 for (&ints, [_]i64{ 7, 3, 9, 11, 1, 0, 10 }) |*constant, value| {
7115 const made = try arith.ConstantOp.createInt(&context, location, half, value);
7116 try block.addOperation(made.op);
7117 constant.* = made.getResult();
7118 }
7119 const seven, const three, const nine, const eleven, const one, const zero, const ten = ints;
7120 const hundred = try arith.MulOp.create(&context, location, ten, ten);
7121 try block.addOperation(hundred.op);
7122
7123 const published = try memref.AtomicStoreOp.create(
7124 &context,
7125 location,
7126 seven,
7127 cells.getResult(),
7128 second.getResult(),
7129 .seq_cst,
7130 );
7131 try block.addOperation(published.op);
7132 const missed = try memref.AtomicCasOp.create(
7133 &context,
7134 location,
7135 three,
7136 nine,
7137 cells.getResult(),
7138 second.getResult(),
7139 half,
7140 );
7141 try block.addOperation(missed.op);
7142 const kept = try memref.AtomicLoadOp.create(
7143 &context,
7144 location,
7145 cells.getResult(),
7146 second.getResult(),
7147 half,
7148 .acquire,
7149 );
7150 try block.addOperation(kept.op);
7151 const matched = try memref.AtomicCasOp.createOrdered(
7152 &context,
7153 location,
7154 seven,
7155 eleven,
7156 cells.getResult(),
7157 second.getResult(),
7158 half,
7159 .release,
7160 );
7161 try block.addOperation(matched.op);
7162 const fence = try memref.FenceOp.create(&context, location, .system, .seq_cst);
7163 try block.addOperation(fence.op);
7164 const written = try memref.AtomicLoadOp.create(
7165 &context,
7166 location,
7167 cells.getResult(),
7168 second.getResult(),
7169 half,
7170 .seq_cst,
7171 );
7172 try block.addOperation(written.op);
7173
7174 const miss_flag = try arith.CmpOp.create(&context, location, .eq, missed.getResult(), three);
7175 try block.addOperation(miss_flag.op);
7176 const miss_bit = try arith.SelectOp.create(
7177 &context,
7178 location,
7179 miss_flag.getResult(),
7180 one,
7181 zero,
7182 );
7183 try block.addOperation(miss_bit.op);
7184 const match_flag = try arith.CmpOp.create(&context, location, .eq, matched.getResult(), seven);
7185 try block.addOperation(match_flag.op);
7186 const match_bit = try arith.SelectOp.create(
7187 &context,
7188 location,
7189 match_flag.getResult(),
7190 one,
7191 zero,
7192 );
7193 try block.addOperation(match_bit.op);
7194
7195 var digits = [_]*ir.Value{
7196 missed.getResult(),
7197 kept.getResult(),
7198 matched.getResult(),
7199 written.getResult(),
7200 };
7201 var packed_value = miss_bit.getResult();
7202 for (&digits) |digit| {
7203 const shifted = try arith.MulOp.create(
7204 &context,
7205 location,
7206 packed_value,
7207 hundred.getResult(),
7208 );
7209 try block.addOperation(shifted.op);
7210 const added = try arith.AddOp.create(&context, location, shifted.getResult(), digit);
7211 try block.addOperation(added.op);
7212 packed_value = added.getResult();
7213 }
7214 const flagged = try arith.MulOp.create(&context, location, packed_value, ten);
7215 try block.addOperation(flagged.op);
7216 const answer_value = try arith.AddOp.create(
7217 &context,
7218 location,
7219 flagged.getResult(),
7220 match_bit.getResult(),
7221 );
7222 try block.addOperation(answer_value.op);
7223 const ret = try dialects.FuncDialect.ReturnOp.create(
7224 &context,
7225 location,
7226 &.{answer_value.getResult()},
7227 );
7228 try block.addOperation(ret.op);
7229
7230 var backend = try Backend.init(allocator, &context, .testing);
7231 defer backend.deinit();
7232 const compiled = try backend.compile(module.op);
7233 const answer = try onlyResult(&backend, compiled, "cas_half", &.{});
7234 try std.testing.expectEqual(@as(i64, 70_707_111), try integerOf(answer));
7235 }
7236
7237 test "x86_64 atomic operations lower to the instructions their orderings need" {
7238 if (!supports_x86_64_backend) return;
7239 const dialects = @import("../../dialects/root.zig");
7240 const arith = dialects.ArithDialect;
7241 const memref = dialects.MemrefDialect;
7242 const allocator = std.testing.allocator;
7243 var context = try ir.Context.init(allocator, ir.Context.Limits.testing);
7244 defer context.deinit(allocator);
7245 try dialects.registerAllDialects(&context);
7246 const location = ir.Location.getUnknown();
7247 const word = try arith.getScalarType(&context, .i64);
7248 const half = try arith.getScalarType(&context, .i32);
7249 const index_type = try arith.getScalarType(&context, .index);
7250 const module = try dialects.BuiltinDialect.ModuleOp.create(&context, location);
7251 var function = try dialects.FuncDialect.FuncOp.create(
7252 &context,
7253 location,
7254 "atomic_windows",
7255 &.{},
7256 &.{},
7257 );
7258 try module.getBodyBlock().addOperation(function.op);
7259 const block = function.getEntryBlock();
7260 const words = try memref.AllocaOp.createStatic(
7261 &context,
7262 location,
7263 try memref.getMemrefType1D(&context, 2, word, .host),
7264 );
7265 try block.addOperation(words.op);
7266 const halves = try memref.AllocaOp.createStatic(
7267 &context,
7268 location,
7269 try memref.getMemrefType1D(&context, 2, half, .host),
7270 );
7271 try block.addOperation(halves.op);
7272 const slot = try arith.ConstantOp.createInt(&context, location, index_type, 1);
7273 try block.addOperation(slot.op);
7274 const w = try arith.ConstantOp.createInt(&context, location, word, 5);
7275 try block.addOperation(w.op);
7276 const h = try arith.ConstantOp.createInt(&context, location, half, 5);
7277 try block.addOperation(h.op);
7278
7279 const operations = [_]*ir.Operation{
7280 (try memref.AtomicLoadOp.create(
7281 &context,
7282 location,
7283 words.getResult(),
7284 slot.getResult(),
7285 word,
7286 .acquire,
7287 )).op,
7288 (try memref.AtomicLoadOp.create(
7289 &context,
7290 location,
7291 halves.getResult(),
7292 slot.getResult(),
7293 half,
7294 .seq_cst,
7295 )).op,
7296 (try memref.AtomicStoreOp.create(
7297 &context,
7298 location,
7299 w.getResult(),
7300 words.getResult(),
7301 slot.getResult(),
7302 .release,
7303 )).op,
7304 (try memref.AtomicStoreOp.create(
7305 &context,
7306 location,
7307 w.getResult(),
7308 words.getResult(),
7309 slot.getResult(),
7310 .seq_cst,
7311 )).op,
7312 (try memref.AtomicStoreOp.create(
7313 &context,
7314 location,
7315 h.getResult(),
7316 halves.getResult(),
7317 slot.getResult(),
7318 .release,
7319 )).op,
7320 (try memref.AtomicStoreOp.create(
7321 &context,
7322 location,
7323 h.getResult(),
7324 halves.getResult(),
7325 slot.getResult(),
7326 .seq_cst,
7327 )).op,
7328 (try memref.AtomicCasOp.create(
7329 &context,
7330 location,
7331 w.getResult(),
7332 w.getResult(),
7333 words.getResult(),
7334 slot.getResult(),
7335 word,
7336 )).op,
7337 (try memref.AtomicCasOp.createOrdered(
7338 &context,
7339 location,
7340 h.getResult(),
7341 h.getResult(),
7342 halves.getResult(),
7343 slot.getResult(),
7344 half,
7345 .acquire,
7346 )).op,
7347 (try memref.FenceOp.create(&context, location, .device, .acquire)).op,
7348 (try memref.FenceOp.create(&context, location, .workgroup, .release)).op,
7349 (try memref.FenceOp.create(&context, location, .system, .acq_rel)).op,
7350 (try memref.FenceOp.create(&context, location, .device, .seq_cst)).op,
7351 };
7352 for (operations) |op| try block.addOperation(op);
7353 const ret = try dialects.FuncDialect.ReturnOp.create(&context, location, &.{});
7354 try block.addOperation(ret.op);
7355
7356 var emitter = x86_64.Emitter.init(allocator);
7357 defer emitter.deinit();
7358 try emitter.emitFunction(function.op);
7359 const code = emitter.code.items;
7360 const windows = [_][]const u8{
7361 &.{ 0x48, 0x8B, 0x04, 0xD1 },
7362 &.{ 0x8B, 0x04, 0x91 },
7363 &.{ 0x4C, 0x89, 0x04, 0xD1 },
7364 &.{ 0x4C, 0x87, 0x04, 0xD1 },
7365 &.{ 0x44, 0x89, 0x04, 0x91 },
7366 &.{ 0x44, 0x87, 0x04, 0x91 },
7367 &.{ 0xF0, 0x4C, 0x0F, 0xB1, 0x04, 0xD1 },
7368 &.{ 0xF0, 0x44, 0x0F, 0xB1, 0x04, 0x91 },
7369 };
7370 var cursor: usize = 0;
7371 for (windows) |window| {
7372 const found = std.mem.indexOfPos(u8, code, cursor, window) orelse
7373 return error.TestExpectedEqual;
7374 cursor = found + window.len;
7375 }
7376 const mfence = [_]u8{ 0x0F, 0xAE, 0xF0 };
7377 try std.testing.expectEqual(@as(usize, 1), std.mem.count(u8, code, &mfence));
7378 try std.testing.expect(std.mem.indexOfPos(u8, code, cursor, &mfence) != null);
7379 }
7380
7381 test "x86_64 scf while condition survives an operation between the compare and the condition" {
7382 if (!supports_x86_64_backend) return;
7383 const dialects = @import("../../dialects/root.zig");
7384 const arith = dialects.ArithDialect;
7385 const scf = dialects.ScfDialect;
7386 const allocator = std.testing.allocator;
7387 var context = try ir.Context.init(allocator, ir.Context.Limits.testing);
7388 defer context.deinit(allocator);
7389 try dialects.registerAllDialects(&context);
7390 const location = ir.Location.getUnknown();
7391 const word = try arith.getScalarType(&context, .i64);
7392 const module = try dialects.BuiltinDialect.ModuleOp.create(&context, location);
7393 var function = try dialects.FuncDialect.FuncOp.create(
7394 &context,
7395 location,
7396 "separated",
7397 &.{},
7398 &.{word},
7399 );
7400 try module.getBodyBlock().addOperation(function.op);
7401 const entry = function.getEntryBlock();
7402 const start = try arith.ConstantOp.createInt(&context, location, word, 0);
7403 try entry.addOperation(start.op);
7404 const loop = try scf.WhileOp.create(&context, location, &.{start.getResult()}, &.{word});
7405 try entry.addOperation(loop.op);
7406 const before = loop.getBeforeBlock();
7407 const limit = try arith.ConstantOp.createInt(&context, location, word, 5);
7408 try before.addOperation(limit.op);
7409 const below = try arith.CmpOp.create(
7410 &context,
7411 location,
7412 .lt,
7413 before.arguments.items[0],
7414 limit.getResult(),
7415 );
7416 try before.addOperation(below.op);
7417 const one = try arith.ConstantOp.createInt(&context, location, word, 1);
7418 try before.addOperation(one.op);
7419 const next = try arith.AddOp.create(
7420 &context,
7421 location,
7422 before.arguments.items[0],
7423 one.getResult(),
7424 );
7425 try before.addOperation(next.op);
7426 const condition = try scf.ConditionOp.create(
7427 &context,
7428 location,
7429 below.getResult(),
7430 &.{next.getResult()},
7431 );
7432 try before.addOperation(condition.op);
7433 const after = loop.getAfterBlock();
7434 const yield = try scf.YieldOp.create(&context, location, &.{after.arguments.items[0]});
7435 try after.addOperation(yield.op);
7436 const ret = try dialects.FuncDialect.ReturnOp.create(
7437 &context,
7438 location,
7439 &.{loop.op.getResult(0).?},
7440 );
7441 try entry.addOperation(ret.op);
7442 var backend = try Backend.init(allocator, &context, .testing);
7443 defer backend.deinit();
7444 const compiled = try backend.compile(module.op);
7445 const answer = try onlyResult(&backend, compiled, "separated", &.{});
7446 try std.testing.expectEqual(@as(i64, 6), try integerOf(answer));
7447 }
7448
7449 const BranchShape = enum { fused, separated };
7450
7451 /// Builds `name(a, b) = if a <pred> b { 1 } else { 0 }`. The separated shape puts an addition
7452 /// between the comparison and the branch, which writes the flags, so the comparison has to keep
7453 /// its materialized boolean.
7454 fn buildPredicateBranch(
7455 context: *ir.Context,
7456 module: @import("../../dialects/root.zig").BuiltinDialect.ModuleOp,
7457 name: []const u8,
7458 predicate: @import("../../dialects/root.zig").arith.CmpPredicate,
7459 shape: BranchShape,
7460 ) !*ir.Operation {
7461 const dialects = @import("../../dialects/root.zig");
7462 const arith = dialects.ArithDialect;
7463 const scf = dialects.ScfDialect;
7464 const location = ir.Location.getUnknown();
7465 const word = try arith.getScalarType(context, .i64);
7466 var function = try dialects.FuncDialect.FuncOp.create(
7467 context,
7468 location,
7469 name,
7470 &.{ word, word },
7471 &.{word},
7472 );
7473 try module.getBodyBlock().addOperation(function.op);
7474 const entry = function.getEntryBlock();
7475 const compare = try arith.CmpOp.create(
7476 context,
7477 location,
7478 predicate,
7479 function.getArgument(0),
7480 function.getArgument(1),
7481 );
7482 try entry.addOperation(compare.op);
7483 var carried = function.getArgument(0);
7484 if (shape == .separated) {
7485 const sum = try arith.AddOp.create(
7486 context,
7487 location,
7488 function.getArgument(0),
7489 function.getArgument(1),
7490 );
7491 try entry.addOperation(sum.op);
7492 carried = sum.getResult();
7493 }
7494 const branch = try scf.IfOp.create(context, location, compare.getResult(), &.{word});
7495 try entry.addOperation(branch.op);
7496 const yes = try arith.ConstantOp.createInt(context, location, word, 1);
7497 try branch.getThenBlock().addOperation(yes.op);
7498 try branch.getThenBlock().addOperation(
7499 (try scf.YieldOp.create(context, location, &.{yes.getResult()})).op,
7500 );
7501 const no = try arith.ConstantOp.createInt(context, location, word, 0);
7502 try branch.getElseBlock().?.addOperation(no.op);
7503 try branch.getElseBlock().?.addOperation(
7504 (try scf.YieldOp.create(context, location, &.{no.getResult()})).op,
7505 );
7506 const answer = try arith.AddOp.create(context, location, branch.op.getResult(0).?, carried);
7507 try entry.addOperation(answer.op);
7508 const difference = try arith.SubOp.create(context, location, answer.getResult(), carried);
7509 try entry.addOperation(difference.op);
7510 const returned = [_]*ir.Value{difference.getResult()};
7511 const ret = try dialects.FuncDialect.ReturnOp.create(context, location, &returned);
7512 try entry.addOperation(ret.op);
7513 return function.op;
7514 }
7515
7516 /// Whether `code` holds a `setcc` (0F 90 through 0F 9F), the instruction that turns flags into
7517 /// a boolean. A byte scan can see one inside an immediate, so a caller asserting absence must own
7518 /// code without such immediates.
7519 fn containsSetcc(code: []const u8) bool {
7520 var at: usize = 0;
7521 while (std.mem.indexOfScalarPos(u8, code, at, 0x0F)) |found| : (at = found + 1) {
7522 if (found + 1 == code.len) return false;
7523 if (code[found + 1] >= 0x90 and code[found + 1] <= 0x9F) return true;
7524 }
7525 return false;
7526 }
7527
7528 fn predicateHolds(
7529 predicate: @import("../../dialects/root.zig").arith.CmpPredicate,
7530 a: i64,
7531 b: i64,
7532 ) bool {
7533 const ua: u64 = @bitCast(a);
7534 const ub: u64 = @bitCast(b);
7535 return switch (predicate) {
7536 .eq => a == b,
7537 .ne => a != b,
7538 .lt, .slt => a < b,
7539 .le, .sle => a <= b,
7540 .gt, .sgt => a > b,
7541 .ge, .sge => a >= b,
7542 .ult => ua < ub,
7543 .ule => ua <= ub,
7544 .ugt => ua > ub,
7545 .uge => ua >= ub,
7546 };
7547 }
7548
7549 test "x86_64 a comparison only a branch reads decides from the flags what its boolean did" {
7550 if (!supports_x86_64_backend) return;
7551 const dialects = @import("../../dialects/root.zig");
7552 const allocator = std.testing.allocator;
7553 var context = try ir.Context.init(allocator, ir.Context.Limits.testing);
7554 defer context.deinit(allocator);
7555 try dialects.registerAllDialects(&context);
7556 const module = try dialects.BuiltinDialect.ModuleOp.create(&context, ir.Location.getUnknown());
7557 const predicates = std.enums.values(dialects.arith.CmpPredicate);
7558 var names: [2][16][24]u8 = undefined;
7559 var name_lengths: [2][16]usize = undefined;
7560 std.debug.assert(predicates.len <= 16);
7561
7562 for (predicates, 0..) |predicate, index| {
7563 for ([_]BranchShape{ .fused, .separated }, 0..) |shape, shape_index| {
7564 const name = try std.fmt.bufPrint(
7565 &names[shape_index][index],
7566 "{s}_{s}",
7567 .{ @tagName(shape), @tagName(predicate) },
7568 );
7569 name_lengths[shape_index][index] = name.len;
7570 const function = try buildPredicateBranch(&context, module, name, predicate, shape);
7571 var emitter = x86_64.Emitter.init(allocator);
7572 defer emitter.deinit();
7573 try emitter.emitFunction(function);
7574 const code = emitter.code.items;
7575 const compare_result = function.getRegion(0).?.getEntryBlock().?.operations.head.?;
7576 const compare_op: *ir.Operation = @ptrCast(@alignCast(compare_result));
7577 const materialized = containsSetcc(code);
7578 switch (shape) {
7579 .fused => {
7580 try std.testing.expect(!materialized);
7581 try std.testing.expect(!emitter.slot_map.contains(compare_op.getResult(0).?));
7582 },
7583 .separated => try std.testing.expect(materialized),
7584 }
7585 }
7586 }
7587
7588 var backend = try Backend.init(allocator, &context, .testing);
7589 defer backend.deinit();
7590 const compiled = try backend.compile(module.op);
7591 const samples = [_]i64{ std.math.minInt(i64), -7, -1, 0, 1, 7, std.math.maxInt(i64) };
7592 for (predicates, 0..) |predicate, index| {
7593 for (0..2) |shape_index| {
7594 const name = names[shape_index][index][0..name_lengths[shape_index][index]];
7595 for (samples) |a| for (samples) |b| {
7596 const arguments = [_]invoke.Value{ .{ .i64 = a }, .{ .i64 = b } };
7597 const answer = try onlyResult(&backend, compiled, name, &arguments);
7598 const expected: i64 = @intFromBool(predicateHolds(predicate, a, b));
7599 try std.testing.expectEqual(expected, try integerOf(answer));
7600 };
7601 }
7602 }
7603 }
7604
7605 fn onlyResult(
7606 backend: *const Backend,
7607 compiled: x86_64.jit.ModuleHandle,
7608 name: []const u8,
7609 args: []const invoke.Value,
7610 ) !invoke.Value {
7611 var results: [1]invoke.Value = undefined;
7612 try backend.runtime.call(compiled, name, args, &results);
7613 return results[0];
7614 }
7615
7616 fn vectorResult(
7617 backend: *const Backend,
7618 compiled: x86_64.jit.ModuleHandle,
7619 name: []const u8,
7620 ) !invoke.Vector {
7621 return switch (try onlyResult(backend, compiled, name, &.{})) {
7622 .vector => |vector| vector,
7623 else => error.TestFailure,
7624 };
7625 }
7626
7627 fn integerOf(value: invoke.Value) !i64 {
7628 return switch (value) {
7629 inline .i8, .i16, .i32, .i64, .u8, .u16, .u32 => |integer| integer,
7630 .u64, .index => |word| @bitCast(word),
7631 .bool => |flag| @intFromBool(flag),
7632 else => error.NoIntResult,
7633 };
7634 }
7635
7636 fn floatOf(value: invoke.Value) !f64 {
7637 return switch (value) {
7638 inline .f32, .f64 => |float| float,
7639 else => error.NoFloatResult,
7640 };
7641 }
7642
7643 /// Builds a module whose one function stores through a `memref.subview`, which
7644 /// is the shape `legalizeMemrefViews` rewrites and therefore allocates for: it
7645 /// parses the subview's shape and stride out of its layout attributes.
7646 fn buildViewModule(ctx: *ir.Context) !*ir.Operation {
7647 const dialects = @import("../../dialects/root.zig");
7648 const loc = ir.Location.getUnknown();
7649
7650 const i64_type = try dialects.ArithDialect.getScalarType(ctx, .i64);
7651 const index_type = try dialects.ArithDialect.getIndexType(ctx);
7652 const source_type = try dialects.MemrefDialect.getMemrefType1D(ctx, 8, i64_type, .host);
7653 const view_type = try dialects.MemrefDialect.getMemrefType1D(ctx, 4, i64_type, .host);
7654
7655 const module = try dialects.BuiltinDialect.ModuleOp.create(ctx, loc);
7656 const module_block = module.getBodyBlock();
7657 var func = try dialects.FuncDialect.FuncOp.create(ctx, loc, "through_view", &.{i64_type}, &.{});
7658 try module_block.addOperation(func.op);
7659
7660 const entry = func.getEntryBlock();
7661 var zero = try dialects.ArithDialect.ConstantOp.createInt(ctx, loc, index_type, 0);
7662 try entry.addOperation(zero.op);
7663 var storage = try dialects.MemrefDialect.AllocaOp.createStatic(ctx, loc, source_type);
7664 try entry.addOperation(storage.op);
7665 var subview = try dialects.MemrefDialect.SubviewOp.create(
7666 ctx,
7667 loc,
7668 storage.getResult(),
7669 view_type,
7670 );
7671 try dialects.MemrefDialect.setLayoutAttrs(subview.op, ctx, .{
7672 .offset = 1,
7673 .shape = &.{4},
7674 .stride = &.{1},
7675 });
7676 try entry.addOperation(subview.op);
7677 const stored = try dialects.MemrefDialect.StoreOp.create(
7678 ctx,
7679 loc,
7680 func.getArgument(0),
7681 subview.getResult(),
7682 zero.getResult(),
7683 );
7684 try entry.addOperation(stored.op);
7685 const ret = try dialects.FuncDialect.ReturnOp.create(ctx, loc, &.{});
7686 try entry.addOperation(ret.op);
7687 return module.op;
7688 }
7689
7690 /// The number of allocations `Backend.init` spends, so a test can fail the one
7691 /// after them. That one is necessarily the first allocation view legalization
7692 /// asks for, because legalization is the first thing `lower` runs.
7693 fn backendInitAllocations(allocator: std.mem.Allocator, ctx: *ir.Context) !usize {
7694 var counter = std.testing.FailingAllocator.init(allocator, .{});
7695 var probe = try Backend.init(counter.allocator(), ctx, .testing);
7696 probe.deinit();
7697 return counter.alloc_index;
7698 }
7699
7700 test "x86_64 view legalization that cannot allocate answers a shortfall and not a refusal" {
7701 const dialects = @import("../../dialects/root.zig");
7702 const allocator = std.testing.allocator;
7703
7704 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
7705 defer ir_ctx.deinit(allocator);
7706 try dialects.registerAllDialects(&ir_ctx);
7707
7708 const for_pass = try buildViewModule(&ir_ctx);
7709 var first = std.testing.FailingAllocator.init(allocator, .{ .fail_index = 0 });
7710 try std.testing.expectError(
7711 error.OutOfMemory,
7712 passes.memref_views.legalizeMemrefViews(for_pass, first.allocator()),
7713 );
7714
7715 const spent = try backendInitAllocations(allocator, &ir_ctx);
7716 const for_backend = try buildViewModule(&ir_ctx);
7717 var failing = std.testing.FailingAllocator.init(allocator, .{ .fail_index = spent });
7718 var backend = try Backend.init(failing.allocator(), &ir_ctx, .testing);
7719 defer backend.deinit();
7720 try std.testing.expectError(error.OutOfMemory, backend.lower(for_backend));
7721 }
7722
7723 test "x86_64 backend init refuses limits below its built-in names and leaks nothing" {
7724 const dialects = @import("../../dialects/root.zig");
7725 const allocator = std.testing.allocator;
7726
7727 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
7728 defer ir_ctx.deinit(allocator);
7729 try dialects.registerAllDialects(&ir_ctx);
7730
7731 const built_in: u32 = 2 + sys.math.symbols.len;
7732 try std.testing.expect(built_in <= Backend.Limits.testing.external_symbols);
7733 var limits: Backend.Limits = .testing;
7734 limits.external_symbols = built_in - 1;
7735 try std.testing.expectError(error.JitRuntimeFull, Backend.init(allocator, &ir_ctx, limits));
7736
7737 limits.external_symbols = built_in;
7738 var backend = try Backend.init(allocator, &ir_ctx, limits);
7739 defer backend.deinit();
7740 try std.testing.expectEqual(built_in, backend.runtime.external_symbols.count());
7741 try std.testing.expectError(
7742 error.JitRuntimeFull,
7743 backend.runtime.registerExternalSymbol("host_extra", 1),
7744 );
7745 try backend.runtime.registerExternalSymbol("malloc", @intFromPtr(&sys.heap.malloc));
7746 try std.testing.expectEqual(built_in, backend.runtime.external_symbols.count());
7747 }
7748
7749 test "x86_64 backend init that cannot allocate fails at every step and then succeeds" {
7750 const dialects = @import("../../dialects/root.zig");
7751 const allocator = std.testing.allocator;
7752
7753 var ir_ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
7754 defer ir_ctx.deinit(allocator);
7755 try dialects.registerAllDialects(&ir_ctx);
7756
7757 const spent = try backendInitAllocations(allocator, &ir_ctx);
7758 try std.testing.expect(spent > 0);
7759 for (0..spent) |fail_index| {
7760 var failing = std.testing.FailingAllocator.init(allocator, .{ .fail_index = fail_index });
7761 try std.testing.expectError(
7762 error.OutOfMemory,
7763 Backend.init(failing.allocator(), &ir_ctx, .testing),
7764 );
7765 try std.testing.expectEqual(failing.allocated_bytes, failing.freed_bytes);
7766 }
7767 var backend = try Backend.init(allocator, &ir_ctx, .testing);
7768 defer backend.deinit();
7769 const built_in: u32 = 2 + sys.math.symbols.len;
7770 try std.testing.expectEqual(built_in, backend.runtime.external_symbols.count());
7771 }