[llvm] [SeparateConstOffsetFromGEP] Extract xor disjoint bits through xor chains (PR #226891)

via llvm-commits llvm-commits at lists.llvm.org
Sun Sep 27 22:41:33 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-transforms

Author: Yuyang Zhang (yuyzhang512)

<details>
<summary>Changes</summary>

extractDisjointBitsFromXor only matched xor(value, constant), so a constant buried under a chain of value xors, such as ((base ^ C) ^ x) ^ y, was never reached and no offset was extracted.

Trace the chain down to its constant leaf and compute the extractable bits as Const & KnownOnes(Xor). A constant bit is known-one in the expression exactly when every other operand contributes a zero to it, so the xor acts as an addition for that bit. For a plain xor this yields the same mask as the previous Const & KnownZeros(Base).

Record the traced nodes in UserChain so the expression is rebuilt around them, and restrict the non-disjoint constant substitution in removeConstOffset to the innermost xor, which is the node that owns the constant leaf.

proof: https://alive2.llvm.org/ce/z/aXWY7E

---
Full diff: https://github.com/llvm/llvm-project/pull/226891.diff


3 Files Affected:

- (modified) llvm/lib/Transforms/Scalar/SeparateConstOffsetFromGEP.cpp (+53-19) 
- (modified) llvm/test/Transforms/SeparateConstOffsetFromGEP/AMDGPU/xor-decompose.ll (+40) 
- (modified) llvm/test/Transforms/SeparateConstOffsetFromGEP/xor-decompose.ll (+108) 


``````````diff
diff --git a/llvm/lib/Transforms/Scalar/SeparateConstOffsetFromGEP.cpp b/llvm/lib/Transforms/Scalar/SeparateConstOffsetFromGEP.cpp
index 477ef6d502c8c..17c7175752569 100644
--- a/llvm/lib/Transforms/Scalar/SeparateConstOffsetFromGEP.cpp
+++ b/llvm/lib/Transforms/Scalar/SeparateConstOffsetFromGEP.cpp
@@ -300,16 +300,21 @@ class ConstantOffsetExtractor {
                     GetElementPtrInst *GEP, Value *Idx);
 
   /// Analyze a xor expression, and identify the bits in the constant operand
-  /// that are disjoint from the base operand's known set bits. For these
+  /// that are disjoint from the other operands' known set bits. For these
   /// disjoint bits, a xor is equivalent to an addition, which allows us to
   /// extract them as constant offsets that can be folded into the immediate
   /// field of addressing operations. The transformation is the following one:
   ///
   ///   Base ^ Const  becomes  (Base ^ NonDisjointBits) + DisjointBits
   ///
-  /// where DisjointBits = Const & KnownZeros(Base) and
+  /// where DisjointBits = Const & KnownOnes(Base ^ Const) and
   ///       NonDisjointBits = Const & ~DisjointBits.
   ///
+  /// A constant bit is known-one in the expression exactly when every other
+  /// operand contributes a zero to it, so for a plain xor this is equivalent to
+  /// Const & KnownZeros(Base). Taking it over the whole expression also covers
+  /// a constant buried under a chain of value xors, e.g. ((Base ^ C) ^ X) ^ Y.
+  ///
   /// Example with ptr having known-zero low bit:
   ///   Original: `xor %ptr, 3`    ; 3 = 0b11
   ///   Analysis: DisjointBits = 3 & KnownZeros(%ptr) = 0b11 & 0b01 = 0b01
@@ -862,8 +867,11 @@ Value *ConstantOffsetExtractor::removeConstOffset(unsigned ChainIndex) {
   // When rewriting xor(TheOther, NextInChain) expressions, the original
   // constant operand is replaced with the non-disjoints bits, which are the
   // non-extractable bits, i.e., those that must remain in the xor (the other
-  // bits have already compounded the GEP offset).
-  if (BO->getOpcode() == Instruction::Xor) {
+  // bits have already compounded the GEP offset). Only the innermost xor of a
+  // chain owns the constant leaf; outer xor nodes rebuild normally around their
+  // value operands.
+  if (BO->getOpcode() == Instruction::Xor &&
+      isa<ConstantInt>(UserChain[ChainIndex - 1])) {
     // The non-disjoint bits are cached in NonDisjointXorConstantBits, which is
     // always up-to-date.
     assert(NonDisjointXorConstantBits &&
@@ -907,26 +915,50 @@ Value *ConstantOffsetExtractor::removeConstOffset(unsigned ChainIndex) {
   return NewBO;
 }
 
+// Descend a chain of value xors, following the xor operand at each step, to the
+// single constant leaf that may be buried under value operands, e.g.
+// ((base ^ C) ^ num0) ^ num1. Records the visited nodes top-down in \p
+// ChainNodes. Returns the constant leaf, or null on a dead end or an ambiguous
+// fork (both operands are xors), in which case nothing is extracted.
+static ConstantInt *
+traceXorConstantLeaf(BinaryOperator *Xor,
+                     SmallVectorImpl<BinaryOperator *> &ChainNodes) {
+  ChainNodes.push_back(Xor);
+  Value *Op0 = Xor->getOperand(0), *Op1 = Xor->getOperand(1);
+  if (auto *CI = dyn_cast<ConstantInt>(Op1))
+    return CI;
+  if (auto *CI = dyn_cast<ConstantInt>(Op0))
+    return CI;
+
+  auto AsXor = [](Value *V) -> BinaryOperator * {
+    auto *BO = dyn_cast<BinaryOperator>(V);
+    return BO && BO->getOpcode() == Instruction::Xor ? BO : nullptr;
+  };
+  BinaryOperator *Xor0 = AsXor(Op0), *Xor1 = AsXor(Op1);
+  if (static_cast<bool>(Xor0) == static_cast<bool>(Xor1))
+    return nullptr;
+  return traceXorConstantLeaf(Xor0 ? Xor0 : Xor1, ChainNodes);
+}
+
 APInt ConstantOffsetExtractor::extractDisjointBitsFromXor(
     BinaryOperator *XorInst) {
   assert(XorInst && XorInst->getOpcode() == Instruction::Xor &&
          "Expected XOR instruction");
 
   unsigned BitWidth = XorInst->getType()->getScalarSizeInBits();
-  Value *BaseOp;
-  ConstantInt *XorConstantOp;
 
-  if (!match(XorInst, m_Xor(m_Value(BaseOp), m_ConstantInt(XorConstantOp))))
+  SmallVector<BinaryOperator *, 8> ChainNodes;
+  ConstantInt *XorConstantOp = traceXorConstantLeaf(XorInst, ChainNodes);
+  if (!XorConstantOp)
     return APInt::getZero(BitWidth);
-
-  const KnownBits BaseKnownBits = computeKnownBits(BaseOp, SQ);
   const APInt &ConstantValue = XorConstantOp->getValue();
 
-  // Compute the disjoint bits, i.e., those bits of the constant operand that
-  // are known-zero in the base. These disjoint bits will contribute to the
-  // final GEP offset. If there are no disjoint bits, there isn't any offset to
-  // extract from the xor.
-  const APInt DisjointBits = ConstantValue & BaseKnownBits.Zero;
+  // A constant bit is extractable as an additive offset only where it is
+  // known-one in the whole expression: there every value operand of the chain
+  // contributes a zero, so the xor behaves like an addition for that bit. (For
+  // a plain xor(base, C) this reduces to C & KnownZeros(base).) If there are no
+  // such bits, there is no offset to extract from the xor.
+  const APInt DisjointBits = ConstantValue & computeKnownBits(XorInst, SQ).One;
   if (DisjointBits.isZero())
     return DisjointBits;
 
@@ -940,12 +972,14 @@ APInt ConstantOffsetExtractor::extractDisjointBitsFromXor(
   NonDisjointXorConstantBits =
       ConstantInt::get(XorInst->getContext(), NonDisjointBits);
 
-  // UserChain maintains a path from the constant up to the GEP index. Push the
-  // xor constant operand, which is the constant leaf of the chain (which is
-  // also what `distributeCastsAndCloneChain` expects). Such a chained operand
-  // is the one to be replaced with the non-disjoint bits, while rebuilding the
-  // xor afterwards. The xor instruction itself is pushed upon returning.
+  // UserChain maintains a path from the constant leaf up to the GEP index. Push
+  // the constant leaf, then the traced xor nodes from innermost outward,
+  // excluding the top node which find() pushes upon returning. For a single xor
+  // this pushes only the constant, matching the non-chained case.
   UserChain.push_back(XorConstantOp);
+  for (BinaryOperator *Node : llvm::reverse(ChainNodes))
+    if (Node != XorInst)
+      UserChain.push_back(Node);
 
   return DisjointBits;
 }
diff --git a/llvm/test/Transforms/SeparateConstOffsetFromGEP/AMDGPU/xor-decompose.ll b/llvm/test/Transforms/SeparateConstOffsetFromGEP/AMDGPU/xor-decompose.ll
index 802e81386d647..176be091009b8 100644
--- a/llvm/test/Transforms/SeparateConstOffsetFromGEP/AMDGPU/xor-decompose.ll
+++ b/llvm/test/Transforms/SeparateConstOffsetFromGEP/AMDGPU/xor-decompose.ll
@@ -418,3 +418,43 @@ entry:
   store <8 x half> %v0, ptr addrspace(3) %ptr, align 16
   ret void
 }
+
+; A chain of value xors feeding an LDS GEP: the disjoint high bit (2048) folds
+; into the addressing offset while the low bit (4) stays in the innermost xor.
+define amdgpu_kernel void @test_chain(ptr addrspace(3) %ptr, i32 %x, i32 %n0) {
+; CHECK-LABEL: define amdgpu_kernel void @test_chain(
+; CHECK-SAME: ptr addrspace(3) [[PTR:%.*]], i32 [[X:%.*]], i32 [[N0:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[BASE:%.*]] = and i32 [[X]], 1023
+; CHECK-NEXT:    [[NUM0:%.*]] = and i32 [[N0]], 1023
+; CHECK-NEXT:    [[A1:%.*]] = xor i32 [[BASE]], 4
+; CHECK-NEXT:    [[B2:%.*]] = xor i32 [[A1]], [[NUM0]]
+; CHECK-NEXT:    [[TMP0:%.*]] = getelementptr half, ptr addrspace(3) [[PTR]], i32 [[B2]]
+; CHECK-NEXT:    [[GEP3:%.*]] = getelementptr i8, ptr addrspace(3) [[TMP0]], i32 4096
+; CHECK-NEXT:    [[V:%.*]] = load <8 x half>, ptr addrspace(3) [[GEP3]], align 16
+; CHECK-NEXT:    store <8 x half> [[V]], ptr addrspace(3) [[PTR]], align 16
+; CHECK-NEXT:    ret void
+;
+; GVN-LABEL: define amdgpu_kernel void @test_chain(
+; GVN-SAME: ptr addrspace(3) [[PTR:%.*]], i32 [[X:%.*]], i32 [[N0:%.*]]) {
+; GVN-NEXT:  [[ENTRY:.*:]]
+; GVN-NEXT:    [[BASE:%.*]] = and i32 [[X]], 1023
+; GVN-NEXT:    [[NUM0:%.*]] = and i32 [[N0]], 1023
+; GVN-NEXT:    [[A1:%.*]] = xor i32 [[BASE]], 4
+; GVN-NEXT:    [[B2:%.*]] = xor i32 [[A1]], [[NUM0]]
+; GVN-NEXT:    [[TMP0:%.*]] = getelementptr half, ptr addrspace(3) [[PTR]], i32 [[B2]]
+; GVN-NEXT:    [[GEP3:%.*]] = getelementptr i8, ptr addrspace(3) [[TMP0]], i32 4096
+; GVN-NEXT:    [[V:%.*]] = load <8 x half>, ptr addrspace(3) [[GEP3]], align 16
+; GVN-NEXT:    store <8 x half> [[V]], ptr addrspace(3) [[PTR]], align 16
+; GVN-NEXT:    ret void
+;
+entry:
+  %base = and i32 %x, 1023
+  %num0 = and i32 %n0, 1023
+  %a = xor i32 %base, 2052
+  %b = xor i32 %a, %num0
+  %gep = getelementptr half, ptr addrspace(3) %ptr, i32 %b
+  %v = load <8 x half>, ptr addrspace(3) %gep, align 16
+  store <8 x half> %v, ptr addrspace(3) %ptr, align 16
+  ret void
+}
diff --git a/llvm/test/Transforms/SeparateConstOffsetFromGEP/xor-decompose.ll b/llvm/test/Transforms/SeparateConstOffsetFromGEP/xor-decompose.ll
index 8ae5cbaf7a9c7..0a10a6e9374fe 100644
--- a/llvm/test/Transforms/SeparateConstOffsetFromGEP/xor-decompose.ll
+++ b/llvm/test/Transforms/SeparateConstOffsetFromGEP/xor-decompose.ll
@@ -161,3 +161,111 @@ entry:
   store i32 0, ptr %gep
   ret void
 }
+
+; The special constant is buried under a chain of value xors. Bit 11 (2048) is
+; known-zero in base, num0 and num1, so it is extracted as the offset while bit 2
+; (4) stays in the innermost xor.
+define ptr @xor_decompose_chain(ptr %p, i32 %x, i32 %n0, i32 %n1) {
+; CHECK-LABEL: define ptr @xor_decompose_chain(
+; CHECK-SAME: ptr [[P:%.*]], i32 [[X:%.*]], i32 [[N0:%.*]], i32 [[N1:%.*]]) {
+; CHECK-NEXT:    [[BASE:%.*]] = and i32 [[X]], 1023
+; CHECK-NEXT:    [[NUM0:%.*]] = and i32 [[N0]], 1023
+; CHECK-NEXT:    [[NUM1:%.*]] = and i32 [[N1]], 1023
+; CHECK-NEXT:    [[TMP1:%.*]] = sext i32 [[NUM1]] to i64
+; CHECK-NEXT:    [[TMP2:%.*]] = sext i32 [[NUM0]] to i64
+; CHECK-NEXT:    [[TMP3:%.*]] = sext i32 [[BASE]] to i64
+; CHECK-NEXT:    [[A1:%.*]] = xor i64 [[TMP3]], 4
+; CHECK-NEXT:    [[B2:%.*]] = xor i64 [[A1]], [[TMP2]]
+; CHECK-NEXT:    [[C3:%.*]] = xor i64 [[B2]], [[TMP1]]
+; CHECK-NEXT:    [[TMP4:%.*]] = shl i64 [[C3]], 2
+; CHECK-NEXT:    [[UGLYGEP:%.*]] = getelementptr i8, ptr [[P]], i64 [[TMP4]]
+; CHECK-NEXT:    [[UGLYGEP4:%.*]] = getelementptr i8, ptr [[UGLYGEP]], i64 8192
+; CHECK-NEXT:    ret ptr [[UGLYGEP4]]
+;
+  %base = and i32 %x, 1023
+  %num0 = and i32 %n0, 1023
+  %num1 = and i32 %n1, 1023
+  %a = xor i32 %base, 2052
+  %b = xor i32 %a, %num0
+  %c = xor i32 %b, %num1
+  %gep = getelementptr i32, ptr %p, i32 %c
+  ret ptr %gep
+}
+
+; Negative: num0 may have bit 11 set, so 2048 is not disjoint from the whole
+; chain and nothing is extracted.
+define ptr @xor_decompose_chain_sibling_not_disjoint(ptr %p, i32 %x, i32 %n0, i32 %n1) {
+; CHECK-LABEL: define ptr @xor_decompose_chain_sibling_not_disjoint(
+; CHECK-SAME: ptr [[P:%.*]], i32 [[X:%.*]], i32 [[N0:%.*]], i32 [[N1:%.*]]) {
+; CHECK-NEXT:    [[BASE:%.*]] = and i32 [[X]], 1023
+; CHECK-NEXT:    [[NUM0:%.*]] = and i32 [[N0]], 4095
+; CHECK-NEXT:    [[NUM1:%.*]] = and i32 [[N1]], 1023
+; CHECK-NEXT:    [[A:%.*]] = xor i32 [[BASE]], 2052
+; CHECK-NEXT:    [[B:%.*]] = xor i32 [[A]], [[NUM0]]
+; CHECK-NEXT:    [[C:%.*]] = xor i32 [[B]], [[NUM1]]
+; CHECK-NEXT:    [[IDXPROM:%.*]] = sext i32 [[C]] to i64
+; CHECK-NEXT:    [[GEP:%.*]] = getelementptr i32, ptr [[P]], i64 [[IDXPROM]]
+; CHECK-NEXT:    ret ptr [[GEP]]
+;
+  %base = and i32 %x, 1023
+  %num0 = and i32 %n0, 4095
+  %num1 = and i32 %n1, 1023
+  %a = xor i32 %base, 2052
+  %b = xor i32 %a, %num0
+  %c = xor i32 %b, %num1
+  %gep = getelementptr i32, ptr %p, i32 %c
+  ret ptr %gep
+}
+
+; The whole constant is disjoint through the chain: both bit 11 and bit 2 are
+; known-zero in every operand, so 2052 is fully extracted and the innermost xor
+; folds away.
+define ptr @xor_decompose_chain_full(ptr %p, i32 %x, i32 %n0, i32 %n1) {
+; CHECK-LABEL: define ptr @xor_decompose_chain_full(
+; CHECK-SAME: ptr [[P:%.*]], i32 [[X:%.*]], i32 [[N0:%.*]], i32 [[N1:%.*]]) {
+; CHECK-NEXT:    [[BASE:%.*]] = and i32 [[X]], 3
+; CHECK-NEXT:    [[NUM0:%.*]] = and i32 [[N0]], 3
+; CHECK-NEXT:    [[NUM1:%.*]] = and i32 [[N1]], 3
+; CHECK-NEXT:    [[TMP1:%.*]] = sext i32 [[NUM1]] to i64
+; CHECK-NEXT:    [[TMP2:%.*]] = sext i32 [[NUM0]] to i64
+; CHECK-NEXT:    [[TMP3:%.*]] = sext i32 [[BASE]] to i64
+; CHECK-NEXT:    [[B2:%.*]] = xor i64 [[TMP3]], [[TMP2]]
+; CHECK-NEXT:    [[C3:%.*]] = xor i64 [[B2]], [[TMP1]]
+; CHECK-NEXT:    [[TMP4:%.*]] = shl i64 [[C3]], 2
+; CHECK-NEXT:    [[UGLYGEP:%.*]] = getelementptr i8, ptr [[P]], i64 [[TMP4]]
+; CHECK-NEXT:    [[UGLYGEP4:%.*]] = getelementptr i8, ptr [[UGLYGEP]], i64 8208
+; CHECK-NEXT:    ret ptr [[UGLYGEP4]]
+;
+  %base = and i32 %x, 3
+  %num0 = and i32 %n0, 3
+  %num1 = and i32 %n1, 3
+  %a = xor i32 %base, 2052
+  %b = xor i32 %a, %num0
+  %c = xor i32 %b, %num1
+  %gep = getelementptr i32, ptr %p, i32 %c
+  ret ptr %gep
+}
+
+; A sext wrapping the chain distributes over the traced xor nodes.
+define ptr @xor_decompose_chain_sext(ptr %p, i16 %x, i16 %n0) {
+; CHECK-LABEL: define ptr @xor_decompose_chain_sext(
+; CHECK-SAME: ptr [[P:%.*]], i16 [[X:%.*]], i16 [[N0:%.*]]) {
+; CHECK-NEXT:    [[BASE:%.*]] = and i16 [[X]], 1023
+; CHECK-NEXT:    [[NUM0:%.*]] = and i16 [[N0]], 1023
+; CHECK-NEXT:    [[TMP1:%.*]] = sext i16 [[NUM0]] to i64
+; CHECK-NEXT:    [[TMP2:%.*]] = sext i16 [[BASE]] to i64
+; CHECK-NEXT:    [[A1:%.*]] = xor i64 [[TMP2]], 4
+; CHECK-NEXT:    [[B2:%.*]] = xor i64 [[A1]], [[TMP1]]
+; CHECK-NEXT:    [[TMP3:%.*]] = shl i64 [[B2]], 2
+; CHECK-NEXT:    [[UGLYGEP:%.*]] = getelementptr i8, ptr [[P]], i64 [[TMP3]]
+; CHECK-NEXT:    [[UGLYGEP3:%.*]] = getelementptr i8, ptr [[UGLYGEP]], i64 8192
+; CHECK-NEXT:    ret ptr [[UGLYGEP3]]
+;
+  %base = and i16 %x, 1023
+  %num0 = and i16 %n0, 1023
+  %a = xor i16 %base, 2052
+  %b = xor i16 %a, %num0
+  %idx = sext i16 %b to i64
+  %gep = getelementptr i32, ptr %p, i64 %idx
+  ret ptr %gep
+}

``````````

</details>


https://github.com/llvm/llvm-project/pull/226891


More information about the llvm-commits mailing list