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 {
45204520fn airSqrt(self: *Self, inst: Air.Inst.Index) !void {
45214521 const un_op = self.air.instructions.items(.data)[inst].un_op;
45224522 const ty = self.air.typeOf(un_op);
4523 const abi_size = @intCast(u32, ty.abiSize(self.target.*));
45234524
45244525 const src_mcv = try self.resolveInst(un_op);
45254526 const dst_mcv = if (src_mcv.isRegister() and self.reuseOperand(inst, un_op, 0, src_mcv))
45264527 src_mcv
45274528 else
45284529 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()) {
4531 .Float => switch (ty.floatBits(self.target.*)) {
4532 32 => .sqrtss,
4533 64 => .sqrtsd,
4534 else => return self.fail("TODO implement airSqrt for {}", .{
4535 ty.fmt(self.bin_file.options.module.?),
4536 }),
4534 const tag = if (@as(?Mir.Inst.Tag, switch (ty.zigTypeTag()) {
4535 .Float => switch (ty.childType().floatBits(self.target.*)) {
4536 32 => if (self.hasFeature(.avx)) .vsqrtss else .sqrtss,
4537 64 => if (self.hasFeature(.avx)) .vsqrtsd else .sqrtsd,
4538 16, 80, 128 => null,
4539 else => unreachable,
45374540 },
4538 else => return self.fail("TODO implement airSqrt for {}", .{
4539 ty.fmt(self.bin_file.options.module.?),
4540 }),
4541 }, ty, dst_mcv, src_mcv);
4541 .Vector => switch (ty.childType().zigTypeTag()) {
4542 .Float => switch (ty.childType().floatBits(self.target.*)) {
4543 32 => switch (ty.vectorLen()) {
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 }
45424586 return self.finishAir(inst, dst_mcv, .{ un_op, .none, .none });
45434587}
45444588
......@@ -9544,85 +9588,92 @@ fn airMulAdd(self: *Self, inst: Air.Inst.Index) !void {
95449588 lock.* = self.register_manager.lockRegAssumeUnused(reg);
95459589 }
95469590
9547 const tag: ?Mir.Inst.Tag =
9591 const tag = if (@as(
9592 ?Mir.Inst.Tag,
95489593 if (mem.eql(u2, &order, &.{ 1, 3, 2 }) or mem.eql(u2, &order, &.{ 3, 1, 2 }))
9549 switch (ty.zigTypeTag()) {
9550 .Float => switch (ty.floatBits(self.target.*)) {
9551 32 => .vfmadd132ss,
9552 64 => .vfmadd132sd,
9553 else => null,
9554 },
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,
9594 switch (ty.zigTypeTag()) {
9595 .Float => switch (ty.floatBits(self.target.*)) {
9596 32 => .vfmadd132ss,
9597 64 => .vfmadd132sd,
9598 16, 80, 128 => null,
9599 else => unreachable,
95689600 },
9569 else => null,
9570 },
9571 else => unreachable,
9572 }
9573 else if (mem.eql(u2, &order, &.{ 2, 1, 3 }) or mem.eql(u2, &order, &.{ 1, 2, 3 }))
9574 switch (ty.zigTypeTag()) {
9575 .Float => switch (ty.floatBits(self.target.*)) {
9576 32 => .vfmadd213ss,
9577 64 => .vfmadd213sd,
9578 else => null,
9579 },
9580 .Vector => switch (ty.childType().zigTypeTag()) {
9581 .Float => switch (ty.childType().floatBits(self.target.*)) {
9582 32 => switch (ty.vectorLen()) {
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,
9601 .Vector => switch (ty.childType().zigTypeTag()) {
9602 .Float => switch (ty.childType().floatBits(self.target.*)) {
9603 32 => switch (ty.vectorLen()) {
9604 1 => .vfmadd132ss,
9605 2...8 => .vfmadd132ps,
9606 else => null,
9607 },
9608 64 => switch (ty.vectorLen()) {
9609 1 => .vfmadd132sd,
9610 2...4 => .vfmadd132pd,
9611 else => null,
9612 },
9613 16, 80, 128 => null,
9614 else => unreachable,
95919615 },
9592 else => null,
9616 else => unreachable,
95939617 },
9594 else => null,
9595 },
9596 else => unreachable,
9597 }
9598 else if (mem.eql(u2, &order, &.{ 2, 3, 1 }) or mem.eql(u2, &order, &.{ 3, 2, 1 }))
9599 switch (ty.zigTypeTag()) {
9600 .Float => switch (ty.floatBits(self.target.*)) {
9601 32 => .vfmadd231ss,
9602 64 => .vfmadd231sd,
9603 else => null,
9604 },
9605 .Vector => switch (ty.childType().zigTypeTag()) {
9606 .Float => switch (ty.childType().floatBits(self.target.*)) {
9607 32 => switch (ty.vectorLen()) {
9608 1 => .vfmadd231ss,
9609 2...8 => .vfmadd231ps,
9610 else => null,
9618 else => unreachable,
9619 }
9620 else if (mem.eql(u2, &order, &.{ 2, 1, 3 }) or mem.eql(u2, &order, &.{ 1, 2, 3 }))
9621 switch (ty.zigTypeTag()) {
9622 .Float => switch (ty.floatBits(self.target.*)) {
9623 32 => .vfmadd213ss,
9624 64 => .vfmadd213sd,
9625 16, 80, 128 => null,
9626 else => unreachable,
9627 },
9628 .Vector => switch (ty.childType().zigTypeTag()) {
9629 .Float => switch (ty.childType().floatBits(self.target.*)) {
9630 32 => switch (ty.vectorLen()) {
9631 1 => .vfmadd213ss,
9632 2...8 => .vfmadd213ps,
9633 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,
96119642 },
9612 64 => switch (ty.vectorLen()) {
9613 1 => .vfmadd231sd,
9614 2...4 => .vfmadd231pd,
9615 else => null,
9643 else => unreachable,
9644 },
9645 else => unreachable,
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,
96169669 },
9617 else => null,
9670 else => unreachable,
96189671 },
9619 else => null,
9620 },
9621 else => null,
9622 }
9623 else
9624 unreachable;
9625 if (tag == null) return self.fail("TODO implement airMulAdd for {}", .{
9672 else => unreachable,
9673 }
9674 else
9675 unreachable,
9676 )) |tag| tag else return self.fail("TODO implement airMulAdd for {}", .{
96269677 ty.fmt(self.bin_file.options.module.?),
96279678 });
96289679
......@@ -9634,14 +9685,14 @@ fn airMulAdd(self: *Self, inst: Air.Inst.Index) !void {
96349685 const mop2_reg = registerAlias(mops[1].getReg().?, abi_size);
96359686 if (mops[2].isRegister())
96369687 try self.asmRegisterRegisterRegister(
9637 tag.?,
9688 tag,
96389689 mop1_reg,
96399690 mop2_reg,
96409691 registerAlias(mops[2].getReg().?, abi_size),
96419692 )
96429693 else
96439694 try self.asmRegisterRegisterMemory(
9644 tag.?,
9695 tag,
96459696 mop1_reg,
96469697 mop2_reg,
96479698 mops[2].mem(Memory.PtrSize.fromSize(abi_size)),
src/arch/x86_64/Encoding.zig+1
......@@ -316,6 +316,7 @@ pub const Mnemonic = enum {
316316 vpsrld, vpsrlq, vpsrlw,
317317 vpunpckhbw, vpunpckhdq, vpunpckhqdq, vpunpckhwd,
318318 vpunpcklbw, vpunpckldq, vpunpcklqdq, vpunpcklwd,
319 vsqrtpd, vsqrtps, vsqrtsd, vsqrtss,
319320 // F16C
320321 vcvtph2ps, vcvtps2ph,
321322 // FMA
src/arch/x86_64/Lower.zig+4
......@@ -212,6 +212,10 @@ pub fn lowerMir(lower: *Lower, index: Mir.Inst.Index) Error!struct {
212212 .vpunpckldq,
213213 .vpunpcklqdq,
214214 .vpunpcklwd,
215 .vsqrtpd,
216 .vsqrtps,
217 .vsqrtsd,
218 .vsqrtss,
215219
216220 .vcvtph2ps,
217221 .vcvtps2ph,
src/arch/x86_64/Mir.zig+8
......@@ -338,6 +338,14 @@ pub const Inst = struct {
338338 vpunpcklqdq,
339339 /// Unpack low data
340340 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
342350 /// Convert 16-bit floating-point values to single-precision floating-point values
343351 vcvtph2ps,
src/arch/x86_64/encodings.zig+15-3
......@@ -869,8 +869,9 @@ pub const table = [_]Entry{
869869
870870 .{ .subss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f, 0x5c }, 0, .none, .sse },
871871
872 .{ .sqrtps, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x0f, 0x51 }, 0, .none, .sse },
873 .{ .sqrtss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f, 0x51 }, 0, .none, .sse },
872 .{ .sqrtps, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x0f, 0x51 }, 0, .none, .sse },
873
874 .{ .sqrtss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f, 0x51 }, 0, .none, .sse },
874875
875876 .{ .ucomiss, .rm, &.{ .xmm, .xmm_m32 }, &.{ 0x0f, 0x2e }, 0, .none, .sse },
876877
......@@ -943,7 +944,8 @@ pub const table = [_]Entry{
943944 .{ .punpcklqdq, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x6c }, 0, .none, .sse2 },
944945
945946 .{ .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
948950 .{ .subsd, .rm, &.{ .xmm, .xmm_m64 }, &.{ 0xf2, 0x0f, 0x5c }, 0, .none, .sse2 },
949951
......@@ -1039,6 +1041,16 @@ pub const table = [_]Entry{
10391041 .{ .vpunpckldq, .rvm, &.{ .xmm, .xmm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x62 }, 0, .vex_128_wig, .avx },
10401042 .{ .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
10421054 // F16C
10431055 .{ .vcvtph2ps, .rm, &.{ .xmm, .xmm_m64 }, &.{ 0x66, 0x0f, 0x38, 0x13 }, 0, .vex_128_w0, .f16c },
10441056 .{ .vcvtph2ps, .rm, &.{ .ymm, .xmm_m128 }, &.{ 0x66, 0x0f, 0x38, 0x13 }, 0, .vex_256_w0, .f16c },