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 }