[llvm] [CodeGen] Add initial multi-def rematerialization support (PR #197580)

Lucas Ramirez via llvm-commits llvm-commits at lists.llvm.org
Thu Jul 16 01:01:20 PDT 2026


https://github.com/lucas-rami updated https://github.com/llvm/llvm-project/pull/197580

>From fc27cbb7de0f4b2e5a043fcc856334e56f59db1a Mon Sep 17 00:00:00 2001
From: Lucas Ramirez <lucas.rami at proton.me>
Date: Mon, 20 Apr 2026 17:40:01 +0000
Subject: [PATCH 1/4] [CodeGen] Add initial multi-def rematerialization support

This significantly improves support for rematerializing registers
with more than one definition. In particular, this includes cases where
different lanes of a register are defined over multiple instructions.

There are still a few restrictions that can hopefully be relaxed in the
future.

- All defining instructions must be part of the same rematerialization
  region.
- No pure user of the register (i.e., an MI that doesn't also defined a
  part of the register) must read the register before its last
  definition.

These constraints ensure that the underlying DAG representation
maintained by the rematerializer is still valid, making this a
relatively incremental improvement.
---
 llvm/include/llvm/CodeGen/Rematerializer.h    | 107 +++--
 llvm/lib/CodeGen/Rematerializer.cpp           | 434 +++++++++++++-----
 llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp   |   2 +-
 llvm/lib/Target/AMDGPU/SIInstrInfo.cpp        |  67 ++-
 .../machine-scheduler-sink-trivial-remats.mir |  42 +-
 llvm/unittests/CodeGen/RematerializerTest.cpp | 254 +++++++++-
 6 files changed, 694 insertions(+), 212 deletions(-)

diff --git a/llvm/include/llvm/CodeGen/Rematerializer.h b/llvm/include/llvm/CodeGen/Rematerializer.h
index e1337cec8d9ed..496b38d4ee66f 100644
--- a/llvm/include/llvm/CodeGen/Rematerializer.h
+++ b/llvm/include/llvm/CodeGen/Rematerializer.h
@@ -29,11 +29,14 @@ namespace llvm {
 ///
 /// At the moment this supports rematerializing registers that meet all of the
 /// following constraints.
-/// 1. The register is virtual and has a single defining instruction.
-/// 2. The single defining instruction is deemed rematerializable by the TII and
-///    doesn't have any physical register use that is both non-constant and
+/// 1. The register is virtual.
+/// 2. The register is defined within a single region---potentially over
+///    multiple MIs---and isn't used by a MI that is not defining part of the
+///    register before its last defining MI.
+/// 3. All defining instructions are deemed rematerializable by the TII and
+///    don't have any physical register use that is both non-constant and
 ///    non-ignorable.
-/// 3. The register has at least one non-debug use that is inside or at a region
+/// 4. The register has at least one non-debug use that is inside or at a region
 ///    boundary (see below for what we consider to be a region).
 ///
 /// Rematerializable registers (represented by \ref Rematerializer::Reg) form a
@@ -89,47 +92,59 @@ namespace llvm {
 /// In its nomenclature, the rematerializer differentiates between "original
 /// registers" (registers that were present when it analyzed the function) and
 /// rematerializations of these original registers. Rematerializations have an
-/// "origin" which is the index of the original regiser they were rematerialized
-/// from (transitivity applies; a rematerialization and all of its own
-/// rematerializations have the same origin). Semantically, only original
-/// registers have rematerializations.
+/// "origin" which is the index of the original register they were
+/// rematerialized from (transitivity applies; a rematerialization and all of
+/// its own rematerializations have the same origin). Semantically, only
+/// original registers have rematerializations.
+///
+/// Dealing with sub-registers is complicated, we have to handle dead-defs,
+/// undef flags, and connected components
 class Rematerializer {
 public:
   /// Index type for rematerializable registers.
   using RegisterIdx = unsigned;
 
-  /// A rematerializable register defined by a single machine instruction.
+  /// A rematerializable register, potentially defined by multiple instructions.
   ///
   /// A rematerializable register has a set of dependencies, which correspond
-  /// to the unique read register operands of its defining instruction and which
-  /// can themselves be rematerializable. Operand indices corresponding to
-  /// unrematerializable dependencies are managed by and queried from the
-  /// rematerializer, whereas rematerializable ones are part of this struct and
-  /// identified through their register index.
+  /// to the unique read register operands of its defining instruction(s) and
+  /// which can themselves be rematerializable. Operands of defining
+  /// instructions corresponding to unrematerializable dependencies are managed
+  /// by and queried from the rematerializer, whereas rematerializable ones are
+  /// part of this struct and identified through their register index.
   ///
   /// A rematerializable register also has an arbitrary number of users in an
   /// arbitrary number of regions, potentially including its own defining
   /// region. When rematerializations lead to operand changes in users, a
   /// register may find itself without any user left, at which point the
-  /// rematerializer deletes it (setting its defining MI to nullptr).
+  /// rematerializer deletes it (emptying \ref Reg::Defs).
   struct Reg {
-    /// Single MI defining the rematerializable register.
-    MachineInstr *DefMI;
-    /// Defining region of \p DefMI.
+    /// All instructions that define the register, in program order.
+    SmallVector<MachineInstr *, 1> Defs;
+    /// Defining region of the register.
     unsigned DefRegion;
     /// The rematerializable register's lane bitmask.
     LaneBitmask Mask;
 
     using RegionUsers = SmallDenseSet<MachineInstr *, 4>;
-    /// Uses of the register, mapped by region.
+    /// Uses of the register, mapped by region. Users that also define a part of
+    /// the register are considered defs and not accounted for here.
     SmallDenseMap<unsigned, RegionUsers, 2> Uses;
+
     /// This register's rematerializable dependencies, one per unique
-    /// rematerializable register operand.
+    /// rematerializable register operand over all definitions.
     SmallVector<RegisterIdx, 2> Dependencies;
 
-    /// Returns the rematerializable register from its defining instruction.
+    MachineInstr *getFirstDef() const { return Defs.front(); }
+    MachineInstr *getLastDef() const { return Defs.back(); }
+
+    /// Returns the rematerializable register from one of its defining
+    /// instructions.
     Register getDefReg() const {
-      assert(DefMI && "defining instruction was deleted");
+      const MachineInstr *DefMI = getFirstDef();
+      assert(DefMI && "defining instruction(s) were deleted");
+      if (!DefMI->getOperand(0).isDef())
+        dbgs() << *DefMI;
       assert(DefMI->getOperand(0).isDef() && "not a register def");
       return DefMI->getOperand(0).getReg();
     }
@@ -144,18 +159,29 @@ class Rematerializer {
       return Uses.size() > 1 || Uses.begin()->first != DefRegion;
     }
 
+    /// Returns the index of \p DefMI in the register's definitions order.
+    /// Returns the number of definitions if \p DefMI is not a definition of the
+    /// register.
+    unsigned getDefIdx(MachineInstr *DefMI) const {
+      return std::distance(Defs.begin(), find(Defs, DefMI));
+    }
+
     /// Returns the first and last user of the register in region \p UseRegion.
     /// If the register has no user in the region, returns a pair of nullptr's.
     LLVM_ABI std::pair<MachineInstr *, MachineInstr *>
     getRegionUseBounds(unsigned UseRegion, const LiveIntervals &LIS) const;
 
-    bool isAlive() const { return DefMI; }
+    bool isAlive() const { return !Defs.empty(); }
 
   private:
     void addUser(MachineInstr *MI, unsigned Region);
     void addUsers(const RegionUsers &NewUsers, unsigned Region);
     void eraseUser(MachineInstr *MI, unsigned Region);
 
+    /// Erases user \p MI from region \p Region if it exists. Returns whether \p
+    /// MI was actually deleted.
+    bool tryEraseUser(MachineInstr *MI, unsigned Region);
+
     friend Rematerializer;
   };
 
@@ -375,12 +401,15 @@ class Rematerializer {
                    MachineBasicBlock::iterator InsertPos,
                    SmallVectorImpl<RegisterIdx> &&Dependencies);
 
-  /// Re-creates a previously deleted register \p RegIdx before \p InsertPos,
-  /// which must be in the register's original defining region. \p DefReg must
-  /// be the original virtual register that \p RegIdx used to define.
-  /// Dependencies are assumed to already exist in the MIR.
+  /// Re-creates each defining instruction of a previously deleted register \p
+  /// RegIdx before each position in \p Positions (one position per defining
+  /// instruction, in the same order). Positions must be in the same region as
+  /// the deleted register, and earlier than all uses of the register in the
+  /// region. \p DefReg must be the original virtual register that \p RegIdx
+  /// used to define. Rematerializable dependencies are assumed to already exist
+  /// in the MIR.
   LLVM_ABI void recreateReg(RegisterIdx RegIdx,
-                            MachineBasicBlock::iterator InsertPos,
+                            ArrayRef<MachineBasicBlock::iterator> Positions,
                             Register DefReg);
 
   /// Transfers all users of register \p FromRegIdx in region \p UseRegion to \p
@@ -424,8 +453,8 @@ class Rematerializer {
 
   LLVM_ABI Printable printDependencyDAG(RegisterIdx RootIdx) const;
   LLVM_ABI Printable printID(RegisterIdx RegIdx) const;
-  LLVM_ABI Printable printRematReg(RegisterIdx RegIdx,
-                                   bool SkipRegions = false) const;
+  LLVM_ABI Printable printRematReg(RegisterIdx RegIdx, bool SkipRegions = false,
+                                   unsigned DefIdx = 0) const;
   LLVM_ABI Printable printRegUsers(RegisterIdx RegIdx) const;
   LLVM_ABI Printable
   printUser(const MachineInstr *MI,
@@ -492,7 +521,7 @@ class Rematerializer {
   void postRematerialization(RegisterIdx ModelRegIdx, RegisterIdx RematRegIdx);
 
   /// Common pre-processing step before deleting a register \p DeleteRegIdx. The
-  /// register's defining instruction must still be alive.
+  /// register must still have alive definitions.
   void preDeletion(RegisterIdx DeleteRegIdx);
 
   /// Extends \p LI over \p Mask to be live at \p UdeIdx.
@@ -527,8 +556,7 @@ class Rematerializer {
 
   /// Determines whether \p MI is considered rematerializable. This further
   /// restricts constraints imposed by the TII on rematerializable instructions,
-  /// requiring for example that the defined register is virtual and only
-  /// defined once.
+  /// requiring for example that the defined register is virtual.
   bool isMIRematerializable(const MachineInstr &MI) const;
 
   /// Implementation of \ref Rematerializer::transferUser that doesn't update
@@ -570,14 +598,14 @@ class LLVM_ABI Rollbacker : public Rematerializer::Listener {
     RegisterIdx Idx;
     /// Original register.
     Register DefReg;
-    /// Original definition of the register. The underlying MI no longer exist
-    /// at rollback time, but may be referenced as re-creation position for
+    /// Original definitions of the register. The underlying MIs no longer exist
+    /// at rollback time, but may be referenced as re-creation positions for
     /// previously deleted registers.
-    MachineInstr *DefMI;
+    SmallVector<MachineInstr *, 1> Defs;
 
     LLVM_ABI DeadReg(RegisterIdx Idx, const Rematerializer &Remater)
         : Idx(Idx), DefReg(Remater.getReg(Idx).getDefReg()),
-          DefMI(Remater.getReg(Idx).DefMI) {}
+          Defs(Remater.getReg(Idx).Defs) {}
   };
 
   /// An insertion position in the MIR, either a MachineInstr* to insert before
@@ -587,8 +615,9 @@ class LLVM_ABI Rollbacker : public Rematerializer::Listener {
   /// Original registers that have been deleted, in order of deletion.
   SmallVector<DeadReg> DeadRegs;
   /// Re-creation positions for all original registers that have been deleted,
-  /// in register deletion order. A position is either a MachineInstr* that
-  /// existed in the MIR at the time the rollbacker was attached to the
+  /// one per defining instruction, in program order for any given register and
+  /// in register deletion order overall. A position is either a MachineInstr*
+  /// that existed in the MIR at the time the rollbacker was attached to the
   /// rematerializer, or a MachineBasicBlock*.
   SmallVector<InsertBeforePos> Positions;
   /// Maps all re-creation positions that exist in \ref Positions to the indices
diff --git a/llvm/lib/CodeGen/Rematerializer.cpp b/llvm/lib/CodeGen/Rematerializer.cpp
index 8db6b0db158e8..0f39db5e24d8e 100644
--- a/llvm/lib/CodeGen/Rematerializer.cpp
+++ b/llvm/lib/CodeGen/Rematerializer.cpp
@@ -18,13 +18,14 @@
 #include "llvm/ADT/SetVector.h"
 #include "llvm/ADT/SmallSet.h"
 #include "llvm/CodeGen/LiveIntervals.h"
+#include "llvm/CodeGen/LiveRangeEdit.h"
 #include "llvm/CodeGen/MachineBasicBlock.h"
 #include "llvm/CodeGen/MachineOperand.h"
 #include "llvm/CodeGen/MachineRegisterInfo.h"
 #include "llvm/CodeGen/Register.h"
 #include "llvm/CodeGen/TargetRegisterInfo.h"
+#include "llvm/MC/LaneBitmask.h"
 #include "llvm/Support/Debug.h"
-#include "llvm/Support/ErrorHandling.h"
 #include <optional>
 
 #define DEBUG_TYPE "rematerializer"
@@ -178,21 +179,50 @@ void Rematerializer::transferUserImpl(RegisterIdx FromRegIdx,
   LLVM_DEBUG(dbgs() << "User transfer from " << printID(FromRegIdx) << " to "
                     << printID(ToRegIdx) << ": " << printUser(&UserMI) << '\n');
 
-  UserMI.substituteRegister(getReg(FromRegIdx).getDefReg(),
-                            getReg(ToRegIdx).getDefReg(), 0, TRI);
+  Register FromReg = getReg(FromRegIdx).getDefReg();
+  UserMI.substituteRegister(FromReg, getReg(ToRegIdx).getDefReg(), 0, TRI);
 
-  // If the user is rematerializable, we must change its dependency to the
-  // new register.
-  if (RegisterIdx UserRegIdx = getDefRegIdx(UserMI); UserRegIdx != NoReg) {
-    // Look for the user's dependency that matches the register.
-    for (RegisterIdx &DepRegIdx : Regs[UserRegIdx].Dependencies) {
-      if (DepRegIdx == FromRegIdx) {
-        DepRegIdx = ToRegIdx;
+  RegisterIdx UserRegIdx = getDefRegIdx(UserMI);
+  if (UserRegIdx == NoReg)
+    return;
+
+  // When the user is rematerializable, we must reflect the change in its
+  // dependencies.
+  Reg &UserReg = Regs[UserRegIdx];
+  SmallVectorImpl<RegisterIdx> &UserDeps = Regs[UserRegIdx].Dependencies;
+  bool IsNewDep = true;
+  if (UserReg.Defs.size() > 1) {
+    // Other defining MIs might already be using the new register.
+    IsNewDep = find(UserDeps, ToRegIdx) == UserDeps.end();
+
+    auto MOIsFromReg = [FromReg](MachineOperand &MO) {
+      return MO.getReg() == FromReg;
+    };
+
+    // If any other defining instruction of the rematerializable user still uses
+    // the original register, we should not remove it from dependencies and may
+    // need to add a new dependency if it is the first time the new register is
+    // used by defining instructions.
+    for (MachineInstr *DefMI : UserReg.Defs) {
+      if (DefMI == &UserMI)
+        continue;
+      if (any_of(DefMI->all_uses(), MOIsFromReg)) {
+        if (IsNewDep)
+          UserDeps.push_back(ToRegIdx);
         return;
       }
     }
-    llvm_unreachable("broken dependency");
   }
+
+  // No other defining instruction has the original register as user. This
+  // either removes a dependency if the new register was previously used, or is
+  // a simple replacement if not.
+  unsigned *FindFromReg = find(UserDeps, FromRegIdx);
+  assert(FindFromReg != UserDeps.end() && "broken dependency");
+  if (IsNewDep)
+    *FindFromReg = ToRegIdx;
+  else
+    UserReg.Dependencies.erase(FindFromReg);
 }
 
 bool Rematerializer::isMOIdenticalAtUses(MachineOperand &MO,
@@ -234,7 +264,7 @@ RegisterIdx Rematerializer::findRematInRegion(RegisterIdx RegIdx,
     if (RematReg.DefRegion != Region || RematReg.Uses.empty())
       continue;
     SlotIndex RematRegSlot =
-        LIS.getInstructionIndex(*RematReg.DefMI).getRegSlot();
+        LIS.getInstructionIndex(*RematReg.getLastDef()).getRegSlot();
     if (RematRegSlot < Before &&
         (BestRegIdx == NoReg || RematRegSlot > BestSlot)) {
       BestSlot = RematRegSlot;
@@ -257,10 +287,17 @@ void Rematerializer::deleteReg(RegisterIdx RootIdx) {
     for (RegisterIdx DepRegIdx : DeleteReg.Dependencies) {
       // All dependencies lose a user (the deleted register).
       Reg &DepReg = Regs[DepRegIdx];
-      DepReg.eraseUser(DeleteReg.DefMI, DeleteReg.DefRegion);
-      if (DepReg.Uses.empty()) {
-        DeleteOrder.push_back(DepRegIdx);
-        DepDAG.push_back(DepRegIdx);
+      for (MachineInstr *DefMI : DeleteReg.Defs) {
+        if (DepReg.tryEraseUser(DefMI, DeleteReg.DefRegion) &&
+            DepReg.Uses.empty()) {
+          // The if condition will only be true at most once for any given
+          // register because, once the dependency no longer has any user,
+          // tryEraseUser will always produce false. We can therefore safely use
+          // vectors instead of sets for determining deletable registers.
+          DeleteOrder.push_back(DepRegIdx);
+          DepDAG.push_back(DepRegIdx);
+          break;
+        }
       }
     }
   } while (!DepDAG.empty());
@@ -269,10 +306,12 @@ void Rematerializer::deleteReg(RegisterIdx RootIdx) {
     preDeletion(RegIdx);
     Reg &DeleteReg = Regs[RegIdx];
     Register DefReg = DeleteReg.getDefReg();
-    LIS.RemoveMachineInstrFromMaps(*DeleteReg.DefMI);
-    DeleteReg.DefMI->eraseFromParent();
-    DeleteReg.DefMI = nullptr;
+    for (MachineInstr *DefMI : reverse(DeleteReg.Defs)) {
+      LIS.RemoveMachineInstrFromMaps(*DefMI);
+      DefMI->eraseFromParent();
+    }
     LIS.removeInterval(DefReg);
+    DeleteReg.Defs.clear();
   }
 
   SmallSet<RegisterIdx, 8> ShrinkRematRegs;
@@ -357,16 +396,27 @@ void Rematerializer::DeadDefDelegate::LRE_WillEraseInstruction(
   // All rematerializable dependencies must be notified.
   Reg &DeleteReg = Remater.Regs[RegIdx];
   for (RegisterIdx DepRegIdx : DeleteReg.Dependencies)
-    Remater.Regs[DepRegIdx].eraseUser(MI, DeleteReg.DefRegion);
-
-  assert(DeleteReg.isAlive() && "register must be alive");
+    Remater.Regs[DepRegIdx].tryEraseUser(MI, DeleteReg.DefRegion);
+
+  // The constraint that no other register reads any intermediate value of a
+  // register defined over multiple MI implies that the live range editor will
+  // either not touch or fully delete rematerializable registers i.e., if this
+  // is called for any defining instruction of a rematerializable register, this
+  // will be called for every definition of the register. Furthermore, def/use
+  // order between defining instructions ensures this will be called from last
+  // definition to first definition. When the last definition / first MI
+  // deletion happens, we want to reflect the deletion in our internal
+  // data-structures and notify any rematerializer listener.
+  if (!DeleteReg.isAlive())
+    return;
+  assert(DeleteReg.getLastDef() == MI && "last def should be deleted first");
   assert(DeleteReg.Uses.empty() && "register should no longer have uses");
 
-  // The live-range editor will delete the defining instruction from the MIR
-  // as well as the register's live-range, so we just need to nullify the def
-  // internally.
+  // The live-reange editor will delete all defining instructions from the MIR
+  // as well as the register's live-range, so we just need to clear out the defs
+  // vector.
   Remater.preDeletion(RegIdx);
-  DeleteReg.DefMI = nullptr;
+  DeleteReg.Defs.clear();
 }
 
 void Rematerializer::preDeletion(RegisterIdx DeleteRegIdx) {
@@ -379,8 +429,11 @@ void Rematerializer::preDeletion(RegisterIdx DeleteRegIdx) {
   // instruction to be the upper region boundary since we don't ever consider
   // them rematerializable.
   MachineBasicBlock::iterator &RegionBegin = Regions[DeleteReg.DefRegion].first;
-  if (RegionBegin == DeleteReg.DefMI)
+  for (MachineInstr *DefMI : DeleteReg.Defs) {
+    if (RegionBegin != DefMI)
+      break;
     ++RegionBegin;
+  }
 
   if (isOriginalRegister(DeleteRegIdx))
     return;
@@ -472,32 +525,76 @@ void Rematerializer::addRegIfRematerializable(
   assert(!SeenRegs[VirtRegIdx] && "register already seen");
   Register DefReg = Register::index2VirtReg(VirtRegIdx);
   SeenRegs.set(VirtRegIdx);
+  Reg RematReg;
 
-  MachineOperand *MO = MRI.getOneDef(DefReg);
-  if (!MO)
-    return;
-  MachineInstr &DefMI = *MO->getParent();
-  if (!isMIRematerializable(DefMI))
-    return;
-  auto DefRegion = MIRegion.find(&DefMI);
-  if (DefRegion == MIRegion.end())
+  // Check that the register's definitions can be rematerialized.
+  SmallPtrSet<MachineInstr *, 1> DefSet;
+  for (MachineOperand &MO : MRI.def_operands(DefReg)) {
+    MachineInstr &DefMI = *MO.getParent();
+    // If a single MI has multiple defs for the same register, we don't need to
+    // redo MI-based checks.
+    if (!DefSet.insert(&DefMI).second)
+      continue;
+
+    // The defining MI must be rematerializable and in the same region as all
+    // other defining MIs.
+    if (!isMIRematerializable(DefMI))
+      return;
+    auto DefRegion = MIRegion.find(&DefMI);
+    if (DefRegion == MIRegion.end())
+      return;
+    if (RematReg.Defs.empty())
+      RematReg.DefRegion = DefRegion->getSecond();
+    else if (RematReg.DefRegion != DefRegion->getSecond())
+      return;
+    RematReg.Defs.push_back(&DefMI);
+  }
+  if (RematReg.Defs.empty())
     return;
 
-  Reg RematReg;
-  RematReg.DefMI = &DefMI;
-  RematReg.DefRegion = DefRegion->second;
-  unsigned SubIdx = DefMI.getOperand(0).getSubReg();
-  RematReg.Mask = SubIdx ? TRI.getSubRegIndexLaneMask(SubIdx)
-                         : MRI.getMaxLaneMaskForVReg(DefReg);
+  // Order defining MIs by slot index.
+  sort(RematReg.Defs, [&](MachineInstr *LHS, MachineInstr *RHS) {
+    return LIS.getInstructionIndex(*LHS) < LIS.getInstructionIndex(*RHS);
+  });
+  // None of the non-first register defintions can be marked undef.
+  for (const MachineInstr *DefMI : drop_begin(RematReg.Defs)) {
+    for (const MachineOperand &DefMO : DefMI->all_defs()) {
+      if (DefMO.getReg() == DefReg && DefMO.isUndef())
+        return;
+    }
+  }
+
+  SlotIndex LastDefSlot = LIS.getInstructionIndex(*RematReg.getLastDef());
+
+  // Set the register's mask to all active lanes after the last def.
+  const LiveInterval &DefLI = LIS.getInterval(DefReg);
+  SlotIndex AfterLastDef = LastDefSlot.getRegSlot();
+  if (DefLI.hasSubRanges()) {
+    for (const LiveInterval::SubRange &SR : DefLI.subranges())
+      if (SR.liveAt(AfterLastDef))
+        RematReg.Mask |= SR.LaneMask;
+  } else {
+    RematReg.Mask = MRI.getMaxLaneMaskForVReg(DefReg);
+  }
 
   // Collect the candidate's direct users, both rematerializable and
-  // unrematerializable. MIs outside provided regions cannot be tracked so the
-  // registers they use are not safely rematerializable.
+  // unrematerializable.
+  const bool MoreThanOneDef = RematReg.Defs.size() > 1;
   for (MachineInstr &UseMI : MRI.use_nodbg_instructions(DefReg)) {
-    if (auto UseRegion = MIRegion.find(&UseMI); UseRegion != MIRegion.end())
-      RematReg.addUser(&UseMI, UseRegion->second);
-    else
+    // We are only interested in users that do not define part of the register.
+    if (DefSet.contains(&UseMI))
+      continue;
+    // MIs outside provided regions cannot be tracked so the registers they use
+    // are not safely rematerializable.
+    auto UseRegion = MIRegion.find(&UseMI);
+    if (UseRegion == MIRegion.end())
       return;
+    // Disallow reads before the last def.
+    if (MoreThanOneDef && RematReg.DefRegion == UseRegion->second &&
+        LastDefSlot > LIS.getInstructionIndex(UseMI))
+      return;
+
+    RematReg.addUser(&UseMI, UseRegion->second);
   }
   if (RematReg.Uses.empty())
     return;
@@ -507,22 +604,39 @@ void Rematerializer::addRegIfRematerializable(
   // it once.
   SmallSetVector<RegisterIdx, 2> RematDeps;
   SmallMapVector<Register, LaneBitmask, 2> UnrematDeps;
-  for (const MachineOperand &MO : DefMI.all_uses()) {
-    Register DepReg = getRegDependency(MO);
-    if (!DepReg)
-      continue;
-    unsigned DepRegIdx = DepReg.virtRegIndex();
-    if (!SeenRegs[DepRegIdx])
-      addRegIfRematerializable(DepRegIdx, MIRegion, SeenRegs);
-    if (auto DepIt = RegToIdx.find(DepReg); DepIt != RegToIdx.end()) {
-      RematDeps.insert(DepIt->second);
-    } else {
-      LaneBitmask &CurrentMask =
-          UnrematDeps.try_emplace(DepReg, LaneBitmask::getNone()).first->second;
-      LaneBitmask Mask = MO.getSubReg()
-                             ? TRI.getSubRegIndexLaneMask(MO.getSubReg())
-                             : MRI.getMaxLaneMaskForVReg(DepReg);
-      CurrentMask |= Mask;
+  for (const MachineInstr *DefMI : RematReg.Defs) {
+    for (const MachineOperand &MO : DefMI->all_uses()) {
+      Register DepReg = getRegDependency(MO);
+      if (!DepReg || DepReg == DefReg)
+        continue;
+      unsigned DepRegIdx = DepReg.virtRegIndex();
+      if (!SeenRegs[DepRegIdx])
+        addRegIfRematerializable(DepRegIdx, MIRegion, SeenRegs);
+      if (auto DepIt = RegToIdx.find(DepReg); DepIt != RegToIdx.end()) {
+        RematDeps.insert(DepIt->second);
+      } else {
+        LaneBitmask &CurrentMask =
+            UnrematDeps.try_emplace(DepReg, LaneBitmask::getNone())
+                .first->second;
+        LaneBitmask Mask = MO.getSubReg()
+                               ? TRI.getSubRegIndexLaneMask(MO.getSubReg())
+                               : MRI.getMaxLaneMaskForVReg(DepReg);
+        CurrentMask |= Mask;
+      }
+    }
+  }
+
+  if (MoreThanOneDef) {
+    // A def of an unrematerializable dependency between the defs of the
+    // register under consideration makes the latter unrematerializable.
+    SlotIndex FirstDefSlot = LIS.getInstructionIndex(*RematReg.getFirstDef());
+    for (const auto &[UnrematDepReg, _] : UnrematDeps) {
+      for (MachineOperand &UnrematMODef : MRI.def_operands(UnrematDepReg)) {
+        MachineInstr &UnrematDefMI = *UnrematMODef.getParent();
+        SlotIndex UnrematDefSlot = LIS.getInstructionIndex(UnrematDefMI);
+        if (UnrematDefSlot > FirstDefSlot || UnrematDefSlot < LastDefSlot)
+          return;
+      }
     }
   }
 
@@ -538,7 +652,6 @@ bool Rematerializer::isMIRematerializable(const MachineInstr &MI) const {
     return false;
 
   assert(MI.getOperand(0).getReg().isVirtual() && "should be virtual");
-  assert(MRI.hasOneDef(MI.getOperand(0).getReg()) && "should have single def");
 
   for (const MachineOperand &MO : MI.all_uses()) {
     // We can't remat physreg uses, unless it is a constant or an ignorable
@@ -555,7 +668,7 @@ bool Rematerializer::isMIRematerializable(const MachineInstr &MI) const {
 
 RegisterIdx Rematerializer::getDefRegIdx(const MachineInstr &MI) const {
   if (!MI.getNumOperands() || !MI.getOperand(0).isReg() ||
-      MI.getOperand(0).readsReg())
+      !MI.getOperand(0).isDef())
     return NoReg;
   Register Reg = MI.getOperand(0).getReg();
   auto UserRegIt = RegToIdx.find(Reg);
@@ -574,6 +687,7 @@ Rematerializer::rematerializeReg(RegisterIdx RegIdx, unsigned UseRegion,
   Reg &FromReg = Regs[RegIdx];
   NewReg.Mask = FromReg.Mask;
   NewReg.DefRegion = UseRegion;
+  NewReg.Defs.reserve(FromReg.Defs.size());
   NewReg.Dependencies = std::move(Dependencies);
 
   // Track rematerialization link between registers. Origins are always
@@ -586,9 +700,10 @@ Rematerializer::rematerializeReg(RegisterIdx RegIdx, unsigned UseRegion,
   // Use the TII to rematerialize the defining instruction with a new defined
   // register.
   Register NewDefReg = MRI.cloneVirtualRegister(FromReg.getDefReg());
-  TII.reMaterialize(*RegionMBB[UseRegion], InsertPos, NewDefReg, 0,
-                    *FromReg.DefMI);
-  NewReg.DefMI = &*std::prev(InsertPos);
+  for (const MachineInstr *DefMI : FromReg.Defs) {
+    TII.reMaterialize(*RegionMBB[UseRegion], InsertPos, NewDefReg, 0, *DefMI);
+    NewReg.Defs.push_back(&*std::prev(InsertPos));
+  }
   RegToIdx.insert({NewDefReg, NewRegIdx});
   postRematerialization(RegIdx, NewRegIdx);
 
@@ -598,13 +713,12 @@ Rematerializer::rematerializeReg(RegisterIdx RegIdx, unsigned UseRegion,
   return NewRegIdx;
 }
 
-void Rematerializer::recreateReg(RegisterIdx RegIdx,
-                                 MachineBasicBlock::iterator InsertPos,
-                                 Register DefReg) {
+void Rematerializer::recreateReg(
+    RegisterIdx RegIdx, ArrayRef<MachineBasicBlock::iterator> Positions,
+    Register DefReg) {
   assert(RegToIdx.contains(DefReg) && "unknown defined register");
   assert(RegToIdx.at(DefReg) == RegIdx && "incorrect defined register");
   assert(!getReg(RegIdx).isAlive() && "register is still alive");
-
   Reg &OriginReg = Regs[RegIdx];
 
   // Re-establish the link between origin and rematerialization if necessary.
@@ -623,11 +737,13 @@ void Rematerializer::recreateReg(RegisterIdx RegIdx,
     assert(getReg(getOriginOf(RegIdx)).isAlive() && "expected alive origin");
     ModelRegIdx = getOriginOf(RegIdx);
   }
-  const MachineInstr &ModelDefMI = *getReg(ModelRegIdx).DefMI;
+  const Reg &ModelReg = getReg(ModelRegIdx);
 
-  TII.reMaterialize(*RegionMBB[OriginReg.DefRegion], InsertPos, DefReg, 0,
-                    ModelDefMI);
-  OriginReg.DefMI = &*std::prev(InsertPos);
+  for (auto [DefMI, InsertPos] : zip_equal(ModelReg.Defs, Positions)) {
+    TII.reMaterialize(*RegionMBB[OriginReg.DefRegion], InsertPos, DefReg, 0,
+                      *DefMI);
+    OriginReg.Defs.push_back(&*std::prev(InsertPos));
+  }
   postRematerialization(ModelRegIdx, RegIdx);
   LLVM_DEBUG(dbgs() << "** Recreated " << printID(RegIdx) << " as "
                     << printRematReg(RegIdx) << '\n');
@@ -637,15 +753,22 @@ void Rematerializer::postRematerialization(RegisterIdx ModelRegIdx,
                                            RegisterIdx RematRegIdx) {
   Reg &ModelReg = Regs[ModelRegIdx], &RematReg = Regs[RematRegIdx];
 
+  SlotIndex UseIdx;
+  for (MachineInstr *DefMI : RematReg.Defs)
+    UseIdx = LIS.InsertMachineInstrInMaps(*DefMI);
+  UseIdx = UseIdx.getRegSlot();
+
   // The rematerialization has no user at this point so its interval will
   // initially be empty.
-  SlotIndex UseIdx = LIS.InsertMachineInstrInMaps(*RematReg.DefMI).getRegSlot();
   LIS.createAndComputeVirtRegInterval(RematReg.getDefReg());
 
   // The start of the new register's region may have changed.
-  MachineBasicBlock::iterator &RegionBegin = Regions[RematReg.DefRegion].first;
-  if (RegionBegin == std::next(MachineBasicBlock::iterator(RematReg.DefMI)))
-    RegionBegin = RematReg.DefMI;
+  MachineInstr &FirstDefMI = *RematReg.getFirstDef();
+  auto &[RegionBegin, RegionEnd] = Regions[RematReg.DefRegion];
+  if (RegionBegin == RegionEnd ||
+      (!RegionBegin->isDebugInstr() && LIS.getInstructionIndex(*RegionBegin) >
+                                           LIS.getInstructionIndex(FirstDefMI)))
+    RegionBegin = FirstDefMI.getIterator();
 
   // Replace dependencies as needed in the rematerialized MI. All dependencies
   // of the latter gain a new user.
@@ -653,15 +776,25 @@ void Rematerializer::postRematerialization(RegisterIdx ModelRegIdx,
   for (const auto &[OldDepRegIdx, NewDepRegIdx] : ZipedDeps) {
     LLVM_DEBUG(dbgs() << "  Dependency: " << printID(OldDepRegIdx) << " -> "
                       << printID(NewDepRegIdx) << '\n');
-
-    Reg &NewDepReg = Regs[NewDepRegIdx];
-    if (OldDepRegIdx != NewDepRegIdx) {
-      Reg &OldDepReg = Regs[OldDepRegIdx];
-      RematReg.DefMI->substituteRegister(OldDepReg.getDefReg(),
-                                         NewDepReg.getDefReg(), 0, TRI);
+    Register OldReg = getReg(OldDepRegIdx).getDefReg();
+    Register NewReg = getReg(NewDepRegIdx).getDefReg();
+
+    SmallVector<MachineInstr *, 2> DefsUsingNewDep;
+    for (MachineInstr *DefMI : RematReg.Defs) {
+      bool NewDefHasReg = false;
+      for (MachineOperand &MO : DefMI->operands()) {
+        if (!MO.isReg() || MO.getReg() != OldReg)
+          continue;
+        NewDefHasReg = true;
+        DefsUsingNewDep.push_back(DefMI);
+        if (OldDepRegIdx != NewDepRegIdx)
+          MO.substVirtReg(NewReg, 0, TRI);
+      }
+      if (NewDefHasReg)
+        Regs[NewDepRegIdx].addUser(DefMI, RematReg.DefRegion);
     }
-    NewDepReg.addUser(RematReg.DefMI, RematReg.DefRegion);
-    extendToNewUsers(NewDepRegIdx, RematReg.DefMI);
+    assert(!DefsUsingNewDep.empty() && "no user of dependency");
+    extendToNewUsers(NewDepRegIdx, DefsUsingNewDep);
   }
 
   // Unrematerializable dependencies always gain a new user after a
@@ -716,12 +849,18 @@ void Rematerializer::extendToNewUsers(RegisterIdx RegIdx,
     extendInterval(LI, RegMask, UseIdx);
   }
 
+  // Rematerializable registers are never read by instructions not defining them
+  // until after their last def, so adding a user to them ensures their last
+  // definition is alive. All potential other definitions are read by the last
+  // definition and are therefore already alive by construction.
   LLVM_DEBUG({
-    if (ExtendReg.DefMI->getOperand(0).isDead())
+    if (ExtendReg.getLastDef()->getOperand(0).isDead())
       dbgs() << "Clearing dead flag for "
-             << printRematReg(RegIdx, /*SkipRegions=*/false) << '\n';
+             << printRematReg(RegIdx, /*SkipRegions=*/false,
+                              /*DefIdx=*/ExtendReg.Defs.size() - 1)
+             << '\n';
   });
-  ExtendReg.DefMI->getOperand(0).setIsDead(false);
+  ExtendReg.getLastDef()->getOperand(0).setIsDead(false);
 }
 
 void Rematerializer::extendInterval(LiveInterval &LI, LaneBitmask Mask,
@@ -841,6 +980,15 @@ void Rematerializer::Reg::eraseUser(MachineInstr *MI, unsigned Region) {
     RUsers.erase(MI);
 }
 
+bool Rematerializer::Reg::tryEraseUser(MachineInstr *MI, unsigned Region) {
+  auto RegionUsers = Uses.find(Region);
+  if (RegionUsers == Uses.end() || !RegionUsers->getSecond().erase(MI))
+    return false;
+  if (RegionUsers->getSecond().empty())
+    Uses.erase(Region);
+  return true;
+}
+
 Printable Rematerializer::printDependencyDAG(RegisterIdx RootIdx) const {
   return Printable([&, RootIdx](raw_ostream &OS) {
     DenseMap<RegisterIdx, unsigned> RegDepths;
@@ -873,19 +1021,17 @@ Printable Rematerializer::printID(RegisterIdx RegIdx) const {
   return Printable([&, RegIdx](raw_ostream &OS) {
     const Reg &PrintReg = getReg(RegIdx);
     OS << '(' << RegIdx << '/';
-    if (!PrintReg.isAlive()) {
+    if (!PrintReg.isAlive())
       OS << "<dead>";
-    } else {
-      OS << printReg(PrintReg.getDefReg(), &TRI,
-                     PrintReg.DefMI->getOperand(0).getSubReg(), &MRI);
-    }
+    else
+      OS << printReg(PrintReg.getDefReg(), &TRI, 0, &MRI);
     OS << ")[" << PrintReg.DefRegion << "]";
   });
 }
 
-Printable Rematerializer::printRematReg(RegisterIdx RegIdx,
-                                        bool SkipRegions) const {
-  return Printable([&, RegIdx, SkipRegions](raw_ostream &OS) {
+Printable Rematerializer::printRematReg(RegisterIdx RegIdx, bool SkipRegions,
+                                        unsigned DefIdx) const {
+  return Printable([&, RegIdx, SkipRegions, DefIdx](raw_ostream &OS) {
     const Reg &PrintReg = getReg(RegIdx);
     OS << printID(RegIdx);
     if (!SkipRegions) {
@@ -929,10 +1075,13 @@ Printable Rematerializer::printRematReg(RegisterIdx RegIdx,
       OS << "] ";
     }
     if (PrintReg.isAlive()) {
-      PrintReg.DefMI->print(OS, /*IsStandalone=*/true, /*SkipOpers=*/false,
-                            /*SkipDebugLoc=*/false, /*AddNewLine=*/false);
+      assert(DefIdx < PrintReg.Defs.size() && "out-of-bound def");
+      MachineInstr &PrintDef = *PrintReg.Defs[DefIdx];
+      OS << "(def. " << DefIdx + 1 << " / " << PrintReg.Defs.size() << ") ";
+      PrintDef.print(OS, /*IsStandalone=*/true, /*SkipOpers=*/false,
+                     /*SkipDebugLoc=*/false, /*AddNewLine=*/false);
       OS << " @ ";
-      LIS.getInstructionIndex(*PrintReg.DefMI).print(OS);
+      LIS.getInstructionIndex(PrintDef).print(OS);
     }
   });
 }
@@ -981,28 +1130,55 @@ void Rollbacker::rematerializerNoteRegWillBeDeleted(
   if (RollingBack)
     return;
 
-  // Find a valid re-creation position after the register's definition.
-  MachineInstr *DefMI = Remater.getReg(RegIdx).DefMI;
-  MachineBasicBlock *ParentMBB = DefMI->getParent();
-  MachineBasicBlock::iterator ValidPos = std::next(DefMI->getIterator());
-  while (ValidPos != ParentMBB->end() && isRollbackableMI(*ValidPos, Remater))
-    ValidPos = std::next(ValidPos);
+  const Rematerializer::Reg &Reg = Remater.getReg(RegIdx);
+  MachineBasicBlock *ParentMBB = Reg.getFirstDef()->getParent();
+  MachineBasicBlock::iterator LastValidPos;
+
+  auto GetNextValidPosAfterDef =
+      [&](unsigned DefIdx) -> MachineBasicBlock::iterator {
+    const MachineInstr *NextDef =
+        DefIdx + 1 < Reg.Defs.size() ? Reg.Defs[DefIdx + 1] : nullptr;
+    MachineBasicBlock::iterator ValidPos =
+        std::next(Reg.Defs[DefIdx]->getIterator());
+
+    while (ValidPos != ParentMBB->end()) {
+      // When there are no valid insert positions between the current and next
+      // definition of the register about to be deleted, the first valid insert
+      // position for the current definition is the same as for the next
+      // definition.
+      const MachineInstr &CandMI = *ValidPos;
+      if (NextDef && &CandMI == NextDef)
+        return LastValidPos;
+      if (!isRollbackableMI(CandMI, Remater))
+        break;
+
+      // Move to the next candidate position.
+      ValidPos = std::next(ValidPos);
+    }
+
+    LastValidPos = ValidPos;
+    return ValidPos;
+  };
 
   if (Remater.isRematerializedRegister(RegIdx)) {
     // Rematerializations will not be re-created. Previously deleted registers
-    // that reference this register's defining instruction as their re-creation
+    // that reference this register's defining instructions as their re-creation
     // position should instead be re-created at a valid position after the
-    // deleted MI.
-    invalidatePosition(DefMI, ValidPos);
+    // deleted MIs.
+    for (unsigned I = Reg.Defs.size(); I > 0; --I)
+      invalidatePosition(Reg.Defs[I - 1], GetNextValidPosAfterDef(I - 1));
     return;
   }
 
-  // Original registers can be re-created. Add a re-creation position for the
+  // Original registers can be re-created. Add a re-creation position for each
   // definition of the rematerializable register.
   DeadRegs.push_back(DeadReg(RegIdx, Remater));
-  const InsertBeforePos InsertPos = makePos(ValidPos, ParentMBB);
-  PosToIdx[InsertPos].insert(Positions.size());
-  Positions.push_back(InsertPos);
+  for (unsigned I = Reg.Defs.size(); I > 0; --I) {
+    const InsertBeforePos InsertPos =
+        makePos(GetNextValidPosAfterDef(I - 1), ParentMBB);
+    PosToIdx[InsertPos].insert(Positions.size());
+    Positions.push_back(InsertPos);
+  }
 }
 
 void Rollbacker::rematerializerNoteMIWillBeDeleted(
@@ -1037,27 +1213,31 @@ void Rollbacker::rollback(Rematerializer &Remater) {
       // It is possible the register was permanently deleted as a consequence of
       // dead-def elimination.
       Rematerializations.erase(Reg.Idx);
-      --PositionIndex;
+      PositionIndex -= Reg.Defs.size();
       continue;
     }
-
     assert(!Remater.getReg(Reg.Idx).isAlive() && "register should be dead");
 
-    // Determine re-creation position for the register's definition.
-    MachineBasicBlock::iterator InsertPosition;
-    InsertBeforePos Pos = Positions[--PositionIndex];
-    if (auto *MBB = dyn_cast<MachineBasicBlock *>(Pos)) {
-      InsertPosition = MBB->end();
-    } else {
-      auto *MI = cast<MachineInstr *>(Pos);
-      InsertPosition = Replacements.lookup_or(MI, MI)->getIterator();
+    // Determine re-creation positions for all the deleted register's defs.
+    SmallVector<MachineBasicBlock::iterator, 1> InsertPositions;
+    for (unsigned I = 0, E = Reg.Defs.size(); I < E; ++I) {
+      InsertBeforePos Pos = Positions[--PositionIndex];
+      if (auto *MBB = dyn_cast<MachineBasicBlock *>(Pos)) {
+        InsertPositions.push_back(MBB->end());
+      } else {
+        auto *MI = cast<MachineInstr *>(Pos);
+        MachineInstr *InsertBeforeMI = Replacements.lookup_or(MI, MI);
+        InsertPositions.push_back(InsertBeforeMI->getIterator());
+      }
     }
 
-    Remater.recreateReg(Reg.Idx, InsertPosition, Reg.DefReg);
+    Remater.recreateReg(Reg.Idx, InsertPositions, Reg.DefReg);
 
     const Rematerializer::Reg &RecreateReg = Remater.getReg(Reg.Idx);
-    if (!Replacements.insert({Reg.DefMI, RecreateReg.DefMI}).second)
-      llvm_unreachable("duplicate deleted MI");
+    for (const auto [OldDef, NewDef] : zip_equal(Reg.Defs, RecreateReg.Defs)) {
+      assert(!Replacements.contains(OldDef) && "duplicate deleted MI");
+      Replacements[OldDef] = NewDef;
+    }
   }
 
   // Rollback rematerializations.
diff --git a/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp b/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
index c86b16573ca2b..59fdd8b8bea23 100644
--- a/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
+++ b/llvm/lib/Target/AMDGPU/GCNSchedStrategy.cpp
@@ -1562,7 +1562,7 @@ bool PreRARematStage::initGCNSchedStage() {
       continue;
     SlotIndex UseIdx = DAG.LIS->getInstructionIndex(*UseMI).getRegSlot(true);
     SlotIndex RefIdx =
-        DAG.LIS->getInstructionIndex(*CandReg.DefMI).getRegSlot(true);
+        DAG.LIS->getInstructionIndex(*CandReg.getLastDef()).getRegSlot(true);
     if (llvm::any_of(CandReg.Dependencies, [&](RegisterIdx DepRegIdx) {
           const Rematerializer::Reg &DepReg = Remater.getReg(DepRegIdx);
           Register DepDefReg = DepReg.getDefReg();
diff --git a/llvm/lib/Target/AMDGPU/SIInstrInfo.cpp b/llvm/lib/Target/AMDGPU/SIInstrInfo.cpp
index 34cf595e8d39b..1ca33ce7def51 100644
--- a/llvm/lib/Target/AMDGPU/SIInstrInfo.cpp
+++ b/llvm/lib/Target/AMDGPU/SIInstrInfo.cpp
@@ -181,7 +181,72 @@ bool SIInstrInfo::isReMaterializableImpl(
       return true;
   }
 
-  return TargetInstrInfo::isReMaterializableImpl(MI);
+  // Everything below copied from TargetInstrInfo::isReMaterializableImpl. The
+  // only difference is that we allow operations that perform read-modify-write
+  // on sub-registers.
+
+  // Remat clients assume operand 0 is the defined register.
+  if (!MI.getNumOperands() || !MI.getOperand(0).isReg())
+    return false;
+  Register DefReg = MI.getOperand(0).getReg();
+
+  const MachineFunction &MF = *MI.getMF();
+
+  // A load from a fixed stack slot can be rematerialized. This may be
+  // redundant with subsequent checks, but it's target-independent,
+  // simple, and a common case.
+  int FrameIdx = 0;
+  if (isLoadFromStackSlot(MI, FrameIdx) &&
+      MF.getFrameInfo().isImmutableObjectIndex(FrameIdx))
+    return true;
+
+  // Avoid instructions obviously unsafe for remat.
+  if (MI.isNotDuplicable() || MI.mayStore() || MI.mayRaiseFPException() ||
+      MI.hasUnmodeledSideEffects())
+    return false;
+
+  // Don't remat inline asm. We have no idea how expensive it is
+  // even if it's side effect free.
+  if (MI.isInlineAsm())
+    return false;
+
+  // Avoid instructions which load from potentially varying memory.
+  if (MI.mayLoad() && !MI.isDereferenceableInvariantLoad())
+    return false;
+
+  const MachineRegisterInfo &MRI = MF.getRegInfo();
+
+  // If any of the registers accessed are non-constant, conservatively assume
+  // the instruction is not rematerializable.
+  for (const MachineOperand &MO : MI.operands()) {
+    if (!MO.isReg())
+      continue;
+    Register Reg = MO.getReg();
+    if (Reg == 0)
+      continue;
+
+    // Check for a well-behaved physical register.
+    if (Reg.isPhysical()) {
+      if (MO.isUse()) {
+        // If the physreg has no defs anywhere, it's just an ambient register
+        // and we can freely move its uses. Alternatively, if it's allocatable,
+        // it could get allocated to something with a def during allocation.
+        if (!MRI.isConstantPhysReg(Reg))
+          return false;
+      } else {
+        // A physreg def. We can't remat it.
+        return false;
+      }
+      continue;
+    }
+
+    // Only allow one virtual-register def.  There may be multiple defs of the
+    // same virtual register, though.
+    if (MO.isDef() && Reg != DefReg)
+      return false;
+  }
+
+  return true;
 }
 
 // Returns true if the result of a VALU instruction depends on exec.
diff --git a/llvm/test/CodeGen/AMDGPU/machine-scheduler-sink-trivial-remats.mir b/llvm/test/CodeGen/AMDGPU/machine-scheduler-sink-trivial-remats.mir
index f9c36192a16a2..c1e1c2f6d5c31 100644
--- a/llvm/test/CodeGen/AMDGPU/machine-scheduler-sink-trivial-remats.mir
+++ b/llvm/test/CodeGen/AMDGPU/machine-scheduler-sink-trivial-remats.mir
@@ -12299,26 +12299,26 @@ body:             |
   ; GFX908-NEXT:   [[DEF28:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
   ; GFX908-NEXT:   [[DEF29:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
   ; GFX908-NEXT:   [[DEF30:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
-  ; GFX908-NEXT:   [[DEF31:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_22:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF24]], implicit $exec, implicit $mode
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_23:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF25]], implicit $exec, implicit $mode
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_24:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF26]], implicit $exec, implicit $mode
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_25:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF27]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_22:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF]].sub0, implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_23:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF24]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_24:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF25]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_25:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF26]], implicit $exec, implicit $mode
   ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_26:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF]].sub2, implicit $exec, implicit $mode
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_27:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF28]], implicit $exec, implicit $mode
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_28:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF29]], implicit $exec, implicit $mode
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_29:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF30]], implicit $exec, implicit $mode
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_30:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF31]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_27:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF27]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_28:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF28]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_29:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF29]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_30:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF30]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[DEF31:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
   ; GFX908-NEXT:   S_BRANCH %bb.1
   ; GFX908-NEXT: {{  $}}
   ; GFX908-NEXT: bb.1:
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_1]], implicit [[V_CVT_I32_F32_e32_6]]
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_31:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF]].sub0, implicit $exec, implicit $mode
-  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_31]], implicit [[V_CVT_I32_F32_e32_26]], implicit [[DEF]].sub0, implicit [[DEF]].sub2
-  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_22]], implicit [[V_CVT_I32_F32_e32_27]], implicit [[DEF24]], implicit [[DEF28]]
-  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_23]], implicit [[V_CVT_I32_F32_e32_28]], implicit [[DEF25]], implicit [[DEF29]]
-  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_24]], implicit [[V_CVT_I32_F32_e32_29]], implicit [[DEF26]], implicit [[DEF30]]
-  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_25]], implicit [[V_CVT_I32_F32_e32_30]], implicit [[DEF27]], implicit [[DEF31]]
+  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_22]], implicit [[V_CVT_I32_F32_e32_26]], implicit [[DEF]].sub0, implicit [[DEF]].sub2
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_31:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF31]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_31]], implicit [[V_CVT_I32_F32_e32_27]], implicit [[DEF31]], implicit [[DEF27]]
+  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_23]], implicit [[V_CVT_I32_F32_e32_28]], implicit [[DEF24]], implicit [[DEF28]]
+  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_24]], implicit [[V_CVT_I32_F32_e32_29]], implicit [[DEF25]], implicit [[DEF29]]
+  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_25]], implicit [[V_CVT_I32_F32_e32_30]], implicit [[DEF26]], implicit [[DEF30]]
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_2]], implicit [[V_CVT_I32_F32_e32_7]]
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_3]], implicit [[V_CVT_I32_F32_e32_8]]
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_4]], implicit [[V_CVT_I32_F32_e32_9]]
@@ -12911,7 +12911,7 @@ body:             |
   ; GFX908-NEXT:   [[DEF30:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
   ; GFX908-NEXT:   [[DEF31:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
   ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_23:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF25]], implicit $exec, implicit $mode
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_24:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF29]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_24:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF]].sub2, implicit $exec, implicit $mode
   ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_25:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF30]], implicit $exec, implicit $mode
   ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_26:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF31]], implicit $exec, implicit $mode
   ; GFX908-NEXT:   undef [[V_CVT_I32_F32_e32_1:%[0-9]+]].sub0:vreg_64 = nofpexcept V_CVT_I32_F32_e32 [[DEF]].sub0, implicit $exec, implicit $mode
