authorgravatar for jacobly@ziglang.orgJacob Young <jacobly@ziglang.org> 2023-10-08 04:04:04-04:00
committergravatar for jacobly@ziglang.orgJacob Young <jacobly@ziglang.org> 2023-10-08 04:41:55-04:00
logb5dedd7c00604bc687aab373a430058d9a97cdf9
tree1f95ea5c9953d200937d12fd8b547ccc623fde4a
parent35c9b717f752ceb1a031e0bcfc237a9c4f3bf6ad

x86_64: implement `@mulAdd` of floats for baseline


2 files changed, 144 insertions(+), 121 deletions(-)

src/arch/x86_64/CodeGen.zig+141-118
......@@ -12721,142 +12721,165 @@ fn airMulAdd(self: *Self, inst: Air.Inst.Index) !void {
1272112721 const extra = self.air.extraData(Air.Bin, pl_op.payload).data;
1272212722 const ty = self.typeOfIndex(inst);
1272312723
12724 if (!self.hasFeature(.fma)) return self.fail("TODO implement airMulAdd for {}", .{ty.fmt(mod)});
12725
1272612724 const ops = [3]Air.Inst.Ref{ extra.lhs, extra.rhs, pl_op.operand };
12727 var mcvs: [3]MCValue = undefined;
12728 var locks = [1]?RegisterManager.RegisterLock{null} ** 3;
12729 defer for (locks) |reg_lock| if (reg_lock) |lock| self.register_manager.unlockReg(lock);
12730 var order = [1]u2{0} ** 3;
12731 var unused = std.StaticBitSet(3).initFull();
12732 for (ops, &mcvs, &locks, 0..) |op, *mcv, *lock, op_i| {
12733 const op_index: u2 = @intCast(op_i);
12734 mcv.* = try self.resolveInst(op);
12735 if (unused.isSet(0) and mcv.isRegister() and self.reuseOperand(inst, op, op_index, mcv.*)) {
12736 order[op_index] = 1;
12737 unused.unset(0);
12738 } else if (unused.isSet(2) and mcv.isMemory()) {
12739 order[op_index] = 3;
12740 unused.unset(2);
12725 const result = result: {
12726 if (switch (ty.scalarType(mod).floatBits(self.target.*)) {
12727 16, 80, 128 => true,
12728 32, 64 => !self.hasFeature(.fma),
12729 else => unreachable,
12730 }) {
12731 if (ty.zigTypeTag(mod) != .Float) return self.fail("TODO implement airMulAdd for {}", .{
12732 ty.fmt(mod),
12733 });
12734
12735 var callee: ["__fma?".len]u8 = undefined;
12736 break :result try self.genCall(.{ .lib = .{
12737 .return_type = ty.toIntern(),
12738 .param_types = &.{ ty.toIntern(), ty.toIntern(), ty.toIntern() },
12739 .callee = std.fmt.bufPrint(&callee, "{s}fma{s}", .{
12740 floatLibcAbiPrefix(ty),
12741 floatLibcAbiSuffix(ty),
12742 }) catch unreachable,
12743 } }, &.{ ty, ty, ty }, &.{
12744 .{ .air_ref = extra.lhs }, .{ .air_ref = extra.rhs }, .{ .air_ref = pl_op.operand },
12745 });
1274112746 }
12742 switch (mcv.*) {
12743 .register => |reg| lock.* = self.register_manager.lockReg(reg),
12744 else => {},
12747
12748 var mcvs: [3]MCValue = undefined;
12749 var locks = [1]?RegisterManager.RegisterLock{null} ** 3;
12750 defer for (locks) |reg_lock| if (reg_lock) |lock| self.register_manager.unlockReg(lock);
12751 var order = [1]u2{0} ** 3;
12752 var unused = std.StaticBitSet(3).initFull();
12753 for (ops, &mcvs, &locks, 0..) |op, *mcv, *lock, op_i| {
12754 const op_index: u2 = @intCast(op_i);
12755 mcv.* = try self.resolveInst(op);
12756 if (unused.isSet(0) and mcv.isRegister() and self.reuseOperand(inst, op, op_index, mcv.*)) {
12757 order[op_index] = 1;
12758 unused.unset(0);
12759 } else if (unused.isSet(2) and mcv.isMemory()) {
12760 order[op_index] = 3;
12761 unused.unset(2);
12762 }
12763 switch (mcv.*) {
12764 .register => |reg| lock.* = self.register_manager.lockReg(reg),
12765 else => {},
12766 }
12767 }
12768 for (&order, &mcvs, &locks) |*mop_index, *mcv, *lock| {
12769 if (mop_index.* != 0) continue;
12770 mop_index.* = 1 + @as(u2, @intCast(unused.toggleFirstSet().?));
12771 if (mop_index.* > 1 and mcv.isRegister()) continue;
12772 const reg = try self.copyToTmpRegister(ty, mcv.*);
12773 mcv.* = .{ .register = reg };
12774 if (lock.*) |old_lock| self.register_manager.unlockReg(old_lock);
12775 lock.* = self.register_manager.lockRegAssumeUnused(reg);
1274512776 }
12746 }
12747 for (&order, &mcvs, &locks) |*mop_index, *mcv, *lock| {
12748 if (mop_index.* != 0) continue;
12749 mop_index.* = 1 + @as(u2, @intCast(unused.toggleFirstSet().?));
12750 if (mop_index.* > 1 and mcv.isRegister()) continue;
12751 const reg = try self.copyToTmpRegister(ty, mcv.*);
12752 mcv.* = .{ .register = reg };
12753 if (lock.*) |old_lock| self.register_manager.unlockReg(old_lock);
12754 lock.* = self.register_manager.lockRegAssumeUnused(reg);
12755 }
1275612777
12757 const mir_tag = @as(?Mir.Inst.FixedTag, if (mem.eql(u2, &order, &.{ 1, 3, 2 }) or
12758 mem.eql(u2, &order, &.{ 3, 1, 2 }))
12759 switch (ty.zigTypeTag(mod)) {
12760 .Float => switch (ty.floatBits(self.target.*)) {
12761 32 => .{ .v_ss, .fmadd132 },
12762 64 => .{ .v_sd, .fmadd132 },
12763 16, 80, 128 => null,
12764 else => unreachable,
12765 },
12766 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {
12767 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {
12768 32 => switch (ty.vectorLen(mod)) {
12769 1 => .{ .v_ss, .fmadd132 },
12770 2...8 => .{ .v_ps, .fmadd132 },
12771 else => null,
12772 },
12773 64 => switch (ty.vectorLen(mod)) {
12774 1 => .{ .v_sd, .fmadd132 },
12775 2...4 => .{ .v_pd, .fmadd132 },
12776 else => null,
12777 },
12778 const mir_tag = @as(?Mir.Inst.FixedTag, if (mem.eql(u2, &order, &.{ 1, 3, 2 }) or
12779 mem.eql(u2, &order, &.{ 3, 1, 2 }))
12780 switch (ty.zigTypeTag(mod)) {
12781 .Float => switch (ty.floatBits(self.target.*)) {
12782 32 => .{ .v_ss, .fmadd132 },
12783 64 => .{ .v_sd, .fmadd132 },
1277812784 16, 80, 128 => null,
1277912785 else => unreachable,
1278012786 },
12781 else => unreachable,
12782 },
12783 else => unreachable,
12784 }
12785 else if (mem.eql(u2, &order, &.{ 2, 1, 3 }) or mem.eql(u2, &order, &.{ 1, 2, 3 }))
12786 switch (ty.zigTypeTag(mod)) {
12787 .Float => switch (ty.floatBits(self.target.*)) {
12788 32 => .{ .v_ss, .fmadd213 },
12789 64 => .{ .v_sd, .fmadd213 },
12790 16, 80, 128 => null,
12791 else => unreachable,
12792 },
12793 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {
12794 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {
12795 32 => switch (ty.vectorLen(mod)) {
12796 1 => .{ .v_ss, .fmadd213 },
12797 2...8 => .{ .v_ps, .fmadd213 },
12798 else => null,
12799 },
12800 64 => switch (ty.vectorLen(mod)) {
12801 1 => .{ .v_sd, .fmadd213 },
12802 2...4 => .{ .v_pd, .fmadd213 },
12803 else => null,
12787 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {
12788 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {
12789 32 => switch (ty.vectorLen(mod)) {
12790 1 => .{ .v_ss, .fmadd132 },
12791 2...8 => .{ .v_ps, .fmadd132 },
12792 else => null,
12793 },
12794 64 => switch (ty.vectorLen(mod)) {
12795 1 => .{ .v_sd, .fmadd132 },
12796 2...4 => .{ .v_pd, .fmadd132 },
12797 else => null,
12798 },
12799 16, 80, 128 => null,
12800 else => unreachable,
1280412801 },
12805 16, 80, 128 => null,
1280612802 else => unreachable,
1280712803 },
1280812804 else => unreachable,
12809 },
12810 else => unreachable,
12811 }
12812 else if (mem.eql(u2, &order, &.{ 2, 3, 1 }) or mem.eql(u2, &order, &.{ 3, 2, 1 }))
12813 switch (ty.zigTypeTag(mod)) {
12814 .Float => switch (ty.floatBits(self.target.*)) {
12815 32 => .{ .v_ss, .fmadd231 },
12816 64 => .{ .v_sd, .fmadd231 },
12817 16, 80, 128 => null,
12818 else => unreachable,
12819 },
12820 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {
12821 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {
12822 32 => switch (ty.vectorLen(mod)) {
12823 1 => .{ .v_ss, .fmadd231 },
12824 2...8 => .{ .v_ps, .fmadd231 },
12825 else => null,
12826 },
12827 64 => switch (ty.vectorLen(mod)) {
12828 1 => .{ .v_sd, .fmadd231 },
12829 2...4 => .{ .v_pd, .fmadd231 },
12830 else => null,
12805 }
12806 else if (mem.eql(u2, &order, &.{ 2, 1, 3 }) or mem.eql(u2, &order, &.{ 1, 2, 3 }))
12807 switch (ty.zigTypeTag(mod)) {
12808 .Float => switch (ty.floatBits(self.target.*)) {
12809 32 => .{ .v_ss, .fmadd213 },
12810 64 => .{ .v_sd, .fmadd213 },
12811 16, 80, 128 => null,
12812 else => unreachable,
12813 },
12814 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {
12815 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {
12816 32 => switch (ty.vectorLen(mod)) {
12817 1 => .{ .v_ss, .fmadd213 },
12818 2...8 => .{ .v_ps, .fmadd213 },
12819 else => null,
12820 },
12821 64 => switch (ty.vectorLen(mod)) {
12822 1 => .{ .v_sd, .fmadd213 },
12823 2...4 => .{ .v_pd, .fmadd213 },
12824 else => null,
12825 },
12826 16, 80, 128 => null,
12827 else => unreachable,
1283112828 },
12829 else => unreachable,
12830 },
12831 else => unreachable,
12832 }
12833 else if (mem.eql(u2, &order, &.{ 2, 3, 1 }) or mem.eql(u2, &order, &.{ 3, 2, 1 }))
12834 switch (ty.zigTypeTag(mod)) {
12835 .Float => switch (ty.floatBits(self.target.*)) {
12836 32 => .{ .v_ss, .fmadd231 },
12837 64 => .{ .v_sd, .fmadd231 },
1283212838 16, 80, 128 => null,
1283312839 else => unreachable,
1283412840 },
12841 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {
12842 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {
12843 32 => switch (ty.vectorLen(mod)) {
12844 1 => .{ .v_ss, .fmadd231 },
12845 2...8 => .{ .v_ps, .fmadd231 },
12846 else => null,
12847 },
12848 64 => switch (ty.vectorLen(mod)) {
12849 1 => .{ .v_sd, .fmadd231 },
12850 2...4 => .{ .v_pd, .fmadd231 },
12851 else => null,
12852 },
12853 16, 80, 128 => null,
12854 else => unreachable,
12855 },
12856 else => unreachable,
12857 },
1283512858 else => unreachable,
12836 },
12837 else => unreachable,
12838 }
12839 else
12840 unreachable) orelse return self.fail("TODO implement airMulAdd for {}", .{ty.fmt(mod)});
12859 }
12860 else
12861 unreachable) orelse return self.fail("TODO implement airMulAdd for {}", .{ty.fmt(mod)});
1284112862
12842 var mops: [3]MCValue = undefined;
12843 for (order, mcvs) |mop_index, mcv| mops[mop_index - 1] = mcv;
12863 var mops: [3]MCValue = undefined;
12864 for (order, mcvs) |mop_index, mcv| mops[mop_index - 1] = mcv;
1284412865
12845 const abi_size: u32 = @intCast(ty.abiSize(mod));
12846 const mop1_reg = registerAlias(mops[0].getReg().?, abi_size);
12847 const mop2_reg = registerAlias(mops[1].getReg().?, abi_size);
12848 if (mops[2].isRegister()) try self.asmRegisterRegisterRegister(
12849 mir_tag,
12850 mop1_reg,
12851 mop2_reg,
12852 registerAlias(mops[2].getReg().?, abi_size),
12853 ) else try self.asmRegisterRegisterMemory(
12854 mir_tag,
12855 mop1_reg,
12856 mop2_reg,
12857 mops[2].mem(Memory.PtrSize.fromSize(abi_size)),
12858 );
12859 return self.finishAir(inst, mops[0], ops);
12866 const abi_size: u32 = @intCast(ty.abiSize(mod));
12867 const mop1_reg = registerAlias(mops[0].getReg().?, abi_size);
12868 const mop2_reg = registerAlias(mops[1].getReg().?, abi_size);
12869 if (mops[2].isRegister()) try self.asmRegisterRegisterRegister(
12870 mir_tag,
12871 mop1_reg,
12872 mop2_reg,
12873 registerAlias(mops[2].getReg().?, abi_size),
12874 ) else try self.asmRegisterRegisterMemory(
12875 mir_tag,
12876 mop1_reg,
12877 mop2_reg,
12878 mops[2].mem(Memory.PtrSize.fromSize(abi_size)),
12879 );
12880 break :result mops[0];
12881 };
12882 return self.finishAir(inst, result, ops);
1286012883}
1286112884
1286212885fn airVaStart(self: *Self, inst: Air.Inst.Index) !void {
test/behavior/muladd.zig+3-3
......@@ -32,11 +32,11 @@ fn testMulAdd() !void {
3232}
3333
3434test "@mulAdd f16" {
35 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
3635 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
3736 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
3837 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
3938 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
39 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;
4040
4141 try comptime testMulAdd16();
4242 try testMulAdd16();
......@@ -50,12 +50,12 @@ fn testMulAdd16() !void {
5050}
5151
5252test "@mulAdd f80" {
53 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
5453 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
5554 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
5655 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
5756 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
5857 if (builtin.zig_backend == .stage2_c and comptime builtin.cpu.arch.isArmOrThumb()) return error.SkipZigTest;
58 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;
5959
6060 try comptime testMulAdd80();
6161 try testMulAdd80();
......@@ -69,12 +69,12 @@ fn testMulAdd80() !void {
6969}
7070
7171test "@mulAdd f128" {
72 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
7372 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
7473 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
7574 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
7675 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
7776 if (builtin.zig_backend == .stage2_c and comptime builtin.cpu.arch.isArmOrThumb()) return error.SkipZigTest;
77 if (builtin.zig_backend == .stage2_x86_64 and builtin.target.ofmt != .elf) return error.SkipZigTest;
7878
7979 try comptime testMulAdd128();
8080 try testMulAdd128();