[llvm] [ValueTracking] Propagate known bits for nvvm.mulhi (PR #228407)
via llvm-commits
llvm-commits at lists.llvm.org
Fri Oct 2 05:12:21 PDT 2026
https://github.com/peterbell10 created https://github.com/llvm/llvm-project/pull/228407
None
>From ce8d526a0a6e2e70724b0a7dbeddeac11cd2638f Mon Sep 17 00:00:00 2001
From: Peter Bell <peterbell10 at openai.com>
Date: Fri, 2 Oct 2026 13:11:33 +0100
Subject: [PATCH] [ValueTracking] Propagate known bits for nvvm.mulhi
---
llvm/lib/Analysis/ValueTracking.cpp | 7 +++
llvm/unittests/Analysis/ValueTrackingTest.cpp | 57 +++++++++++++++++++
2 files changed, 64 insertions(+)
diff --git a/llvm/lib/Analysis/ValueTracking.cpp b/llvm/lib/Analysis/ValueTracking.cpp
index 1580571a014518..829ad49e918ae3 100644
--- a/llvm/lib/Analysis/ValueTracking.cpp
+++ b/llvm/lib/Analysis/ValueTracking.cpp
@@ -58,6 +58,7 @@
#include "llvm/IR/Intrinsics.h"
#include "llvm/IR/IntrinsicsAArch64.h"
#include "llvm/IR/IntrinsicsAMDGPU.h"
+#include "llvm/IR/IntrinsicsNVPTX.h"
#include "llvm/IR/IntrinsicsRISCV.h"
#include "llvm/IR/IntrinsicsX86.h"
#include "llvm/IR/LLVMContext.h"
@@ -2235,6 +2236,9 @@ static void computeKnownBitsFromOperator(const Operator *I,
Known &= Known2.anyextOrTrunc(BitWidth);
break;
}
+ case Intrinsic::nvvm_mulhi_s:
+ case Intrinsic::nvvm_mulhi_i:
+ case Intrinsic::nvvm_mulhi_ll:
case Intrinsic::x86_sse2_pmulh_w:
case Intrinsic::x86_avx2_pmulh_w:
case Intrinsic::x86_avx512_pmulh_w_512:
@@ -2242,6 +2246,9 @@ static void computeKnownBitsFromOperator(const Operator *I,
computeKnownBits(I->getOperand(1), DemandedElts, Known2, Q, Depth + 1);
Known = KnownBits::mulhs(Known, Known2);
break;
+ case Intrinsic::nvvm_mulhi_us:
+ case Intrinsic::nvvm_mulhi_ui:
+ case Intrinsic::nvvm_mulhi_ull:
case Intrinsic::x86_sse2_pmulhu_w:
case Intrinsic::x86_avx2_pmulhu_w:
case Intrinsic::x86_avx512_pmulhu_w_512:
diff --git a/llvm/unittests/Analysis/ValueTrackingTest.cpp b/llvm/unittests/Analysis/ValueTrackingTest.cpp
index 0634f43f3e79c8..f69da84222298d 100644
--- a/llvm/unittests/Analysis/ValueTrackingTest.cpp
+++ b/llvm/unittests/Analysis/ValueTrackingTest.cpp
@@ -18,6 +18,7 @@
#include "llvm/IR/InstIterator.h"
#include "llvm/IR/Instructions.h"
#include "llvm/IR/IntrinsicInst.h"
+#include "llvm/IR/IntrinsicsNVPTX.h"
#include "llvm/IR/LLVMContext.h"
#include "llvm/IR/Module.h"
#include "llvm/Support/ErrorHandling.h"
@@ -1514,6 +1515,62 @@ TEST_F(ComputeKnownBitsTest, ComputeKnownMulBits) {
expectKnownBits(/*zero*/ 95u, /*one*/ 32u);
}
+TEST_F(ComputeKnownBitsTest, NVVMMulHi) {
+ for (auto [ID, IsSigned] : {std::pair{Intrinsic::nvvm_mulhi_s, true},
+ {Intrinsic::nvvm_mulhi_i, true},
+ {Intrinsic::nvvm_mulhi_ll, true},
+ {Intrinsic::nvvm_mulhi_us, false},
+ {Intrinsic::nvvm_mulhi_ui, false},
+ {Intrinsic::nvvm_mulhi_ull, false}}) {
+ Module M("test", Context);
+ Function *MulHi = Intrinsic::getOrInsertDeclaration(&M, ID);
+ SCOPED_TRACE(MulHi->getName().str());
+ Function *F = Function::Create(MulHi->getFunctionType(),
+ Function::ExternalLinkage, "test", M);
+ IRBuilder<> B(BasicBlock::Create(Context, "entry", F));
+ Type *Ty = MulHi->getReturnType();
+ unsigned BitWidth = Ty->getIntegerBitWidth();
+ APInt Zero = APInt::getZero(BitWidth);
+ APInt AllOnes = APInt::getAllOnes(BitWidth);
+ Value *X = F->getArg(0);
+ Value *Y = F->getArg(1);
+
+ auto Check = [&](Value *LHS, Value *RHS, const APInt &ExpectedZero,
+ const APInt &ExpectedOne) {
+ Value *Product = B.CreateCall(MulHi, {LHS, RHS});
+ KnownBits Known = computeKnownBits(Product, M.getDataLayout());
+ EXPECT_FALSE(Known.hasConflict());
+ EXPECT_EQ(Known.Zero, ExpectedZero);
+ EXPECT_EQ(Known.One, ExpectedOne);
+ };
+
+ Check(X, Y, Zero, Zero);
+ Check(X, ConstantInt::get(Ty, 0), AllOnes, Zero);
+ Check(ConstantInt::get(Ty, 0), Y, AllOnes, Zero);
+
+ // Two half-width operands have no set bits in the high half of the product.
+ APInt HalfMask = APInt::getLowBitsSet(BitWidth, BitWidth / 2);
+ Check(B.CreateAnd(X, HalfMask), B.CreateAnd(Y, HalfMask), AllOnes, Zero);
+
+ // Known leading zeros from both operands propagate to the high product.
+ APInt Mask = APInt::getLowBitsSet(BitWidth, BitWidth - 2);
+ Check(B.CreateAnd(X, Mask), B.CreateAnd(Y, Mask),
+ APInt::getHighBitsSet(BitWidth, 4), Zero);
+
+ // The low two bits of the high product are 01, even with unknown sign bits.
+ APInt Bit = APInt::getOneBitSet(BitWidth, BitWidth / 2);
+ Value *LHS = B.CreateOr(B.CreateShl(X, BitWidth / 2 + 2), Bit);
+ Value *RHS = B.CreateOr(B.CreateShl(Y, BitWidth / 2 + 2), Bit);
+ Check(LHS, RHS, APInt(BitWidth, 2), APInt(BitWidth, 1));
+
+ // Sign extension and zero extension give different high halves for -1 * 2.
+ APInt Expected = IsSigned ? AllOnes : APInt(BitWidth, 1);
+ Check(ConstantInt::get(Ty, AllOnes), ConstantInt::get(Ty, 2), ~Expected,
+ Expected);
+ B.CreateRet(ConstantInt::get(Ty, 0));
+ }
+}
+
TEST_F(ComputeKnownFPClassTest, SelectPos0) {
parseAssembly(
"define float @test(i1 %cond) {\n"
More information about the llvm-commits
mailing list