[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