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 {...@@ -12721,142 +12721,165 @@ fn airMulAdd(self: *Self, inst: Air.Inst.Index) !void {
12721 const extra = self.air.extraData(Air.Bin, pl_op.payload).data;12721 const extra = self.air.extraData(Air.Bin, pl_op.payload).data;
12722 const ty = self.typeOfIndex(inst);12722 const ty = self.typeOfIndex(inst);
1272312723
12724 if (!self.hasFeature(.fma)) return self.fail("TODO implement airMulAdd for {}", .{ty.fmt(mod)});
12725
12726 const ops = [3]Air.Inst.Ref{ extra.lhs, extra.rhs, pl_op.operand };12724 const ops = [3]Air.Inst.Ref{ extra.lhs, extra.rhs, pl_op.operand };
12727 var mcvs: [3]MCValue = undefined;12725 const result = result: {
12728 var locks = [1]?RegisterManager.RegisterLock{null} ** 3;12726 if (switch (ty.scalarType(mod).floatBits(self.target.*)) {
12729 defer for (locks) |reg_lock| if (reg_lock) |lock| self.register_manager.unlockReg(lock);12727 16, 80, 128 => true,
12730 var order = [1]u2{0} ** 3;12728 32, 64 => !self.hasFeature(.fma),
12731 var unused = std.StaticBitSet(3).initFull();12729 else => unreachable,
12732 for (ops, &mcvs, &locks, 0..) |op, *mcv, *lock, op_i| {12730 }) {
12733 const op_index: u2 = @intCast(op_i);12731 if (ty.zigTypeTag(mod) != .Float) return self.fail("TODO implement airMulAdd for {}", .{
12734 mcv.* = try self.resolveInst(op);12732 ty.fmt(mod),
12735 if (unused.isSet(0) and mcv.isRegister() and self.reuseOperand(inst, op, op_index, mcv.*)) {12733 });
12736 order[op_index] = 1;12734
12737 unused.unset(0);12735 var callee: ["__fma?".len]u8 = undefined;
12738 } else if (unused.isSet(2) and mcv.isMemory()) {12736 break :result try self.genCall(.{ .lib = .{
12739 order[op_index] = 3;12737 .return_type = ty.toIntern(),
12740 unused.unset(2);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 });
12741 }12746 }
12742 switch (mcv.*) {12747
12743 .register => |reg| lock.* = self.register_manager.lockReg(reg),12748 var mcvs: [3]MCValue = undefined;
12744 else => {},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);
12745 }12776 }
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 }) or12778 const mir_tag = @as(?Mir.Inst.FixedTag, if (mem.eql(u2, &order, &.{ 1, 3, 2 }) or
12758 mem.eql(u2, &order, &.{ 3, 1, 2 }))12779 mem.eql(u2, &order, &.{ 3, 1, 2 }))
12759 switch (ty.zigTypeTag(mod)) {12780 switch (ty.zigTypeTag(mod)) {
12760 .Float => switch (ty.floatBits(self.target.*)) {12781 .Float => switch (ty.floatBits(self.target.*)) {
12761 32 => .{ .v_ss, .fmadd132 },12782 32 => .{ .v_ss, .fmadd132 },
12762 64 => .{ .v_sd, .fmadd132 },12783 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 16, 80, 128 => null,12784 16, 80, 128 => null,
12779 else => unreachable,12785 else => unreachable,
12780 },12786 },
12781 else => unreachable,12787 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {
12782 },12788 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {
12783 else => unreachable,12789 32 => switch (ty.vectorLen(mod)) {
12784 }12790 1 => .{ .v_ss, .fmadd132 },
12785 else if (mem.eql(u2, &order, &.{ 2, 1, 3 }) or mem.eql(u2, &order, &.{ 1, 2, 3 }))12791 2...8 => .{ .v_ps, .fmadd132 },
12786 switch (ty.zigTypeTag(mod)) {12792 else => null,
12787 .Float => switch (ty.floatBits(self.target.*)) {12793 },
12788 32 => .{ .v_ss, .fmadd213 },12794 64 => switch (ty.vectorLen(mod)) {
12789 64 => .{ .v_sd, .fmadd213 },12795 1 => .{ .v_sd, .fmadd132 },
12790 16, 80, 128 => null,12796 2...4 => .{ .v_pd, .fmadd132 },
12791 else => unreachable,12797 else => null,
12792 },12798 },
12793 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {12799 16, 80, 128 => null,
12794 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {12800 else => unreachable,
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,
12804 },12801 },
12805 16, 80, 128 => null,
12806 else => unreachable,12802 else => unreachable,
12807 },12803 },
12808 else => unreachable,12804 else => unreachable,
12809 },12805 }
12810 else => unreachable,12806 else if (mem.eql(u2, &order, &.{ 2, 1, 3 }) or mem.eql(u2, &order, &.{ 1, 2, 3 }))
12811 }12807 switch (ty.zigTypeTag(mod)) {
12812 else if (mem.eql(u2, &order, &.{ 2, 3, 1 }) or mem.eql(u2, &order, &.{ 3, 2, 1 }))12808 .Float => switch (ty.floatBits(self.target.*)) {
12813 switch (ty.zigTypeTag(mod)) {12809 32 => .{ .v_ss, .fmadd213 },
12814 .Float => switch (ty.floatBits(self.target.*)) {12810 64 => .{ .v_sd, .fmadd213 },
12815 32 => .{ .v_ss, .fmadd231 },12811 16, 80, 128 => null,
12816 64 => .{ .v_sd, .fmadd231 },12812 else => unreachable,
12817 16, 80, 128 => null,12813 },
12818 else => unreachable,12814 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {
12819 },12815 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {
12820 .Vector => switch (ty.childType(mod).zigTypeTag(mod)) {12816 32 => switch (ty.vectorLen(mod)) {
12821 .Float => switch (ty.childType(mod).floatBits(self.target.*)) {12817 1 => .{ .v_ss, .fmadd213 },
12822 32 => switch (ty.vectorLen(mod)) {12818 2...8 => .{ .v_ps, .fmadd213 },
12823 1 => .{ .v_ss, .fmadd231 },12819 else => null,
12824 2...8 => .{ .v_ps, .fmadd231 },12820 },
12825 else => null,12821 64 => switch (ty.vectorLen(mod)) {
12826 },12822 1 => .{ .v_sd, .fmadd213 },
12827 64 => switch (ty.vectorLen(mod)) {12823 2...4 => .{ .v_pd, .fmadd213 },
12828 1 => .{ .v_sd, .fmadd231 },12824 else => null,
12829 2...4 => .{ .v_pd, .fmadd231 },12825 },
12830 else => null,12826 16, 80, 128 => null,
12827 else => unreachable,
12831 },12828 },
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 },
12832 16, 80, 128 => null,12838 16, 80, 128 => null,
12833 else => unreachable,12839 else => unreachable,
12834 },12840 },
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 },
12835 else => unreachable,12858 else => unreachable,
12836 },12859 }
12837 else => unreachable,12860 else
12838 }12861 unreachable) orelse return self.fail("TODO implement airMulAdd for {}", .{ty.fmt(mod)});
12839 else
12840 unreachable) orelse return self.fail("TODO implement airMulAdd for {}", .{ty.fmt(mod)});
1284112862
12842 var mops: [3]MCValue = undefined;12863 var mops: [3]MCValue = undefined;
12843 for (order, mcvs) |mop_index, mcv| mops[mop_index - 1] = mcv;12864 for (order, mcvs) |mop_index, mcv| mops[mop_index - 1] = mcv;
1284412865
12845 const abi_size: u32 = @intCast(ty.abiSize(mod));12866 const abi_size: u32 = @intCast(ty.abiSize(mod));
12846 const mop1_reg = registerAlias(mops[0].getReg().?, abi_size);12867 const mop1_reg = registerAlias(mops[0].getReg().?, abi_size);
12847 const mop2_reg = registerAlias(mops[1].getReg().?, abi_size);12868 const mop2_reg = registerAlias(mops[1].getReg().?, abi_size);
12848 if (mops[2].isRegister()) try self.asmRegisterRegisterRegister(12869 if (mops[2].isRegister()) try self.asmRegisterRegisterRegister(
12849 mir_tag,12870 mir_tag,
12850 mop1_reg,12871 mop1_reg,
12851 mop2_reg,12872 mop2_reg,
12852 registerAlias(mops[2].getReg().?, abi_size),12873 registerAlias(mops[2].getReg().?, abi_size),
12853 ) else try self.asmRegisterRegisterMemory(12874 ) else try self.asmRegisterRegisterMemory(
12854 mir_tag,12875 mir_tag,
12855 mop1_reg,12876 mop1_reg,
12856 mop2_reg,12877 mop2_reg,
12857 mops[2].mem(Memory.PtrSize.fromSize(abi_size)),12878 mops[2].mem(Memory.PtrSize.fromSize(abi_size)),
12858 );12879 );
12859 return self.finishAir(inst, mops[0], ops);12880 break :result mops[0];
12881 };
12882 return self.finishAir(inst, result, ops);
12860}12883}
1286112884
12862fn airVaStart(self: *Self, inst: Air.Inst.Index) !void {12885fn airVaStart(self: *Self, inst: Air.Inst.Index) !void {
test/behavior/muladd.zig+3-3
...@@ -32,11 +32,11 @@ fn testMulAdd() !void {...@@ -32,11 +32,11 @@ fn testMulAdd() !void {
32}32}
3333
34test "@mulAdd f16" {34test "@mulAdd f16" {
35 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
36 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO35 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
37 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO36 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
38 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO37 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
39 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;38 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
41 try comptime testMulAdd16();41 try comptime testMulAdd16();
42 try testMulAdd16();42 try testMulAdd16();
...@@ -50,12 +50,12 @@ fn testMulAdd16() !void {...@@ -50,12 +50,12 @@ fn testMulAdd16() !void {
50}50}
5151
52test "@mulAdd f80" {52test "@mulAdd f80" {
53 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
54 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO53 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
55 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO54 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
56 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO55 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
57 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;56 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
58 if (builtin.zig_backend == .stage2_c and comptime builtin.cpu.arch.isArmOrThumb()) return error.SkipZigTest;57 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
60 try comptime testMulAdd80();60 try comptime testMulAdd80();
61 try testMulAdd80();61 try testMulAdd80();
...@@ -69,12 +69,12 @@ fn testMulAdd80() !void {...@@ -69,12 +69,12 @@ fn testMulAdd80() !void {
69}69}
7070
71test "@mulAdd f128" {71test "@mulAdd f128" {
72 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest; // TODO
73 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO72 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
74 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO73 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest; // TODO
75 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO74 if (builtin.zig_backend == .stage2_sparc64) return error.SkipZigTest; // TODO
76 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;75 if (builtin.zig_backend == .stage2_spirv64) return error.SkipZigTest;
77 if (builtin.zig_backend == .stage2_c and comptime builtin.cpu.arch.isArmOrThumb()) return error.SkipZigTest;76 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
79 try comptime testMulAdd128();79 try comptime testMulAdd128();
80 try testMulAdd128();80 try testMulAdd128();