From d10ae00c2130a0226195cfd4671adce51d2081a4 Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 22:06:15 +0200 Subject: [PATCH 01/13] separate x86 floating point register allocation --- AGENTS.md | 2 + tests/test_cfg.cpp | 37 ++++++++++ tests/test_insts.cpp | 44 ++++++++++++ x86gen.hpp | 158 ++++++++++++++++++++++++++----------------- x86insts.inc.hpp | 2 +- 5 files changed, 180 insertions(+), 63 deletions(-) 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/tests/test_cfg.cpp b/tests/test_cfg.cpp index b2866cd..c13d156 100644 --- a/tests/test_cfg.cpp +++ b/tests/test_cfg.cpp @@ -39,6 +39,43 @@ int main(int argc, char** argv) { data.output(input); }); + for (Type type : {Type::Float32, Type::Float64}) { + std::string suffix = type == Type::Float32 ? "float32" : "float64"; + suite.diff_test("float_block_argument_" + suffix).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(builder.build_add_f(cont->arg(0), cont->arg(0))); + }); + suite.diff_test("float_swap_loop_" + suffix).run([type](Builder& builder, TestData& data) { + Block* header = builder.build_block({Type::Int64, type, type}); + Block* body = builder.build_block(); + Block* end = builder.build_block(); + Value* a = data.input(type); + Value* b = data.input(type); + builder.build_jump(header, {builder.build_const(Type::Int64, 0), a, b}); + builder.move_to_end(header); + builder.build_branch(builder.build_lt_u(header->arg(0), builder.build_const(Type::Int64, 3)), body, end); + builder.move_to_end(body); + builder.build_jump(header, { + builder.build_add(header->arg(0), builder.build_const(Type::Int64, 1)), + 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..f1d8297 100644 --- a/tests/test_insts.cpp +++ b/tests/test_insts.cpp @@ -456,7 +456,51 @@ void test_call(DiffTestSuite& suite) { }); } +uint64_t test_call_clobber_xmm() { + asm volatile("xorps %%xmm0, %%xmm0\n\txorps %%xmm15, %%xmm15" ::: "xmm0", "xmm15"); + return 42; +} + void test_binop_f(DiffTestSuite& suite) { + for (Type type : {Type::Float32, Type::Float64}) { + std::string suffix = type == Type::Float32 ? "float32" : "float64"; + suite.diff_test("float_spills_" + suffix).aot(false).run([type](Builder& builder, TestData& data) { + std::vector values; + for (size_t index = 0; index < 24; index++) { + values.push_back(data.input(type)); + } + Value* integer = data.input(Type::Int64); + for (Value* value : values) { + data.output(builder.build_add_f(value, value)); + } + data.output(integer); + }); + suite.diff_test("float_across_call_" + suffix).aot(false).interpreter(false).run([type](Builder& builder, TestData& data) { + std::vector values; + for (size_t index = 0; index < 16; index++) { + values.push_back(data.input(type)); + } + Value* callee = builder.build_const(Type::Ptr, (uint64_t)(void*) test_call_clobber_xmm); + data.output(builder.build_call(callee, Type::Int64, std::vector(), CallConv::Default)); + for (Value* value : values) { + data.output(value); + } + }); + suite.diff_test("mixed_register_classes_" + suffix).run([type](Builder& builder, TestData& data) { + std::vector floats; + std::vector integers; + for (size_t index = 0; index < 10; 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(builder.build_add_f(floats[index], floats[index])); + data.output(builder.build_lt_f_o(floats[index], floats[0])); + data.output(builder.build_add(integers[index], 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))); \ diff --git a/x86gen.hpp b/x86gen.hpp index 44f549a..2d15a87 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -29,6 +29,8 @@ namespace metajit { enum class Kind { Invalid, Virtual, Physical }; + + enum class Class { GP, FP }; private: Kind _kind = Kind::Invalid; size_t _id = 0; @@ -40,6 +42,17 @@ namespace metajit { return Reg(Kind::Physical, id); } + static constexpr Reg xmm(size_t id) { return phys(16 + id); } + + Class reg_class() const { + assert(is_physical()); + return _id < 16 ? Class::GP : Class::FP; + } + + static constexpr uint32_t mask(Class reg_class) { + return reg_class == Class::GP ? 0xffff : 0xffff0000; + } + static constexpr Reg X86_RAX() { return phys(0); } static constexpr Reg X86_RCX() { return phys(1); } static constexpr Reg X86_RDX() { return phys(2); } @@ -100,7 +113,7 @@ namespace metajit { stream << "v" << _id; break; case Kind::Physical: - stream << "p" << _id; + stream << (_id < 16 ? "p" : "xmm") << (_id < 16 ? _id : _id - 16); break; } } @@ -559,6 +572,7 @@ namespace metajit { }; struct VRegInfo { + Reg::Class reg_class = Reg::Class::GP; Reg fixed; Interval interval; Reg current_reg; @@ -612,15 +626,21 @@ namespace metajit { } } - Reg vreg() { + static Reg::Class reg_class(Type type) { + return type == Type::Float32 || type == Type::Float64 ? Reg::Class::FP : Reg::Class::GP; + } + + Reg vreg(Reg::Class reg_class = Reg::Class::GP) { size_t id = _vreg_info.size(); _vreg_info.emplace_back(); + _vreg_info.back().reg_class = reg_class; return Reg::virt(id); } Reg fix_to_preg(Reg vreg, Reg preg) { assert(vreg.is_virtual()); VRegInfo& info = _vreg_info[vreg.id()]; + assert(info.reg_class == preg.reg_class()); info.fixed = preg; return vreg; } @@ -638,7 +658,7 @@ namespace metajit { Reg vreg(Value* value) { if (dynmatch(Const, constant, value)) { - Reg reg = vreg(); + Reg reg = vreg(reg_class(value->type())); switch (constant->type()) { case Type::Bool: case Type::Int8: _builder.mov8_imm(reg, constant->value()); break; @@ -681,7 +701,7 @@ namespace metajit { } else if (value->is_named()) { NamedValue* named = (NamedValue*) value; if (_vregs.at(named).is_invalid()) { - _vregs[named] = vreg(); + _vregs[named] = vreg(reg_class(value->type())); } return _vregs.at(named); } else { @@ -690,6 +710,17 @@ namespace metajit { } } + void copy(Reg dst, Reg src) { + Reg::Class dst_class = dst.is_virtual() ? _vreg_info[dst.id()].reg_class : dst.reg_class(); + Reg::Class src_class = src.is_virtual() ? _vreg_info[src.id()].reg_class : src.reg_class(); + assert(dst_class == src_class); + if (dst_class == Reg::Class::FP) { + _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)) { @@ -775,11 +806,11 @@ namespace metajit { void isel(Inst* inst, Block* block) { if (dynmatch(FreezeInst, freeze, inst)) { - _builder.mov64(vreg(inst), vreg(freeze->arg(0))); + copy(vreg(inst), vreg(freeze->arg(0))); } else if (dynmatch(PromoteInst, promote, inst)) { - _builder.mov64(vreg(inst), vreg(promote->arg(0))); + copy(vreg(inst), vreg(promote->arg(0))); } else if (dynmatch(AssumeConstInst, assume_const, inst)) { - _builder.mov64(vreg(inst), vreg(assume_const->arg(0))); + copy(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; @@ -1275,11 +1306,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(reg_class(arg->type())); + copy(copies[arg->index()], vreg(jump->arg(arg->index()))); } for (Arg* arg : jump->block()->args()) { - _builder.mov64(vreg(arg), copies[arg->index()]); + copy(vreg(arg), copies[arg->index()]); } _builder.jmp(_blocks[jump->block()->name()]); } else if (dynmatch(ExitInst, exit, inst)) { @@ -1419,15 +1450,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 +1472,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 +1484,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,31 +1498,33 @@ 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(Reg::Class reg_class) { + uint32_t free = _free & Reg::mask(reg_class); + if (free == 0) { return Reg(); } else { - return Reg::phys(__builtin_ctz(_free)); + return Reg::phys(__builtin_ctz(free)); } } - Reg get_lru() { + Reg get_lru(Reg::Class reg_class) { size_t min_index = 0; size_t min_value = ~size_t(0); - for (size_t it = 0; it < _lru.size(); it++) { + for (size_t it = reg_class == Reg::Class::GP ? 0 : 16; + it < (reg_class == Reg::Class::GP ? 16 : 32); it++) { if (_lru[it] < min_value) { min_value = _lru[it]; min_index = it; @@ -1518,7 +1551,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 +1596,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); + copy(free_reg, preg); reg_file.free(preg); info.current_reg = free_reg; reg_file.set(free_reg, vreg); @@ -1574,13 +1607,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 == Reg::Class::FP) { + _builder.movsd_mem(mem, preg); + } else { + _builder.mov64_mem(mem, preg); + } reg_file.free(preg); info.current_reg = Reg(); } @@ -1592,17 +1624,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); + copy(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 == Reg::Class::FP) { + _builder.movsd(preg, mem); + } else { + _builder.mov64(preg, mem); + } } info.current_reg = preg; reg_file.set(preg, vreg); @@ -1627,9 +1658,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 +1711,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) { @@ -1806,9 +1842,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 +1986,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 +2016,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 +2092,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 = Reg::mask(_vreg_info[reg.id()].reg_class); + 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 +2128,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 +2178,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) { @@ -2358,10 +2392,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..47d328e 100644 --- a/x86insts.inc.hpp +++ b/x86insts.inc.hpp @@ -195,7 +195,7 @@ 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(); }) -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(); }) From 64b448a5824119544ff4e029bbca573327f42432 Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 22:16:16 +0200 Subject: [PATCH 02/13] address floating point register allocation review --- jitir.tmpl.hpp | 22 ++++++++----- tests/test_cfg.cpp | 5 ++- tests/test_insts.cpp | 7 ++--- x86gen.hpp | 74 +++++++++++++++++++++----------------------- 4 files changed, 55 insertions(+), 53 deletions(-) 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/tests/test_cfg.cpp b/tests/test_cfg.cpp index c13d156..57ac515 100644 --- a/tests/test_cfg.cpp +++ b/tests/test_cfg.cpp @@ -40,8 +40,7 @@ int main(int argc, char** argv) { }); for (Type type : {Type::Float32, Type::Float64}) { - std::string suffix = type == Type::Float32 ? "float32" : "float64"; - suite.diff_test("float_block_argument_" + suffix).run([type](Builder& builder, TestData& data) { + 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}); @@ -56,7 +55,7 @@ int main(int argc, char** argv) { builder.move_to_end(cont); data.output(builder.build_add_f(cont->arg(0), cont->arg(0))); }); - suite.diff_test("float_swap_loop_" + suffix).run([type](Builder& builder, TestData& data) { + suite.diff_test(std::string("float_swap_loop_") + to_string(type)).run([type](Builder& builder, TestData& data) { Block* header = builder.build_block({Type::Int64, type, type}); Block* body = builder.build_block(); Block* end = builder.build_block(); diff --git a/tests/test_insts.cpp b/tests/test_insts.cpp index f1d8297..2959eae 100644 --- a/tests/test_insts.cpp +++ b/tests/test_insts.cpp @@ -463,8 +463,7 @@ uint64_t test_call_clobber_xmm() { void test_binop_f(DiffTestSuite& suite) { for (Type type : {Type::Float32, Type::Float64}) { - std::string suffix = type == Type::Float32 ? "float32" : "float64"; - suite.diff_test("float_spills_" + suffix).aot(false).run([type](Builder& builder, TestData& data) { + 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 < 24; index++) { values.push_back(data.input(type)); @@ -475,7 +474,7 @@ void test_binop_f(DiffTestSuite& suite) { } data.output(integer); }); - suite.diff_test("float_across_call_" + suffix).aot(false).interpreter(false).run([type](Builder& builder, TestData& data) { + suite.diff_test(std::string("float_across_call_") + to_string(type)).aot(false).interpreter(false).run([type](Builder& builder, TestData& data) { std::vector values; for (size_t index = 0; index < 16; index++) { values.push_back(data.input(type)); @@ -486,7 +485,7 @@ void test_binop_f(DiffTestSuite& suite) { data.output(value); } }); - suite.diff_test("mixed_register_classes_" + suffix).run([type](Builder& builder, TestData& data) { + 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 < 10; index++) { diff --git a/x86gen.hpp b/x86gen.hpp index 2d15a87..4850d07 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -29,8 +29,6 @@ namespace metajit { enum class Kind { Invalid, Virtual, Physical }; - - enum class Class { GP, FP }; private: Kind _kind = Kind::Invalid; size_t _id = 0; @@ -42,17 +40,6 @@ namespace metajit { return Reg(Kind::Physical, id); } - static constexpr Reg xmm(size_t id) { return phys(16 + id); } - - Class reg_class() const { - assert(is_physical()); - return _id < 16 ? Class::GP : Class::FP; - } - - static constexpr uint32_t mask(Class reg_class) { - return reg_class == Class::GP ? 0xffff : 0xffff0000; - } - static constexpr Reg X86_RAX() { return phys(0); } static constexpr Reg X86_RCX() { return phys(1); } static constexpr Reg X86_RDX() { return phys(2); } @@ -113,7 +100,7 @@ namespace metajit { stream << "v" << _id; break; case Kind::Physical: - stream << (_id < 16 ? "p" : "xmm") << (_id < 16 ? _id : _id - 16); + stream << "p" << _id; break; } } @@ -543,6 +530,17 @@ namespace metajit { Timer peephole; }; private: + enum class RegClass { GP, FP }; + + static RegClass reg_class(Reg preg) { + assert(preg.is_physical()); + return preg.id() < 16 ? RegClass::GP : RegClass::FP; + } + + static constexpr uint32_t reg_mask(RegClass reg_class) { + return reg_class == RegClass::GP ? 0xffff : 0xffff0000; + } + struct Interval { size_t min = 0; size_t max = 0; @@ -572,7 +570,7 @@ namespace metajit { }; struct VRegInfo { - Reg::Class reg_class = Reg::Class::GP; + RegClass reg_class = RegClass::GP; Reg fixed; Interval interval; Reg current_reg; @@ -626,11 +624,11 @@ namespace metajit { } } - static Reg::Class reg_class(Type type) { - return type == Type::Float32 || type == Type::Float64 ? Reg::Class::FP : Reg::Class::GP; + static RegClass reg_class(Type type) { + return type == Type::Float32 || type == Type::Float64 ? RegClass::FP : RegClass::GP; } - Reg vreg(Reg::Class reg_class = Reg::Class::GP) { + Reg vreg(RegClass reg_class = RegClass::GP) { size_t id = _vreg_info.size(); _vreg_info.emplace_back(); _vreg_info.back().reg_class = reg_class; @@ -640,7 +638,7 @@ namespace metajit { Reg fix_to_preg(Reg vreg, Reg preg) { assert(vreg.is_virtual()); VRegInfo& info = _vreg_info[vreg.id()]; - assert(info.reg_class == preg.reg_class()); + assert(info.reg_class == reg_class(preg)); info.fixed = preg; return vreg; } @@ -710,11 +708,11 @@ namespace metajit { } } - void copy(Reg dst, Reg src) { - Reg::Class dst_class = dst.is_virtual() ? _vreg_info[dst.id()].reg_class : dst.reg_class(); - Reg::Class src_class = src.is_virtual() ? _vreg_info[src.id()].reg_class : src.reg_class(); + void move(Reg dst, Reg src) { + RegClass dst_class = dst.is_virtual() ? _vreg_info[dst.id()].reg_class : reg_class(dst); + RegClass src_class = src.is_virtual() ? _vreg_info[src.id()].reg_class : reg_class(src); assert(dst_class == src_class); - if (dst_class == Reg::Class::FP) { + if (dst_class == RegClass::FP) { _builder.movsd(dst, src); } else { _builder.mov64(dst, src); @@ -806,11 +804,11 @@ namespace metajit { void isel(Inst* inst, Block* block) { if (dynmatch(FreezeInst, freeze, inst)) { - copy(vreg(inst), vreg(freeze->arg(0))); + move(vreg(inst), vreg(freeze->arg(0))); } else if (dynmatch(PromoteInst, promote, inst)) { - copy(vreg(inst), vreg(promote->arg(0))); + move(vreg(inst), vreg(promote->arg(0))); } else if (dynmatch(AssumeConstInst, assume_const, inst)) { - copy(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; @@ -1307,10 +1305,10 @@ namespace metajit { lwir::Span copies = _builder.alloc_regs(jump->block()->args().size()); for (Arg* arg : jump->block()->args()) { copies[arg->index()] = vreg(reg_class(arg->type())); - copy(copies[arg->index()], vreg(jump->arg(arg->index()))); + move(copies[arg->index()], vreg(jump->arg(arg->index()))); } for (Arg* arg : jump->block()->args()) { - copy(vreg(arg), copies[arg->index()]); + move(vreg(arg), copies[arg->index()]); } _builder.jmp(_blocks[jump->block()->name()]); } else if (dynmatch(ExitInst, exit, inst)) { @@ -1511,8 +1509,8 @@ namespace metajit { return (_max_free & (1u << preg.id())) == 0; } - Reg get_free_reg(Reg::Class reg_class) { - uint32_t free = _free & Reg::mask(reg_class); + Reg get_free_reg(RegClass reg_class) { + uint32_t free = _free & reg_mask(reg_class); if (free == 0) { return Reg(); } else { @@ -1520,11 +1518,11 @@ namespace metajit { } } - Reg get_lru(Reg::Class reg_class) { + Reg get_lru(RegClass reg_class) { size_t min_index = 0; size_t min_value = ~size_t(0); - for (size_t it = reg_class == Reg::Class::GP ? 0 : 16; - it < (reg_class == Reg::Class::GP ? 16 : 32); it++) { + for (uint32_t mask = reg_mask(reg_class); mask != 0; mask &= mask - 1) { + size_t it = __builtin_ctz(mask); if (_lru[it] < min_value) { min_value = _lru[it]; min_index = it; @@ -1599,7 +1597,7 @@ namespace metajit { 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 - copy(free_reg, preg); + move(free_reg, preg); reg_file.free(preg); info.current_reg = free_reg; reg_file.set(free_reg, vreg); @@ -1608,7 +1606,7 @@ namespace metajit { info.stack_offset = _stack_offset_alloc.alloc(); } X86Inst::Mem mem(Reg::X86_RSP(), (int32_t) info.stack_offset); - if (info.reg_class == Reg::Class::FP) { + if (info.reg_class == RegClass::FP) { _builder.movsd_mem(mem, preg); } else { _builder.mov64_mem(mem, preg); @@ -1624,12 +1622,12 @@ namespace metajit { VRegInfo& info = _vreg_info[vreg.id()]; if (info.current_reg.is_physical()) { // No need to unspill, just move from current reg - copy(preg, info.current_reg); + move(preg, info.current_reg); reg_file.free(info.current_reg); } else { assert(info.stack_offset != ~size_t(0)); X86Inst::Mem mem(Reg::X86_RSP(), (int32_t) info.stack_offset); - if (info.reg_class == Reg::Class::FP) { + if (info.reg_class == RegClass::FP) { _builder.movsd(preg, mem); } else { _builder.mov64(preg, mem); @@ -2092,7 +2090,7 @@ namespace metajit { Reg reg = order.at(it); assert(reg.is_virtual()); - uint32_t free_mask = Reg::mask(_vreg_info[reg.id()].reg_class); + uint32_t free_mask = reg_mask(_vreg_info[reg.id()].reg_class); free_mask &= ~(1u << Reg::X86_RSP().id()); free_mask &= ~(1u << Reg::X86_RBP().id()); for (Reg conflict : conflicts.at(reg.id())) { From 71df1aac20c304bdbb2026ceb687a16e4a5debc6 Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 22:29:20 +0200 Subject: [PATCH 03/13] support floating point call arguments and returns --- tests/test_cfg.cpp | 1 + tests/test_insts.cpp | 119 +++++++++++++++++++++++++++++++++++++++++++ x86gen.hpp | 49 ++++++++++-------- 3 files changed, 149 insertions(+), 20 deletions(-) diff --git a/tests/test_cfg.cpp b/tests/test_cfg.cpp index 57ac515..ab3b741 100644 --- a/tests/test_cfg.cpp +++ b/tests/test_cfg.cpp @@ -55,6 +55,7 @@ int main(int argc, char** argv) { builder.move_to_end(cont); data.output(builder.build_add_f(cont->arg(0), 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::Int64, type, type}); Block* body = builder.build_block(); diff --git a/tests/test_insts.cpp b/tests/test_insts.cpp index 2959eae..576b96c 100644 --- a/tests/test_insts.cpp +++ b/tests/test_insts.cpp @@ -456,6 +456,121 @@ void test_call(DiffTestSuite& suite) { }); } +#define fp_call_targets(prefix, attrs) \ + template attrs F prefix##_identity(F value) { return value; } \ + template attrs F prefix##_zero() { return (F) -3.25; } \ + template attrs uint64_t prefix##_integer(F value, uint64_t integer) { \ + return (uint64_t) value + integer; \ + } \ + template attrs void prefix##_void(F value, F* out) { *out = value; } \ + template attrs F prefix##_mixed( \ + uint64_t i0, F f0, uint64_t i1, double f1, uint64_t i2, float f2, \ + uint64_t i3, F f3, uint64_t i4, F f4, uint64_t i5, F f5, F f6, F f7) { \ + return (F) (i0 + 2 * i1 + 3 * i2 + 4 * i3 + 5 * i4 + 6 * i5) + \ + f0 + (F) (2 * f1) + (F) (3 * f2) + 4 * f3 + 5 * f4 + 6 * f5 + 7 * f6 + 8 * f7; \ + } + +fp_call_targets(test_call_fp_default, __attribute__((noinline))) +fp_call_targets(test_call_fp_preserve_none, __attribute__((preserve_none, noinline))) + +#undef fp_call_targets + +template +void test_call_fp(DiffTestSuite& suite, Type type) { + for (CallConv call_conv : {CallConv::Default, CallConv::PreserveNone}) { + std::string name = std::string("call_fp_") + to_string(type) + + (call_conv == CallConv::Default ? "_default" : "_preserve_none"); + uint64_t identity = call_conv == CallConv::Default ? + (uint64_t)(void*) test_call_fp_default_identity : (uint64_t)(void*) test_call_fp_preserve_none_identity; + uint64_t zero = call_conv == CallConv::Default ? + (uint64_t)(void*) test_call_fp_default_zero : (uint64_t)(void*) test_call_fp_preserve_none_zero; + uint64_t integer = call_conv == CallConv::Default ? + (uint64_t)(void*) test_call_fp_default_integer : (uint64_t)(void*) test_call_fp_preserve_none_integer; + uint64_t void_target = call_conv == CallConv::Default ? + (uint64_t)(void*) test_call_fp_default_void : (uint64_t)(void*) test_call_fp_preserve_none_void; + uint64_t mixed = call_conv == CallConv::Default ? + (uint64_t)(void*) test_call_fp_default_mixed : (uint64_t)(void*) test_call_fp_preserve_none_mixed; + + suite.diff_test(name + "_identity").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { + Value* value = data.input(type); + Value* callee = builder.build_const(Type::Ptr, identity); + Value* first = builder.build_call(callee, type, {value}, call_conv); + Value* second = builder.build_call(callee, type, {first}, call_conv); + data.output(value); + data.output(first); + data.output(second); + }); + + suite.diff_test(name + "_zero_args").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { + data.output(builder.build_call(builder.build_const(Type::Ptr, zero), type, std::vector(), call_conv)); + }); + + suite.diff_test(name + "_constant").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { + Value* value = builder.build_int_to_float_s(builder.build_const(Type::Int64, 17), type); + data.output(builder.build_call(builder.build_const(Type::Ptr, identity), type, {value}, call_conv)); + }); + + suite.diff_test(name + "_special_values").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { + std::vector values = type == Type::Float32 ? + std::vector{0, 0x80000000, 0x7f800000, 0xff800000, 0x7fc00001, 1} : + std::vector{0, 0x8000000000000000ULL, 0x7ff0000000000000ULL, + 0xfff0000000000000ULL, 0x7ff8000000000001ULL, 1}; + Value* callee = builder.build_const(Type::Ptr, identity); + for (uint64_t bits : values) { + data.output(builder.build_call(callee, type, {builder.build_const(type, bits)}, call_conv)); + } + }); + + suite.diff_test(name + "_integer_return").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { + Value* value = builder.build_int_to_float_s(data.input(RandomRange(Type::Int64, 0, 100)), type); + Value* i = data.input(Type::Int64); + data.output(builder.build_call(builder.build_const(Type::Ptr, integer), Type::Int64, {value, i}, call_conv)); + data.output(value); + data.output(i); + }); + + suite.diff_test(name + "_void_return").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { + Value* value = data.input(type); + Value* out = builder.build_alloca(builder.build_const(Type::Int64, 8), 8); + builder.build_call(builder.build_const(Type::Ptr, void_target), Type::Void, {value, out}, call_conv); + data.output(builder.build_load(out, type, LoadFlags::None, AliasingGroup(0), 0)); + data.output(value); + }); + + suite.diff_test(name + "_mixed_pressure").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { + std::vector floats; + std::vector integers; + for (size_t index = 0; index < 24; index++) { + floats.push_back(builder.build_int_to_float_s(data.input(RandomRange(Type::Int64, 1, 32)), type)); + } + for (size_t index = 0; index < 16; index++) { + integers.push_back(data.input(RandomRange(Type::Int64, 1, 32))); + } + std::vector args; + for (size_t index = 0; index < 6; index++) { + args.push_back(integers[index]); + Type arg_type = index == 1 ? Type::Float64 : index == 2 ? Type::Float32 : type; + args.push_back(arg_type == type ? floats[index] : builder.build_resize_f(floats[index], arg_type)); + } + args.push_back(floats[6]); + args.push_back(floats[7]); + Value* callee = builder.build_const(Type::Ptr, mixed); + Value* first = builder.build_call(callee, type, args, call_conv); + args[1] = first; + args[13] = first; + Value* second = builder.build_call(callee, type, args, call_conv); + data.output(first); + data.output(second); + for (Value* value : floats) { + data.output(value); + } + for (Value* value : integers) { + data.output(value); + } + }); + } +} + uint64_t test_call_clobber_xmm() { asm volatile("xorps %%xmm0, %%xmm0\n\txorps %%xmm15, %%xmm15" ::: "xmm0", "xmm15"); return 42; @@ -474,6 +589,7 @@ void test_binop_f(DiffTestSuite& suite) { } data.output(integer); }); + suite.diff_test(std::string("float_across_call_") + to_string(type)).aot(false).interpreter(false).run([type](Builder& builder, TestData& data) { std::vector values; for (size_t index = 0; index < 16; index++) { @@ -485,6 +601,7 @@ void test_binop_f(DiffTestSuite& suite) { 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; @@ -619,6 +736,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 4850d07..d741e1d 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -454,6 +454,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::phys(16), Reg::phys(17), Reg::phys(18), Reg::phys(19), + Reg::phys(20), Reg::phys(21), Reg::phys(22), Reg::phys(23) + }; public: CallConvInfo(CallConv call_conv) { switch (call_conv) { @@ -472,12 +477,19 @@ namespace metajit { } } - const lwir::Span& args() const { return _arg_regs; } + lwir::Span args(Type type = Type::Int64) const { + if (type == Type::Float32 || type == Type::Float64) { + return lwir::Span(fp_arg_regs, sizeof(fp_arg_regs) / sizeof(fp_arg_regs[0])); + } + return _arg_regs; + } const lwir::Span& preserved() const { return _preserved_regs; } - Reg arg(size_t index) const { return _arg_regs.at(index); } + Reg arg(size_t index, Type type = Type::Int64) const { return args(type).at(index); } Reg preserved(size_t index) const { return _preserved_regs.at(index); } - Reg ret() const { return _ret_reg; } + Reg ret(Type type = Type::Int64) const { + return type == Type::Float32 || type == Type::Float64 ? Reg::phys(16) : _ret_reg; + } // TODO: Optimize @@ -489,15 +501,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 { @@ -1237,22 +1240,26 @@ 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 gp_count = 0; + size_t fp_count = 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(); + size_t index = reg_class(type) == RegClass::FP ? fp_count++ : gp_count++; + assert(index < info.args(type).size() && "Call with too many register arguments"); + Reg arg_reg = fix_to_preg(vreg(reg_class(type)), info.arg(index, type)); + move(arg_reg, vreg(call->arg(it))); args[it - 1] = arg_reg; } - Reg ret_reg = fix_to_preg(vreg(), info.ret()); + Reg ret_reg = fix_to_preg(vreg(reg_class(call->type())), info.ret(call->type())); 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(); @@ -1794,7 +1801,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); } @@ -1816,8 +1823,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; From 6cb8069dc3e8052c1d89dd5e5f0855d77b7a2a1c Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 22:35:19 +0200 Subject: [PATCH 04/13] test floating point loads and stores --- tests/test_insts.cpp | 62 ++++++++++++++++++++++++++++++++++++++++++++ x86gen.hpp | 9 ++++--- 2 files changed, 67 insertions(+), 4 deletions(-) diff --git a/tests/test_insts.cpp b/tests/test_insts.cpp index 576b96c..7b45490 100644 --- a/tests/test_insts.cpp +++ b/tests/test_insts.cpp @@ -287,6 +287,67 @@ void test_alloca(DiffTestSuite& suite) { }); } +void test_load_store_f(DiffTestSuite& suite) { + for (Type type : {Type::Float32, Type::Float64}) { + Type bits_type = type == Type::Float32 ? Type::Int32 : Type::Int64; + std::string name = std::string("load_store_") + to_string(type); + + suite.diff_test(name + "_roundtrip").run([=](Builder& builder, TestData& data) { + Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 24), 8); + Value* sentinel = builder.build_const(Type::Int64, 0x123456789abcdef0ULL); + for (int32_t offset : {0, 8, 16}) { + builder.build_store(ptr, sentinel, AliasingGroup(0), offset); + } + Value* first = data.input(type); + Value* second = data.input(type); + builder.build_store(ptr, first, AliasingGroup(0), 8); + data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 8)); + data.output(builder.build_load(ptr, Type::Int64, LoadFlags::None, AliasingGroup(0), 8)); + builder.build_store(ptr, second, AliasingGroup(0), 8); + data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 8)); + for (int32_t offset : {0, 8, 16}) { + data.output(builder.build_load(ptr, Type::Int64, LoadFlags::None, AliasingGroup(0), offset)); + } + }); + + suite.diff_test(name + "_offsets").run([=](Builder& builder, TestData& data) { + Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 512), 8); + for (int32_t offset : {1, 127, 257}) { + Value* value = data.input(type); + builder.build_store(ptr, value, AliasingGroup(0), offset); + Value* shifted = builder.build_add_ptr(ptr, builder.build_const(Type::Int64, offset + 3)); + data.output(builder.build_load(shifted, type, LoadFlags::None, AliasingGroup(0), -3)); + data.output(builder.build_load(ptr, bits_type, LoadFlags::None, AliasingGroup(0), offset)); + } + }); + + suite.diff_test(name + "_constants").run([=](Builder& builder, TestData& data) { + Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 8), 8); + for (uint64_t bits : {uint64_t(0), uint64_t(1), uint64_t(1) << (type_size(type) * 8 - 1), type_mask(type)}) { + builder.build_store(ptr, builder.build_const(type, bits), AliasingGroup(0), 0); + data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 0)); + data.output(builder.build_load(ptr, bits_type, LoadFlags::None, AliasingGroup(0), 0)); + } + }); + + suite.diff_test(name + "_pressure").aot(false).run([=](Builder& builder, TestData& data) { + Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 24 * 8), 8); + std::vector values; + for (int32_t index = 0; index < 24; index++) { + values.push_back(data.input(type)); + } + for (int32_t index = 0; index < 24; index++) { + builder.build_store(ptr, values[index], AliasingGroup(0), index * 8); + } + for (int32_t index = 0; index < 24; index++) { + data.output(values[index]); + data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), index * 8)); + data.output(builder.build_load(ptr, bits_type, LoadFlags::None, AliasingGroup(0), index * 8)); + } + }); + } +} + void test_call(DiffTestSuite& suite) { suite.diff_test("call_preserve_none").aot(false).interpreter(false).run([](Builder& builder, TestData& data) { Value* a = data.input(Type::Int64); @@ -735,6 +796,7 @@ int main(int argc, char** argv) { test_freeze(suite); test_assume_const(suite); test_alloca(suite); + test_load_store_f(suite); test_call(suite); test_call_fp(suite, Type::Float32); test_call_fp(suite, Type::Float64); diff --git a/x86gen.hpp b/x86gen.hpp index d741e1d..1861516 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -56,6 +56,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); @@ -456,8 +457,8 @@ namespace metajit { }; static constexpr Reg fp_arg_regs[] = { - Reg::phys(16), Reg::phys(17), Reg::phys(18), Reg::phys(19), - Reg::phys(20), Reg::phys(21), Reg::phys(22), Reg::phys(23) + 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) { @@ -488,7 +489,7 @@ namespace metajit { Reg arg(size_t index, Type type = Type::Int64) const { return args(type).at(index); } Reg preserved(size_t index) const { return _preserved_regs.at(index); } Reg ret(Type type = Type::Int64) const { - return type == Type::Float32 || type == Type::Float64 ? Reg::phys(16) : _ret_reg; + return type == Type::Float32 || type == Type::Float64 ? Reg::X86_XMM(0) : _ret_reg; } // TODO: Optimize @@ -537,7 +538,7 @@ namespace metajit { static RegClass reg_class(Reg preg) { assert(preg.is_physical()); - return preg.id() < 16 ? RegClass::GP : RegClass::FP; + return preg.id() < Reg::X86_XMM(0).id() ? RegClass::GP : RegClass::FP; } static constexpr uint32_t reg_mask(RegClass reg_class) { From 3322bc19226c8aec8db727782a8b0a7d71a9b58a Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 22:54:51 +0200 Subject: [PATCH 05/13] simplify floating point tests and calling convention metadata --- tests/test_cfg.cpp | 11 ++- tests/test_insts.cpp | 225 ++++++++++--------------------------------- x86gen.hpp | 83 +++++++++++----- 3 files changed, 113 insertions(+), 206 deletions(-) diff --git a/tests/test_cfg.cpp b/tests/test_cfg.cpp index ab3b741..ad4b6fd 100644 --- a/tests/test_cfg.cpp +++ b/tests/test_cfg.cpp @@ -53,21 +53,22 @@ int main(int argc, char** argv) { builder.move_to_end(b); builder.build_jump(cont, {value_b}); builder.move_to_end(cont); - data.output(builder.build_add_f(cont->arg(0), cont->arg(0))); + 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::Int64, type, type}); + 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); - builder.build_jump(header, {builder.build_const(Type::Int64, 0), a, b}); + Value* cond = data.input(Type::Bool); + builder.build_jump(header, {cond, a, b}); builder.move_to_end(header); - builder.build_branch(builder.build_lt_u(header->arg(0), builder.build_const(Type::Int64, 3)), body, end); + builder.build_branch(header->arg(0), body, end); builder.move_to_end(body); builder.build_jump(header, { - builder.build_add(header->arg(0), builder.build_const(Type::Int64, 1)), + builder.build_const(Type::Bool, false), header->arg(2), header->arg(1) }); builder.move_to_end(end); diff --git a/tests/test_insts.cpp b/tests/test_insts.cpp index 7b45490..e91a050 100644 --- a/tests/test_insts.cpp +++ b/tests/test_insts.cpp @@ -289,61 +289,25 @@ void test_alloca(DiffTestSuite& suite) { void test_load_store_f(DiffTestSuite& suite) { for (Type type : {Type::Float32, Type::Float64}) { - Type bits_type = type == Type::Float32 ? Type::Int32 : Type::Int64; + Type bits_type; + if (type == Type::Float32) { + bits_type = Type::Int32; + } else { + bits_type = Type::Int64; + } std::string name = std::string("load_store_") + to_string(type); suite.diff_test(name + "_roundtrip").run([=](Builder& builder, TestData& data) { - Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 24), 8); - Value* sentinel = builder.build_const(Type::Int64, 0x123456789abcdef0ULL); - for (int32_t offset : {0, 8, 16}) { - builder.build_store(ptr, sentinel, AliasingGroup(0), offset); - } - Value* first = data.input(type); - Value* second = data.input(type); - builder.build_store(ptr, first, AliasingGroup(0), 8); - data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 8)); - data.output(builder.build_load(ptr, Type::Int64, LoadFlags::None, AliasingGroup(0), 8)); - builder.build_store(ptr, second, AliasingGroup(0), 8); - data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 8)); - for (int32_t offset : {0, 8, 16}) { - data.output(builder.build_load(ptr, Type::Int64, LoadFlags::None, AliasingGroup(0), offset)); - } - }); - - suite.diff_test(name + "_offsets").run([=](Builder& builder, TestData& data) { - Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 512), 8); - for (int32_t offset : {1, 127, 257}) { - Value* value = data.input(type); - builder.build_store(ptr, value, AliasingGroup(0), offset); - Value* shifted = builder.build_add_ptr(ptr, builder.build_const(Type::Int64, offset + 3)); - data.output(builder.build_load(shifted, type, LoadFlags::None, AliasingGroup(0), -3)); - data.output(builder.build_load(ptr, bits_type, LoadFlags::None, AliasingGroup(0), offset)); - } - }); - - suite.diff_test(name + "_constants").run([=](Builder& builder, TestData& data) { + Value* value = data.input(type); Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 8), 8); - for (uint64_t bits : {uint64_t(0), uint64_t(1), uint64_t(1) << (type_size(type) * 8 - 1), type_mask(type)}) { - builder.build_store(ptr, builder.build_const(type, bits), AliasingGroup(0), 0); - data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 0)); - data.output(builder.build_load(ptr, bits_type, LoadFlags::None, AliasingGroup(0), 0)); - } - }); - suite.diff_test(name + "_pressure").aot(false).run([=](Builder& builder, TestData& data) { - Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 24 * 8), 8); - std::vector values; - for (int32_t index = 0; index < 24; index++) { - values.push_back(data.input(type)); - } - for (int32_t index = 0; index < 24; index++) { - builder.build_store(ptr, values[index], AliasingGroup(0), index * 8); - } - for (int32_t index = 0; index < 24; index++) { - data.output(values[index]); - data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), index * 8)); - data.output(builder.build_load(ptr, bits_type, LoadFlags::None, AliasingGroup(0), index * 8)); - } + builder.build_store(ptr, builder.build_const(Type::Int64, ~uint64_t(0)), AliasingGroup(0), 0); + builder.build_store(ptr, value, AliasingGroup(0), 0); + data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 0)); + data.output(builder.build_load(ptr, Type::Int64, LoadFlags::None, AliasingGroup(0), 0)); + + builder.build_store(ptr, builder.build_const(type, 1), AliasingGroup(0), 0); + data.output(builder.build_load(ptr, bits_type, LoadFlags::None, AliasingGroup(0), 0)); }); } } @@ -517,147 +481,57 @@ void test_call(DiffTestSuite& suite) { }); } -#define fp_call_targets(prefix, attrs) \ - template attrs F prefix##_identity(F value) { return value; } \ - template attrs F prefix##_zero() { return (F) -3.25; } \ - template attrs uint64_t prefix##_integer(F value, uint64_t integer) { \ - return (uint64_t) value + integer; \ - } \ - template attrs void prefix##_void(F value, F* out) { *out = value; } \ - template attrs F prefix##_mixed( \ - uint64_t i0, F f0, uint64_t i1, double f1, uint64_t i2, float f2, \ - uint64_t i3, F f3, uint64_t i4, F f4, uint64_t i5, F f5, F f6, F f7) { \ - return (F) (i0 + 2 * i1 + 3 * i2 + 4 * i3 + 5 * i4 + 6 * i5) + \ - f0 + (F) (2 * f1) + (F) (3 * f2) + 4 * f3 + 5 * f4 + 6 * f5 + 7 * f6 + 8 * f7; \ - } - -fp_call_targets(test_call_fp_default, __attribute__((noinline))) -fp_call_targets(test_call_fp_preserve_none, __attribute__((preserve_none, noinline))) +template +__attribute__((noinline)) F test_call_fp_default(uint64_t i, float a, double b) { + return (F) i + (F) a + (F) b; +} -#undef fp_call_targets +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::string name = std::string("call_fp_") + to_string(type) + - (call_conv == CallConv::Default ? "_default" : "_preserve_none"); - uint64_t identity = call_conv == CallConv::Default ? - (uint64_t)(void*) test_call_fp_default_identity : (uint64_t)(void*) test_call_fp_preserve_none_identity; - uint64_t zero = call_conv == CallConv::Default ? - (uint64_t)(void*) test_call_fp_default_zero : (uint64_t)(void*) test_call_fp_preserve_none_zero; - uint64_t integer = call_conv == CallConv::Default ? - (uint64_t)(void*) test_call_fp_default_integer : (uint64_t)(void*) test_call_fp_preserve_none_integer; - uint64_t void_target = call_conv == CallConv::Default ? - (uint64_t)(void*) test_call_fp_default_void : (uint64_t)(void*) test_call_fp_preserve_none_void; - uint64_t mixed = call_conv == CallConv::Default ? - (uint64_t)(void*) test_call_fp_default_mixed : (uint64_t)(void*) test_call_fp_preserve_none_mixed; - - suite.diff_test(name + "_identity").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { - Value* value = data.input(type); - Value* callee = builder.build_const(Type::Ptr, identity); - Value* first = builder.build_call(callee, type, {value}, call_conv); - Value* second = builder.build_call(callee, type, {first}, call_conv); - data.output(value); - data.output(first); - data.output(second); - }); - - suite.diff_test(name + "_zero_args").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { - data.output(builder.build_call(builder.build_const(Type::Ptr, zero), type, std::vector(), call_conv)); - }); - - suite.diff_test(name + "_constant").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { - Value* value = builder.build_int_to_float_s(builder.build_const(Type::Int64, 17), type); - data.output(builder.build_call(builder.build_const(Type::Ptr, identity), type, {value}, call_conv)); - }); - - suite.diff_test(name + "_special_values").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { - std::vector values = type == Type::Float32 ? - std::vector{0, 0x80000000, 0x7f800000, 0xff800000, 0x7fc00001, 1} : - std::vector{0, 0x8000000000000000ULL, 0x7ff0000000000000ULL, - 0xfff0000000000000ULL, 0x7ff8000000000001ULL, 1}; - Value* callee = builder.build_const(Type::Ptr, identity); - for (uint64_t bits : values) { - data.output(builder.build_call(callee, type, {builder.build_const(type, bits)}, call_conv)); + 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); } - }); - - suite.diff_test(name + "_integer_return").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { - Value* value = builder.build_int_to_float_s(data.input(RandomRange(Type::Int64, 0, 100)), type); - Value* i = data.input(Type::Int64); - data.output(builder.build_call(builder.build_const(Type::Ptr, integer), Type::Int64, {value, i}, call_conv)); - data.output(value); - data.output(i); - }); - - suite.diff_test(name + "_void_return").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { - Value* value = data.input(type); - Value* out = builder.build_alloca(builder.build_const(Type::Int64, 8), 8); - builder.build_call(builder.build_const(Type::Ptr, void_target), Type::Void, {value, out}, call_conv); - data.output(builder.build_load(out, type, LoadFlags::None, AliasingGroup(0), 0)); - data.output(value); - }); + Value* first = builder.build_call(callee, type, {i, a, b}, call_conv); + Value* second = builder.build_call(callee, type, {i, a, b}, call_conv); - suite.diff_test(name + "_mixed_pressure").aot(false).interpreter(false).run([=](Builder& builder, TestData& data) { - std::vector floats; - std::vector integers; - for (size_t index = 0; index < 24; index++) { - floats.push_back(builder.build_int_to_float_s(data.input(RandomRange(Type::Int64, 1, 32)), type)); - } - for (size_t index = 0; index < 16; index++) { - integers.push_back(data.input(RandomRange(Type::Int64, 1, 32))); - } - std::vector args; - for (size_t index = 0; index < 6; index++) { - args.push_back(integers[index]); - Type arg_type = index == 1 ? Type::Float64 : index == 2 ? Type::Float32 : type; - args.push_back(arg_type == type ? floats[index] : builder.build_resize_f(floats[index], arg_type)); - } - args.push_back(floats[6]); - args.push_back(floats[7]); - Value* callee = builder.build_const(Type::Ptr, mixed); - Value* first = builder.build_call(callee, type, args, call_conv); - args[1] = first; - args[13] = first; - Value* second = builder.build_call(callee, type, args, call_conv); data.output(first); data.output(second); - for (Value* value : floats) { - data.output(value); - } - for (Value* value : integers) { - data.output(value); - } + data.output(i); + data.output(a); + data.output(b); }); } } -uint64_t test_call_clobber_xmm() { - asm volatile("xorps %%xmm0, %%xmm0\n\txorps %%xmm15, %%xmm15" ::: "xmm0", "xmm15"); - return 42; -} - 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 < 24; index++) { - values.push_back(data.input(type)); - } - Value* integer = data.input(Type::Int64); - for (Value* value : values) { - data.output(builder.build_add_f(value, value)); - } - data.output(integer); - }); - - suite.diff_test(std::string("float_across_call_") + to_string(type)).aot(false).interpreter(false).run([type](Builder& builder, TestData& data) { - std::vector values; - for (size_t index = 0; index < 16; index++) { + for (size_t index = 0; index < 17; index++) { values.push_back(data.input(type)); } - Value* callee = builder.build_const(Type::Ptr, (uint64_t)(void*) test_call_clobber_xmm); - data.output(builder.build_call(callee, Type::Int64, std::vector(), CallConv::Default)); for (Value* value : values) { data.output(value); } @@ -666,14 +540,13 @@ void test_binop_f(DiffTestSuite& suite) { 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 < 10; index++) { + 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(builder.build_add_f(floats[index], floats[index])); - data.output(builder.build_lt_f_o(floats[index], floats[0])); - data.output(builder.build_add(integers[index], integers[index])); + data.output(floats[index]); + data.output(integers[index]); } }); } diff --git a/x86gen.hpp b/x86gen.hpp index 1861516..f927dd3 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -24,6 +24,8 @@ #include "jitir.hpp" namespace metajit { + enum class RegClass { Int, Float }; + class Reg { public: enum class Kind { @@ -437,6 +439,7 @@ namespace metajit { lwir::Span _arg_regs; lwir::Span _preserved_regs; Reg _ret_reg; + Reg _fp_ret_reg = Reg::X86_XMM(0); static constexpr Reg preserve_none_arg_regs[] = { Reg::X86_R12(), Reg::X86_R13(), Reg::X86_R14(), Reg::X86_R15(), @@ -478,18 +481,22 @@ namespace metajit { } } - lwir::Span args(Type type = Type::Int64) const { - if (type == Type::Float32 || type == Type::Float64) { + lwir::Span args(RegClass reg_class = RegClass::Int) const { + if (reg_class == RegClass::Float) { return lwir::Span(fp_arg_regs, sizeof(fp_arg_regs) / sizeof(fp_arg_regs[0])); } return _arg_regs; } const lwir::Span& preserved() const { return _preserved_regs; } - Reg arg(size_t index, Type type = Type::Int64) const { return args(type).at(index); } + Reg arg(size_t index, RegClass reg_class = RegClass::Int) const { return args(reg_class).at(index); } Reg preserved(size_t index) const { return _preserved_regs.at(index); } - Reg ret(Type type = Type::Int64) const { - return type == Type::Float32 || type == Type::Float64 ? Reg::X86_XMM(0) : _ret_reg; + Reg ret(RegClass reg_class = RegClass::Int) const { + if (reg_class == RegClass::Float) { + return _fp_ret_reg; + } else { + return _ret_reg; + } } // TODO: Optimize @@ -534,15 +541,29 @@ namespace metajit { Timer peephole; }; private: - enum class RegClass { GP, FP }; - static RegClass reg_class(Reg preg) { assert(preg.is_physical()); - return preg.id() < Reg::X86_XMM(0).id() ? RegClass::GP : RegClass::FP; + if (preg.id() < Reg::X86_XMM(0).id()) { + return RegClass::Int; + } else { + return RegClass::Float; + } + } + + static RegClass reg_class(Type type) { + if (type == Type::Float32 || type == Type::Float64) { + return RegClass::Float; + } else { + return RegClass::Int; + } } static constexpr uint32_t reg_mask(RegClass reg_class) { - return reg_class == RegClass::GP ? 0xffff : 0xffff0000; + if (reg_class == RegClass::Int) { + return 0xffff; + } else { + return 0xffff0000; + } } struct Interval { @@ -574,7 +595,7 @@ namespace metajit { }; struct VRegInfo { - RegClass reg_class = RegClass::GP; + RegClass reg_class = RegClass::Int; Reg fixed; Interval interval; Reg current_reg; @@ -628,11 +649,7 @@ namespace metajit { } } - static RegClass reg_class(Type type) { - return type == Type::Float32 || type == Type::Float64 ? RegClass::FP : RegClass::GP; - } - - Reg vreg(RegClass reg_class = RegClass::GP) { + Reg vreg(RegClass reg_class = RegClass::Int) { size_t id = _vreg_info.size(); _vreg_info.emplace_back(); _vreg_info.back().reg_class = reg_class; @@ -713,10 +730,20 @@ namespace metajit { } void move(Reg dst, Reg src) { - RegClass dst_class = dst.is_virtual() ? _vreg_info[dst.id()].reg_class : reg_class(dst); - RegClass src_class = src.is_virtual() ? _vreg_info[src.id()].reg_class : reg_class(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::FP) { + if (dst_class == RegClass::Float) { _builder.movsd(dst, src); } else { _builder.mov64(dst, src); @@ -1246,15 +1273,21 @@ namespace metajit { size_t gp_count = 0; size_t fp_count = 0; for (size_t it = 1; it < call->args().size(); it++) { - Type type = call->arg(it)->type(); - size_t index = reg_class(type) == RegClass::FP ? fp_count++ : gp_count++; - assert(index < info.args(type).size() && "Call with too many register arguments"); - Reg arg_reg = fix_to_preg(vreg(reg_class(type)), info.arg(index, type)); + RegClass arg_class = reg_class(call->arg(it)->type()); + size_t index; + if (arg_class == RegClass::Float) { + index = fp_count++; + } else { + index = gp_count++; + } + assert(index < info.args(arg_class).size() && "Call with too many register arguments"); + Reg arg_reg = fix_to_preg(vreg(arg_class), info.arg(index, arg_class)); move(arg_reg, vreg(call->arg(it))); args[it - 1] = arg_reg; } - Reg ret_reg = fix_to_preg(vreg(reg_class(call->type())), info.ret(call->type())); + RegClass ret_class = reg_class(call->type()); + Reg ret_reg = fix_to_preg(vreg(ret_class), 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); @@ -1614,7 +1647,7 @@ namespace metajit { info.stack_offset = _stack_offset_alloc.alloc(); } X86Inst::Mem mem(Reg::X86_RSP(), (int32_t) info.stack_offset); - if (info.reg_class == RegClass::FP) { + if (info.reg_class == RegClass::Float) { _builder.movsd_mem(mem, preg); } else { _builder.mov64_mem(mem, preg); @@ -1635,7 +1668,7 @@ namespace metajit { } else { assert(info.stack_offset != ~size_t(0)); X86Inst::Mem mem(Reg::X86_RSP(), (int32_t) info.stack_offset); - if (info.reg_class == RegClass::FP) { + if (info.reg_class == RegClass::Float) { _builder.movsd(preg, mem); } else { _builder.mov64(preg, mem); From f5ef819797a7942bec339c9a9e8bded615aaa68f Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 22:59:29 +0200 Subject: [PATCH 06/13] fix floating point select in x86 backend --- tests/test_insts.cpp | 2 ++ x86gen.hpp | 13 +++++++++++-- x86insts.inc.hpp | 1 + 3 files changed, 14 insertions(+), 2 deletions(-) diff --git a/tests/test_insts.cpp b/tests/test_insts.cpp index e91a050..6f9dce0 100644 --- a/tests/test_insts.cpp +++ b/tests/test_insts.cpp @@ -104,6 +104,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) { diff --git a/x86gen.hpp b/x86gen.hpp index f927dd3..3814e34 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -849,8 +849,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 (reg_class(select->type()) == RegClass::Float) { + Reg res = vreg(); + Reg then = vreg(); + _builder.movq_to_int(res, vreg(select->arg(2))); + _builder.movq_to_int(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))); diff --git a/x86insts.inc.hpp b/x86insts.inc.hpp index 47d328e..411cd7f 100644 --- a/x86insts.inc.hpp +++ b/x86insts.inc.hpp @@ -194,6 +194,7 @@ 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(MovQToInt, movq_to_int, mov_mem_usedef, true, { byte(0x66); rex_w(); byte(0x0f); byte(0x7e); 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(); }) From d500dbf3043fea177ec109cce3b7432d86f1a1a9 Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 23:05:29 +0200 Subject: [PATCH 07/13] store register class masks and reuse load store tests --- tests/test_insts.cpp | 50 +++++++------------------------------ x86gen.hpp | 59 +++++++++++++++++++++++--------------------- 2 files changed, 40 insertions(+), 69 deletions(-) diff --git a/tests/test_insts.cpp b/tests/test_insts.cpp index 6f9dce0..070ade7 100644 --- a/tests/test_insts.cpp +++ b/tests/test_insts.cpp @@ -253,21 +253,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); @@ -289,31 +283,6 @@ void test_alloca(DiffTestSuite& suite) { }); } -void test_load_store_f(DiffTestSuite& suite) { - for (Type type : {Type::Float32, Type::Float64}) { - Type bits_type; - if (type == Type::Float32) { - bits_type = Type::Int32; - } else { - bits_type = Type::Int64; - } - std::string name = std::string("load_store_") + to_string(type); - - suite.diff_test(name + "_roundtrip").run([=](Builder& builder, TestData& data) { - Value* value = data.input(type); - Value* ptr = builder.build_alloca(builder.build_const(Type::Int64, 8), 8); - - builder.build_store(ptr, builder.build_const(Type::Int64, ~uint64_t(0)), AliasingGroup(0), 0); - builder.build_store(ptr, value, AliasingGroup(0), 0); - data.output(builder.build_load(ptr, type, LoadFlags::None, AliasingGroup(0), 0)); - data.output(builder.build_load(ptr, Type::Int64, LoadFlags::None, AliasingGroup(0), 0)); - - builder.build_store(ptr, builder.build_const(type, 1), AliasingGroup(0), 0); - data.output(builder.build_load(ptr, bits_type, LoadFlags::None, AliasingGroup(0), 0)); - }); - } -} - void test_call(DiffTestSuite& suite) { suite.diff_test("call_preserve_none").aot(false).interpreter(false).run([](Builder& builder, TestData& data) { Value* a = data.input(Type::Int64); @@ -671,7 +640,6 @@ int main(int argc, char** argv) { test_freeze(suite); test_assume_const(suite); test_alloca(suite); - test_load_store_f(suite); test_call(suite); test_call_fp(suite, Type::Float32); test_call_fp(suite, Type::Float64); diff --git a/x86gen.hpp b/x86gen.hpp index 3814e34..dfbfaa4 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -24,7 +24,18 @@ #include "jitir.hpp" namespace metajit { - enum class RegClass { Int, Float }; + 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: @@ -481,18 +492,18 @@ namespace metajit { } } - lwir::Span args(RegClass reg_class = RegClass::Int) const { - if (reg_class == RegClass::Float) { + lwir::Span args(RegClass reg_class = RegClass::X86_INT()) const { + if (reg_class == RegClass::X86_FLOAT()) { return lwir::Span(fp_arg_regs, sizeof(fp_arg_regs) / sizeof(fp_arg_regs[0])); } return _arg_regs; } const lwir::Span& preserved() const { return _preserved_regs; } - Reg arg(size_t index, RegClass reg_class = RegClass::Int) const { return args(reg_class).at(index); } + Reg arg(size_t index, RegClass reg_class = RegClass::X86_INT()) const { return args(reg_class).at(index); } Reg preserved(size_t index) const { return _preserved_regs.at(index); } - Reg ret(RegClass reg_class = RegClass::Int) const { - if (reg_class == RegClass::Float) { + Reg ret(RegClass reg_class = RegClass::X86_INT()) const { + if (reg_class == RegClass::X86_FLOAT()) { return _fp_ret_reg; } else { return _ret_reg; @@ -544,25 +555,17 @@ namespace metajit { static RegClass reg_class(Reg preg) { assert(preg.is_physical()); if (preg.id() < Reg::X86_XMM(0).id()) { - return RegClass::Int; + return RegClass::X86_INT(); } else { - return RegClass::Float; + return RegClass::X86_FLOAT(); } } static RegClass reg_class(Type type) { if (type == Type::Float32 || type == Type::Float64) { - return RegClass::Float; - } else { - return RegClass::Int; - } - } - - static constexpr uint32_t reg_mask(RegClass reg_class) { - if (reg_class == RegClass::Int) { - return 0xffff; + return RegClass::X86_FLOAT(); } else { - return 0xffff0000; + return RegClass::X86_INT(); } } @@ -595,7 +598,7 @@ namespace metajit { }; struct VRegInfo { - RegClass reg_class = RegClass::Int; + RegClass reg_class = RegClass::X86_INT(); Reg fixed; Interval interval; Reg current_reg; @@ -649,7 +652,7 @@ namespace metajit { } } - Reg vreg(RegClass reg_class = RegClass::Int) { + 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; @@ -743,7 +746,7 @@ namespace metajit { src_class = reg_class(src); } assert(dst_class == src_class); - if (dst_class == RegClass::Float) { + if (dst_class == RegClass::X86_FLOAT()) { _builder.movsd(dst, src); } else { _builder.mov64(dst, src); @@ -849,7 +852,7 @@ namespace metajit { default: assert(false && "Unsupported pointer conversion type"); } } else if (dynmatch(SelectInst, select, inst)) { - if (reg_class(select->type()) == RegClass::Float) { + if (reg_class(select->type()) == RegClass::X86_FLOAT()) { Reg res = vreg(); Reg then = vreg(); _builder.movq_to_int(res, vreg(select->arg(2))); @@ -1284,7 +1287,7 @@ namespace metajit { for (size_t it = 1; it < call->args().size(); it++) { RegClass arg_class = reg_class(call->arg(it)->type()); size_t index; - if (arg_class == RegClass::Float) { + if (arg_class == RegClass::X86_FLOAT()) { index = fp_count++; } else { index = gp_count++; @@ -1560,7 +1563,7 @@ namespace metajit { } Reg get_free_reg(RegClass reg_class) { - uint32_t free = _free & reg_mask(reg_class); + uint32_t free = _free & reg_class.mask(); if (free == 0) { return Reg(); } else { @@ -1571,7 +1574,7 @@ namespace metajit { Reg get_lru(RegClass reg_class) { size_t min_index = 0; size_t min_value = ~size_t(0); - for (uint32_t mask = reg_mask(reg_class); mask != 0; mask &= mask - 1) { + 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]; @@ -1656,7 +1659,7 @@ namespace metajit { info.stack_offset = _stack_offset_alloc.alloc(); } X86Inst::Mem mem(Reg::X86_RSP(), (int32_t) info.stack_offset); - if (info.reg_class == RegClass::Float) { + if (info.reg_class == RegClass::X86_FLOAT()) { _builder.movsd_mem(mem, preg); } else { _builder.mov64_mem(mem, preg); @@ -1677,7 +1680,7 @@ namespace metajit { } else { assert(info.stack_offset != ~size_t(0)); X86Inst::Mem mem(Reg::X86_RSP(), (int32_t) info.stack_offset); - if (info.reg_class == RegClass::Float) { + if (info.reg_class == RegClass::X86_FLOAT()) { _builder.movsd(preg, mem); } else { _builder.mov64(preg, mem); @@ -2142,7 +2145,7 @@ namespace metajit { Reg reg = order.at(it); assert(reg.is_virtual()); - uint32_t free_mask = reg_mask(_vreg_info[reg.id()].reg_class); + 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())) { From 549a998f41d445a11dca80a0e98171f4773758ab Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 23:08:31 +0200 Subject: [PATCH 08/13] handle floating point entry arguments in both allocators --- tests/test_cfg.cpp | 24 ++++++++++++++++++++++++ x86gen.hpp | 28 ++++++++++++++++++++++++---- 2 files changed, 48 insertions(+), 4 deletions(-) diff --git a/tests/test_cfg.cpp b/tests/test_cfg.cpp index ad4b6fd..0c1c4c2 100644 --- a/tests/test_cfg.cpp +++ b/tests/test_cfg.cpp @@ -24,6 +24,30 @@ int main(int argc, char** argv) { DiffTestSuite suite("tests/output/test_cfg", argc, argv); + for (Type type : {Type::Float32, Type::Float64}) { + suite.test(std::string("float_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); + + X86CodeGen codegen(section, {Reg::X86_R12(), Reg::X86_R13()}, mode); + using Func = void(* [[clang::preserve_none]])(uint64_t, uint64_t*); + uint64_t result = 0; + ((Func) codegen.deploy())(0x12345678, &result); + unittest_assert(result == 0x12345678); + 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++) { diff --git a/x86gen.hpp b/x86gen.hpp index dfbfaa4..64d6526 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -450,7 +450,7 @@ namespace metajit { lwir::Span _arg_regs; lwir::Span _preserved_regs; Reg _ret_reg; - Reg _fp_ret_reg = Reg::X86_XMM(0); + Reg _fp_ret_reg; static constexpr Reg preserve_none_arg_regs[] = { Reg::X86_R12(), Reg::X86_R13(), Reg::X86_R14(), Reg::X86_R15(), @@ -481,11 +481,13 @@ 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"); @@ -1800,10 +1802,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 = Reg::virt(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); @@ -2319,7 +2322,13 @@ namespace metajit { _vregs.init(_section); for (Arg* arg : _section->entry()->args()) { - fix_to_preg(vreg(arg), input_pregs[arg->index()]); + Reg preg = input_pregs[arg->index()]; + _vregs[arg] = fix_to_preg(vreg(reg_class(preg)), preg); + } + for (Arg* arg : _section->entry()->args()) { + if (!(reg_class(arg->type()) == reg_class(input_pregs[arg->index()]))) { + _vregs[arg] = vreg(reg_class(arg->type())); + } } // We create one extra block for pseudo_use instructions after loops @@ -2332,6 +2341,17 @@ namespace metajit { memory_deps(); with_timer(isel, isel()); + _builder.move_to_begin(_blocks[0]); + for (Arg* arg : _section->entry()->args()) { + Reg input = Reg::virt(arg->index()); + if (!(input == vreg(arg))) { + if (reg_class(arg->type()) == RegClass::X86_FLOAT()) { + _builder.movq(vreg(arg), input); + } else { + _builder.movq_to_int(vreg(arg), input); + } + } + } autoname_insts(); if (_mode == Mode::JIT) { From 2a35ac84dc4ad4250fa4428ef496417367b3d4fe Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 23:14:23 +0200 Subject: [PATCH 09/13] compare floating point equality by bits in both backends --- llvmgen.hpp | 7 ++++--- tests/test_insts.cpp | 13 +++++++++++++ x86gen.hpp | 19 ++++++++++++++++--- 3 files changed, 33 insertions(+), 6 deletions(-) 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_insts.cpp b/tests/test_insts.cpp index 070ade7..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) { diff --git a/x86gen.hpp b/x86gen.hpp index 64d6526..7be137f 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -564,7 +564,7 @@ namespace metajit { } static RegClass reg_class(Type type) { - if (type == Type::Float32 || type == Type::Float64) { + if (is_float(type)) { return RegClass::X86_FLOAT(); } else { return RegClass::X86_INT(); @@ -792,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_int(bits_a, vreg(a)); + _builder.movq_to_int(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())) { @@ -854,7 +867,7 @@ namespace metajit { default: assert(false && "Unsupported pointer conversion type"); } } else if (dynmatch(SelectInst, select, inst)) { - if (reg_class(select->type()) == RegClass::X86_FLOAT()) { + if (is_float(select->type())) { Reg res = vreg(); Reg then = vreg(); _builder.movq_to_int(res, vreg(select->arg(2))); @@ -2345,7 +2358,7 @@ namespace metajit { for (Arg* arg : _section->entry()->args()) { Reg input = Reg::virt(arg->index()); if (!(input == vreg(arg))) { - if (reg_class(arg->type()) == RegClass::X86_FLOAT()) { + if (is_float(arg->type())) { _builder.movq(vreg(arg), input); } else { _builder.movq_to_int(vreg(arg), input); From 2835b969beeae47c8e0e069ef242da55cc372d9b Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 23:18:30 +0200 Subject: [PATCH 10/13] clarify call argument counter names --- x86gen.hpp | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/x86gen.hpp b/x86gen.hpp index 7be137f..dd30aa8 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -1297,15 +1297,15 @@ namespace metajit { assert(call->arg_count() >= 1); lwir::Span args = _builder.alloc_regs(call->args().size() - 1); - size_t gp_count = 0; - size_t fp_count = 0; + size_t int_arg_count = 0; + size_t float_arg_count = 0; for (size_t it = 1; it < call->args().size(); it++) { RegClass arg_class = reg_class(call->arg(it)->type()); size_t index; if (arg_class == RegClass::X86_FLOAT()) { - index = fp_count++; + index = float_arg_count++; } else { - index = gp_count++; + index = int_arg_count++; } assert(index < info.args(arg_class).size() && "Call with too many register arguments"); Reg arg_reg = fix_to_preg(vreg(arg_class), info.arg(index, arg_class)); From 98698ebf07ab0c3ff99dce45c5b632e54eed5e1d Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 23:21:32 +0200 Subject: [PATCH 11/13] simplify typed virtual registers and call argument access --- x86gen.hpp | 38 ++++++++++++++++++-------------------- 1 file changed, 18 insertions(+), 20 deletions(-) diff --git a/x86gen.hpp b/x86gen.hpp index dd30aa8..1b18c16 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -494,15 +494,12 @@ namespace metajit { } } - lwir::Span args(RegClass reg_class = RegClass::X86_INT()) const { - if (reg_class == RegClass::X86_FLOAT()) { - return lwir::Span(fp_arg_regs, sizeof(fp_arg_regs) / sizeof(fp_arg_regs[0])); - } - 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, RegClass reg_class = RegClass::X86_INT()) const { return args(reg_class).at(index); } Reg preserved(size_t index) const { return _preserved_regs.at(index); } Reg ret(RegClass reg_class = RegClass::X86_INT()) const { if (reg_class == RegClass::X86_FLOAT()) { @@ -661,6 +658,8 @@ namespace metajit { 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()]; @@ -682,7 +681,7 @@ namespace metajit { Reg vreg(Value* value) { if (dynmatch(Const, constant, value)) { - Reg reg = vreg(reg_class(value->type())); + Reg reg = vreg(value->type()); switch (constant->type()) { case Type::Bool: case Type::Int8: _builder.mov8_imm(reg, constant->value()); break; @@ -725,7 +724,7 @@ namespace metajit { } else if (value->is_named()) { NamedValue* named = (NamedValue*) value; if (_vregs.at(named).is_invalid()) { - _vregs[named] = vreg(reg_class(value->type())); + _vregs[named] = vreg(value->type()); } return _vregs.at(named); } else { @@ -1297,24 +1296,23 @@ namespace metajit { assert(call->arg_count() >= 1); lwir::Span args = _builder.alloc_regs(call->args().size() - 1); - size_t int_arg_count = 0; - size_t float_arg_count = 0; + size_t int_index = 0; + size_t float_index = 0; for (size_t it = 1; it < call->args().size(); it++) { - RegClass arg_class = reg_class(call->arg(it)->type()); - size_t index; - if (arg_class == RegClass::X86_FLOAT()) { - index = float_arg_count++; + Type type = call->arg(it)->type(); + Reg preg; + if (is_float(type)) { + preg = info.float_arg(float_index++); } else { - index = int_arg_count++; + preg = info.int_arg(int_index++); } - assert(index < info.args(arg_class).size() && "Call with too many register arguments"); - Reg arg_reg = fix_to_preg(vreg(arg_class), info.arg(index, arg_class)); + Reg arg_reg = fix_to_preg(vreg(type), preg); move(arg_reg, vreg(call->arg(it))); args[it - 1] = arg_reg; } RegClass ret_class = reg_class(call->type()); - Reg ret_reg = fix_to_preg(vreg(ret_class), info.ret(ret_class)); + 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); @@ -1372,7 +1370,7 @@ 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(reg_class(arg->type())); + copies[arg->index()] = vreg(arg->type()); move(copies[arg->index()], vreg(jump->arg(arg->index()))); } for (Arg* arg : jump->block()->args()) { @@ -2340,7 +2338,7 @@ namespace metajit { } for (Arg* arg : _section->entry()->args()) { if (!(reg_class(arg->type()) == reg_class(input_pregs[arg->index()]))) { - _vregs[arg] = vreg(reg_class(arg->type())); + _vregs[arg] = vreg(arg->type()); } } From c12d4b80c0c4ab45cc8934243e76618bcc31c9fd Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 23:25:55 +0200 Subject: [PATCH 12/13] track input vregs explicitly and test both entry move directions --- tests/test_cfg.cpp | 19 +++++++++++++------ x86gen.hpp | 15 +++++++++------ 2 files changed, 22 insertions(+), 12 deletions(-) diff --git a/tests/test_cfg.cpp b/tests/test_cfg.cpp index 0c1c4c2..febebea 100644 --- a/tests/test_cfg.cpp +++ b/tests/test_cfg.cpp @@ -24,8 +24,8 @@ int main(int argc, char** argv) { DiffTestSuite suite("tests/output/test_cfg", argc, argv); - for (Type type : {Type::Float32, Type::Float64}) { - suite.test(std::string("float_entry_argument_") + to_string(type)).run([type]() { + 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; @@ -38,11 +38,18 @@ int main(int argc, char** argv) { section->autoname(); section->set_ordering(BlockOrdering::Natural); - X86CodeGen codegen(section, {Reg::X86_R12(), Reg::X86_R13()}, mode); - using Func = void(* [[clang::preserve_none]])(uint64_t, uint64_t*); + uint64_t bits = 0x123456789abcdef0; uint64_t result = 0; - ((Func) codegen.deploy())(0x12345678, &result); - unittest_assert(result == 0x12345678); + 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; } }); diff --git a/x86gen.hpp b/x86gen.hpp index 1b18c16..0d856b6 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -615,6 +615,7 @@ namespace metajit { NameMap _memory_deps; NameMap _vregs; + lwir::Span _input_vregs; std::vector _vreg_info; #ifdef METAJIT_STATS @@ -1594,6 +1595,7 @@ namespace metajit { min_index = it; } } + assert(min_value != ~size_t(0)); return Reg::phys(min_index); } @@ -1813,7 +1815,7 @@ namespace metajit { std::fill(initial_state, initial_state + reg_file.size(), Reg()); for (Arg* arg : _section->entry()->args()) { - Reg input = Reg::virt(arg->index()); + 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"); @@ -2331,13 +2333,14 @@ namespace metajit { _memory_deps.init(_section); _vregs.init(_section); + _input_vregs = _builder.alloc_regs(_section->entry()->args().size()); for (Arg* arg : _section->entry()->args()) { Reg preg = input_pregs[arg->index()]; - _vregs[arg] = fix_to_preg(vreg(reg_class(preg)), preg); - } - for (Arg* arg : _section->entry()->args()) { - if (!(reg_class(arg->type()) == reg_class(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()); } } @@ -2354,7 +2357,7 @@ namespace metajit { with_timer(isel, isel()); _builder.move_to_begin(_blocks[0]); for (Arg* arg : _section->entry()->args()) { - Reg input = Reg::virt(arg->index()); + Reg input = _input_vregs.at(arg->index()); if (!(input == vreg(arg))) { if (is_float(arg->type())) { _builder.movq(vreg(arg), input); From aff0fee710251e83a1f557e098a51fca1e7161b0 Mon Sep 17 00:00:00 2001 From: Can Lehmann Date: Wed, 30 Sep 2026 23:35:05 +0200 Subject: [PATCH 13/13] rename movq to general purpose register instruction --- x86gen.hpp | 10 +++++----- x86insts.inc.hpp | 2 +- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/x86gen.hpp b/x86gen.hpp index 0d856b6..7cef25c 100644 --- a/x86gen.hpp +++ b/x86gen.hpp @@ -795,8 +795,8 @@ namespace metajit { if (is_float(a->type())) { Reg bits_a = vreg(); Reg bits_b = vreg(); - _builder.movq_to_int(bits_a, vreg(a)); - _builder.movq_to_int(bits_b, vreg(b)); + _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 { @@ -870,8 +870,8 @@ namespace metajit { if (is_float(select->type())) { Reg res = vreg(); Reg then = vreg(); - _builder.movq_to_int(res, vreg(select->arg(2))); - _builder.movq_to_int(then, vreg(select->arg(1))); + _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 { @@ -2362,7 +2362,7 @@ namespace metajit { if (is_float(arg->type())) { _builder.movq(vreg(arg), input); } else { - _builder.movq_to_int(vreg(arg), input); + _builder.movq_to_gp(vreg(arg), input); } } } diff --git a/x86insts.inc.hpp b/x86insts.inc.hpp index 411cd7f..53b4d50 100644 --- a/x86insts.inc.hpp +++ b/x86insts.inc.hpp @@ -194,7 +194,7 @@ 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(MovQToInt, movq_to_int, mov_mem_usedef, true, { byte(0x66); rex_w(); byte(0x0f); byte(0x7e); 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, { rex_opt(); byte(0x0f); byte(0x2e); modrm(); }) binop_x86_inst(UComISD, ucomisd, binop_usedef, true, { byte(0x66); rex_opt(); byte(0x0f); byte(0x2e); modrm(); })