diff --git a/ir/instr.cpp b/ir/instr.cpp index 20613c968..5748c53ba 100644 --- a/ir/instr.cpp +++ b/ir/instr.cpp @@ -189,6 +189,8 @@ void BinOp::print(ostream &os) const { case Clmul: str = "clmul "; break; case PExt: str = "pext "; break; case PDep: str = "pdep "; break; + case UMulH: str = "umulh "; break; + case SMulH: str = "smulh "; break; } os << getName() << " = " << str; @@ -491,6 +493,18 @@ StateValue BinOp::toSMT(State &s) const { return {a.pdep(b), ap && bp}; }; break; + case UMulH: + fn = [&](auto &a, auto &ap, auto &b, auto &bp) -> StateValue { + auto bw = a.bits(); + return {(a.zext(bw) * b.zext(bw)).extract(2*bw - 1, bw), ap && bp}; + }; + break; + case SMulH: + fn = [&](auto &a, auto &ap, auto &b, auto &bp) -> StateValue { + auto bw = a.bits(); + return {(a.sext(bw) * b.sext(bw)).extract(2*bw - 1, bw), ap && bp}; + }; + break; } function(const expr&, const expr&, const expr&, diff --git a/ir/instr.h b/ir/instr.h index d9fa8ea51..ffa92d7bd 100644 --- a/ir/instr.h +++ b/ir/instr.h @@ -41,7 +41,7 @@ class BinOp final : public Instr { SAdd_Overflow, UAdd_Overflow, SSub_Overflow, USub_Overflow, SMul_Overflow, UMul_Overflow, And, Or, Xor, Cttz, Ctlz, UMin, UMax, SMin, SMax, Abs, - UCmp, SCmp, Clmul, PExt, PDep }; + UCmp, SCmp, Clmul, PExt, PDep, UMulH, SMulH }; enum Flags { None = 0, NSW = 1 << 0, NUW = 1 << 1, Exact = 1 << 2, Disjoint = 1 << 3 }; private: diff --git a/llvm_util/llvm2alive.cpp b/llvm_util/llvm2alive.cpp index b7e821f99..c79ddf3c1 100644 --- a/llvm_util/llvm2alive.cpp +++ b/llvm_util/llvm2alive.cpp @@ -843,7 +843,9 @@ class llvm2alive_ : public llvm::InstVisitor> { case llvm::Intrinsic::scmp: case llvm::Intrinsic::clmul: case llvm::Intrinsic::pext: - case llvm::Intrinsic::pdep: { + case llvm::Intrinsic::pdep: + case llvm::Intrinsic::umulh: + case llvm::Intrinsic::smulh: { PARSE_BINOP(); addNoundefAssumes(i, {a, b}); BinOp::Op op; @@ -872,6 +874,8 @@ class llvm2alive_ : public llvm::InstVisitor> { case llvm::Intrinsic::clmul: op = BinOp::Clmul; break; case llvm::Intrinsic::pext: op = BinOp::PExt; break; case llvm::Intrinsic::pdep: op = BinOp::PDep; break; + case llvm::Intrinsic::umulh: op = BinOp::UMulH; break; + case llvm::Intrinsic::smulh: op = BinOp::SMulH; break; default: UNREACHABLE(); } ret = make_unique(*ty, value_name(i), *a, *b, op); diff --git a/tests/alive-tv/usmulh.srctgt.ll b/tests/alive-tv/usmulh.srctgt.ll new file mode 100644 index 000000000..fa0e723ab --- /dev/null +++ b/tests/alive-tv/usmulh.srctgt.ll @@ -0,0 +1,55 @@ +define i16 @src_smulh_ashr(i8 %a, i8 %b) { + %ea = sext i8 %a to i16 + %eb = sext i8 %b to i16 + %mul = mul i16 %ea, %eb + %shr = ashr i16 %mul, 8 + ret i16 %shr +} + +define i16 @tgt_smulh_ashr(i8 %a, i8 %b) { + %res = call i8 @llvm.smulh.i8(i8 %a, i8 %b) + %ext = sext i8 %res to i16 + ret i16 %ext +} + +define i16 @src_smulh_lshr(i8 %a, i8 %b) { + %ea = sext i8 %a to i16 + %eb = sext i8 %b to i16 + %mul = mul i16 %ea, %eb + %shr = lshr i16 %mul, 8 + ret i16 %shr +} + +define i16 @tgt_smulh_lshr(i8 %a, i8 %b) { + %res = call i8 @llvm.smulh.i8(i8 %a, i8 %b) + %ext = zext i8 %res to i16 + ret i16 %ext +} + +define i16 @src_umulh_ashr(i8 %a, i8 %b) { + %ea = zext i8 %a to i16 + %eb = zext i8 %b to i16 + %mul = mul i16 %ea, %eb + %shr = ashr i16 %mul, 8 + ret i16 %shr +} + +define i16 @tgt_umulh_ashr(i8 %a, i8 %b) { + %res = call i8 @llvm.umulh.i8(i8 %a, i8 %b) + %ext = sext i8 %res to i16 + ret i16 %ext +} + +define i16 @src_umulh_lshr(i8 %a, i8 %b) { + %ea = zext i8 %a to i16 + %eb = zext i8 %b to i16 + %mul = mul i16 %ea, %eb + %shr = lshr i16 %mul, 8 + ret i16 %shr +} + +define i16 @tgt_umulh_lshr(i8 %a, i8 %b) { + %res = call i8 @llvm.umulh.i8(i8 %a, i8 %b) + %ext = zext i8 %res to i16 + ret i16 %ext +} diff --git a/tests/unit/usmulh.opt b/tests/unit/usmulh.opt new file mode 100644 index 000000000..ba2de45c6 --- /dev/null +++ b/tests/unit/usmulh.opt @@ -0,0 +1,29 @@ +Name: umulh constant1 +%r = umulh i4 1, i4 2 + => +%r = 0 + +Name: umulh constant2 +%r = umulh i4 5, i4 6 + => +%r = 1 + +Name: umulh constant2 +%r = umulh i4 4, i4 10 + => +%r = 2 + +Name: smulh constant1 +%r = smulh i4 1, i4 2 + => +%r = 0 + +Name: smulh constant2 +%r = smulh i4 5, i4 6 + => +%r = 1 + +Name: smulh constant2 +%r = smulh i4 4, i4 10 + => +%r = -2 diff --git a/tools/alive_lexer.re b/tools/alive_lexer.re index bd7ed4f54..c173437fb 100644 --- a/tools/alive_lexer.re +++ b/tools/alive_lexer.re @@ -289,6 +289,8 @@ space+ { "clmul" { return CLMUL; } "pext" { return PEXT; } "pdep" { return PDEP; } +"umulh" { return UMULH; } +"smulh" { return SMULH; } "oeq" { return OEQ; } "ogt" { return OGT; } "oge" { return OGE; } diff --git a/tools/alive_parser.cpp b/tools/alive_parser.cpp index 5bbca9c65..9a6309f72 100644 --- a/tools/alive_parser.cpp +++ b/tools/alive_parser.cpp @@ -713,6 +713,8 @@ static unsigned parse_binop_flags(token op_token) { case CLMUL: case PEXT: case PDEP: + case UMULH: + case SMULH: return BinOp::None; default: UNREACHABLE(); @@ -788,6 +790,8 @@ static unique_ptr parse_binop(string_view name, token op_token) { case CLMUL: op = BinOp::Clmul; break; case PEXT: op = BinOp::PExt; break; case PDEP: op = BinOp::PDep; break; + case UMULH: op = BinOp::UMulH; break; + case SMULH: op = BinOp::SMulH; break; default: UNREACHABLE(); } @@ -1282,6 +1286,8 @@ static unique_ptr parse_instr(string_view name) { case CLMUL: case PEXT: case PDEP: + case UMULH: + case SMULH: return parse_binop(name, t); case FADD: case FSUB: diff --git a/tools/tokens.h b/tools/tokens.h index a62871f41..b9daa23fd 100644 --- a/tools/tokens.h +++ b/tools/tokens.h @@ -139,6 +139,7 @@ TOKEN(SITOFP) TOKEN(SLE) TOKEN(SLT) TOKEN(SMUL_OVERFLOW) +TOKEN(SMULH) TOKEN(SREM) TOKEN(SSUB_OVERFLOW) TOKEN(SSUB_SAT) @@ -158,6 +159,7 @@ TOKEN(UITOFP) TOKEN(ULE) TOKEN(ULT) TOKEN(UMUL_OVERFLOW) +TOKEN(UMULH) TOKEN(UMIN) TOKEN(UMAX) TOKEN(SMIN) diff --git a/tv/tv.cpp b/tv/tv.cpp index dfaaddd99..a860b1f2d 100644 --- a/tv/tv.cpp +++ b/tv/tv.cpp @@ -545,7 +545,7 @@ bool is_terminate_pass(const llvm::StringRef &pass0) { } -struct TVPass : public llvm::detail::PassInfoMixin { +struct TVPass : public llvm::RequiredPassInfoMixin { static string batched_pass_begin_name; static bool batch_started; // # of run passes when batching is enabled @@ -681,7 +681,8 @@ void runTVPass(Ty &M) { tv.run(M, get_TLI); } -struct ClangTVFinalizePass : public llvm::detail::PassInfoMixin { +struct ClangTVFinalizePass + : public llvm::RequiredPassInfoMixin { llvm::PreservedAnalyses run(llvm::Module &M, llvm::ModuleAnalysisManager &AM) { if (is_clangtv) {