@@ -12920,13 +12920,13 @@ body:             |
   ; GFX908-NEXT: bb.1:
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_2]], implicit [[V_CVT_I32_F32_e32_7]]
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_3]], implicit [[V_CVT_I32_F32_e32_8]]
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_27:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF]].sub2, implicit $exec, implicit $mode
-  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_1]].sub0, implicit [[V_CVT_I32_F32_e32_27]], implicit [[DEF]].sub0, implicit [[DEF]].sub2
-  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_28:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF28]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_1]].sub0, implicit [[V_CVT_I32_F32_e32_24]], implicit [[DEF]].sub0, implicit [[DEF]].sub2
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_27:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF28]], implicit $exec, implicit $mode
+  ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_28:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF29]], implicit $exec, implicit $mode
   ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_29:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF26]], implicit $exec, implicit $mode
   ; GFX908-NEXT:   [[V_CVT_I32_F32_e32_30:%[0-9]+]]:vgpr_32 = nofpexcept V_CVT_I32_F32_e32 [[DEF27]], implicit $exec, implicit $mode
-  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_23]], implicit [[V_CVT_I32_F32_e32_28]], implicit [[DEF1]], implicit [[DEF28]]
-  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_23]], implicit [[V_CVT_I32_F32_e32_24]], implicit [[DEF25]], implicit [[DEF29]]
+  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_23]], implicit [[V_CVT_I32_F32_e32_27]], implicit [[DEF1]], implicit [[DEF28]]
+  ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_23]], implicit [[V_CVT_I32_F32_e32_28]], implicit [[DEF25]], implicit [[DEF29]]
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_29]], implicit [[V_CVT_I32_F32_e32_25]], implicit [[DEF26]], implicit [[DEF30]]
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_30]], implicit [[V_CVT_I32_F32_e32_26]], implicit [[DEF27]], implicit [[DEF31]]
   ; GFX908-NEXT:   S_NOP 0, implicit [[V_CVT_I32_F32_e32_4]], implicit [[V_CVT_I32_F32_e32_9]]
