diff --git a/AGENTS.md b/AGENTS.md index 6c175db..63b74f0 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -64,6 +64,8 @@ metajit.cpp uses a generating extension for tracing. ## Coding Guidelines - Do not add code comments. +- NEVER push directly to main under any circumstances. +- Start commit messages with a lowercase letter. - Name feature branches `-`, using the author's first name followed by a dash and a short description with words separated by dashes. - Never edit the generated jitir.hpp and jitir_llvmapi.hpp files directly. Instead, edit the corresponding template files jitir.tmpl.hpp and jitir_llvmapi.tmpl.hpp. The instructions are specified in the jitir.py generator script. - Never edit any files in tests/output. They are just output files from the unit tests used to debug failing test cases. They are also not golden tests; in fact, they are ignored by Git. diff --git a/jitir.tmpl.hpp b/jitir.tmpl.hpp index 8fca68f..c2cdac0 100644 --- a/jitir.tmpl.hpp +++ b/jitir.tmpl.hpp @@ -222,15 +222,21 @@ namespace metajit { }; } +namespace metajit { + inline const char* to_string(Type type) { + static const char* names[] = { + "Void", + "Bool", + "Int8", "Int16", "Int32", "Int64", + "Float32", "Float64", + "Ptr" + }; + return names[(size_t) type]; + } +} + std::ostream& operator<<(std::ostream& stream, metajit::Type type) { - static const char* names[] = { - "Void", - "Bool", - "Int8", "Int16", "Int32", "Int64", - "Float32", "Float64", - "Ptr" - }; - stream << names[(size_t) type]; + stream << metajit::to_string(type); return stream; } diff --git a/llvmgen.hpp b/llvmgen.hpp index c7e9ba7..a09899b 100644 --- a/llvmgen.hpp +++ b/llvmgen.hpp @@ -256,9 +256,10 @@ namespace metajit { return call_inst; } else if (dynmatch(EqInst, eq, inst)) { if (is_float(eq->arg(0)->type())) { - return _builder.CreateFCmpUEQ( - emit_arg(eq->arg(0)), - emit_arg(eq->arg(1)) + llvm::Type* bits_type = _builder.getIntNTy(type_size(eq->arg(0)->type()) * 8); + return _builder.CreateICmpEQ( + _builder.CreateBitCast(emit_arg(eq->arg(0)), bits_type), + _builder.CreateBitCast(emit_arg(eq->arg(1)), bits_type) ); } else { return _builder.CreateICmpEQ( diff --git a/tests/test_cfg.cpp b/tests/test_cfg.cpp index b2866cd..febebea 100644 --- a/tests/test_cfg.cpp +++ b/tests/test_cfg.cpp @@ -24,6 +24,37 @@ int main(int argc, char** argv) { DiffTestSuite suite("tests/output/test_cfg", argc, argv); + for (Type type : {Type::Float32, Type::Float64, Type::Int32, Type::Int64}) { + suite.test(std::string("entry_argument_") + to_string(type)).run([type]() { + for (auto mode : {X86CodeGen::Mode::JIT, X86CodeGen::Mode::AOT}) { + Context context; + Allocator allocator; + Section* section = new Section(context, allocator); + Builder builder(section); + Block* entry = builder.build_block({type, Type::Ptr}); + builder.move_to_end(entry); + builder.build_store(entry->arg(1), entry->arg(0), AliasingGroup(0), 0); + builder.build_exit(); + section->autoname(); + section->set_ordering(BlockOrdering::Natural); + + uint64_t bits = 0x123456789abcdef0; + uint64_t result = 0; + if (is_float(type)) { + X86CodeGen codegen(section, {Reg::X86_R12(), Reg::X86_R13()}, mode); + using Func = void(* [[clang::preserve_none]])(uint64_t, uint64_t*); + ((Func) codegen.deploy())(bits, &result); + } else { + X86CodeGen codegen(section, {Reg::X86_XMM(0), Reg::X86_R12()}, mode); + using Func = void(* [[clang::preserve_none]])(double, uint64_t*); + ((Func) codegen.deploy())(bit_cast(bits), &result); + } + unittest_assert(result == (bits & type_mask(type))); + delete section; + } + }); + } + suite.diff_test("entry_argument_spilled_before_first_use").aot(false).run([](Builder& builder, TestData& data) { std::vector values; for (size_t index = 0; index < 32; index++) { @@ -39,6 +70,44 @@ int main(int argc, char** argv) { data.output(input); }); + for (Type type : {Type::Float32, Type::Float64}) { + suite.diff_test(std::string("float_block_argument_") + to_string(type)).run([type](Builder& builder, TestData& data) { + Block* a = builder.build_block(); + Block* b = builder.build_block(); + Block* cont = builder.build_block({type}); + Value* cond = data.input(Type::Bool); + Value* value_a = data.input(type); + Value* value_b = data.input(type); + builder.build_branch(cond, a, b); + builder.move_to_end(a); + builder.build_jump(cont, {value_a}); + builder.move_to_end(b); + builder.build_jump(cont, {value_b}); + builder.move_to_end(cont); + data.output(cont->arg(0)); + }); + + suite.diff_test(std::string("float_swap_loop_") + to_string(type)).run([type](Builder& builder, TestData& data) { + Block* header = builder.build_block({Type::Bool, type, type}); + Block* body = builder.build_block(); + Block* end = builder.build_block(); + Value* a = data.input(type); + Value* b = data.input(type); + Value* cond = data.input(Type::Bool); + builder.build_jump(header, {cond, a, b}); + builder.move_to_end(header); + builder.build_branch(header->arg(0), body, end); + builder.move_to_end(body); + builder.build_jump(header, { + builder.build_const(Type::Bool, false), + header->arg(2), header->arg(1) + }); + builder.move_to_end(end); + data.output(header->arg(1)); + data.output(header->arg(2)); + }); + } + suite.diff_test("branch").run([](Builder& builder, TestData& data) { Block* a = builder.build_block(); Block* b = builder.build_block(); diff --git a/tests/test_insts.cpp b/tests/test_insts.cpp index 5ccd176..373ab73 100644 --- a/tests/test_insts.cpp +++ b/tests/test_insts.cpp @@ -44,8 +44,21 @@ void test_binop(DiffTestSuite& suite) { binop(xor, true) binop(eq, false) + binop_type(eq, Float32) + binop_type(eq, Float64) binop(lt_u, false) binop(lt_s, false) + + for (Type type : {Type::Float32, Type::Float64}) { + suite.diff_test(std::string("eq_bits_") + to_string(type)).run([type](Builder& builder, TestData& data) { + Value* zero = data.input(RandomRange(type, 0, 0)); + uint64_t sign = uint64_t(1) << (type_size(type) * 8 - 1); + Value* negative_zero = data.input(RandomRange(type, sign, sign)); + Value* nan = data.input(RandomRange(type, type_mask(type), type_mask(type))); + data.output(builder.build_eq(zero, negative_zero)); + data.output(builder.build_eq(nan, nan)); + }); + } } void test_shift(DiffTestSuite& suite) { @@ -104,6 +117,8 @@ void test_select(DiffTestSuite& suite) { select_type(Int16) select_type(Int32) select_type(Int64) + select_type(Float32) + select_type(Float64) } void test_resize(DiffTestSuite& suite) { @@ -251,21 +266,15 @@ void test_assume_const(DiffTestSuite& suite) { } void test_alloca(DiffTestSuite& suite) { - suite.diff_test("alloca_store_load_Int32").run([](Builder& builder, TestData& data) { - Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 4), 4); - Value* val = data.input(Type::Int32); - builder.build_store(ptr, val, AliasingGroup(0), 0); - Value* loaded = builder.build_load(ptr, Type::Int32, LoadFlags::None, AliasingGroup(0), 0); - data.output(loaded); - }); - - suite.diff_test("alloca_store_load_Int64").run([](Builder& builder, TestData& data) { - Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 8), 8); - Value* val = data.input(Type::Int64); - builder.build_store(ptr, val, AliasingGroup(0), 0); - Value* loaded = builder.build_load(ptr, Type::Int64, LoadFlags::None, AliasingGroup(0), 0); - data.output(loaded); - }); + for (Type type : {Type::Int32, Type::Int64, Type::Float32, Type::Float64}) { + suite.diff_test(std::string("alloca_store_load_") + to_string(type)).run([=](Builder& builder, TestData& data) { + Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, type_size(type)), type_size(type)); + Value* val = data.input(type); + builder.build_store(ptr, val, AliasingGroup(0), 0); + Value* loaded = builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 0); + data.output(loaded); + }); + } suite.diff_test("alloca_multiple_stores").run([](Builder& builder, TestData& data) { Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 8), 8); @@ -456,7 +465,76 @@ void test_call(DiffTestSuite& suite) { }); } +template +__attribute__((noinline)) F test_call_fp_default(uint64_t i, float a, double b) { + return (F) i + (F) a + (F) b; +} + +template +__attribute__((preserve_none, noinline)) F test_call_fp_preserve_none(uint64_t i, float a, double b) { + return (F) i + (F) a + (F) b; +} + +template +void test_call_fp(DiffTestSuite& suite, Type type) { + for (CallConv call_conv : {CallConv::Default, CallConv::PreserveNone}) { + std::ostringstream name; + name << "call_fp_" << type; + if (call_conv == CallConv::Default) { + name << "_default"; + } else { + name << "_preserve_none"; + } + + suite.diff_test(name.str()).aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { + Value* i = data.input(RandomRange(Type::Int64, 0, 100)); + Value* a = data.input(Type::Float32); + Value* b = data.input(Type::Float64); + + Value* callee; + if (call_conv == CallConv::Default) { + callee = builder.build_const(Type::Ptr, (uint64_t)(void*) test_call_fp_default); + } else { + callee = builder.build_const(Type::Ptr, (uint64_t)(void*) test_call_fp_preserve_none); + } + Value* first = builder.build_call(callee, type, {i, a, b}, call_conv); + Value* second = builder.build_call(callee, type, {i, a, b}, call_conv); + + data.output(first); + data.output(second); + data.output(i); + data.output(a); + data.output(b); + }); + } +} + void test_binop_f(DiffTestSuite& suite) { + for (Type type : {Type::Float32, Type::Float64}) { + suite.diff_test(std::string("float_spills_") + to_string(type)).aot(false).run([type](Builder& builder, TestData& data) { + std::vector values; + for (size_t index = 0; index < 17; index++) { + values.push_back(data.input(type)); + } + for (Value* value : values) { + data.output(value); + } + }); + + suite.diff_test(std::string("mixed_register_classes_") + to_string(type)).run([type](Builder& builder, TestData& data) { + std::vector floats; + std::vector integers; + for (size_t index = 0; index < 8; index++) { + floats.push_back(data.input(type)); + integers.push_back(data.input(Type::Int64)); + } + for (size_t index = 0; index < floats.size(); index++) { + data.output(floats[index]); + data.output(integers[index]); + } + }); + } + #define binop_f_type(name, type) \ suite.diff_test(#name "_" #type).run([](Builder& builder, TestData& data) { \ data.output(builder.build_##name(data.input(Type::type), data.input(Type::type))); \ @@ -576,6 +654,8 @@ int main(int argc, char** argv) { test_assume_const(suite); test_alloca(suite); test_call(suite); + test_call_fp(suite, Type::Float32); + test_call_fp(suite, Type::Float64); test_binop_f(suite); test_convert_f(suite); test_ptr_to_int(suite); diff --git a/x86gen.hpp b/x86gen.hpp index 44f549a..7cef25c 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -24,6 +24,19 @@ #include "jitir.hpp" namespace metajit { + class RegClass { + uint32_t _mask; + public: + constexpr explicit RegClass(uint32_t mask = 0): _mask(mask) {} + + static constexpr RegClass X86_INT() { return RegClass(0xffff); } + static constexpr RegClass X86_FLOAT() { return RegClass(0xffff0000); } + + constexpr uint32_t mask() const { return _mask; } + + constexpr bool operator==(RegClass other) const { return _mask == other._mask; } + }; + class Reg { public: enum class Kind { @@ -56,6 +69,7 @@ namespace metajit { static constexpr Reg X86_R13() { return phys(13); } static constexpr Reg X86_R14() { return phys(14); } static constexpr Reg X86_R15() { return phys(15); } + static constexpr Reg X86_XMM(size_t index) { return phys(16 + index); } static constexpr Reg virt(size_t id) { return Reg(Kind::Virtual, id); @@ -436,6 +450,7 @@ namespace metajit { lwir::Span _arg_regs; lwir::Span _preserved_regs; Reg _ret_reg; + Reg _fp_ret_reg; static constexpr Reg preserve_none_arg_regs[] = { Reg::X86_R12(), Reg::X86_R13(), Reg::X86_R14(), Reg::X86_R15(), @@ -454,6 +469,11 @@ namespace metajit { Reg::X86_RBX(), Reg::X86_RBP(), Reg::X86_R12(), Reg::X86_R13(), Reg::X86_R14(), Reg::X86_R15() }; + + static constexpr Reg fp_arg_regs[] = { + Reg::X86_XMM(0), Reg::X86_XMM(1), Reg::X86_XMM(2), Reg::X86_XMM(3), + Reg::X86_XMM(4), Reg::X86_XMM(5), Reg::X86_XMM(6), Reg::X86_XMM(7) + }; public: CallConvInfo(CallConv call_conv) { switch (call_conv) { @@ -461,23 +481,33 @@ namespace metajit { _arg_regs = lwir::Span(preserve_none_arg_regs, sizeof(preserve_none_arg_regs) / sizeof(preserve_none_arg_regs[0])); _preserved_regs = lwir::Span(preserve_none_preserved_regs, sizeof(preserve_none_preserved_regs) / sizeof(preserve_none_preserved_regs[0])); _ret_reg = Reg::X86_RAX(); + _fp_ret_reg = Reg::X86_XMM(0); break; case CallConv::Default: _arg_regs = lwir::Span(default_arg_regs, sizeof(default_arg_regs) / sizeof(default_arg_regs[0])); _preserved_regs = lwir::Span(default_preserved_regs, sizeof(default_preserved_regs) / sizeof(default_preserved_regs[0])); _ret_reg = Reg::X86_RAX(); + _fp_ret_reg = Reg::X86_XMM(0); break; default: assert(false && "Unsupported calling convention"); } } - const lwir::Span& args() const { return _arg_regs; } + Reg int_arg(size_t index) const { return _arg_regs.at(index); } + Reg float_arg(size_t index) const { + return lwir::Span(fp_arg_regs, sizeof(fp_arg_regs) / sizeof(fp_arg_regs[0])).at(index); + } const lwir::Span& preserved() const { return _preserved_regs; } - Reg arg(size_t index) const { return _arg_regs.at(index); } Reg preserved(size_t index) const { return _preserved_regs.at(index); } - Reg ret() const { return _ret_reg; } + Reg ret(RegClass reg_class = RegClass::X86_INT()) const { + if (reg_class == RegClass::X86_FLOAT()) { + return _fp_ret_reg; + } else { + return _ret_reg; + } + } // TODO: Optimize @@ -489,15 +519,6 @@ namespace metajit { } return false; } - - bool is_arg(Reg reg, size_t arg_count) const { - for (size_t it = 0; it < arg_count && it < _arg_regs.size(); it++) { - if (reg == _arg_regs.at(it)) { - return true; - } - } - return false; - } }; class SymbolMap { @@ -530,6 +551,23 @@ namespace metajit { Timer peephole; }; private: + static RegClass reg_class(Reg preg) { + assert(preg.is_physical()); + if (preg.id() < Reg::X86_XMM(0).id()) { + return RegClass::X86_INT(); + } else { + return RegClass::X86_FLOAT(); + } + } + + static RegClass reg_class(Type type) { + if (is_float(type)) { + return RegClass::X86_FLOAT(); + } else { + return RegClass::X86_INT(); + } + } + struct Interval { size_t min = 0; size_t max = 0; @@ -559,6 +597,7 @@ namespace metajit { }; struct VRegInfo { + RegClass reg_class = RegClass::X86_INT(); Reg fixed; Interval interval; Reg current_reg; @@ -576,6 +615,7 @@ namespace metajit { NameMap _memory_deps; NameMap _vregs; + lwir::Span _input_vregs; std::vector _vreg_info; #ifdef METAJIT_STATS @@ -612,15 +652,19 @@ namespace metajit { } } - Reg vreg() { + Reg vreg(RegClass reg_class = RegClass::X86_INT()) { size_t id = _vreg_info.size(); _vreg_info.emplace_back(); + _vreg_info.back().reg_class = reg_class; return Reg::virt(id); } + Reg vreg(Type type) { return vreg(reg_class(type)); } + Reg fix_to_preg(Reg vreg, Reg preg) { assert(vreg.is_virtual()); VRegInfo& info = _vreg_info[vreg.id()]; + assert(info.reg_class == reg_class(preg)); info.fixed = preg; return vreg; } @@ -638,7 +682,7 @@ namespace metajit { Reg vreg(Value* value) { if (dynmatch(Const, constant, value)) { - Reg reg = vreg(); + Reg reg = vreg(value->type()); switch (constant->type()) { case Type::Bool: case Type::Int8: _builder.mov8_imm(reg, constant->value()); break; @@ -681,7 +725,7 @@ namespace metajit { } else if (value->is_named()) { NamedValue* named = (NamedValue*) value; if (_vregs.at(named).is_invalid()) { - _vregs[named] = vreg(); + _vregs[named] = vreg(value->type()); } return _vregs.at(named); } else { @@ -690,6 +734,27 @@ namespace metajit { } } + void move(Reg dst, Reg src) { + RegClass dst_class; + if (dst.is_virtual()) { + dst_class = _vreg_info[dst.id()].reg_class; + } else { + dst_class = reg_class(dst); + } + RegClass src_class; + if (src.is_virtual()) { + src_class = _vreg_info[src.id()].reg_class; + } else { + src_class = reg_class(src); + } + assert(dst_class == src_class); + if (dst_class == RegClass::X86_FLOAT()) { + _builder.movsd(dst, src); + } else { + _builder.mov64(dst, src); + } + } + void build_add(Reg dst, Value* a, Value* b) { X86Inst::Mem mem; if (dynmatch(Const, constant_b, b)) { @@ -727,6 +792,19 @@ namespace metajit { } void build_cmp(Value* a, Value* b) { + if (is_float(a->type())) { + Reg bits_a = vreg(); + Reg bits_b = vreg(); + _builder.movq_to_gp(bits_a, vreg(a)); + _builder.movq_to_gp(bits_b, vreg(b)); + if (a->type() == Type::Float32) { + _builder.cmp32(bits_a, bits_b); + } else { + _builder.cmp64(bits_a, bits_b); + } + return; + } + if (dynmatch(Const, constant_b, b)) { if (is_sext_imm32(constant_b)) { switch (type_size(a->type())) { @@ -775,11 +853,11 @@ namespace metajit { void isel(Inst* inst, Block* block) { if (dynmatch(FreezeInst, freeze, inst)) { - _builder.mov64(vreg(inst), vreg(freeze->arg(0))); + move(vreg(inst), vreg(freeze->arg(0))); } else if (dynmatch(PromoteInst, promote, inst)) { - _builder.mov64(vreg(inst), vreg(promote->arg(0))); + move(vreg(inst), vreg(promote->arg(0))); } else if (dynmatch(AssumeConstInst, assume_const, inst)) { - _builder.mov64(vreg(inst), vreg(assume_const->arg(0))); + move(vreg(inst), vreg(assume_const->arg(0))); } else if (dynmatch(PtrToIntInst, ptr_to_int, inst)) { switch (type_size(ptr_to_int->type())) { case 1: _builder.movzx8to64(vreg(inst), vreg(ptr_to_int->arg(0))); break; @@ -789,8 +867,17 @@ namespace metajit { default: assert(false && "Unsupported pointer conversion type"); } } else if (dynmatch(SelectInst, select, inst)) { - _builder.mov64(vreg(inst), vreg(select->arg(2))); - build_cmov(vreg(inst), select->cond(), vreg(select->arg(1))); + if (is_float(select->type())) { + Reg res = vreg(); + Reg then = vreg(); + _builder.movq_to_gp(res, vreg(select->arg(2))); + _builder.movq_to_gp(then, vreg(select->arg(1))); + build_cmov(res, select->cond(), then); + _builder.movq(vreg(inst), res); + } else { + _builder.mov64(vreg(inst), vreg(select->arg(2))); + build_cmov(vreg(inst), select->cond(), vreg(select->arg(1))); + } } else if (dynmatch(ResizeUInst, resize_u, inst)) { if (resize_u->arg(0)->type() == Type::Bool) { _builder.mov64(vreg(inst), vreg(resize_u->arg(0))); @@ -1208,22 +1295,31 @@ namespace metajit { CallConvInfo info(call->call_conv()); assert(call->arg_count() >= 1); - assert(call->arg_count() - 1 <= info.args().size() && "Call with too many register arguments"); lwir::Span args = _builder.alloc_regs(call->args().size() - 1); + size_t int_index = 0; + size_t float_index = 0; for (size_t it = 1; it < call->args().size(); it++) { - Reg arg_reg = fix_to_preg(vreg(), info.arg(it - 1)); - _builder.mov64(arg_reg, vreg(call->arg(it))); + Type type = call->arg(it)->type(); + Reg preg; + if (is_float(type)) { + preg = info.float_arg(float_index++); + } else { + preg = info.int_arg(int_index++); + } + Reg arg_reg = fix_to_preg(vreg(type), preg); + move(arg_reg, vreg(call->arg(it))); args[it - 1] = arg_reg; } - Reg ret_reg = fix_to_preg(vreg(), info.ret()); + RegClass ret_class = reg_class(call->type()); + Reg ret_reg = fix_to_preg(vreg(call->type()), info.ret(ret_class)); Reg callee_reg = fix_to_preg(vreg(), Reg::X86_R10()); _builder.mov64(callee_reg, vreg(call->callee())); _builder.call(callee_reg, ret_reg, call->call_conv(), args); if (call->type() != Type::Void) { - _builder.mov64(vreg(call), ret_reg); + move(vreg(call), ret_reg); } _stack_offset_alloc.require_call_alignment(); @@ -1275,11 +1371,11 @@ namespace metajit { } else if (dynmatch(JumpInst, jump, inst)) { lwir::Span copies = _builder.alloc_regs(jump->block()->args().size()); for (Arg* arg : jump->block()->args()) { - copies[arg->index()] = vreg(); - _builder.mov64(copies[arg->index()], vreg(jump->arg(arg->index()))); + copies[arg->index()] = vreg(arg->type()); + move(copies[arg->index()], vreg(jump->arg(arg->index()))); } for (Arg* arg : jump->block()->args()) { - _builder.mov64(vreg(arg), copies[arg->index()]); + move(vreg(arg), copies[arg->index()]); } _builder.jmp(_blocks[jump->block()->name()]); } else if (dynmatch(ExitInst, exit, inst)) { @@ -1419,15 +1515,15 @@ namespace metajit { class RegFileState { private: std::vector _regs; - uint16_t _free = 0xffff; - uint16_t _max_free = 0xffff; + uint32_t _free = 0xffffffff; + uint32_t _max_free = 0xffffffff; std::vector _lru; size_t _lru_count = 0; public: RegFileState() { - _regs.resize(16, Reg()); - _lru.resize(16, 0); + _regs.resize(32, Reg()); + _lru.resize(32, 0); disable(Reg::X86_RSP()); disable(Reg::X86_RBP()); @@ -1441,7 +1537,7 @@ namespace metajit { void disable(Reg preg) { assert(preg.is_physical()); - _max_free &= ~(1 << preg.id()); + _max_free &= ~(1u << preg.id()); _lru[preg.id()] = ~size_t(0); } @@ -1453,7 +1549,7 @@ namespace metajit { void set(Reg preg, Reg vreg) { assert(preg.is_physical() && vreg.is_virtual()); _regs[preg.id()] = vreg; - _free &= ~(1 << preg.id()); + _free &= ~(1u << preg.id()); } void touch(Reg preg) { @@ -1467,36 +1563,39 @@ namespace metajit { void free(Reg preg) { assert(preg.is_physical()); _regs[preg.id()] = Reg(); - _free |= (1 << preg.id()) & _max_free; + _free |= (1u << preg.id()) & _max_free; } bool is_free(Reg preg) const { assert(preg.is_physical()); - return (_free & (1 << preg.id())) != 0; + return (_free & (1u << preg.id())) != 0; } bool is_disabled(Reg preg) const { assert(preg.is_physical()); - return (_max_free & (1 << preg.id())) == 0; + return (_max_free & (1u << preg.id())) == 0; } - Reg get_free_reg() { - if (_free == 0) { + Reg get_free_reg(RegClass reg_class) { + uint32_t free = _free & reg_class.mask(); + if (free == 0) { return Reg(); } else { - return Reg::phys(__builtin_ctz(_free)); + return Reg::phys(__builtin_ctz(free)); } } - Reg get_lru() { + Reg get_lru(RegClass reg_class) { size_t min_index = 0; size_t min_value = ~size_t(0); - for (size_t it = 0; it < _lru.size(); it++) { + for (uint32_t mask = reg_class.mask(); mask != 0; mask &= mask - 1) { + size_t it = __builtin_ctz(mask); if (_lru[it] < min_value) { min_value = _lru[it]; min_index = it; } } + assert(min_value != ~size_t(0)); return Reg::phys(min_index); } @@ -1518,7 +1617,7 @@ namespace metajit { for (size_t it = 0; it < _regs.size(); it++) { _regs[it] = state[it]; if (!_regs[it].is_invalid()) { - _free &= ~(1 << it); + _free &= ~(1u << it); } } assert_invariant(); @@ -1563,10 +1662,10 @@ namespace metajit { Reg vreg = reg_file[preg]; if (vreg.is_virtual()) { VRegInfo& info = _vreg_info[vreg.id()]; - Reg free_reg = reg_file.get_free_reg(); + Reg free_reg = reg_file.get_free_reg(info.reg_class); if (allow_spill_to_reg && free_reg.is_physical()) { // No need to spill, just move to free reg - _builder.mov64(free_reg, preg); + move(free_reg, preg); reg_file.free(preg); info.current_reg = free_reg; reg_file.set(free_reg, vreg); @@ -1574,13 +1673,12 @@ namespace metajit { if (info.stack_offset == ~size_t(0)) { info.stack_offset = _stack_offset_alloc.alloc(); } - _builder.mov64_mem( - X86Inst::Mem( - Reg::X86_RSP(), - (int32_t) info.stack_offset - ), - preg - ); + X86Inst::Mem mem(Reg::X86_RSP(), (int32_t) info.stack_offset); + if (info.reg_class == RegClass::X86_FLOAT()) { + _builder.movsd_mem(mem, preg); + } else { + _builder.mov64_mem(mem, preg); + } reg_file.free(preg); info.current_reg = Reg(); } @@ -1592,17 +1690,16 @@ namespace metajit { VRegInfo& info = _vreg_info[vreg.id()]; if (info.current_reg.is_physical()) { // No need to unspill, just move from current reg - _builder.mov64(preg, info.current_reg); + move(preg, info.current_reg); reg_file.free(info.current_reg); } else { assert(info.stack_offset != ~size_t(0)); - _builder.mov64( - preg, - X86Inst::Mem( - Reg::X86_RSP(), - (int32_t) info.stack_offset - ) - ); + X86Inst::Mem mem(Reg::X86_RSP(), (int32_t) info.stack_offset); + if (info.reg_class == RegClass::X86_FLOAT()) { + _builder.movsd(preg, mem); + } else { + _builder.mov64(preg, mem); + } } info.current_reg = preg; reg_file.set(preg, vreg); @@ -1627,9 +1724,15 @@ namespace metajit { } } + static bool is_reg_mov(X86Inst* inst) { + return (inst->kind() == X86Inst::Kind::Mov64 || + inst->kind() == X86Inst::Kind::MovSS || + inst->kind() == X86Inst::Kind::MovSD) && + std::holds_alternative(inst->rm()); + } + bool is_foldable_mov(X86Inst* inst) { - if (inst->kind() == X86Inst::Kind::Mov64 && - std::holds_alternative(inst->rm())) { + if (is_reg_mov(inst)) { Reg src = std::get(inst->rm()); Reg dst = inst->reg(); assert(src.is_virtual() && dst.is_virtual()); @@ -1674,8 +1777,7 @@ namespace metajit { // Since they are used for block arguments, the register may // be def-only even if the instruction is not the first in // the register's live interval. - if (inst->kind() == X86Inst::Kind::Mov64 && - std::holds_alternative(inst->rm())) { + if (is_reg_mov(inst)) { Reg src = std::get(inst->rm()); Reg dst = inst->reg(); if (reg != src && reg == dst) { @@ -1713,10 +1815,11 @@ namespace metajit { std::fill(initial_state, initial_state + reg_file.size(), Reg()); for (Arg* arg : _section->entry()->args()) { - VRegInfo& info = _vreg_info[vreg(arg).id()]; + Reg input = _input_vregs.at(arg->index()); + VRegInfo& info = _vreg_info[input.id()]; info.interval.incl(0); assert(info.fixed.is_physical() && "Entry arguments must be in fixed registers"); - initial_state[info.fixed.id()] = vreg(arg); + initial_state[info.fixed.id()] = input; } _blocks[0]->set_regalloc(initial_state); @@ -1760,7 +1863,7 @@ namespace metajit { Reg preg = Reg::phys(it); if (!reg_file.is_free(preg) && !info.is_preserved(preg) && - !info.is_arg(preg, data->args.size()) && + std::find(data->args.begin(), data->args.end(), reg_file[preg]) == data->args.end() && reg_file[preg] != std::get(inst->rm())) { spill(reg_file, preg, false); } @@ -1782,8 +1885,10 @@ namespace metajit { } } - reg_file.set(info.ret(), data->ret); - reg_file.touch(info.ret()); + VRegInfo& ret_info = _vreg_info[data->ret.id()]; + ret_info.current_reg = ret_info.fixed; + reg_file.set(ret_info.fixed, data->ret); + reg_file.touch(ret_info.fixed); it++; continue; @@ -1806,9 +1911,9 @@ namespace metajit { inst->visit_regs([&](Reg reg) { VRegInfo& info = _vreg_info[reg.id()]; if (info.current_reg.is_invalid() && !info.fixed.is_physical()) { - Reg preg = reg_file.get_free_reg(); + Reg preg = reg_file.get_free_reg(info.reg_class); if (!preg.is_physical()) { - preg = reg_file.get_lru(); + preg = reg_file.get_lru(info.reg_class); } spill_and_unspill(reg_file, preg, reg, is_def_only(reg, inst)); } @@ -1950,8 +2055,7 @@ namespace metajit { incoming[target->name()].insert(block); } - if (inst->kind() == X86Inst::Kind::Mov64 && - std::holds_alternative(inst->rm())) { + if (is_reg_mov(inst)) { Reg src = std::get(inst->rm()); Reg dst = inst->reg(); if (src.is_virtual() && dst.is_virtual()) { @@ -1981,7 +2085,7 @@ namespace metajit { for (Reg reg : live) { for (Reg other : live) { - if (reg != other) { + if (reg != other && _vreg_info[reg.id()].reg_class == _vreg_info[other.id()].reg_class) { conflicts.at(reg.id()).insert(other); } } @@ -2057,14 +2161,14 @@ namespace metajit { Reg reg = order.at(it); assert(reg.is_virtual()); - uint16_t free_mask = 0xffff; - free_mask &= ~(1 << Reg::X86_RSP().id()); - free_mask &= ~(1 << Reg::X86_RBP().id()); + uint32_t free_mask = _vreg_info[reg.id()].reg_class.mask(); + free_mask &= ~(1u << Reg::X86_RSP().id()); + free_mask &= ~(1u << Reg::X86_RBP().id()); for (Reg conflict : conflicts.at(reg.id())) { if (std::holds_alternative(assigned.at(conflict.id()))) { Reg assigned_reg = std::get(assigned.at(conflict.id())); if (assigned_reg.is_physical()) { - free_mask &= ~(1 << assigned_reg.id()); + free_mask &= ~(1u << assigned_reg.id()); } } } @@ -2093,19 +2197,19 @@ namespace metajit { if (std::holds_alternative(assigned.at(merge_reg.id())) && std::get(assigned.at(merge_reg.id())).is_physical()) { Reg assigned_merge_reg = std::get(assigned.at(merge_reg.id())); - if ((free_mask & (1 << assigned_merge_reg.id())) != 0) { + if ((free_mask & (1u << assigned_merge_reg.id())) != 0) { assigned.at(reg.id()) = assigned_merge_reg; merged = true; break; } } - uint16_t merge_free_mask = free_mask; + uint32_t merge_free_mask = free_mask; for (Reg conflict : conflicts.at(merge_reg.id())) { if (std::holds_alternative(assigned.at(conflict.id()))) { Reg assigned_conflict = std::get(assigned.at(conflict.id())); if (assigned_conflict.is_physical()) { - merge_free_mask &= ~(1 << assigned_conflict.id()); + merge_free_mask &= ~(1u << assigned_conflict.id()); } } } @@ -2143,8 +2247,7 @@ namespace metajit { } }); - if (inst->kind() == X86Inst::Kind::Mov64 && - std::holds_alternative(inst->rm())) { + if (is_reg_mov(inst)) { Reg src = std::get(inst->rm()); Reg dst = inst->reg(); if (src == dst) { @@ -2230,9 +2333,16 @@ namespace metajit { _memory_deps.init(_section); _vregs.init(_section); + _input_vregs = _builder.alloc_regs(_section->entry()->args().size()); for (Arg* arg : _section->entry()->args()) { - fix_to_preg(vreg(arg), input_pregs[arg->index()]); + Reg preg = input_pregs[arg->index()]; + Reg input = fix_to_preg(vreg(reg_class(preg)), preg); + _input_vregs[arg->index()] = input; + _vregs[arg] = input; + if (!(reg_class(arg->type()) == reg_class(preg))) { + _vregs[arg] = vreg(arg->type()); + } } // We create one extra block for pseudo_use instructions after loops @@ -2245,6 +2355,17 @@ namespace metajit { memory_deps(); with_timer(isel, isel()); + _builder.move_to_begin(_blocks[0]); + for (Arg* arg : _section->entry()->args()) { + Reg input = _input_vregs.at(arg->index()); + if (!(input == vreg(arg))) { + if (is_float(arg->type())) { + _builder.movq(vreg(arg), input); + } else { + _builder.movq_to_gp(vreg(arg), input); + } + } + } autoname_insts(); if (_mode == Mode::JIT) { @@ -2358,10 +2479,10 @@ namespace metajit { auto rex_opt = [&]() { bool need_rex = false; - if (reg.id() >= 8) { + if ((reg.id() & 15) >= 8) { need_rex = true; } else if (std::holds_alternative(rm)) { - if (std::get(rm).id() >= 8) { + if ((std::get(rm).id() & 15) >= 8) { need_rex = true; } } else if (std::holds_alternative(rm)) { diff --git a/x86insts.inc.hpp b/x86insts.inc.hpp index 43c3d3a..53b4d50 100644 --- a/x86insts.inc.hpp +++ b/x86insts.inc.hpp @@ -194,8 +194,9 @@ rev_binop_x86_inst(MovSDMem, movsd_mem, mov_mem_usedef, true, { byte(0xf2); rex_ binop_x86_inst(MovD, movd, mov_usedef, false, { byte(0x66); rex_opt(); byte(0x0f); byte(0x6e); modrm(); }) binop_x86_inst(MovQ, movq, mov_usedef, true, { byte(0x66); rex_w(); byte(0x0f); byte(0x6e); modrm(); }) +rev_binop_x86_inst(MovQToGP, movq_to_gp, mov_mem_usedef, true, { byte(0x66); rex_w(); byte(0x0f); byte(0x7e); modrm(); }) -binop_x86_inst(UComISS, ucomiss, binop_usedef, false, { byte(0x0f); byte(0x2e); modrm(); }) +binop_x86_inst(UComISS, ucomiss, binop_usedef, false, { rex_opt(); byte(0x0f); byte(0x2e); modrm(); }) binop_x86_inst(UComISD, ucomisd, binop_usedef, true, { byte(0x66); rex_opt(); byte(0x0f); byte(0x2e); modrm(); }) binop_x86_inst(CvtSS2SD, cvtss2sd, mov_usedef, true, { byte(0xf3); rex_opt(); byte(0x0f); byte(0x5a); modrm(); })