[llvm] [ARM] Combine four byte add stores into SWAR word ops (PR #203153)

via llvm-commits llvm-commits at lists.llvm.org
Mon Aug 3 18:57:44 PDT 2026


https://github.com/Sightem updated https://github.com/llvm/llvm-project/pull/203153

>From 8d53a6f23f03068660a183752f2978de80e5556f Mon Sep 17 00:00:00 2001
From: Sightem <sightem at national.shitposting.agency>
Date: Wed, 10 Jun 2026 18:42:53 -0700
Subject: [PATCH] [ARM] Combine four byte add stores into SWAR word ops

Recognize four consecutive i8 stores of dst[i] + src[i] in the ARM
DAG combiner and replace them with an i32 load/load/store plus the
masked add SWAR sequence when unaligned word accesses are legal and fast.

This avoids the register-heavy scalar byte sequence on targets without
vector support while preserving strict align codegen
---
 llvm/lib/Target/ARM/ARMISelLowering.cpp  | 320 +++++++++++++++++++++++
 llvm/test/CodeGen/ARM/merge-byte-adds.ll |  46 ++++
 2 files changed, 366 insertions(+)
 create mode 100644 llvm/test/CodeGen/ARM/merge-byte-adds.ll

diff --git a/llvm/lib/Target/ARM/ARMISelLowering.cpp b/llvm/lib/Target/ARM/ARMISelLowering.cpp
index e08a4f5ebda08..9e7be175798ec 100644
--- a/llvm/lib/Target/ARM/ARMISelLowering.cpp
+++ b/llvm/lib/Target/ARM/ARMISelLowering.cpp
@@ -613,6 +613,9 @@ ARMTargetLowering::ARMTargetLowering(const TargetMachine &TM_,
     setTargetDAGCombine({ISD::BRCOND, ISD::BR_CC});
   }
 
+  if (!Subtarget->isThumb1Only())
+    setTargetDAGCombine(ISD::STORE);
+
   if (Subtarget->hasNEON()) {
     addDRTypeForNEON(MVT::v2f32);
     addDRTypeForNEON(MVT::v8i8);
@@ -16909,6 +16912,320 @@ static SDValue PerformExtractFpToIntStores(StoreSDNode *St, SelectionDAG &DAG) {
   return Store;
 }
 
+struct ByteAddStoreMatch {
+  StoreSDNode *Store = nullptr;
+  LoadSDNode *DstLoad = nullptr;
+  LoadSDNode *SrcLoad = nullptr;
+  int64_t Offset = 0;
+};
+
+static LoadSDNode *matchByteAddLoad(SDValue V, EVT AddVT) {
+  V = peekThroughBitcasts(V);
+  auto *Ld = dyn_cast<LoadSDNode>(V);
+  if (!Ld || Ld->getMemoryVT() != MVT::i8 || Ld->getValueType(0) != AddVT ||
+      !Ld->hasNUsesOfValue(1, 0) || !Ld->isSimple() || Ld->isIndexed())
+    return nullptr;
+  return Ld;
+}
+
+static bool byteLoadMatchesStore(LoadSDNode *Ld, StoreSDNode *St,
+                                 SelectionDAG &DAG) {
+  int64_t Offset = 0;
+  BaseIndexOffset StorePtr = BaseIndexOffset::match(St, DAG);
+  BaseIndexOffset LoadPtr = BaseIndexOffset::match(Ld, DAG);
+  return StorePtr.equalBaseIndex(LoadPtr, DAG, Offset) && Offset == 0;
+}
+
+static bool matchByteAddStore(StoreSDNode *St, SelectionDAG &DAG,
+                              ByteAddStoreMatch &Match) {
+  if (!St->isSimple() || St->isIndexed() || St->getMemoryVT() != MVT::i8)
+    return false;
+
+  SDValue StoredVal = peekThroughBitcasts(St->getValue());
+  EVT AddVT = StoredVal.getValueType();
+  if (StoredVal.getOpcode() != ISD::ADD || !StoredVal.hasOneUse() ||
+      (AddVT != MVT::i8 && AddVT != MVT::i32))
+    return false;
+
+  LoadSDNode *LHS = matchByteAddLoad(StoredVal.getOperand(0), AddVT);
+  LoadSDNode *RHS = matchByteAddLoad(StoredVal.getOperand(1), AddVT);
+  if (!LHS || !RHS)
+    return false;
+
+  LoadSDNode *DstLoad = nullptr;
+  LoadSDNode *SrcLoad = nullptr;
+  if (byteLoadMatchesStore(LHS, St, DAG)) {
+    DstLoad = LHS;
+    SrcLoad = RHS;
+  } else if (byteLoadMatchesStore(RHS, St, DAG)) {
+    DstLoad = RHS;
+    SrcLoad = LHS;
+  } else {
+    return false;
+  }
+
+  Match.Store = St;
+  Match.DstLoad = DstLoad;
+  Match.SrcLoad = SrcLoad;
+  return true;
+}
+
+static SDValue getMergedStoreChain(ArrayRef<ByteAddStoreMatch> Stores,
+                                   SelectionDAG &DAG) {
+  SmallVector<SDValue, 4> Chains;
+  SmallPtrSet<const SDNode *, 4> Visited;
+  SDLoc DL(Stores[0].Store);
+
+  for (const ByteAddStoreMatch &M : Stores)
+    Visited.insert(M.Store);
+
+  for (const ByteAddStoreMatch &M : Stores)
+    if (Visited.insert(M.Store->getChain().getNode()).second)
+      Chains.push_back(M.Store->getChain());
+
+  assert(!Chains.empty() && "Expected at least one chain");
+  return DAG.getTokenFactor(DL, Chains);
+}
+
+static bool checkByteAddStoreDependencies(ArrayRef<ByteAddStoreMatch> Stores,
+                                          SDNode *RootNode) {
+  SmallPtrSet<const SDNode *, 32> Visited;
+  SmallVector<const SDNode *, 8> Worklist;
+
+  Worklist.push_back(RootNode);
+  while (!Worklist.empty()) {
+    const SDNode *N = Worklist.pop_back_val();
+    if (!Visited.insert(N).second)
+      continue;
+    if (N->getOpcode() == ISD::TokenFactor)
+      for (SDValue Op : N->ops())
+        Worklist.push_back(Op.getNode());
+  }
+
+  unsigned Max = 1024 + Visited.size();
+  for (const ByteAddStoreMatch &M : Stores)
+    for (const SDValue &Op : M.Store->op_values())
+      Worklist.push_back(Op.getNode());
+
+  for (const ByteAddStoreMatch &M : Stores)
+    if (SDNode::hasPredecessorHelper(M.Store, Visited, Worklist, Max))
+      return false;
+  return true;
+}
+
+static SDValue
+PerformByteAddStoreCombine(StoreSDNode *St,
+                           TargetLowering::DAGCombinerInfo &DCI) {
+  SelectionDAG &DAG = DCI.DAG;
+  LLVMContext &Context = *DAG.getContext();
+  const DataLayout &DL = DAG.getDataLayout();
+  const TargetLowering &TLI = DAG.getTargetLoweringInfo();
+  constexpr unsigned NumElem = 4;
+  EVT JointVT = MVT::i32;
+
+  ByteAddStoreMatch FirstMatch;
+  if (!matchByteAddStore(St, DAG, FirstMatch))
+    return SDValue();
+
+  BaseIndexOffset BasePtr = BaseIndexOffset::match(St, DAG);
+  if (BasePtr.getBase().isUndef())
+    return SDValue();
+
+  SDNode *RootNode = St->getChain().getNode();
+  bool RootWasLoad = isa<LoadSDNode>(RootNode);
+  if (RootWasLoad)
+    RootNode = cast<LoadSDNode>(RootNode)->getChain().getNode();
+
+  SmallVector<ByteAddStoreMatch, 8> Candidates;
+  Candidates.push_back(FirstMatch);
+  auto TryAddCandidate = [&](StoreSDNode *Other) {
+    if (Other == St)
+      return;
+    if (St->isNonTemporal() != Other->isNonTemporal())
+      return;
+
+    BaseIndexOffset Ptr = BaseIndexOffset::match(Other, DAG);
+    int64_t Offset = 0;
+    if (!BasePtr.equalBaseIndex(Ptr, DAG, Offset))
+      return;
+
+    ByteAddStoreMatch Match;
+    if (!matchByteAddStore(Other, DAG, Match))
+      return;
+    Match.Offset = Offset;
+    Candidates.push_back(Match);
+  };
+
+  auto TryChainUse = [&](SDUse &Use) {
+    if (Use.getOperandNo() != 0)
+      return;
+    if (auto *OtherStore = dyn_cast<StoreSDNode>(Use.getUser()))
+      TryAddCandidate(OtherStore);
+  };
+
+  if (RootWasLoad) {
+    for (SDUse &Use : RootNode->uses()) {
+      if (Use.getOperandNo() != 0)
+        continue;
+      if (auto *OtherLoad = dyn_cast<LoadSDNode>(Use.getUser())) {
+        for (SDUse &LoadUse : OtherLoad->uses())
+          TryChainUse(LoadUse);
+      } else {
+        TryChainUse(Use);
+      }
+    }
+  } else {
+    for (SDUse &Use : RootNode->uses())
+      TryChainUse(Use);
+  }
+
+  llvm::sort(Candidates,
+             [](const ByteAddStoreMatch &LHS, const ByteAddStoreMatch &RHS) {
+               return LHS.Offset < RHS.Offset;
+             });
+
+  for (unsigned Start = 0; Start + NumElem <= Candidates.size(); ++Start) {
+    ArrayRef<ByteAddStoreMatch> Stores(Candidates.data() + Start, NumElem);
+    if (!llvm::any_of(
+            Stores, [St](const ByteAddStoreMatch &M) { return M.Store == St; }))
+      continue;
+
+    int64_t StartOffset = Stores[0].Offset;
+    bool Consecutive = true;
+    for (unsigned I = 1; I != NumElem; ++I)
+      if (Stores[I].Offset - StartOffset != static_cast<int64_t>(I)) {
+        Consecutive = false;
+        break;
+      }
+    if (!Consecutive)
+      continue;
+
+    BaseIndexOffset SrcBasePtr;
+    SDValue DstLoadChain;
+    SDValue SrcLoadChain;
+    bool DstDereferenceable = true;
+    bool SrcDereferenceable = true;
+    bool MatchFailed = false;
+
+    for (unsigned I = 0; I != NumElem; ++I) {
+      const ByteAddStoreMatch &M = Stores[I];
+      int64_t LaneOffset = M.Offset - StartOffset;
+      BaseIndexOffset SrcPtr = BaseIndexOffset::match(M.SrcLoad, DAG);
+
+      if (I == 0) {
+        SrcBasePtr = SrcPtr;
+        DstLoadChain = M.DstLoad->getChain();
+        SrcLoadChain = M.SrcLoad->getChain();
+      } else {
+        int64_t SrcOffset = 0;
+        if (M.DstLoad->getChain() != DstLoadChain ||
+            M.SrcLoad->getChain() != SrcLoadChain ||
+            !SrcBasePtr.equalBaseIndex(SrcPtr, DAG, SrcOffset) ||
+            SrcOffset != LaneOffset ||
+            M.DstLoad->isNonTemporal() != Stores[0].DstLoad->isNonTemporal() ||
+            M.SrcLoad->isNonTemporal() != Stores[0].SrcLoad->isNonTemporal()) {
+          MatchFailed = true;
+          break;
+        }
+      }
+
+      DstDereferenceable &= M.DstLoad->isDereferenceable();
+      SrcDereferenceable &= M.SrcLoad->isDereferenceable();
+    }
+    if (MatchFailed)
+      continue;
+
+    StoreSDNode *FirstStore = Stores[0].Store;
+    LoadSDNode *FirstDstLoad = Stores[0].DstLoad;
+    LoadSDNode *FirstSrcLoad = Stores[0].SrcLoad;
+    unsigned FirstStoreAS = FirstStore->getAddressSpace();
+    unsigned IsFastSt = 0;
+    unsigned IsFastDstLd = 0;
+    unsigned IsFastSrcLd = 0;
+    if (!TLI.allowsMemoryAccess(Context, DL, JointVT,
+                                *FirstStore->getMemOperand(), &IsFastSt) ||
+        !IsFastSt ||
+        !TLI.allowsMemoryAccess(Context, DL, JointVT,
+                                *FirstDstLoad->getMemOperand(), &IsFastDstLd) ||
+        !IsFastDstLd ||
+        !TLI.allowsMemoryAccess(Context, DL, JointVT,
+                                *FirstSrcLoad->getMemOperand(), &IsFastSrcLd) ||
+        !IsFastSrcLd)
+      continue;
+
+    if (!checkByteAddStoreDependencies(Stores, RootNode))
+      continue;
+
+    SDLoc LoadDL(FirstDstLoad);
+    SDLoc StoreDL(FirstStore);
+    SDValue NewStoreChain = getMergedStoreChain(Stores, DAG);
+    DCI.AddToWorklist(NewStoreChain.getNode());
+
+    MachineMemOperand::Flags DstLdMMOFlags =
+        DstDereferenceable ? MachineMemOperand::MODereferenceable
+                           : MachineMemOperand::MONone;
+    if (FirstDstLoad->isNonTemporal())
+      DstLdMMOFlags |= MachineMemOperand::MONonTemporal;
+
+    MachineMemOperand::Flags SrcLdMMOFlags =
+        SrcDereferenceable ? MachineMemOperand::MODereferenceable
+                           : MachineMemOperand::MONone;
+    if (FirstSrcLoad->isNonTemporal())
+      SrcLdMMOFlags |= MachineMemOperand::MONonTemporal;
+
+    MachineMemOperand::Flags StMMOFlags = FirstStore->isNonTemporal()
+                                              ? MachineMemOperand::MONonTemporal
+                                              : MachineMemOperand::MONone;
+
+    SDValue DstWord =
+        DAG.getLoad(JointVT, LoadDL, FirstDstLoad->getChain(),
+                    FirstDstLoad->getBasePtr(), FirstDstLoad->getPointerInfo(),
+                    FirstDstLoad->getAlign(), DstLdMMOFlags);
+    SDValue SrcWord =
+        DAG.getLoad(JointVT, LoadDL, FirstSrcLoad->getChain(),
+                    FirstSrcLoad->getBasePtr(), FirstSrcLoad->getPointerInfo(),
+                    FirstSrcLoad->getAlign(), SrcLdMMOFlags);
+
+    SDValue Mask7F = DAG.getConstant(APInt(32, 0x7f7f7f7f), LoadDL, JointVT);
+    SDValue Mask80 = DAG.getConstant(APInt(32, 0x80808080), LoadDL, JointVT);
+    SDValue DstLo = DAG.getNode(ISD::AND, LoadDL, JointVT, DstWord, Mask7F);
+    SDValue SrcLo = DAG.getNode(ISD::AND, LoadDL, JointVT, SrcWord, Mask7F);
+    SDValue Lo = DAG.getNode(ISD::ADD, LoadDL, JointVT, DstLo, SrcLo);
+    SDValue HiXor = DAG.getNode(ISD::XOR, LoadDL, JointVT, DstWord, SrcWord);
+    SDValue Hi = DAG.getNode(ISD::AND, LoadDL, JointVT, HiXor, Mask80);
+    SDValue Result = DAG.getNode(ISD::XOR, StoreDL, JointVT, Lo, Hi);
+    SDValue NewStore = DAG.getStore(
+        NewStoreChain, StoreDL, Result, FirstStore->getBasePtr(),
+        MachinePointerInfo(FirstStoreAS), FirstStore->getAlign(), StMMOFlags);
+
+    SmallPtrSet<LoadSDNode *, 8> ReplacedLoads;
+    for (const ByteAddStoreMatch &M : Stores)
+      if (ReplacedLoads.insert(M.DstLoad).second)
+        DAG.ReplaceAllUsesOfValueWith(SDValue(M.DstLoad, 1),
+                                      SDValue(DstWord.getNode(), 1));
+    for (const ByteAddStoreMatch &M : Stores)
+      if (ReplacedLoads.insert(M.SrcLoad).second)
+        DAG.ReplaceAllUsesOfValueWith(SDValue(M.SrcLoad, 1),
+                                      SDValue(SrcWord.getNode(), 1));
+
+    for (const ByteAddStoreMatch &M : Stores) {
+      SDValue Val = M.Store->getOperand(1);
+      DCI.CombineTo(M.Store, NewStore);
+      if (Val->use_empty())
+        DCI.recursivelyDeleteUnusedNodes(Val.getNode());
+    }
+
+    DCI.AddToWorklist(DstWord.getNode());
+    DCI.AddToWorklist(SrcWord.getNode());
+    DCI.AddToWorklist(Result.getNode());
+    DCI.AddToWorklist(NewStore.getNode());
+
+    return SDValue(St, 0);
+  }
+
+  return SDValue();
+}
+
 /// PerformSTORECombine - Target-specific dag combine xforms for
 /// ISD::STORE.
 static SDValue PerformSTORECombine(SDNode *N,
@@ -16936,6 +17253,9 @@ static SDValue PerformSTORECombine(SDNode *N,
       return NewToken;
   }
 
+  if (SDValue ByteAddStore = PerformByteAddStoreCombine(St, DCI))
+    return ByteAddStore;
+
   if (!ISD::isNormalStore(St))
     return SDValue();
 
diff --git a/llvm/test/CodeGen/ARM/merge-byte-adds.ll b/llvm/test/CodeGen/ARM/merge-byte-adds.ll
new file mode 100644
index 0000000000000..8261bec0f8038
--- /dev/null
+++ b/llvm/test/CodeGen/ARM/merge-byte-adds.ll
@@ -0,0 +1,46 @@
+; RUN: llc -mtriple=thumbv8m.main-none-eabi -mcpu=cortex-m33 -mattr=-dsp,-mve,-fpregs < %s | FileCheck %s --check-prefix=FAST --implicit-check-not=ldrb --implicit-check-not=strb
+; RUN: llc -mtriple=thumbv8m.main-none-eabi -mcpu=cortex-m33 -mattr=-dsp,-mve,-fpregs,+strict-align < %s | FileCheck %s --check-prefix=STRICT --implicit-check-not=bic
+
+define void @add4(ptr noalias %dst, ptr noalias readonly %src) {
+; FAST-LABEL: add4:
+; FAST: ldr{{(\.w)?}} {{.*}}, [r1]
+; FAST: ldr{{(\.w)?}} {{.*}}, [r0]
+; FAST: bic{{(\.w)?}} {{.*}}, {{.*}}, #-2139062144
+; FAST: bic{{(\.w)?}} {{.*}}, {{.*}}, #-2139062144
+; FAST: eor
+; FAST: add
+; FAST: bic{{(\.w)?}} {{.*}}, {{.*}}, #2139062143
+; FAST: eor
+; FAST: str{{(\.w)?}} {{.*}}, [r0]
+; FAST: bx lr
+; STRICT-LABEL: add4:
+; STRICT: ldrb
+; STRICT: strb
+entry:
+  %d0 = load i8, ptr %dst, align 1
+  %s0 = load i8, ptr %src, align 1
+  %a0 = add i8 %s0, %d0
+  store i8 %a0, ptr %dst, align 1
+
+  %dst1 = getelementptr inbounds i8, ptr %dst, i32 1
+  %src1 = getelementptr inbounds i8, ptr %src, i32 1
+  %d1 = load i8, ptr %dst1, align 1
+  %s1 = load i8, ptr %src1, align 1
+  %a1 = add i8 %s1, %d1
+  store i8 %a1, ptr %dst1, align 1
+
+  %dst2 = getelementptr inbounds i8, ptr %dst, i32 2
+  %src2 = getelementptr inbounds i8, ptr %src, i32 2
+  %d2 = load i8, ptr %dst2, align 1
+  %s2 = load i8, ptr %src2, align 1
+  %a2 = add i8 %s2, %d2
+  store i8 %a2, ptr %dst2, align 1
+
+  %dst3 = getelementptr inbounds i8, ptr %dst, i32 3
+  %src3 = getelementptr inbounds i8, ptr %src, i32 3
+  %d3 = load i8, ptr %dst3, align 1
+  %s3 = load i8, ptr %src3, align 1
+  %a3 = add i8 %s3, %d3
+  store i8 %a3, ptr %dst3, align 1
+  ret void
+}



More information about the llvm-commits mailing list