diff --git a/llvm/unittests/CodeGen/RematerializerTest.cpp b/llvm/unittests/CodeGen/RematerializerTest.cpp
index 85e13ea19d578..8605a856bf759 100644
--- a/llvm/unittests/CodeGen/RematerializerTest.cpp
+++ b/llvm/unittests/CodeGen/RematerializerTest.cpp
@@ -191,6 +191,11 @@ body:             |
 #define EXPECT_NUM_USERS(RegIdx, N)                                            \
   EXPECT_EQ(RW.getNumUsers(RegIdx), static_cast<unsigned>(N))
 
+/// Expects that register RegIdx in the rematerializer has a total of N
+/// dependencies.
+#define EXPECT_NUM_DEPENDENCIES(RegIdx, N)                                     \
+  EXPECT_EQ(RW->getReg(RegIdx).Dependencies.size(), static_cast<unsigned>(N))
+
 /// Expects that register RegIdx in the rematerializer has no users.
 #define EXPECT_NO_USERS(RegIdx) EXPECT_NUM_USERS(RegIdx, 0)
 
@@ -508,21 +513,217 @@ TEST_F(RematerializerTest, SubRegRematSupport) {
     Rematerializer::DependencyReuseInfo DRI;
 
     const unsigned MBB0 = 0, MBB1 = 1;
-    const RegisterIdx Cst2 = 0;
+    const RegisterIdx Cst01 = 0, Cst2 = 1, Cst99 = 2;
 
-    // - %01 is not rematerializable because it has multiplie definitions.
     // - %34 is not rematerializable because it is defined over multiple
     // regions.
     // - %56 is not rematerializable because the second defining MI is
     // unrematerializable due to the implicit def.
     // - %78 is not rematerializable because it is read by an MI not defining it
     // before its last definition.
-    // - %99 is not rematerializable because it has multiplie definitions.
+    EXPECT_EQ(RW->getNumRegs(), 3U);
+
+    auto CheckBasicRemat = [&](RegisterIdx RegIdx,
+                               unsigned NumExpectDefs) -> void {
+      Rematerializer::DependencyReuseInfo DRI;
+      EXPECT_EQ(RW->getReg(RegIdx).Defs.size(), NumExpectDefs);
+      const RegisterIdx Remat = RW->rematerializeToRegion(RegIdx, MBB1, DRI);
+      RW.moveMIs(MBB0, MBB1, NumExpectDefs);
+      ASSERT_REGION_SIZES();
+      EXPECT_REMAT(Remat, RegIdx, MBB1, 1);
+    };
+
+    CheckBasicRemat(Cst01, 2);
+    CheckBasicRemat(Cst2, 1);
+    CheckBasicRemat(Cst99, 2);
+  });
+}
+
+/// Checks that the user transfer logic works correctly when different defining
+/// MIs of the same rematerializable register start dependening on different
+/// versions (original and rematerialized) of the same register.
+TEST_F(RematerializerTest, SubRegUserTransfer) {
+  StringRef MIRBody = R"MIR(
+  bb.0:
+    undef %01.sub0:sreg_64 = S_MOV_B32 0
+    %01.sub1:sreg_64 = S_MOV_B32 1    
+    
+  bb.1:
+    undef %23.sub0:sreg_64 = S_MOV_B32 %01.sub0
+    %23.sub1:sreg_64 = S_MOV_B32 %01.sub1
+    S_NOP 0, implicit %23
+  
+    S_ENDPGM 0
+)MIR";
+  rematerializerTest(MIRBody, [](RematerializerWrapper &RW) {
+    Rematerializer::DependencyReuseInfo DRI;
+    Rollbacker Rollback;
+    RW->addListener(&Rollback);
+
+    const unsigned MBB1 = 1;
+    const RegisterIdx Cst01 = 0, Cst23 = 1;
+    EXPECT_EQ(RW->getReg(Cst01).Defs.size(), 2U);
+    EXPECT_EQ(RW->getReg(Cst23).Defs.size(), 2U);
+    MachineInstr *Cst23FirstDef = RW->getReg(Cst23).Defs[0];
+    MachineInstr *Cst23SecondDef = RW->getReg(Cst23).Defs[1];
+
+    // Create a rematerialization of %01 just before %23.
+    const RegisterIdx RematCst01 =
+        RW->rematerializeToPos(Cst01, MBB1, Cst23FirstDef, DRI);
+    EXPECT_NUM_USERS(Cst01, 2);
+    EXPECT_NUM_USERS(RematCst01, 0);
+    EXPECT_NUM_USERS(Cst23, 1);
+    EXPECT_NUM_DEPENDENCIES(Cst23, 1);
+
+    // Have the first def of %23 use the rematerialization of %01 (the second
+    // def still uses %01). This transfers a user to the rematerialization of
+    // %01 and adds the rematerialization of %01 as a rematerializable
+    // dependency to %23.
+    RW->transferUser(Cst01, RematCst01, MBB1, *Cst23FirstDef);
+    EXPECT_NUM_USERS(Cst01, 1);
+    EXPECT_NUM_USERS(RematCst01, 1);
+    EXPECT_NUM_USERS(Cst23, 1);
+    EXPECT_NUM_DEPENDENCIES(Cst23, 2);
+
+    // Have the second def of %23 use the rematerialization of %01 as well. This
+    // transfers a user to the rematerialization of %01 and removes %01 as a
+    // rematerializable dependency of %23.
+    RW->transferUser(Cst01, RematCst01, MBB1, *Cst23SecondDef);
+    EXPECT_NUM_USERS(Cst01, 0);
+    EXPECT_NUM_USERS(RematCst01, 2);
+    EXPECT_NUM_DEPENDENCIES(Cst23, 1);
+
+    // Rollback should restore everything to its original state.
+    Rollback.rollback(*RW);
+    EXPECT_NUM_USERS(Cst01, 2);
+    EXPECT_NUM_USERS(RematCst01, 0);
+    EXPECT_NUM_USERS(Cst23, 1);
+    EXPECT_NUM_DEPENDENCIES(Cst23, 1);
+  });
+}
+
+TEST_F(RematerializerTest, SubRegRollback) {
+  StringRef MIRBody = R"MIR(
+  bb.0:
+    undef %01.sub0:sreg_64 = S_MOV_B32 0
+    %unremat0:vgpr_32 = nofpexcept V_CVT_I32_F64_e32 0, implicit $exec, implicit $mode, implicit-def $m0
+    %01.sub1:sreg_64 = S_MOV_B32 1    
+    %unremat1:vgpr_32 = nofpexcept V_CVT_I32_F64_e32 1, implicit $exec, implicit $mode, implicit-def $m0
+  
+  bb.1:
+    undef %23.sub0:sreg_64 = S_MOV_B32 2
+    %23.sub1:sreg_64 = S_MOV_B32 3
+
+  bb.2:
+    undef %45.sub0:sreg_64 = S_MOV_B32 4
+    undef %67.sub0:sreg_64 = S_MOV_B32 6
+    %45.sub1:sreg_64 = S_MOV_B32 5
+    %67.sub1:sreg_64 = S_MOV_B32 7
+
+  bb.3:
+    S_NOP 0, implicit %01, implicit %23, implicit %45, implicit %67
+    S_NOP 0, implicit %unremat0, implicit %unremat1 
+    S_ENDPGM 0
+)MIR";
+  rematerializerTest(MIRBody, [](RematerializerWrapper &RW) {
+    Rematerializer::DependencyReuseInfo DRI;
+    Rollbacker Rollback;
+    RW->addListener(&Rollback);
+
+    const unsigned MBB0 = 0, MBB1 = 1, MBB2 = 2, MBB3 = 3;
+    const RegisterIdx Cst01 = 0, Cst23 = 1, Cst45 = 2, Cst67 = 3;
+
+    EXPECT_EQ(RW->getReg(Cst01).Defs.size(), 2U);
+    EXPECT_EQ(RW->getReg(Cst23).Defs.size(), 2U);
+    EXPECT_EQ(RW->getReg(Cst45).Defs.size(), 2U);
+    EXPECT_EQ(RW->getReg(Cst67).Defs.size(), 2U);
+
+    auto GetNextMI = [&](MachineInstr *MI) -> MachineInstr * {
+      return &*std::next(MI->getIterator());
+    };
+
+    auto GetDefMI = [&](RegisterIdx RegIdx, unsigned DefIdx) -> MachineInstr * {
+      return RW->getReg(RegIdx).Defs[DefIdx];
+    };
+
+    // Rematerialize and rollback %01.
+    MachineInstr *Unremat0 = GetNextMI(GetDefMI(Cst01, 0));
+    MachineInstr *Unremat1 = GetNextMI(GetDefMI(Cst01, 1));
+    const RegisterIdx RematCst01 =
+        RW->rematerializeToRegion(Cst01, MBB3, DRI.clear());
+    RW.moveMIs(MBB0, MBB3, 2);
+    ASSERT_REGION_SIZES();
+    EXPECT_REMAT(RematCst01, Cst01, MBB3, 1);
+
+    // Rollback must re-create MIs in the same order.
+    Rollback.rollback(*RW);
+    RW.moveMIs(MBB3, MBB0, 2);
+    ASSERT_REGION_SIZES();
+    EXPECT_EQ(Unremat0, GetNextMI(GetDefMI(Cst01, 0)));
+    EXPECT_EQ(Unremat1, GetNextMI(GetDefMI(Cst01, 1)));
+
+    // Rematerialize and rollback %23.
+    MachineBasicBlock::iterator EndOfMBB1 =
+        std::next(GetDefMI(Cst23, 1)->getIterator());
+    const RegisterIdx RematCst23 =
+        RW->rematerializeToRegion(Cst23, MBB3, DRI.clear());
+    RW.moveMIs(MBB1, MBB3, 2);
+    ASSERT_REGION_SIZES();
+    EXPECT_REMAT(RematCst23, Cst23, MBB3, 1);
 
-    RegisterIdx RematCst2 = RW->rematerializeToRegion(Cst2, MBB1, DRI);
-    RW.moveMIs(MBB0, MBB1, 1);
+    // Rollback must re-create MIs in the same order.
+    Rollback.rollback(*RW);
+    RW.moveMIs(MBB3, MBB1, 2);
+    ASSERT_REGION_SIZES();
+    MachineInstr *Cst23Def0 = GetDefMI(Cst23, 0);
+    MachineInstr *Cst23Def1 = GetDefMI(Cst23, 1);
+    EXPECT_EQ(Cst23Def1, GetNextMI(Cst23Def0));
+    EXPECT_EQ(EndOfMBB1, std::next(Cst23Def1->getIterator()));
+
+    // Rematerialize and rollback %45 and %67.
+    MachineBasicBlock::iterator EndOfMBB2 =
+        std::next(GetDefMI(Cst67, 1)->getIterator());
+    const RegisterIdx RematCst45 =
+        RW->rematerializeToRegion(Cst45, MBB3, DRI.clear());
+    const RegisterIdx RematCst67 =
+        RW->rematerializeToRegion(Cst67, MBB3, DRI.clear());
+    RW.moveMIs(MBB2, MBB3, 4);
+    ASSERT_REGION_SIZES();
+    EXPECT_REMAT(RematCst45, Cst45, MBB3, 1);
+    EXPECT_REMAT(RematCst67, Cst67, MBB3, 1);
+
+    // Rollback must re-create MIs in the same order.
+    Rollback.rollback(*RW);
+    RW.moveMIs(MBB3, MBB2, 4);
     ASSERT_REGION_SIZES();
-    EXPECT_REMAT(RematCst2, Cst2, MBB1, 1);
+    MachineInstr *Cst45Def0 = GetDefMI(Cst45, 0);
+    MachineInstr *Cst67Def0 = GetDefMI(Cst67, 0);
+    MachineInstr *Cst45Def1 = GetDefMI(Cst45, 1);
+    MachineInstr *Cst67Def1 = GetDefMI(Cst67, 1);
+    EXPECT_EQ(Cst67Def0, GetNextMI(Cst45Def0));
+    EXPECT_EQ(Cst45Def1, GetNextMI(Cst67Def0));
+    EXPECT_EQ(Cst67Def1, GetNextMI(Cst45Def1));
+    EXPECT_EQ(EndOfMBB2, std::next(Cst67Def1->getIterator()));
+  });
+}
+
+/// Checks that instructions which use a rematerializable register as their
+/// first operand (here the KILL pseudo) are not treated as defining
+/// instructions for that register.
+TEST_F(RematerializerTest, FirstOperandNotDef) {
+  StringRef MIRBody = R"MIR(
+  bb.0:
+    undef %0.sub0:sgpr_64 = S_MOV_B32 0
+    KILL %0
+    S_ENDPGM 0
+)MIR";
+  rematerializerTest(MIRBody, [](RematerializerWrapper &RW) {
+    Rematerializer::DependencyReuseInfo DRI;
+
+    const RegisterIdx Cst0 = 0;
+    EXPECT_EQ(RW->getNumRegs(), 1U);
+    EXPECT_EQ(RW->getReg(Cst0).Defs.size(), 1U);
+    EXPECT_NUM_USERS(Cst0, 1);
   });
 }
 
