[llvm] [InstCombine][NFC] Use uint64_t for TruncRatio to avoid overflow (PR #225150)

via llvm-commits llvm-commits at lists.llvm.org
Fri Sep 25 07:56:29 PDT 2026


https://github.com/ahradwan2-public updated https://github.com/llvm/llvm-project/pull/225150

>From ebc50e5dbba9d678a74f7a6be08580e36b7d6991 Mon Sep 17 00:00:00 2001
From: Ahmed Radwan <ahradwan at amd.com>
Date: Mon, 21 Sep 2026 11:32:45 -0400
Subject: [PATCH 1/2] [InstCombine][NFC] Use uint64_t for TruncRatio to avoid
 32-bit overflow of BitCastNumElts in foldVecExtTruncToExtElt

---
 llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp | 2 +-
 1 file changed, 1 insertion(+), 1 deletion(-)

diff --git a/llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp b/llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp
index 97defc5e3ddd3..2784fd7893f70 100644
--- a/llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp
+++ b/llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp
@@ -754,7 +754,7 @@ static Instruction *foldVecExtTruncToExtElt(TruncInst &Trunc,
   // A badly fit destination size would result in an invalid cast.
   unsigned SrcBits = SrcType->getScalarSizeInBits();
   unsigned DstBits = DstType->getScalarSizeInBits();
-  unsigned TruncRatio = SrcBits / DstBits;
+  uint64_t TruncRatio = SrcBits / DstBits;
   if ((SrcBits % DstBits) != 0)
     return nullptr;
 

>From af926608173aac582865077f6dbd2f27b1ff46d2 Mon Sep 17 00:00:00 2001
From: Ahmed Radwan <ahradwan at amd.com>
Date: Fri, 25 Sep 2026 10:56:04 -0400
Subject: [PATCH 2/2] [InstCombine][NFC] Replace assert with a bail condition
 and add lit test

---
 .../Transforms/InstCombine/InstCombineCasts.cpp |  8 +++++---
 .../InstCombine/trunc-extractelement.ll         | 17 +++++++++++++++++
 2 files changed, 22 insertions(+), 3 deletions(-)

diff --git a/llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp b/llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp
index 2784fd7893f70..e3ccf6788b673 100644
--- a/llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp
+++ b/llvm/lib/Transforms/InstCombine/InstCombineCasts.cpp
@@ -771,6 +771,11 @@ static Instruction *foldVecExtTruncToExtElt(TruncInst &Trunc,
   auto VecElts = VecOpTy->getElementCount();
 
   uint64_t BitCastNumElts = VecElts.getKnownMinValue() * TruncRatio;
+  // Computed in 64-bit above to avoid a 32-bit overflow. Bail out if the
+  // element count exceeds IntegerType::MAX_INT_BITS, as we cannot create a
+  // wider vector type.
+  if (BitCastNumElts > IntegerType::MAX_INT_BITS)
+    return nullptr;
   // Make sure we don't overflow in the calculation of the new index.
   // (VecOpIdx + 1) * TruncRatio should not overflow.
   if (Cst->uge(std::numeric_limits<uint64_t>::max() / TruncRatio))
@@ -796,9 +801,6 @@ static Instruction *foldVecExtTruncToExtElt(TruncInst &Trunc,
                                               : (NewIdx + IdxOfs);
   }
 
-  assert(BitCastNumElts <= std::numeric_limits<uint32_t>::max() &&
-         "overflow 32-bits");
-
   auto *BitCastTo =
       VectorType::get(DstType, BitCastNumElts, VecElts.isScalable());
   Value *BitCast = IC.Builder.CreateBitCast(VecOp, BitCastTo);
diff --git a/llvm/test/Transforms/InstCombine/trunc-extractelement.ll b/llvm/test/Transforms/InstCombine/trunc-extractelement.ll
index 074b5b35d7f2d..a4a59ee4b7cb7 100644
--- a/llvm/test/Transforms/InstCombine/trunc-extractelement.ll
+++ b/llvm/test/Transforms/InstCombine/trunc-extractelement.ll
@@ -330,3 +330,20 @@ entry:
   %trunc = trunc i32 %ext to i16
   ret i16 %trunc
 }
+
+; Do not fold if the bitcast vector's element count (NumElts * TruncRatio)
+; exceeds IntegerType::MAX_INT_BITS. Here 536870912 (2^29) * 8 = 2^32, which
+; also wraps to 0 when the product is computed in 32 bits; computing it in
+; 64 bits and bailing out avoids forming an invalid vector type.
+define i8 @test_bitcast_numelts_overflow(<536870912 x i64> %vec) {
+; ANY-LABEL: @test_bitcast_numelts_overflow(
+; ANY-NEXT:  entry:
+; ANY-NEXT:    [[EXT:%.*]] = extractelement <536870912 x i64> [[VEC:%.*]], i64 0
+; ANY-NEXT:    [[TRUNC:%.*]] = trunc i64 [[EXT]] to i8
+; ANY-NEXT:    ret i8 [[TRUNC]]
+;
+entry:
+  %ext = extractelement <536870912 x i64> %vec, i64 0
+  %trunc = trunc i64 %ext to i8
+  ret i8 %trunc
+}



More information about the llvm-commits mailing list