[llvm] [AMDGPU][Codegen] Legalize 16bit register class in sdag (PR #220275)

Matt Arsenault via llvm-commits llvm-commits at lists.llvm.org
Thu Sep 10 13:56:52 PDT 2026


================
@@ -4781,6 +4781,283 @@ bool AMDGPUDAGToDAGISel::isUniformLoad(const SDNode *N) const {
                ->isMemOpHasNoClobberedMemOperand(N)));
 }
 
+const TargetRegisterClass *
+AMDGPUDAGToDAGISel::inferDefRegClass(SDNode *N) const {
+  const SIRegisterInfo *TRI = Subtarget->getRegisterInfo();
+  if (!N->isMachineOpcode()) {
+    switch (N->getOpcode()) {
+    case ISD::CopyFromReg: {
+      Register Reg = cast<RegisterSDNode>(N->getOperand(1))->getReg();
+      if (Reg.isPhysical())
+        return TRI->getPhysRegBaseClass(Reg);
+      return CurDAG->getMachineFunction().getRegInfo().getRegClass(Reg);
+    }
+    }
+    return nullptr;
+  }
+
+  switch (N->getMachineOpcode()) {
+  case TargetOpcode::COPY: {
+    return inferDefRegClass(N->getOperand(0).getNode());
+  }
+  case TargetOpcode::EXTRACT_SUBREG: {
+    const TargetRegisterClass *SrcRC =
+        inferDefRegClass(N->getOperand(0).getNode());
+    if (!SrcRC)
+      return nullptr;
+    unsigned SubIdx = cast<ConstantSDNode>(N->getOperand(1))->getZExtValue();
+    return Subtarget->getRegisterInfo()->getSubRegisterClass(SrcRC, SubIdx);
+  }
+  case TargetOpcode::REG_SEQUENCE: {
+    unsigned RCID = N->getConstantOperandVal(0);
+    return Subtarget->getRegisterInfo()->getRegClass(RCID);
+  }
+  case TargetOpcode::COPY_TO_REGCLASS: {
+    unsigned RCID = cast<ConstantSDNode>(N->getOperand(1))->getZExtValue();
+    return TRI->getRegClass(RCID);
+  }
+  default:
+    const MCInstrDesc &Desc = TII->get(N->getMachineOpcode());
+    return TII->getRegClass(Desc, 0);
+  }
+}
+
+// Create a sreg32 from a vgpr16 in true16 mode
+static SDValue createVGPR16ToSGPR32(SDValue VReg16, SDLoc DL, EVT VT,
+                                    llvm::SelectionDAG *CurDAG) {
+  SDValue SubIdx0 = CurDAG->getTargetConstant(AMDGPU::lo16, DL, MVT::i32);
+  SDValue SubIdx1 = CurDAG->getTargetConstant(AMDGPU::hi16, DL, MVT::i32);
+  SDValue Undef = SDValue(
+      CurDAG->getMachineNode(TargetOpcode::IMPLICIT_DEF, DL, MVT::i16), 0);
+  SDValue VRegRCImm =
+      CurDAG->getTargetConstant(AMDGPU::VGPR_32RegClassID, DL, MVT::i32);
+  const SDValue Ops[] = {VRegRCImm, VReg16, SubIdx0, Undef, SubIdx1};
+  SDValue RegSeq = SDValue(
+      CurDAG->getMachineNode(TargetOpcode::REG_SEQUENCE, DL, MVT::i32, Ops), 0);
+  SDValue SRegRCImm =
+      CurDAG->getTargetConstant(AMDGPU::SGPR_32RegClassID, DL, MVT::i32);
+  return SDValue(CurDAG->getMachineNode(AMDGPU::COPY_TO_REGCLASS, DL, VT,
+                                        RegSeq, SRegRCImm),
+                 0);
+}
+
+// Create a vgpr16 from a sreg32 in true16 mode
+static SDValue createSGPR32ToVGPR16(SDValue SReg32, SDValue LoHi16, SDLoc DL,
+                                    EVT VT, llvm::SelectionDAG *CurDAG) {
+  SDValue VRegRCImm =
+      CurDAG->getTargetConstant(AMDGPU::VGPR_32RegClassID, DL, MVT::i32);
+  SDValue VReg32 = SDValue(CurDAG->getMachineNode(AMDGPU::COPY_TO_REGCLASS, DL,
+                                                  VT, SReg32, VRegRCImm),
+                           0);
+  return SDValue(CurDAG->getMachineNode(TargetOpcode::EXTRACT_SUBREG, DL, VT,
+                                        VReg32, LoHi16),
+                 0);
+}
+
+// Due to missing of sgpr16 class, 16bit value could be in vgpr16/sgpr32.
+// Check and legalize 16bit Register/SubregIdx in true16 mode includuing:
+// 1. 16bit register def-use chain that requires a fix (i.e. sgpr32->vgpr16)
+// 2. extract_subreg lo/hi16
+// by inserting REG_SEQUENCE and COPY_TO_REGCLASS
+// Legalization expected to be done from top-down
+bool AMDGPUDAGToDAGISel::Legalize16BitRegClass(SDNode *N) {
+  // This check is required for r600 and older targets
+  if (!CurDAG->getTarget().getTargetTriple().isAMDGCN() ||
+      !Subtarget->useRealTrue16Insts())
+    return false;
+
+  const SIRegisterInfo *TRI = Subtarget->getRegisterInfo();
+
+  EVT VT = N->getValueType(0);
+  // Check only 16bit types
+  if (VT != MVT::i16 && VT != MVT::f16 && VT != MVT::bf16)
+    return false;
+
+  const TargetRegisterClass *DstRC = nullptr;
+  SDLoc DL(N);
+
+  // EXTRACT_SUBREG Src, Lo16/Hi16
+  // 1. Src is SGPR:
+  // t0 = EXTRACT_SUBREG Src, sub0
+  // t1 = COPY_TO_REGCLASS t0, VGPR32
+  // t2 = EXTRACT_SUBREG t1, Lo16/Hi16
+  // User is SGPR => select t0
+  // User is VGPR => select t2
+  // 2. Src is VGPR:
+  // t0 = EXTRACT_SUBREG Src, Lo16/Hi16
+  // t1 = EXTRACT_SUBREG Src, sub0
+  // t2 = COPY_TO_REGCLASS t1, SGPR32
+  // User is VGPR => select t0
+  // User is SGPR => select t2
+  if (N->isMachineOpcode() &&
+      N->getMachineOpcode() == TargetOpcode::EXTRACT_SUBREG) {
+    unsigned SubIdx = cast<ConstantSDNode>(N->getOperand(1))->getZExtValue();
+    // Only check lo/hi16 subregidx
+    if (TRI->getSubRegIdxSize(SubIdx) != 16)
+      return false;
+
+    SDNode *Src = N->getOperand(0).getNode();
+    const TargetRegisterClass *SrcRC = inferDefRegClass(Src);
+    if (!SrcRC)
+      return false;
+
+    SmallVector<std::pair<SDNode *, unsigned>, 4> UserSGPR;
+    SmallVector<std::pair<SDNode *, unsigned>, 4> UserVGPR;
+
+    // Check Src/User regclass of extract_subreg
+    for (SDNode::use_iterator UI = N->use_begin(), UE = N->use_end(); UI != UE;
+         ++UI) {
+      SDNode *User = UI->getUser();
+      unsigned OperandNo = UI->getOperandNo();
+
+      const TargetRegisterClass *UserRC = getOperandRegClass(User, OperandNo);
+      if (!UserRC)
+        continue;
+
+      if (!TRI->getCommonSubClass(UserRC, &AMDGPU::SGPR_32RegClass))
+        UserVGPR.emplace_back(User, OperandNo);
+      else if (!TRI->getCommonSubClass(UserRC, &AMDGPU::VGPR_16RegClass))
+        UserSGPR.emplace_back(User, OperandNo);
+      else
+        TRI->isSGPRClass(SrcRC) ? UserSGPR.emplace_back(User, OperandNo)
+                                : UserVGPR.emplace_back(User, OperandNo);
+    }
+
+    SDValue NewValue;
+    if (TRI->isSGPRClass(SrcRC)) {
+      // SGPR extract_subreg with lo/hi16 is illegal
+      SDValue SReg32;
+      if (TRI->getSubClassWithSubReg(SrcRC, AMDGPU::sub0)) {
+        SDValue SubIdx = CurDAG->getTargetConstant(AMDGPU::sub0, DL, MVT::i32);
+        SReg32 =
+            SDValue(CurDAG->getMachineNode(TargetOpcode::EXTRACT_SUBREG, DL, VT,
+                                           SDValue(Src, 0), SubIdx),
+                    0);
+      } else
+        SReg32 = SDValue(Src, 0);
+
+      if (UserVGPR.size()) {
+        // t0: sgpr_xx = ...
+        // ... = extract_subreg t0, lo/hi16
+        // to
+        // t0: sgpr_xx = ...
+        // t1: sgpr_32 = extract_subreg t0, sub0
+        // t2: vgpr_32 = COPY_TO_VGPR32 t1
+        // t3  ... = extract_subreg t2, lo/hi16
+        // ... = t3: vgpr_16
+        NewValue =
+            createSGPR32ToVGPR16(SReg32, N->getOperand(1), DL, VT, CurDAG);
+        for (auto &[User, OperandNo] : UserVGPR) {
+          SmallVector<SDValue, 8> NewOps(User->op_begin(), User->op_end());
+          NewOps[OperandNo] = NewValue;
+          CurDAG->UpdateNodeOperands(User, NewOps);
+        }
+      }
+
+      if (UserSGPR.size()) {
+        // t0: sgpr_xx = ...
+        // ... = extract_subreg t0, lo/hi16
+        // to
+        // t0: sgpr_xx = ...
+        // t1: sgpr_32 = extract_subreg t0, sub0
+        // ... = t1: sgpr_32
+        for (auto &[User, OperandNo] : UserSGPR) {
+          SmallVector<SDValue, 8> NewOps(User->op_begin(), User->op_end());
+          NewOps[OperandNo] = SReg32;
+          CurDAG->UpdateNodeOperands(User, NewOps);
+        }
+      }
+      return true;
+    } else {
+
+      if (!UserSGPR.size())
+        return false;
+
+      // t0: vgpr_xx = ...
+      // t1: vgpr_16 = extract_subreg t0, lo/hi16
+      // to
+      // t0: vgpr_xx = ...
+      // t1: vgpr_32 = extract_subreg t0, sub0
+      // t2: sreg_32 = COPY_REGCLASS t1
+      // ... = t3: sreg_32
+      SDValue VReg32;
+      if (TRI->getSubClassWithSubReg(SrcRC, AMDGPU::sub0)) {
+        SDValue SubIdx = CurDAG->getTargetConstant(AMDGPU::sub0, DL, MVT::i32);
+        VReg32 =
+            SDValue(CurDAG->getMachineNode(TargetOpcode::EXTRACT_SUBREG, DL, VT,
+                                           SDValue(Src, 0), SubIdx),
+                    0);
+      } else
+        VReg32 = SDValue(Src, 0);
+      SDValue RCImm =
+          CurDAG->getTargetConstant(AMDGPU::SGPR_32RegClassID, DL, MVT::i32);
+      NewValue = SDValue(CurDAG->getMachineNode(AMDGPU::COPY_TO_REGCLASS, DL,
+                                                VT, VReg32, RCImm),
+                         0);
+      for (auto &[User, OperandNo] : UserSGPR) {
+        SmallVector<SDValue, 8> NewOps(User->op_begin(), User->op_end());
+        NewOps[OperandNo] = NewValue;
+        CurDAG->UpdateNodeOperands(User, NewOps);
+      }
+    }
+    return true;
+  }
+
+  // Check def register class
+  DstRC = inferDefRegClass(N);
+
+  if (!DstRC)
+    return false;
+
+  SmallVector<std::pair<SDNode *, unsigned>, 4> ToFix;
+
+  bool IsSGPR32 = TRI->getCommonSubClass(DstRC, &AMDGPU::SGPR_32RegClass);
+  bool IsVGPR16 = TRI->getCommonSubClass(DstRC, &AMDGPU::VGPR_16RegClass);
+
+  // Fix user:
+  // 1. Def is SGPR32, Use is VGPR16:
+  // t0 = COPY_TO_REGCLASS Def, VGPR32
+  // Use = EXTRACT_SUBREG t0, Lo16/Hi16
+  // 2. Def is VGPR16, Use is SGPR32:
+  // t0 = REG_SEQUENCE Def, lo16, undef, hi16
+  // Use = COPY_TO_REGCLASS t0, SGPR32
+  for (SDNode::use_iterator UI = N->use_begin(), UE = N->use_end(); UI != UE;
----------------
arsenm wrote:

I'm weary of use scans 

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


More information about the llvm-commits mailing list