@@ -621,8 +822,9 @@ TEST_F(RematerializerTest, DeadDefCascadeDeletion) {
 
         Rematerializer::DependencyReuseInfo DRI;
         const unsigned MBB0 = 0, MBB1 = 1, MBB2 = 2;
-        const RegisterIdx Cst1Die = 1, AddDie = 2, Cst2 = 3, Add = 4;
-        ASSERT_EQ(RW->getNumRegs(), 5U);
+        const RegisterIdx Cst1Die = 1, AddDie = 2, MultidefDontDie = 3,
+                          Cst2 = 4, Add = 5;
+        ASSERT_EQ(RW->getNumRegs(), 6U);
 
         // Rematerialize %addDie along with %cst0Die right after %cst2.
         RW->rematerializeToRegion(AddDie, MBB1, DRI.reuse(Cst1Die));
@@ -630,7 +832,8 @@ TEST_F(RematerializerTest, DeadDefCascadeDeletion) {
 
         // %cst2 and %add are moved to their using region.
         RW->rematerializeToRegion(Cst2, MBB2, DRI.clear());
-        RW->rematerializeToRegion(Add, MBB2, DRI.clear());
+        RW->rematerializeToRegion(Add, MBB2,
+                                  DRI.clear().reuse(MultidefDontDie));
         RW.moveMIs(MBB1, MBB2, 2);
 
         // The rematerialization of %add makes %multidef.sub1 become a dead def.
@@ -648,7 +851,8 @@ TEST_F(RematerializerTest, DeadDefCascadeDeletion) {
         // deleted. It should be re-created at the beginning of its block, as it
         // was initially.
         Rollback.rollback(*RW);
-        EXPECT_EQ(RW->getReg(Cst2).DefMI, &*RW.MF.getBlockNumbered(1)->begin());
+        EXPECT_EQ(RW->getReg(Cst2).getFirstDef(),
+                  &*RW.MF.getBlockNumbered(1)->begin());
         RW.moveMIs(MBB2, MBB1, 2);
         ASSERT_REGION_SIZES();
       },
@@ -766,10 +970,10 @@ TEST_F(RematerializerTest, RollbackInvalidInsertPos) {
       RW.moveMIs(MBB1, MBB0, 3);
       ASSERT_REGION_SIZES();
 
-      MachineInstr *DefCst0 = RW->getReg(Cst0).DefMI;
-      MachineInstr *DefCst1 = RW->getReg(Cst1).DefMI;
-      MachineInstr *DefCst2 = RW->getReg(Cst2).DefMI;
-      MachineInstr *DefCst3 = RW->getReg(Cst3).DefMI;
+      MachineInstr *DefCst0 = RW->getReg(Cst0).getFirstDef();
+      MachineInstr *DefCst1 = RW->getReg(Cst1).getFirstDef();
+      MachineInstr *DefCst2 = RW->getReg(Cst2).getFirstDef();
+      MachineInstr *DefCst3 = RW->getReg(Cst3).getFirstDef();
       EXPECT_EQ(GetNextMI(DefCst0), DefCst1);
       EXPECT_EQ(GetNextMI(DefCst1), DefCst2);
       EXPECT_EQ(GetNextMI(DefCst2), DefCst3);
@@ -849,22 +1053,25 @@ TEST_F(RematerializerTest, RollbackNextPosIsRemat) {
     // This rematerialization is created right after %2, which is later
     // rematerialized. It is *not* recorded by the rollbacker.
     RegisterIdx RematCst0 = RW->rematerializeToRegion(Cst0, MBB1, DRI.clear());
-    ExpectSeq(RW->getReg(Cst2).DefMI, RW->getReg(RematCst0).DefMI);
-    ExpectSeq(RW->getReg(RematCst0).DefMI, Nop1);
+    ExpectSeq(RW->getReg(Cst2).getFirstDef(),
+              RW->getReg(RematCst0).getFirstDef());
+    ExpectSeq(RW->getReg(RematCst0).getFirstDef(), Nop1);
 
     RW->addListener(&Rollback);
 
     // This rematerialization is created right after %3, which is later
     // rematerialized. It is recorded by the rollbacker.
     RegisterIdx RematCst1 = RW->rematerializeToRegion(Cst1, MBB2, DRI.clear());
-    ExpectSeq(RW->getReg(Cst3).DefMI, RW->getReg(RematCst1).DefMI);
-    ExpectSeq(RW->getReg(RematCst1).DefMI, Nop2);
+    ExpectSeq(RW->getReg(Cst3).getFirstDef(),
+              RW->getReg(RematCst1).getFirstDef());
+    ExpectSeq(RW->getReg(RematCst1).getFirstDef(), Nop2);
 
     RegisterIdx RematCst2 = RW->rematerializeToRegion(Cst2, MBB3, DRI.clear());
     RegisterIdx RematCst3 = RW->rematerializeToRegion(Cst3, MBB3, DRI.clear());
 
-    ExpectSeq(RW->getReg(RematCst2).DefMI, RW->getReg(RematCst3).DefMI);
-    ExpectSeq(RW->getReg(RematCst3).DefMI, Nop3);
+    ExpectSeq(RW->getReg(RematCst2).getFirstDef(),
+              RW->getReg(RematCst3).getFirstDef());
+    ExpectSeq(RW->getReg(RematCst3).getFirstDef(), Nop3);
 
     // After rollback, %2 and %3 should be re-created at the beginning of their
     // respective original region.
@@ -872,11 +1079,12 @@ TEST_F(RematerializerTest, RollbackNextPosIsRemat) {
 
     // The rematerialization of %0 was not recorded so isn't rolled back, %2 is
     // re-created right before it.
-    ExpectSeq(RW->getReg(Cst2).DefMI, RW->getReg(RematCst0).DefMI);
-    ExpectSeq(RW->getReg(RematCst0).DefMI, Nop1);
+    ExpectSeq(RW->getReg(Cst2).getFirstDef(),
+              RW->getReg(RematCst0).getFirstDef());
+    ExpectSeq(RW->getReg(RematCst0).getFirstDef(), Nop1);
 
     // The rematerialization of %1 was recorded so is rolled back, %3 is
     // re-created before the S_NOP in its region.
-    ExpectSeq(RW->getReg(Cst3).DefMI, Nop2);
+    ExpectSeq(RW->getReg(Cst3).getFirstDef(), Nop2);
   });
 }

>From 1d418f9d5031223fcddb00b6f212604ae987c804 Mon Sep 17 00:00:00 2001
From: Lucas Ramirez <lucas.rami at proton.me>
Date: Tue, 9 Jun 2026 14:34:53 +0000
Subject: [PATCH 2/4] Use is_contained

---
 llvm/lib/CodeGen/Rematerializer.cpp | 2 +-
 1 file changed, 1 insertion(+), 1 deletion(-)

diff --git a/llvm/lib/CodeGen/Rematerializer.cpp b/llvm/lib/CodeGen/Rematerializer.cpp
index 0f39db5e24d8e..b08066dfec036 100644
--- a/llvm/lib/CodeGen/Rematerializer.cpp
+++ b/llvm/lib/CodeGen/Rematerializer.cpp
@@ -193,7 +193,7 @@ void Rematerializer::transferUserImpl(RegisterIdx FromRegIdx,
   bool IsNewDep = true;
   if (UserReg.Defs.size() > 1) {
     // Other defining MIs might already be using the new register.
-    IsNewDep = find(UserDeps, ToRegIdx) == UserDeps.end();
+    IsNewDep = !is_contained(UserDeps, ToRegIdx);
 
     auto MOIsFromReg = [FromReg](MachineOperand &MO) {
       return MO.getReg() == FromReg;

>From e44a8d4910d0817b04d16e060f34a3e0575b49ce Mon Sep 17 00:00:00 2001
From: Lucas Ramirez <lucas.rami at proton.me>
Date: Wed, 15 Jul 2026 13:12:06 +0000
Subject: [PATCH 3/4] Stray debug + inline closure

---
 llvm/include/llvm/CodeGen/Rematerializer.h |  5 +----
 llvm/lib/CodeGen/Rematerializer.cpp        | 14 ++++++--------
 2 files changed, 7 insertions(+), 12 deletions(-)

diff --git a/llvm/include/llvm/CodeGen/Rematerializer.h b/llvm/include/llvm/CodeGen/Rematerializer.h
index 496b38d4ee66f..ca611ad77f25c 100644
--- a/llvm/include/llvm/CodeGen/Rematerializer.h
+++ b/llvm/include/llvm/CodeGen/Rematerializer.h
@@ -142,10 +142,7 @@ class Rematerializer {
     /// instructions.
     Register getDefReg() const {
       const MachineInstr *DefMI = getFirstDef();
-      assert(DefMI && "defining instruction(s) were deleted");
-      if (!DefMI->getOperand(0).isDef())
-        dbgs() << *DefMI;
-      assert(DefMI->getOperand(0).isDef() && "not a register def");
+      assert(DefMI && DefMI->getOperand(0).isDef() && "not a register def");
       return DefMI->getOperand(0).getReg();
     }
 
diff --git a/llvm/lib/CodeGen/Rematerializer.cpp b/llvm/lib/CodeGen/Rematerializer.cpp
index b08066dfec036..29480f7015e9e 100644
--- a/llvm/lib/CodeGen/Rematerializer.cpp
+++ b/llvm/lib/CodeGen/Rematerializer.cpp
@@ -195,10 +195,6 @@ void Rematerializer::transferUserImpl(RegisterIdx FromRegIdx,
     // Other defining MIs might already be using the new register.
     IsNewDep = !is_contained(UserDeps, ToRegIdx);
 
-    auto MOIsFromReg = [FromReg](MachineOperand &MO) {
-      return MO.getReg() == FromReg;
-    };
-
     // If any other defining instruction of the rematerializable user still uses
     // the original register, we should not remove it from dependencies and may
     // need to add a new dependency if it is the first time the new register is
@@ -206,10 +202,12 @@ void Rematerializer::transferUserImpl(RegisterIdx FromRegIdx,
     for (MachineInstr *DefMI : UserReg.Defs) {
       if (DefMI == &UserMI)
         continue;
-      if (any_of(DefMI->all_uses(), MOIsFromReg)) {
-        if (IsNewDep)
-          UserDeps.push_back(ToRegIdx);
-        return;
+      for (const MachineOperand &MO : DefMI->all_uses()) {
+        if (MO.getReg() == FromReg) {
+          if (IsNewDep)
+            UserDeps.push_back(ToRegIdx);
+          return;
+        }
       }
     }
   }

>From c7127935b4c0e0ec2ffcf1c80f2a80b3ed4cd59b Mon Sep 17 00:00:00 2001
From: Lucas Ramirez <lucas.rami at proton.me>
Date: Wed, 15 Jul 2026 13:47:26 +0000
Subject: [PATCH 4/4] Fix failing test

---
 .../CodeGen/AMDGPU/sched_mfma_rewrite_copies.mir     | 12 ++++++------
 1 file changed, 6 insertions(+), 6 deletions(-)

diff --git a/llvm/test/CodeGen/AMDGPU/sched_mfma_rewrite_copies.mir b/llvm/test/CodeGen/AMDGPU/sched_mfma_rewrite_copies.mir
index bd642d0d88c7b..0362fcd0e4b6f 100644
--- a/llvm/test/CodeGen/AMDGPU/sched_mfma_rewrite_copies.mir
+++ b/llvm/test/CodeGen/AMDGPU/sched_mfma_rewrite_copies.mir
@@ -5277,14 +5277,14 @@ body:             |
   ; CHECK-NEXT:   [[DEF13:%[0-9]+]]:vreg_128_align2 = IMPLICIT_DEF
   ; CHECK-NEXT:   [[DEF14:%[0-9]+]]:vreg_64_align2 = IMPLICIT_DEF
   ; CHECK-NEXT:   [[DEF15:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
-  ; CHECK-NEXT:   undef [[AV_MOV_:%[0-9]+]].sub0:vreg_128_align2 = AV_MOV_B32_IMM_PSEUDO 0, implicit $exec
-  ; CHECK-NEXT:   [[AV_MOV_:%[0-9]+]].sub1:vreg_128_align2 = AV_MOV_B32_IMM_PSEUDO 0, implicit $exec
-  ; CHECK-NEXT:   [[AV_MOV_:%[0-9]+]].sub2:vreg_128_align2 = AV_MOV_B32_IMM_PSEUDO 0, implicit $exec
-  ; CHECK-NEXT:   [[AV_MOV_:%[0-9]+]].sub3:vreg_128_align2 = AV_MOV_B32_IMM_PSEUDO 0, implicit $exec
   ; CHECK-NEXT: {{  $}}
   ; CHECK-NEXT: bb.1:
   ; CHECK-NEXT:   successors: %bb.2(0x80000000)
   ; CHECK-NEXT: {{  $}}
+  ; CHECK-NEXT:   undef [[AV_MOV_:%[0-9]+]].sub0:vreg_128_align2 = AV_MOV_B32_IMM_PSEUDO 0, implicit $exec
+  ; CHECK-NEXT:   [[AV_MOV_:%[0-9]+]].sub1:vreg_128_align2 = AV_MOV_B32_IMM_PSEUDO 0, implicit $exec
+  ; CHECK-NEXT:   [[AV_MOV_:%[0-9]+]].sub2:vreg_128_align2 = AV_MOV_B32_IMM_PSEUDO 0, implicit $exec
+  ; CHECK-NEXT:   [[AV_MOV_:%[0-9]+]].sub3:vreg_128_align2 = AV_MOV_B32_IMM_PSEUDO 0, implicit $exec
   ; CHECK-NEXT:   [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_:%[0-9]+]]:vreg_128_align2 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 [[DEF11]], [[DEF12]], [[AV_MOV_]], 4, 4, [[DEF14]].sub0, [[DEF15]], 0, 0, implicit $mode, implicit $exec
   ; CHECK-NEXT:   [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_1:%[0-9]+]]:vreg_128_align2 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 [[DEF11]], [[DEF12]], [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_]], 4, 4, [[DEF14]].sub0, [[DEF15]], 0, 0, implicit $mode, implicit $exec
   ; CHECK-NEXT:   [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_2:%[0-9]+]]:vreg_128_align2 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 [[DEF11]], [[DEF12]], [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_1]], 4, 4, [[DEF14]].sub0, [[DEF15]], 0, 0, implicit $mode, implicit $exec
@@ -5377,12 +5377,12 @@ body:             |
   ; CHECK-NEXT:   [[DEF13:%[0-9]+]]:vreg_128_align2 = IMPLICIT_DEF
   ; CHECK-NEXT:   [[DEF14:%[0-9]+]]:vreg_64_align2 = IMPLICIT_DEF
   ; CHECK-NEXT:   [[DEF15:%[0-9]+]]:vgpr_32 = IMPLICIT_DEF
-  ; CHECK-NEXT:   undef [[AV_MOV_:%[0-9]+]].sub0_sub1:vreg_128_align2 = AV_MOV_B64_IMM_PSEUDO 0, implicit $exec
-  ; CHECK-NEXT:   [[AV_MOV_:%[0-9]+]].sub2_sub3:vreg_128_align2 = AV_MOV_B64_IMM_PSEUDO 0, implicit $exec
   ; CHECK-NEXT: {{  $}}
   ; CHECK-NEXT: bb.1:
   ; CHECK-NEXT:   successors: %bb.2(0x80000000)
   ; CHECK-NEXT: {{  $}}
+  ; CHECK-NEXT:   undef [[AV_MOV_:%[0-9]+]].sub0_sub1:vreg_128_align2 = AV_MOV_B64_IMM_PSEUDO 0, implicit $exec
+  ; CHECK-NEXT:   [[AV_MOV_:%[0-9]+]].sub2_sub3:vreg_128_align2 = AV_MOV_B64_IMM_PSEUDO 0, implicit $exec
   ; CHECK-NEXT:   [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_:%[0-9]+]]:vreg_128_align2 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 [[DEF11]], [[DEF12]], [[AV_MOV_]], 4, 4, [[DEF14]].sub0, [[DEF15]], 0, 0, implicit $mode, implicit $exec
   ; CHECK-NEXT:   [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_1:%[0-9]+]]:vreg_128_align2 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 [[DEF11]], [[DEF12]], [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_]], 4, 4, [[DEF14]].sub0, [[DEF15]], 0, 0, implicit $mode, implicit $exec
   ; CHECK-NEXT:   [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_2:%[0-9]+]]:vreg_128_align2 = contract nofpexcept V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64 [[DEF11]], [[DEF12]], [[V_MFMA_SCALE_F32_16X16X128_F8F6F4_f4_f4_vgprcd_e64_1]], 4, 4, [[DEF14]].sub0, [[DEF15]], 0, 0, implicit $mode, implicit $exec



More information about the llvm-commits mailing list