1 //===- RISCVTargetTransformInfo.h - RISC-V specific TTI ---------*- C++ -*-===// 2 // 3 // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. 4 // See https://llvm.org/LICENSE.txt for license information. 5 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception 6 // 7 //===----------------------------------------------------------------------===// 8 /// \file 9 /// This file defines a TargetTransformInfo::Concept conforming object specific 10 /// to the RISC-V target machine. It uses the target's detailed information to 11 /// provide more precise answers to certain TTI queries, while letting the 12 /// target independent and default TTI implementations handle the rest. 13 /// 14 //===----------------------------------------------------------------------===// 15 16 #ifndef LLVM_LIB_TARGET_RISCV_RISCVTARGETTRANSFORMINFO_H 17 #define LLVM_LIB_TARGET_RISCV_RISCVTARGETTRANSFORMINFO_H 18 19 #include "RISCVSubtarget.h" 20 #include "RISCVTargetMachine.h" 21 #include "llvm/Analysis/IVDescriptors.h" 22 #include "llvm/Analysis/TargetTransformInfo.h" 23 #include "llvm/CodeGen/BasicTTIImpl.h" 24 #include "llvm/IR/Function.h" 25 26 namespace llvm { 27 28 class RISCVTTIImpl : public BasicTTIImplBase<RISCVTTIImpl> { 29 using BaseT = BasicTTIImplBase<RISCVTTIImpl>; 30 using TTI = TargetTransformInfo; 31 32 friend BaseT; 33 34 const RISCVSubtarget *ST; 35 const RISCVTargetLowering *TLI; 36 37 const RISCVSubtarget *getST() const { return ST; } 38 const RISCVTargetLowering *getTLI() const { return TLI; } 39 40 unsigned getMaxVLFor(VectorType *Ty); 41 public: 42 explicit RISCVTTIImpl(const RISCVTargetMachine *TM, const Function &F) 43 : BaseT(TM, F.getParent()->getDataLayout()), ST(TM->getSubtargetImpl(F)), 44 TLI(ST->getTargetLowering()) {} 45 46 InstructionCost getIntImmCost(const APInt &Imm, Type *Ty, 47 TTI::TargetCostKind CostKind); 48 InstructionCost getIntImmCostInst(unsigned Opcode, unsigned Idx, 49 const APInt &Imm, Type *Ty, 50 TTI::TargetCostKind CostKind, 51 Instruction *Inst = nullptr); 52 InstructionCost getIntImmCostIntrin(Intrinsic::ID IID, unsigned Idx, 53 const APInt &Imm, Type *Ty, 54 TTI::TargetCostKind CostKind); 55 56 TargetTransformInfo::PopcntSupportKind getPopcntSupport(unsigned TyWidth); 57 58 bool shouldExpandReduction(const IntrinsicInst *II) const; 59 bool supportsScalableVectors() const { return ST->hasVInstructions(); } 60 bool emitGetActiveLaneMask() const { return ST->hasVInstructions(); } 61 Optional<unsigned> getMaxVScale() const; 62 Optional<unsigned> getVScaleForTuning() const; 63 64 TypeSize getRegisterBitWidth(TargetTransformInfo::RegisterKind K) const; 65 66 unsigned getRegUsageForType(Type *Ty); 67 68 InstructionCost getMaskedMemoryOpCost(unsigned Opcode, Type *Src, 69 Align Alignment, unsigned AddressSpace, 70 TTI::TargetCostKind CostKind); 71 72 void getUnrollingPreferences(Loop *L, ScalarEvolution &SE, 73 TTI::UnrollingPreferences &UP, 74 OptimizationRemarkEmitter *ORE); 75 76 void getPeelingPreferences(Loop *L, ScalarEvolution &SE, 77 TTI::PeelingPreferences &PP); 78 79 unsigned getMinVectorRegisterBitWidth() const { 80 return ST->useRVVForFixedLengthVectors() ? 16 : 0; 81 } 82 83 InstructionCost getSpliceCost(VectorType *Tp, int Index); 84 InstructionCost getShuffleCost(TTI::ShuffleKind Kind, VectorType *Tp, 85 ArrayRef<int> Mask, int Index, 86 VectorType *SubTp, 87 ArrayRef<const Value *> Args = None); 88 89 InstructionCost getIntrinsicInstrCost(const IntrinsicCostAttributes &ICA, 90 TTI::TargetCostKind CostKind); 91 92 InstructionCost getGatherScatterOpCost(unsigned Opcode, Type *DataTy, 93 const Value *Ptr, bool VariableMask, 94 Align Alignment, 95 TTI::TargetCostKind CostKind, 96 const Instruction *I); 97 98 InstructionCost getCastInstrCost(unsigned Opcode, Type *Dst, Type *Src, 99 TTI::CastContextHint CCH, 100 TTI::TargetCostKind CostKind, 101 const Instruction *I = nullptr); 102 103 InstructionCost getMinMaxReductionCost(VectorType *Ty, VectorType *CondTy, 104 bool IsUnsigned, 105 TTI::TargetCostKind CostKind); 106 107 InstructionCost getArithmeticReductionCost(unsigned Opcode, VectorType *Ty, 108 Optional<FastMathFlags> FMF, 109 TTI::TargetCostKind CostKind); 110 111 bool isElementTypeLegalForScalableVector(Type *Ty) const { 112 return TLI->isLegalElementTypeForRVV(Ty); 113 } 114 115 bool isLegalMaskedLoadStore(Type *DataType, Align Alignment) { 116 if (!ST->hasVInstructions()) 117 return false; 118 119 // Only support fixed vectors if we know the minimum vector size. 120 if (isa<FixedVectorType>(DataType) && !ST->useRVVForFixedLengthVectors()) 121 return false; 122 123 // Don't allow elements larger than the ELEN. 124 // FIXME: How to limit for scalable vectors? 125 if (isa<FixedVectorType>(DataType) && 126 DataType->getScalarSizeInBits() > ST->getELEN()) 127 return false; 128 129 if (Alignment < 130 DL.getTypeStoreSize(DataType->getScalarType()).getFixedSize()) 131 return false; 132 133 return TLI->isLegalElementTypeForRVV(DataType->getScalarType()); 134 } 135 136 bool isLegalMaskedLoad(Type *DataType, Align Alignment) { 137 return isLegalMaskedLoadStore(DataType, Alignment); 138 } 139 bool isLegalMaskedStore(Type *DataType, Align Alignment) { 140 return isLegalMaskedLoadStore(DataType, Alignment); 141 } 142 143 bool isLegalMaskedGatherScatter(Type *DataType, Align Alignment) { 144 if (!ST->hasVInstructions()) 145 return false; 146 147 // Only support fixed vectors if we know the minimum vector size. 148 if (isa<FixedVectorType>(DataType) && !ST->useRVVForFixedLengthVectors()) 149 return false; 150 151 // Don't allow elements larger than the ELEN. 152 // FIXME: How to limit for scalable vectors? 153 if (isa<FixedVectorType>(DataType) && 154 DataType->getScalarSizeInBits() > ST->getELEN()) 155 return false; 156 157 if (Alignment < 158 DL.getTypeStoreSize(DataType->getScalarType()).getFixedSize()) 159 return false; 160 161 return TLI->isLegalElementTypeForRVV(DataType->getScalarType()); 162 } 163 164 bool isLegalMaskedGather(Type *DataType, Align Alignment) { 165 return isLegalMaskedGatherScatter(DataType, Alignment); 166 } 167 bool isLegalMaskedScatter(Type *DataType, Align Alignment) { 168 return isLegalMaskedGatherScatter(DataType, Alignment); 169 } 170 171 bool forceScalarizeMaskedGather(VectorType *VTy, Align Alignment) { 172 // Scalarize masked gather for RV64 if EEW=64 indices aren't supported. 173 return ST->is64Bit() && !ST->hasVInstructionsI64(); 174 } 175 176 bool forceScalarizeMaskedScatter(VectorType *VTy, Align Alignment) { 177 // Scalarize masked scatter for RV64 if EEW=64 indices aren't supported. 178 return ST->is64Bit() && !ST->hasVInstructionsI64(); 179 } 180 181 /// \returns How the target needs this vector-predicated operation to be 182 /// transformed. 183 TargetTransformInfo::VPLegalization 184 getVPLegalizationStrategy(const VPIntrinsic &PI) const { 185 using VPLegalization = TargetTransformInfo::VPLegalization; 186 return VPLegalization(VPLegalization::Legal, VPLegalization::Legal); 187 } 188 189 bool isLegalToVectorizeReduction(const RecurrenceDescriptor &RdxDesc, 190 ElementCount VF) const { 191 if (!VF.isScalable()) 192 return true; 193 194 Type *Ty = RdxDesc.getRecurrenceType(); 195 if (!TLI->isLegalElementTypeForRVV(Ty)) 196 return false; 197 198 switch (RdxDesc.getRecurrenceKind()) { 199 case RecurKind::Add: 200 case RecurKind::FAdd: 201 case RecurKind::And: 202 case RecurKind::Or: 203 case RecurKind::Xor: 204 case RecurKind::SMin: 205 case RecurKind::SMax: 206 case RecurKind::UMin: 207 case RecurKind::UMax: 208 case RecurKind::FMin: 209 case RecurKind::FMax: 210 return true; 211 default: 212 return false; 213 } 214 } 215 216 unsigned getMaxInterleaveFactor(unsigned VF) { 217 // If the loop will not be vectorized, don't interleave the loop. 218 // Let regular unroll to unroll the loop. 219 return VF == 1 ? 1 : ST->getMaxInterleaveFactor(); 220 } 221 222 enum RISCVRegisterClass { GPRRC, FPRRC, VRRC }; 223 unsigned getNumberOfRegisters(unsigned ClassID) const { 224 switch (ClassID) { 225 case RISCVRegisterClass::GPRRC: 226 // 31 = 32 GPR - x0 (zero register) 227 // FIXME: Should we exclude fixed registers like SP, TP or GP? 228 return 31; 229 case RISCVRegisterClass::FPRRC: 230 if (ST->hasStdExtF()) 231 return 32; 232 return 0; 233 case RISCVRegisterClass::VRRC: 234 // Although there are 32 vector registers, v0 is special in that it is the 235 // only register that can be used to hold a mask. 236 // FIXME: Should we conservatively return 31 as the number of usable 237 // vector registers? 238 return ST->hasVInstructions() ? 32 : 0; 239 } 240 llvm_unreachable("unknown register class"); 241 } 242 243 unsigned getRegisterClassForType(bool Vector, Type *Ty = nullptr) const { 244 if (Vector) 245 return RISCVRegisterClass::VRRC; 246 if (!Ty) 247 return RISCVRegisterClass::GPRRC; 248 249 Type *ScalarTy = Ty->getScalarType(); 250 if ((ScalarTy->isHalfTy() && ST->hasStdExtZfh()) || 251 (ScalarTy->isFloatTy() && ST->hasStdExtF()) || 252 (ScalarTy->isDoubleTy() && ST->hasStdExtD())) { 253 return RISCVRegisterClass::FPRRC; 254 } 255 256 return RISCVRegisterClass::GPRRC; 257 } 258 259 const char *getRegisterClassName(unsigned ClassID) const { 260 switch (ClassID) { 261 case RISCVRegisterClass::GPRRC: 262 return "RISCV::GPRRC"; 263 case RISCVRegisterClass::FPRRC: 264 return "RISCV::FPRRC"; 265 case RISCVRegisterClass::VRRC: 266 return "RISCV::VRRC"; 267 } 268 llvm_unreachable("unknown register class"); 269 } 270 }; 271 272 } // end namespace llvm 273 274 #endif // LLVM_LIB_TARGET_RISCV_RISCVTARGETTRANSFORMINFO_H 275