authorgravatar for jacobly@ziglang.orgJacob Young <jacobly@ziglang.org> 2023-05-07 05:01:37-04:00
committergravatar for jacobly@ziglang.orgJacob Young <jacobly@ziglang.org> 2023-05-08 07:36:20-04:00
logea957c4cff77f045108863cb5552b3511cb455c1
treea40ee9152c2c8b830e32ee81f0d7df39d2a3681a
parent5c5da179fb930c9d8be9366a851eb4a36f4044f1

x86_64: implement `@sqrt` for `f16` scalars and vectors


3 files changed, 109 insertions(+), 52 deletions(-)

src/arch/x86_64/CodeGen.zig+107-49
...@@ -4531,59 +4531,117 @@ fn airSqrt(self: *Self, inst: Air.Inst.Index) !void {...@@ -4531,59 +4531,117 @@ fn airSqrt(self: *Self, inst: Air.Inst.Index) !void {
4531 const dst_lock = self.register_manager.lockReg(dst_reg);4531 const dst_lock = self.register_manager.lockReg(dst_reg);
4532 defer if (dst_lock) |lock| self.register_manager.unlockReg(lock);4532 defer if (dst_lock) |lock| self.register_manager.unlockReg(lock);
45334533
4534 const tag = if (@as(?Mir.Inst.Tag, switch (ty.zigTypeTag()) {4534 const result: MCValue = result: {
4535 .Float => switch (ty.childType().floatBits(self.target.*)) {4535 const tag = if (@as(?Mir.Inst.Tag, switch (ty.zigTypeTag()) {
4536 32 => if (self.hasFeature(.avx)) .vsqrtss else .sqrtss,4536 .Float => switch (ty.floatBits(self.target.*)) {
4537 64 => if (self.hasFeature(.avx)) .vsqrtsd else .sqrtsd,4537 16 => if (self.hasFeature(.f16c)) {
4538 16, 80, 128 => null,4538 const mat_src_reg = if (src_mcv.isRegister())
4539 else => unreachable,4539 src_mcv.getReg().?
4540 },4540 else
4541 .Vector => switch (ty.childType().zigTypeTag()) {4541 try self.copyToTmpRegister(ty, src_mcv);
4542 .Float => switch (ty.childType().floatBits(self.target.*)) {4542 try self.asmRegisterRegister(.vcvtph2ps, dst_reg, mat_src_reg.to128());
4543 32 => switch (ty.vectorLen()) {4543 try self.asmRegisterRegisterRegister(.vsqrtss, dst_reg, dst_reg, dst_reg);
4544 1 => if (self.hasFeature(.avx)) .vsqrtss else .sqrtss,4544 try self.asmRegisterRegisterImmediate(
4545 2...4 => if (self.hasFeature(.avx)) .vsqrtps else .sqrtps,4545 .vcvtps2ph,
4546 5...8 => if (self.hasFeature(.avx)) .vsqrtps else null,4546 dst_reg,
4547 else => null,4547 dst_reg,
4548 },4548 Immediate.u(0b1_00),
4549 64 => switch (ty.vectorLen()) {4549 );
4550 1 => if (self.hasFeature(.avx)) .vsqrtsd else .sqrtsd,4550 break :result dst_mcv;
4551 2 => if (self.hasFeature(.avx)) .vsqrtpd else .sqrtpd,4551 } else null,
4552 3...4 => if (self.hasFeature(.avx)) .vsqrtpd else null,4552 32 => if (self.hasFeature(.avx)) .vsqrtss else .sqrtss,
4553 else => null,4553 64 => if (self.hasFeature(.avx)) .vsqrtsd else .sqrtsd,
4554 80, 128 => null,
4555 else => unreachable,
4556 },
4557 .Vector => switch (ty.childType().zigTypeTag()) {
4558 .Float => switch (ty.childType().floatBits(self.target.*)) {
4559 16 => if (self.hasFeature(.f16c)) switch (ty.vectorLen()) {
4560 1 => {
4561 const mat_src_reg = if (src_mcv.isRegister())
4562 src_mcv.getReg().?
4563 else
4564 try self.copyToTmpRegister(ty, src_mcv);
4565 try self.asmRegisterRegister(.vcvtph2ps, dst_reg, mat_src_reg.to128());
4566 try self.asmRegisterRegisterRegister(.vsqrtss, dst_reg, dst_reg, dst_reg);
4567 try self.asmRegisterRegisterImmediate(
4568 .vcvtps2ph,
4569 dst_reg,
4570 dst_reg,
4571 Immediate.u(0b1_00),
4572 );
4573 break :result dst_mcv;
4574 },
4575 2...8 => {
4576 const wide_reg = registerAlias(dst_reg, abi_size * 2);
4577 if (src_mcv.isRegister()) try self.asmRegisterRegister(
4578 .vcvtph2ps,
4579 wide_reg,
4580 src_mcv.getReg().?.to128(),
4581 ) else try self.asmRegisterMemory(
4582 .vcvtph2ps,
4583 wide_reg,
4584 src_mcv.mem(Memory.PtrSize.fromSize(
4585 @intCast(u32, @divExact(wide_reg.bitSize(), 16)),
4586 )),
4587 );
4588 try self.asmRegisterRegister(.vsqrtps, wide_reg, wide_reg);
4589 try self.asmRegisterRegisterImmediate(
4590 .vcvtps2ph,
4591 dst_reg,
4592 wide_reg,
4593 Immediate.u(0b1_00),
4594 );
4595 break :result dst_mcv;
4596 },
4597 else => null,
4598 } else null,
4599 32 => switch (ty.vectorLen()) {
4600 1 => if (self.hasFeature(.avx)) .vsqrtss else .sqrtss,
4601 2...4 => if (self.hasFeature(.avx)) .vsqrtps else .sqrtps,
4602 5...8 => if (self.hasFeature(.avx)) .vsqrtps else null,
4603 else => null,
4604 },
4605 64 => switch (ty.vectorLen()) {
4606 1 => if (self.hasFeature(.avx)) .vsqrtsd else .sqrtsd,
4607 2 => if (self.hasFeature(.avx)) .vsqrtpd else .sqrtpd,
4608 3...4 => if (self.hasFeature(.avx)) .vsqrtpd else null,
4609 else => null,
4610 },
4611 80, 128 => null,
4612 else => unreachable,
4554 },4613 },
4555 16, 80, 128 => null,
4556 else => unreachable,4614 else => unreachable,
4557 },4615 },
4558 else => unreachable,4616 else => unreachable,
4559 },4617 })) |tag| tag else return self.fail("TODO implement airSqrt for {}", .{
4560 else => unreachable,4618 ty.fmt(self.bin_file.options.module.?),
4561 })) |tag| tag else return self.fail("TODO implement airSqrt for {}", .{4619 });
4562 ty.fmt(self.bin_file.options.module.?),4620 switch (tag) {
4563 });4621 .vsqrtss, .vsqrtsd => if (src_mcv.isRegister()) try self.asmRegisterRegisterRegister(
4564 switch (tag) {4622 tag,
4565 .vsqrtss, .vsqrtsd => if (src_mcv.isRegister()) try self.asmRegisterRegisterRegister(4623 dst_reg,
4566 tag,4624 dst_reg,
4567 dst_reg,4625 registerAlias(src_mcv.getReg().?, abi_size),
4568 dst_reg,4626 ) else try self.asmRegisterRegisterMemory(
4569 registerAlias(src_mcv.getReg().?, abi_size),4627 tag,
4570 ) else try self.asmRegisterRegisterMemory(4628 dst_reg,
4571 tag,4629 dst_reg,
4572 dst_reg,4630 src_mcv.mem(Memory.PtrSize.fromSize(abi_size)),
4573 dst_reg,4631 ),
4574 src_mcv.mem(Memory.PtrSize.fromSize(abi_size)),4632 else => if (src_mcv.isRegister()) try self.asmRegisterRegister(
4575 ),4633 tag,
4576 else => if (src_mcv.isRegister()) try self.asmRegisterRegister(4634 dst_reg,
4577 tag,4635 registerAlias(src_mcv.getReg().?, abi_size),
4578 dst_reg,4636 ) else try self.asmRegisterMemory(
4579 registerAlias(src_mcv.getReg().?, abi_size),4637 tag,
4580 ) else try self.asmRegisterMemory(4638 dst_reg,
4581 tag,4639 src_mcv.mem(Memory.PtrSize.fromSize(abi_size)),
4582 dst_reg,4640 ),
4583 src_mcv.mem(Memory.PtrSize.fromSize(abi_size)),4641 }
4584 ),4642 break :result dst_mcv;
4585 }4643 };
4586 return self.finishAir(inst, dst_mcv, .{ un_op, .none, .none });4644 return self.finishAir(inst, result, .{ un_op, .none, .none });
4587}4645}
45884646
4589fn airUnaryMath(self: *Self, inst: Air.Inst.Index) !void {4647fn airUnaryMath(self: *Self, inst: Air.Inst.Index) !void {
src/arch/x86_64/encodings.zig+2-2
...@@ -1047,9 +1047,9 @@ pub const table = [_]Entry{...@@ -1047,9 +1047,9 @@ pub const table = [_]Entry{
1047 .{ .vsqrtps, .rm, &.{ .xmm, .xmm_m128 }, &.{ 0x0f, 0x51 }, 0, .vex_128_wig, .avx },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 },1048 .{ .vsqrtps, .rm, &.{ .ymm, .ymm_m256 }, &.{ 0x0f, 0x51 }, 0, .vex_256_wig, .avx },
10491049
1050 .{ .vsqrtsd, .rvm, &.{ .xmm, .xmm, .xmm_m64 }, &.{ 0xf2, 0x0f }, 0, .vex_lig_wig, .avx },1050 .{ .vsqrtsd, .rvm, &.{ .xmm, .xmm, .xmm_m64 }, &.{ 0xf2, 0x0f, 0x51 }, 0, .vex_lig_wig, .avx },
10511051
1052 .{ .vsqrtss, .rvm, &.{ .xmm, .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f }, 0, .vex_lig_wig, .avx },1052 .{ .vsqrtss, .rvm, &.{ .xmm, .xmm, .xmm_m32 }, &.{ 0xf3, 0x0f, 0x51 }, 0, .vex_lig_wig, .avx },
10531053
1054 // F16C1054 // F16C
1055 .{ .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 },
test/behavior/floatop.zig-1
...@@ -135,7 +135,6 @@ fn testSqrt() !void {...@@ -135,7 +135,6 @@ fn testSqrt() !void {
135135
136test "@sqrt with vectors" {136test "@sqrt with vectors" {
137 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO137 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
138 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
139 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO138 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
140 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO139 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
141140