authorgravatar for jacobly@ziglang.orgJacob Young <jacobly@ziglang.org> 2023-05-07 03:47:56-04:00
committergravatar for jacobly@ziglang.orgJacob Young <jacobly@ziglang.org> 2023-05-08 07:36:20-04:00
log5c5da179fb930c9d8be9366a851eb4a36f4044f1
tree693460d65399c1c3dd443f70aec71fe711920652
parent05580b9453e4ae2d9b62fe4178651937d8b73989

x86_64: implement `@sqrt` for vectors


5 files changed, 164 insertions(+), 88 deletions(-)

src/arch/x86_64/CodeGen.zig+136-85
...@@ -4520,25 +4520,69 @@ fn airRound(self: *Self, inst: Air.Inst.Index, mode: Immediate) !void {...@@ -4520,25 +4520,69 @@ fn airRound(self: *Self, inst: Air.Inst.Index, mode: Immediate) !void {
4520fn airSqrt(self: *Self, inst: Air.Inst.Index) !void {4520fn airSqrt(self: *Self, inst: Air.Inst.Index) !void {
4521 const un_op = self.air.instructions.items(.data)[inst].un_op;4521 const un_op = self.air.instructions.items(.data)[inst].un_op;
4522 const ty = self.air.typeOf(un_op);4522 const ty = self.air.typeOf(un_op);
4523 const abi_size = @intCast(u32, ty.abiSize(self.target.*));
45234524
4524 const src_mcv = try self.resolveInst(un_op);4525 const src_mcv = try self.resolveInst(un_op);
4525 const dst_mcv = if (src_mcv.isRegister() and self.reuseOperand(inst, un_op, 0, src_mcv))4526 const dst_mcv = if (src_mcv.isRegister() and self.reuseOperand(inst, un_op, 0, src_mcv))
4526 src_mcv4527 src_mcv
4527 else4528 else
4528 try self.copyToRegisterWithInstTracking(inst, ty, src_mcv);4529 try self.copyToRegisterWithInstTracking(inst, ty, src_mcv);
4530 const dst_reg = registerAlias(dst_mcv.getReg().?, abi_size);
4531 const dst_lock = self.register_manager.lockReg(dst_reg);
4532 defer if (dst_lock) |lock| self.register_manager.unlockReg(lock);
45294533
4530 try self.genBinOpMir(switch (ty.zigTypeTag()) {4534 const tag = if (@as(?Mir.Inst.Tag, switch (ty.zigTypeTag()) {
4531 .Float => switch (ty.floatBits(self.target.*)) {4535 .Float => switch (ty.childType().floatBits(self.target.*)) {
4532 32 => .sqrtss,4536 32 => if (self.hasFeature(.avx)) .vsqrtss else .sqrtss,
4533 64 => .sqrtsd,4537 64 => if (self.hasFeature(.avx)) .vsqrtsd else .sqrtsd,
4534 else => return self.fail("TODO implement airSqrt for {}", .{4538 16, 80, 128 => null,
4535 ty.fmt(self.bin_file.options.module.?),4539 else => unreachable,
4536 }),
4537 },4540 },
4538 else => return self.fail("TODO implement airSqrt for {}", .{4541 .Vector => switch (ty.childType().zigTypeTag()) {
4539 ty.fmt(self.bin_file.options.module.?),4542 .Float => switch (ty.childType().floatBits(self.target.*)) {
4540 }),4543 32 => switch (ty.vectorLen()) {
4541 }, ty, dst_mcv, src_mcv);4544 1 => if (self.hasFeature(.avx)) .vsqrtss else .sqrtss,
4545 2...4 => if (self.hasFeature(.avx)) .vsqrtps else .sqrtps,
4546 5...8 => if (self.hasFeature(.avx)) .vsqrtps else null,
4547 else => null,
4548 },
4549 64 => switch (ty.vectorLen()) {
4550 1 => if (self.hasFeature(.avx)) .vsqrtsd else .sqrtsd,
4551 2 => if (self.hasFeature(.avx)) .vsqrtpd else .sqrtpd,
4552 3...4 => if (self.hasFeature(.avx)) .vsqrtpd else null,
4553 else => null,
4554 },
4555 16, 80, 128 => null,
4556 else => unreachable,
4557 },
4558 else => unreachable,
4559 },
4560 else => unreachable,
4561 })) |tag| tag else return self.fail("TODO implement airSqrt for {}", .{
4562 ty.fmt(self.bin_file.options.module.?),
4563 });
4564 switch (tag) {
4565 .vsqrtss, .vsqrtsd => if (src_mcv.isRegister()) try self.asmRegisterRegisterRegister(
4566 tag,
4567 dst_reg,
4568 dst_reg,
4569 registerAlias(src_mcv.getReg().?, abi_size),
4570 ) else try self.asmRegisterRegisterMemory(
4571 tag,
4572 dst_reg,
4573 dst_reg,
4574 src_mcv.mem(Memory.PtrSize.fromSize(abi_size)),
4575 ),
4576 else => if (src_mcv.isRegister()) try self.asmRegisterRegister(
4577 tag,
4578 dst_reg,
4579 registerAlias(src_mcv.getReg().?, abi_size),
4580 ) else try self.asmRegisterMemory(
4581 tag,
4582 dst_reg,
4583 src_mcv.mem(Memory.PtrSize.fromSize(abi_size)),
4584 ),
4585 }
4542 return self.finishAir(inst, dst_mcv, .{ un_op, .none, .none });4586 return self.finishAir(inst, dst_mcv, .{ un_op, .none, .none });
4543}4587}
45444588
...@@ -9544,85 +9588,92 @@ fn airMulAdd(self: *Self, inst: Air.Inst.Index) !void {...@@ -9544,85 +9588,92 @@ fn airMulAdd(self: *Self, inst: Air.Inst.Index) !void {
9544 lock.* = self.register_manager.lockRegAssumeUnused(reg);9588 lock.* = self.register_manager.lockRegAssumeUnused(reg);
9545 }9589 }
95469590
9547 const tag: ?Mir.Inst.Tag =9591 const tag = if (@as(
9592 ?Mir.Inst.Tag,
9548 if (mem.eql(u2, &order, &.{ 1, 3, 2 }) or mem.eql(u2, &order, &.{ 3, 1, 2 }))9593 if (mem.eql(u2, &order, &.{ 1, 3, 2 }) or mem.eql(u2, &order, &.{ 3, 1, 2 }))
9549 switch (ty.zigTypeTag()) {9594 switch (ty.zigTypeTag()) {
9550 .Float => switch (ty.floatBits(self.target.*)) {9595 .Float => switch (ty.floatBits(self.target.*)) {
9551 32 => .vfmadd132ss,9596 32 => .vfmadd132ss,
9552 64 => .vfmadd132sd,9597 64 => .vfmadd132sd,
9553 else => null,9598 16, 80, 128 => null,
9554 },9599 else => unreachable,
9555 .Vector => switch (ty.childType().zigTypeTag()) {
9556 .Float => switch (ty.childType().floatBits(self.target.*)) {
9557 32 => switch (ty.vectorLen()) {
9558 1 => .vfmadd132ss,
9559 2...8 => .vfmadd132ps,
9560 else => null,
9561 },
9562 64 => switch (ty.vectorLen()) {
9563 1 => .vfmadd132sd,
9564 2...4 => .vfmadd132pd,
9565 else => null,
9566 },
9567 else => null,
9568 },9600 },
9569 else => null,9601 .Vector => switch (ty.childType().zigTypeTag()) {
9570 },9602 .Float => switch (ty.childType().floatBits(self.target.*)) {
9571 else => unreachable,9603 32 => switch (ty.vectorLen()) {
9572 }9604 1 => .vfmadd132ss,
9573 else if (mem.eql(u2, &order, &.{ 2, 1, 3 }) or mem.eql(u2, &order, &.{ 1, 2, 3 }))9605 2...8 => .vfmadd132ps,
9574 switch (ty.zigTypeTag()) {9606 else => null,
9575 .Float => switch (ty.floatBits(self.target.*)) {9607 },
9576 32 => .vfmadd213ss,9608 64 => switch (ty.vectorLen()) {
9577 64 => .vfmadd213sd,9609 1 => .vfmadd132sd,
9578 else => null,9610 2...4 => .vfmadd132pd,
9579 },9611 else => null,
9580 .Vector => switch (ty.childType().zigTypeTag()) {9612 },
9581 .Float => switch (ty.childType().floatBits(self.target.*)) {9613 16, 80, 128 => null,
9582 32 => switch (ty.vectorLen()) {9614 else => unreachable,
9583 1 => .vfmadd213ss,
9584 2...8 => .vfmadd213ps,
9585 else => null,
9586 },
9587 64 => switch (ty.vectorLen()) {
9588 1 => .vfmadd213sd,
9589 2...4 => .vfmadd213pd,
9590 else => null,
9591 },9615 },
9592 else => null,9616 else => unreachable,
9593 },9617 },
9594 else => null,9618 else => unreachable,
9595 },9619 }
9596 else => unreachable,9620 else if (mem.eql(u2, &order, &.{ 2, 1, 3 }) or mem.eql(u2, &order, &.{ 1, 2, 3 }))
9597 }9621 switch (ty.zigTypeTag()) {
9598 else if (mem.eql(u2, &order, &.{ 2, 3, 1 }) or mem.eql(u2, &order, &.{ 3, 2, 1 }))9622 .Float => switch (ty.floatBits(self.target.*)) {
9599 switch (ty.zigTypeTag()) {9623 32 => .vfmadd213ss,
9600 .Float => switch (ty.floatBits(self.target.*)) {9624 64 => .vfmadd213sd,
9601 32 => .vfmadd231ss,9625 16, 80, 128 => null,
9602 64 => .vfmadd231sd,9626 else => unreachable,
9603 else => null,9627 },
9604 },9628 .Vector => switch (ty.childType().zigTypeTag()) {
9605 .Vector => switch (ty.childType().zigTypeTag()) {9629 .Float => switch (ty.childType().floatBits(self.target.*)) {
9606 .Float => switch (ty.childType().floatBits(self.target.*)) {9630 32 => switch (ty.vectorLen()) {
9607 32 => switch (ty.vectorLen()) {9631 1 => .vfmadd213ss,
9608 1 => .vfmadd231ss,9632 2...8 => .vfmadd213ps,
9609 2...8 => .vfmadd231ps,9633 else => null,
9610 else => null,9634 },
9635 64 => switch (ty.vectorLen()) {
9636 1 => .vfmadd213sd,
9637 2...4 => .vfmadd213pd,
9638 else => null,
9639 },
9640 16, 80, 128 => null,
9641 else => unreachable,
9611 },9642 },
9612 64 => switch (ty.vectorLen()) {9643 else => unreachable,
9613 1 => .vfmadd231sd,9644 },
9614 2...4 => .vfmadd231pd,9645 else => unreachable,
9615 else => null,9646 }
9647 else if (mem.eql(u2, &order, &.{ 2, 3, 1 }) or mem.eql(u2, &order, &.{ 3, 2, 1 }))
9648 switch (ty.zigTypeTag()) {
9649 .Float => switch (ty.floatBits(self.target.*)) {
9650 32 => .vfmadd231ss,
9651 64 => .vfmadd231sd,
9652 16, 80, 128 => null,
9653 else => unreachable,
9654 },
9655 .Vector => switch (ty.childType().zigTypeTag()) {
9656 .Float => switch (ty.childType().floatBits(self.target.*)) {
9657 32 => switch (ty.vectorLen()) {
9658 1 => .vfmadd231ss,
9659 2...8 => .vfmadd231ps,
9660 else => null,
9661 },
9662 64 => switch (ty.vectorLen()) {
9663 1 => .vfmadd231sd,
9664 2...4 => .vfmadd231pd,
9665 else => null,
9666 },
9667 16, 80, 128 => null,
9668 else => unreachable,
9616 },9669 },
9617 else => null,9670 else => unreachable,
9618 },9671 },
9619 else => null,9672 else => unreachable,
9620 },9673 }
9621 else => null,9674 else
9622 }9675 unreachable,
9623 else9676 )) |tag| tag else return self.fail("TODO implement airMulAdd for {}", .{
9624 unreachable;
9625 if (tag == null) return self.fail("TODO implement airMulAdd for {}", .{
9626 ty.fmt(self.bin_file.options.module.?),9677 ty.fmt(self.bin_file.options.module.?),
9627 });9678 });
96289679
...@@ -9634,14 +9685,14 @@ fn airMulAdd(self: *Self, inst: Air.Inst.Index) !void {...@@ -9634,14 +9685,14 @@ fn airMulAdd(self: *Self, inst: Air.Inst.Index) !void {
9634 const mop2_reg = registerAlias(mops[1].getReg().?, abi_size);9685 const mop2_reg = registerAlias(mops[1].getReg().?, abi_size);
9635 if (mops[2].isRegister())9686 if (mops[2].isRegister())
9636 try self.asmRegisterRegisterRegister(9687 try self.asmRegisterRegisterRegister(
9637 tag.?,9688 tag,
9638 mop1_reg,9689 mop1_reg,
9639 mop2_reg,9690 mop2_reg,
9640 registerAlias(mops[2].getReg().?, abi_size),9691 registerAlias(mops[2].getReg().?, abi_size),
9641 )9692 )
9642 else9693 else
9643 try self.asmRegisterRegisterMemory(9694 try self.asmRegisterRegisterMemory(
9644 tag.?,9695 tag,
9645 mop1_reg,9696 mop1_reg,
9646 mop2_reg,9697 mop2_reg,
9647 mops[2].mem(Memory.PtrSize.fromSize(abi_size)),9698 mops[2].mem(Memory.PtrSize.fromSize(abi_size)),
src/arch/x86_64/Encoding.zig+1
...@@ -316,6 +316,7 @@ pub const Mnemonic = enum {...@@ -316,6 +316,7 @@ pub const Mnemonic = enum {
316 vpsrld, vpsrlq, vpsrlw,316 vpsrld, vpsrlq, vpsrlw,
317 vpunpckhbw, vpunpckhdq, vpunpckhqdq, vpunpckhwd,317 vpunpckhbw, vpunpckhdq, vpunpckhqdq, vpunpckhwd,
318 vpunpcklbw, vpunpckldq, vpunpcklqdq, vpunpcklwd,318 vpunpcklbw, vpunpckldq, vpunpcklqdq, vpunpcklwd,
319 vsqrtpd, vsqrtps, vsqrtsd, vsqrtss,
319 // F16C320 // F16C
320 vcvtph2ps, vcvtps2ph,321 vcvtph2ps, vcvtps2ph,
321 // FMA322 // FMA
src/arch/x86_64/Lower.zig+4
...@@ -212,6 +212,10 @@ pub fn lowerMir(lower: *Lower, index: Mir.Inst.Index) Error!struct {...@@ -212,6 +212,10 @@ pub fn lowerMir(lower: *Lower, index: Mir.Inst.Index) Error!struct {
212 .vpunpckldq,212 .vpunpckldq,
213 .vpunpcklqdq,213 .vpunpcklqdq,
214 .vpunpcklwd,214 .vpunpcklwd,
215 .vsqrtpd,
216 .vsqrtps,
217 .vsqrtsd,
218 .vsqrtss,
215219
216 .vcvtph2ps,220 .vcvtph2ps,
217 .vcvtps2ph,221 .vcvtps2ph,
src/arch/x86_64/Mir.zig+8
...@@ -338,6 +338,14 @@ pub const Inst = struct {...@@ -338,6 +338,14 @@ pub const Inst = struct {
338 vpunpcklqdq,338 vpunpcklqdq,
339 /// Unpack low data339 /// Unpack low data
340 vpunpcklwd,340 vpunpcklwd,
341 /// Square root of packed double-precision floating-point value
342 vsqrtpd,
343 /// Square root of packed single-precision floating-point value
344 vsqrtps,
345 /// Square root of scalar double-precision floating-point value
346 vsqrtsd,
347 /// Square root of scalar single-precision floating-point value
348 vsqrtss,
341349
342 /// Convert 16-bit floating-point values to single-precision floating-point values350 /// Convert 16-bit floating-point values to single-precision floating-point values
343 vcvtph2ps,351 vcvtph2ps,
src/arch/x86_64/encodings.zig+15-3
...@@ -869,8 +869,9 @@ pub const table = [_]Entry{...@@ -869,8 +869,9 @@ pub const table = [_]Entry{
869869
870 .{ .subss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f, 0x5c }, 0, .none, .sse },870 .{ .subss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f, 0x5c }, 0, .none, .sse },
871871
872 .{ .sqrtps, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x0f, 0x51 }, 0, .none, .sse },872 .{ .sqrtps, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x0f, 0x51 }, 0, .none, .sse },
873 .{ .sqrtss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f, 0x51 }, 0, .none, .sse },873
874 .{ .sqrtss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f, 0x51 }, 0, .none, .sse },
874875
875 .{ .ucomiss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0x0f, 0x2e }, 0, .none, .sse },876 .{ .ucomiss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0x0f, 0x2e }, 0, .none, .sse },
876877
...@@ -943,7 +944,8 @@ pub const table = [_]Entry{...@@ -943,7 +944,8 @@ pub const table = [_]Entry{
943 .{ .punpcklqdq, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x6c }, 0, .none, .sse2 },944 .{ .punpcklqdq, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x6c }, 0, .none, .sse2 },
944945
945 .{ .sqrtpd, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x51 }, 0, .none, .sse2 },946 .{ .sqrtpd, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x51 }, 0, .none, .sse2 },
946 .{ .sqrtsd, .rm, &.{ .xmm, .xmm_m64 }, &.{ 0xf2, 0x0f, 0x51 }, 0, .none, .sse2 },947
948 .{ .sqrtsd, .rm, &.{ .xmm, .xmm_m64 }, &.{ 0xf2, 0x0f, 0x51 }, 0, .none, .sse2 },
947949
948 .{ .subsd, .rm, &.{ .xmm, .xmm_m64 }, &.{ 0xf2, 0x0f, 0x5c }, 0, .none, .sse2 },950 .{ .subsd, .rm, &.{ .xmm, .xmm_m64 }, &.{ 0xf2, 0x0f, 0x5c }, 0, .none, .sse2 },
949951
...@@ -1039,6 +1041,16 @@ pub const table = [_]Entry{...@@ -1039,6 +1041,16 @@ pub const table = [_]Entry{
1039 .{ .vpunpckldq, .rvm, &.{ .xmm, .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x62 }, 0, .vex_128_wig, .avx },1041 .{ .vpunpckldq, .rvm, &.{ .xmm, .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x62 }, 0, .vex_128_wig, .avx },
1040 .{ .vpunpcklqdq, .rvm, &.{ .xmm, .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x6c }, 0, .vex_128_wig, .avx },1042 .{ .vpunpcklqdq, .rvm, &.{ .xmm, .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x6c }, 0, .vex_128_wig, .avx },
10411043
1044 .{ .vsqrtpd, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x51 }, 0, .vex_128_wig, .avx },
1045 .{ .vsqrtpd, .rm, &.{ .ymm, .ymm_m256 }, &.{ 0x66, 0x0f, 0x51 }, 0, .vex_256_wig, .avx },
1046
1047 .{ .vsqrtps, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x0f, 0x51 }, 0, .vex_128_wig, .avx },
1048 .{ .vsqrtps, .rm, &.{ .ymm, .ymm_m256 }, &.{ 0x0f, 0x51 }, 0, .vex_256_wig, .avx },
1049
1050 .{ .vsqrtsd, .rvm, &.{ .xmm, .xmm, .xmm_m64 }, &.{ 0xf2, 0x0f }, 0, .vex_lig_wig, .avx },
1051
1052 .{ .vsqrtss, .rvm, &.{ .xmm, .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f }, 0, .vex_lig_wig, .avx },
1053
1042 // F16C1054 // F16C
1043 .{ .vcvtph2ps, .rm, &.{ .xmm, .xmm_m64 }, &.{ 0x66, 0x0f, 0x38, 0x13 }, 0, .vex_128_w0, .f16c },1055 .{ .vcvtph2ps, .rm, &.{ .xmm, .xmm_m64 }, &.{ 0x66, 0x0f, 0x38, 0x13 }, 0, .vex_128_w0, .f16c },
1044 .{ .vcvtph2ps, .rm, &.{ .ymm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x38, 0x13 }, 0, .vex_256_w0, .f16c },1056 .{ .vcvtph2ps, .rm, &.{ .ymm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x38, 0x13 }, 0, .vex_256_w0, .f16c },