[llvm] [SeparateConstOffsetFromGEP] Extract xor disjoint bits through xor chains (PR #226891)
Yuyang Zhang via llvm-commits
llvm-commits at lists.llvm.org
Sun Sep 27 22:41:02 PDT 2026
https://github.com/yuyzhang512 created https://github.com/llvm/llvm-project/pull/226891
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
>From 15b081ad04f796d19205dec68f19b995d57b9737 Mon Sep 17 00:00:00 2001
From: yuyzhang512 <yuyzhang at amd.com>
Date: Mon, 28 Sep 2026 05:32:23 +0000
Subject: [PATCH] [SeparateConstOffsetFromGEP] Extract xor disjoint bits
through xor chains
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.
---
.../Scalar/SeparateConstOffsetFromGEP.cpp | 72 +++++++++---
.../AMDGPU/xor-decompose.ll | 40 +++++++
.../xor-decompose.ll | 108 ++++++++++++++++++
3 files changed, 201 insertions(+), 19 deletions(-)
diff --git a/llvm/lib/Transforms/Scalar/SeparateConstOffsetFromGEP.cpp b/llvm/lib/Transforms/Scalar/SeparateConstOffsetFromGEP.cpp
index 477ef6d502c8cc..17c71757525690 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 802e81386d6476..176be091009b8a 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 8ae5cbaf7a9c7a..0a10a6e9374fe4 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
+}
More information about the llvm-commits
mailing list