[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