lib/choir/src/backends/x64/cast.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const ir = @import("../../core/root.zig");
  3 const dialects = @import("../../dialects/root.zig");
  4 const encoding = @import("encoding.zig");
  5 const labels = @import("labels.zig");
  6 const registers = @import("registers/root.zig");
  7 const slot_layout = @import("slots.zig");
  8 
  9 const ArithDialect = dialects.arith.ArithDialect;
 10 
 11 pub const Kind = enum {
 12     bool_,
 13     i8_,
 14     u8_,
 15     i16_,
 16     i32_,
 17     u32_,
 18     i64_,
 19     u64_,
 20     index_,
 21     f32_,
 22     f64_,
 23 
 24     pub fn fromTypeName(name: []const u8) ?Kind {
 25         const scalar_kind = dialects.arith.scalarKindFromTypeName(name) orelse return null;
 26         return switch (scalar_kind) {
 27             .bool => .bool_,
 28             .i8 => .i8_,
 29             .u8 => .u8_,
 30             .i16 => .i16_,
 31             .i32 => .i32_,
 32             .u32 => .u32_,
 33             .i64 => .i64_,
 34             .u64 => .u64_,
 35             .index => .index_,
 36             .f32 => .f32_,
 37             .f64 => .f64_,
 38             .u16, .f16, .bf16 => null,
 39         };
 40     }
 41 
 42     pub fn isFloat(self: Kind) bool {
 43         return self == .f32_ or self == .f64_;
 44     }
 45 
 46     pub fn isSignedInt(self: Kind) bool {
 47         return switch (self) {
 48             .i8_, .i16_, .i32_, .i64_ => true,
 49             else => false,
 50         };
 51     }
 52 
 53     pub fn isUnsigned64Int(self: Kind) bool {
 54         return self == .u64_ or self == .index_;
 55     }
 56 
 57     pub fn isUnsigned32Int(self: Kind) bool {
 58         return self == .u32_;
 59     }
 60 
 61     pub fn intBits(self: Kind) u8 {
 62         return switch (self) {
 63             .i8_, .u8_ => 8,
 64             .i16_ => 16,
 65             .i32_ => 32,
 66             .u32_ => 32,
 67             .i64_ => 64,
 68             .u64_ => 64,
 69             .index_ => 64,
 70             else => 0,
 71         };
 72     }
 73 };
 74 
 75 pub fn emitCast(emitter: anytype, op: *ir.Operation) !void {
 76     const result = op.getResult(0) orelse return error.MissingResult;
 77     const operand = op.getOperand(0) orelse return error.MissingOperand;
 78 
 79     const result_slot = try emitter.slotFor(result);
 80     const operand_slot = try emitter.slotFor(operand);
 81 
 82     const result_name = result.type.getDialectTypeName() orelse return error.UnsupportedType;
 83     const operand_name = operand.type.getDialectTypeName() orelse return error.UnsupportedType;
 84 
 85     const result_kind = Kind.fromTypeName(result_name) orelse return error.UnsupportedType;
 86     const operand_kind = Kind.fromTypeName(operand_name) orelse return error.UnsupportedType;
 87 
 88     if (result_kind == operand_kind) {
 89         if (result_kind.isFloat()) {
 90             try emitter.loadSlotXmm(operand_slot, .xmm0);
 91             try emitter.storeSlotXmm(result_slot, .xmm0);
 92         } else {
 93             try emitter.loadInto(operand, .rax);
 94             try emitter.storeFrom(result, .rax);
 95         }
 96         return;
 97     }
 98 
 99     if (operand_kind == .bool_ and result_kind != .bool_) {
100         if (result_kind.isFloat()) {
101             try emitter.loadInto(operand, .rax);
102             switch (result_kind) {
103                 .f32_ => try emitter.emitEncoding(encoding.cvtsi2ssReg64(.xmm0, .rax)),
104                 .f64_ => try emitter.emitEncoding(encoding.cvtsi2sdReg64(.xmm0, .rax)),
105                 else => unreachable,
106             }
107             try emitter.storeSlotXmm(result_slot, .xmm0);
108         } else {
109             try emitter.loadIntoUnsigned(operand, .rax);
110             try emitter.storeFrom(result, .rax);
111         }
112         return;
113     }
114 
115     if (result_kind == .bool_ and operand_kind != .bool_) {
116         if (operand_kind.isFloat()) {
117             try emitter.loadSlotXmm(operand_slot, .xmm0);
118             try emitter.emitEncoding(encoding.xorps(.xmm1, .{ .reg = .xmm1 }));
119             if (operand_kind == .f32_) {
120                 try emitter.emitEncoding(encoding.ucomiss(.xmm0, .{ .reg = .xmm1 }));
121             } else {
122                 try emitter.emitEncoding(encoding.ucomisd(.xmm0, .{ .reg = .xmm1 }));
123             }
124             try emitter.emitEncoding(encoding.setcc(.rax, .np));
125             try emitter.emitEncoding(encoding.setcc(.rcx, .ne));
126             try emitter.emitEncoding(encoding.andRegReg(.rax, .rcx));
127             try emitter.emitEncoding(encoding.movzxReg8(.rax, .rax));
128         } else {
129             try emitter.loadIntoUnsigned(operand, .rax);
130             try emitter.emitEncoding(encoding.cmpRegImm(.rax, 0));
131             try emitter.emitEncoding(encoding.setcc(.rax, .ne));
132             try emitter.emitEncoding(encoding.movzxReg8(.rax, .rax));
133         }
134         try emitter.storeFrom(result, .rax);
135         return;
136     }
137 
138     if (!result_kind.isFloat() and !operand_kind.isFloat()) {
139         const operand_signed = operand_kind.isSignedInt();
140         const result_bits = result_kind.intBits();
141         const operand_bits = operand_kind.intBits();
142 
143         if (operand_kind == .i32_ and (result_kind == .i64_ or result_kind.isUnsigned64Int())) {
144             try emitter.loadInto(operand, .rax);
145             try emitter.emitEncoding(encoding.movsxdRegReg(.rax, .rax));
146             try emitter.storeFrom(result, .rax);
147             return;
148         }
149 
150         if (operand_signed) {
151             try emitter.loadInto(operand, .rax);
152         } else {
153             try emitter.loadIntoUnsigned(operand, .rax);
154         }
155 
156         _ = result_bits;
157         _ = operand_bits;
158         try emitter.storeFrom(result, .rax);
159         return;
160     }
161 
162     if (result_kind.isFloat() and operand_kind.isFloat()) {
163         try emitter.loadSlotXmm(operand_slot, .xmm0);
164         if (operand_kind == .f32_ and result_kind == .f64_) {
165             try emitter.emitEncoding(encoding.cvtss2sd(.xmm0, .{ .reg = .xmm0 }));
166         } else if (operand_kind == .f64_ and result_kind == .f32_) {
167             try emitter.emitEncoding(encoding.cvtsd2ss(.xmm0, .{ .reg = .xmm0 }));
168         } else {
169             return error.UnsupportedType;
170         }
171         try emitter.storeSlotXmm(result_slot, .xmm0);
172         return;
173     }
174 
175     if (result_kind.isFloat() and !operand_kind.isFloat()) {
176         const operand_bits = operand_kind.intBits();
177         if (operand_kind.isSignedInt()) {
178             try emitter.loadSlot(operand_slot, .rax);
179         } else {
180             try emitter.loadSlotUnsigned(operand_slot, .rax);
181         }
182 
183         if (operand_kind.isUnsigned64Int()) {
184             var high_label = labels.Label{};
185             defer high_label.deinit(emitter.allocator);
186             var done_label = labels.Label{};
187             defer done_label.deinit(emitter.allocator);
188 
189             try emitter.emitEncoding(encoding.cmpRegImm(.rax, 0));
190             try emitter.emitJccLabel(.s, &high_label);
191 
192             switch (result_kind) {
193                 .f32_ => try emitter.emitEncoding(encoding.cvtsi2ssReg64(.xmm0, .rax)),
194                 .f64_ => try emitter.emitEncoding(encoding.cvtsi2sdReg64(.xmm0, .rax)),
195                 else => unreachable,
196             }
197             try emitter.emitJmpLabel(&done_label);
198 
199             try emitter.bindLabel(&high_label);
200             try emitter.emitEncoding(encoding.movRegReg(.rcx, .rax));
201             try emitter.emitEncoding(encoding.shrRegImm(.rcx, 1));
202             try emitter.emitEncoding(encoding.andRegImm(.rax, 1));
203             try emitter.emitEncoding(encoding.orRegReg(.rcx, .rax));
204             switch (result_kind) {
205                 .f32_ => {
206                     try emitter.emitEncoding(encoding.cvtsi2ssReg64(.xmm0, .rcx));
207                     try emitter.emitEncoding(encoding.addss(.xmm0, .{ .reg = .xmm0 }));
208                 },
209                 .f64_ => {
210                     try emitter.emitEncoding(encoding.cvtsi2sdReg64(.xmm0, .rcx));
211                     try emitter.emitEncoding(encoding.addsd(.xmm0, .{ .reg = .xmm0 }));
212                 },
213                 else => unreachable,
214             }
215 
216             try emitter.bindLabel(&done_label);
217             try emitter.storeSlotXmm(result_slot, .xmm0);
218             return;
219         }
220 
221         switch (result_kind) {
222             .f32_ => {
223                 if (operand_bits <= 32 and operand_kind.isSignedInt()) {
224                     try emitter.emitEncoding(encoding.cvtsi2ssReg32(.xmm0, .rax));
225                 } else {
226                     try emitter.emitEncoding(encoding.cvtsi2ssReg64(.xmm0, .rax));
227                 }
228             },
229             .f64_ => {
230                 if (operand_bits <= 32 and operand_kind.isSignedInt()) {
231                     try emitter.emitEncoding(encoding.cvtsi2sdReg32(.xmm0, .rax));
232                 } else {
233                     try emitter.emitEncoding(encoding.cvtsi2sdReg64(.xmm0, .rax));
234                 }
235             },
236             else => unreachable,
237         }
238         try emitter.storeSlotXmm(result_slot, .xmm0);
239         return;
240     }
241 
242     if (operand_kind.isFloat() and !result_kind.isFloat()) {
243         try emitter.loadSlotXmm(operand_slot, .xmm0);
244 
245         if (result_kind.isUnsigned32Int()) {
246             var low_label = labels.Label{};
247             defer low_label.deinit(emitter.allocator);
248             var done_label = labels.Label{};
249             defer done_label.deinit(emitter.allocator);
250 
251             switch (operand_kind) {
252                 .f32_ => {
253                     try emitter.emitEncoding(encoding.movRegImm64(.rcx, 0x4F000000));
254                     try emitter.emitEncoding(encoding.movdXmmFromReg32(.xmm1, .rcx));
255                     try emitter.emitEncoding(encoding.ucomiss(.xmm0, .{ .reg = .xmm1 }));
256                 },
257                 .f64_ => {
258                     try emitter.emitEncoding(encoding.movRegImm64(.rcx, 0x41E0000000000000));
259                     try emitter.emitEncoding(encoding.movqXmmFromReg64(.xmm1, .rcx));
260                     try emitter.emitEncoding(encoding.ucomisd(.xmm0, .{ .reg = .xmm1 }));
261                 },
262                 else => unreachable,
263             }
264 
265             try emitter.emitJccLabel(.b, &low_label);
266 
267             switch (operand_kind) {
268                 .f32_ => {
269                     try emitter.emitEncoding(encoding.subss(.xmm0, .{ .reg = .xmm1 }));
270                     try emitter.emitEncoding(encoding.cvttss2siReg32(.rax, .{ .reg = .xmm0 }));
271                 },
272                 .f64_ => {
273                     try emitter.emitEncoding(encoding.subsd(.xmm0, .{ .reg = .xmm1 }));
274                     try emitter.emitEncoding(encoding.cvttsd2siReg32(.rax, .{ .reg = .xmm0 }));
275                 },
276                 else => unreachable,
277             }
278             try emitter.emitEncoding(encoding.addRegImm(.rax, std.math.minInt(i32)));
279             try emitter.emitJmpLabel(&done_label);
280 
281             try emitter.bindLabel(&low_label);
282             switch (operand_kind) {
283                 .f32_ => try emitter.emitEncoding(encoding.cvttss2siReg32(.rax, .{ .reg = .xmm0 })),
284                 .f64_ => try emitter.emitEncoding(encoding.cvttsd2siReg32(.rax, .{ .reg = .xmm0 })),
285                 else => unreachable,
286             }
287 
288             try emitter.bindLabel(&done_label);
289             try emitter.storeSlot(result_slot, .rax);
290             return;
291         }
292 
293         if (result_kind.isUnsigned64Int()) {
294             var low_label = labels.Label{};
295             defer low_label.deinit(emitter.allocator);
296             var done_label = labels.Label{};
297             defer done_label.deinit(emitter.allocator);
298 
299             switch (operand_kind) {
300                 .f32_ => {
301                     try emitter.emitEncoding(encoding.movRegImm64(.rcx, 0x5F000000));
302                     try emitter.emitEncoding(encoding.movdXmmFromReg32(.xmm1, .rcx));
303                     try emitter.emitEncoding(encoding.ucomiss(.xmm0, .{ .reg = .xmm1 }));
304                 },
305                 .f64_ => {
306                     try emitter.emitEncoding(encoding.movRegImm64(.rcx, 0x43E0000000000000));
307                     try emitter.emitEncoding(encoding.movqXmmFromReg64(.xmm1, .rcx));
308                     try emitter.emitEncoding(encoding.ucomisd(.xmm0, .{ .reg = .xmm1 }));
309                 },
310                 else => unreachable,
311             }
312 
313             try emitter.emitJccLabel(.b, &low_label);
314 
315             switch (operand_kind) {
316                 .f32_ => {
317                     try emitter.emitEncoding(encoding.subss(.xmm0, .{ .reg = .xmm1 }));
318                     try emitter.emitEncoding(encoding.cvttss2siReg64(.rax, .{ .reg = .xmm0 }));
319                 },
320                 .f64_ => {
321                     try emitter.emitEncoding(encoding.subsd(.xmm0, .{ .reg = .xmm1 }));
322                     try emitter.emitEncoding(encoding.cvttsd2siReg64(.rax, .{ .reg = .xmm0 }));
323                 },
324                 else => unreachable,
325             }
326             try emitter.emitEncoding(encoding.movRegImm64(.rcx, 0x8000000000000000));
327             try emitter.emitEncoding(encoding.addRegReg(.rax, .rcx));
328             try emitter.emitJmpLabel(&done_label);
329 
330             try emitter.bindLabel(&low_label);
331             switch (operand_kind) {
332                 .f32_ => try emitter.emitEncoding(encoding.cvttss2siReg64(.rax, .{ .reg = .xmm0 })),
333                 .f64_ => try emitter.emitEncoding(encoding.cvttsd2siReg64(.rax, .{ .reg = .xmm0 })),
334                 else => unreachable,
335             }
336 
337             try emitter.bindLabel(&done_label);
338             try emitter.storeSlot(result_slot, .rax);
339             return;
340         }
341 
342         const result_bits = result_kind.intBits();
343         switch (operand_kind) {
344             .f32_ => {
345                 if (result_bits <= 32) {
346                     try emitter.emitEncoding(encoding.cvttss2siReg32(.rax, .{ .reg = .xmm0 }));
347                 } else {
348                     try emitter.emitEncoding(encoding.cvttss2siReg64(.rax, .{ .reg = .xmm0 }));
349                 }
350             },
351             .f64_ => {
352                 if (result_bits <= 32) {
353                     try emitter.emitEncoding(encoding.cvttsd2siReg32(.rax, .{ .reg = .xmm0 }));
354                 } else {
355                     try emitter.emitEncoding(encoding.cvttsd2siReg64(.rax, .{ .reg = .xmm0 }));
356                 }
357             },
358             else => unreachable,
359         }
360         try emitter.storeSlot(result_slot, .rax);
361         return;
362     }
363 
364     return error.UnsupportedType;
365 }
366 
367 pub fn emitBitcast(emitter: anytype, op: *ir.Operation) !void {
368     const result = op.getResult(0) orelse return error.MissingResult;
369     const operand = op.getOperand(0) orelse return error.MissingOperand;
370 
371     const result_slot = try emitter.slotFor(result);
372     const operand_slot = try emitter.slotFor(operand);
373 
374     const result_name = result.type.getDialectTypeName() orelse return error.UnsupportedType;
375     const operand_name = operand.type.getDialectTypeName() orelse return error.UnsupportedType;
376 
377     const result_kind = Kind.fromTypeName(result_name) orelse return error.UnsupportedType;
378     const operand_kind = Kind.fromTypeName(operand_name) orelse return error.UnsupportedType;
379     if (result_kind == .bool_ or operand_kind == .bool_) return error.UnsupportedType;
380     if (result_kind == .index_ or operand_kind == .index_) return error.UnsupportedType;
381 
382     if (std.mem.eql(u8, result_name, "arith.f16") or
383         std.mem.eql(u8, result_name, "arith.bf16") or
384         std.mem.eql(u8, operand_name, "arith.f16") or
385         std.mem.eql(u8, operand_name, "arith.bf16")) return error.UnsupportedType;
386 
387     if (result_slot.width != operand_slot.width) return error.UnsupportedType;
388 
389     if (operand_kind.isFloat()) {
390         if (emitter.xmmHome(operand)) |home| {
391             if (operand_slot.width == 32) {
392                 try emitter.emitEncoding(encoding.movdReg32FromXmm(.rax, home));
393             } else {
394                 try emitter.emitEncoding(encoding.movqReg64FromXmm(.rax, home));
395             }
396         } else {
397             try emitter.loadSlotUnsigned(operand_slot, .rax);
398         }
399     } else {
400         try emitter.loadIntoUnsigned(operand, .rax);
401     }
402 
403     if (result_kind.isFloat()) {
404         if (emitter.xmmHomeForPhase(result, .definition)) |home| {
405             if (result_slot.width == 32) {
406                 try emitter.emitEncoding(encoding.movdXmmFromReg32(home, .rax));
407             } else {
408                 try emitter.emitEncoding(encoding.movqXmmFromReg64(home, .rax));
409             }
410         } else {
411             try emitter.storeSlot(result_slot, .rax);
412         }
413     } else {
414         try emitter.storeFrom(result, .rax);
415     }
416 }
417 
418 const CopyRecorder = struct {
419     allocator: std.mem.Allocator,
420     input: *ir.Value,
421     result: *ir.Value,
422     input_slot: slot_layout.Slot,
423     result_slot: slot_layout.Slot,
424     loads: usize = 0,
425     stores: usize = 0,
426 
427     pub fn slotFor(self: *CopyRecorder, value: *ir.Value) !slot_layout.Slot {
428         if (value == self.input) return self.input_slot;
429         if (value == self.result) return self.result_slot;
430         return error.MissingSlot;
431     }
432 
433     pub fn loadSlot(self: *CopyRecorder, slot: slot_layout.Slot, reg: registers.GPR) !void {
434         _ = slot;
435         _ = reg;
436         self.loads += 1;
437     }
438 
439     pub fn loadSlotUnsigned(self: *CopyRecorder, slot: slot_layout.Slot, reg: registers.GPR) !void {
440         try self.loadSlot(slot, reg);
441     }
442 
443     pub fn loadInto(self: *CopyRecorder, value: *ir.Value, reg: registers.GPR) !void {
444         try self.loadSlot(try self.slotFor(value), reg);
445     }
446 
447     pub fn loadIntoUnsigned(self: *CopyRecorder, value: *ir.Value, reg: registers.GPR) !void {
448         try self.loadSlot(try self.slotFor(value), reg);
449     }
450 
451     pub fn storeFrom(self: *CopyRecorder, value: *ir.Value, reg: registers.GPR) !void {
452         try self.storeSlot(try self.slotFor(value), reg);
453     }
454 
455     pub fn loadSlotXmm(self: *CopyRecorder, slot: slot_layout.Slot, reg: registers.XMM) !void {
456         _ = slot;
457         _ = reg;
458         self.loads += 1;
459     }
460 
461     pub fn storeSlot(self: *CopyRecorder, slot: slot_layout.Slot, reg: registers.GPR) !void {
462         _ = slot;
463         _ = reg;
464         self.stores += 1;
465     }
466 
467     pub fn storeSlotXmm(self: *CopyRecorder, slot: slot_layout.Slot, reg: registers.XMM) !void {
468         _ = slot;
469         _ = reg;
470         self.stores += 1;
471     }
472 
473     pub fn emitEncoding(_: *CopyRecorder, _: encoding.Encoding) !void {}
474 
475     pub fn emitJccLabel(_: *CopyRecorder, _: encoding.Condition, _: *labels.Label) !void {}
476 
477     pub fn emitJmpLabel(_: *CopyRecorder, _: *labels.Label) !void {}
478 
479     pub fn bindLabel(_: *CopyRecorder, _: *labels.Label) !void {}
480 };
481 
482 test "x86_64 cast owner emits integer identity copy" {
483     var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
484     defer ctx.deinit(std.testing.allocator);
485     try @import("../../dialects/root.zig").registerAllDialects(&ctx);
486 
487     const loc = ir.Location.getUnknown();
488     const i32_type = try ArithDialect.getScalarType(&ctx, .i32);
489     var constant = try ArithDialect.ConstantOp.createInt(&ctx, loc, i32_type, 7);
490     var cast = try ArithDialect.CastOp.create(&ctx, loc, constant.getResult(), i32_type);
491 
492     var recorder = CopyRecorder{
493         .allocator = std.testing.allocator,
494         .input = constant.getResult(),
495         .result = cast.getResult(),
496         .input_slot = .{ .offset = -8, .width = 32, .ext = .signed },
497         .result_slot = .{ .offset = -16, .width = 32, .ext = .signed },
498     };
499 
500     try emitCast(&recorder, cast.op);
501 
502     try std.testing.expectEqual(@as(usize, 1), recorder.loads);
503     try std.testing.expectEqual(@as(usize, 1), recorder.stores);
504 }
505 
506 test "x86_64 cast owner classifies scalar types" {
507     try std.testing.expectEqual(Kind.i64_, Kind.fromTypeName("arith.i64").?);
508     try std.testing.expectEqual(Kind.u32_, Kind.fromTypeName("arith.u32").?);
509     try std.testing.expectEqual(Kind.u64_, Kind.fromTypeName("arith.u64").?);
510     try std.testing.expectEqual(Kind.f32_, Kind.fromTypeName("arith.f32").?);
511     try std.testing.expect(Kind.fromTypeName("arith.vec4xf32") == null);
512 }