LLVM 24.0.0git
SPIRVCombinerHelper.cpp
Go to the documentation of this file.
1//===-- SPIRVCombinerHelper.cpp -------------------------------------------===//
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
10#include "SPIRVGlobalRegistry.h"
11#include "SPIRVUtils.h"
15#include "llvm/IR/IntrinsicsSPIRV.h"
17
18using namespace llvm;
19using namespace MIPatternMatch;
20
26
27/// This match is part of a combine that
28/// rewrites X / length(X) to normalize(X)
29/// (vXf32 (g_fdiv
30/// (vXf32 X)
31/// (vXf32 splat
32/// (f32 (g_intrinsic length (vXf32 X))))))
33/// ->
34/// (vXf32 (g_intrinsic normalize (vXf32 X)))
35///
37 Register NumeratorReg = MI.getOperand(1).getReg();
38 Register DivisorReg = MI.getOperand(2).getReg();
39
40 // Match the divisor as a splat of length, inserted into lane 0.
41 MachineInstr *ShuffleInstr = MRI.getVRegDef(DivisorReg);
42 if (ShuffleInstr->getOpcode() != TargetOpcode::G_SHUFFLE_VECTOR)
43 return false;
44 if (!all_of(cast<GShuffleVector>(ShuffleInstr)->getMask(),
45 [](int M) { return M == 0; }))
46 return false;
47
48 MachineInstr *InsertInstr =
49 MRI.getVRegDef(ShuffleInstr->getOperand(1).getReg());
50 if (!isSpvIntrinsic(*InsertInstr, Intrinsic::spv_insertelt))
51 return false;
52 if (!mi_match(InsertInstr->getOperand(4).getReg(), MRI, m_ZeroInt()))
53 return false;
54
55 MachineInstr *LengthInstr =
56 MRI.getVRegDef(InsertInstr->getOperand(3).getReg());
57 if (!isSpvIntrinsic(*LengthInstr, Intrinsic::spv_length))
58 return false;
59
60 // Check that length's argument is the same as the numerator.
61 return LengthInstr->getOperand(2).getReg() == NumeratorReg;
62}
63
65 // Extract the operand for X from the match criteria.
66 Register NumeratorReg = MI.getOperand(1).getReg();
67 Register ResultReg = MI.getOperand(0).getReg();
68
69 Builder.setInstrAndDebugLoc(MI);
70 Builder.buildIntrinsic(Intrinsic::spv_normalize, ResultReg)
71 .addUse(NumeratorReg);
72
73 MI.eraseFromParent();
74}
75
76/// This match is part of a combine that
77/// rewrites select(fcmp(dot(I, Ng), 0), N, -N) to faceforward(N, I, Ng)
78/// (vXf32 (g_select
79/// (g_fcmp
80/// (g_intrinsic dot(vXf32 I) (vXf32 Ng)
81/// 0)
82/// (vXf32 N)
83/// (vXf32 g_fneg (vXf32 N))))
84/// ->
85/// (vXf32 (g_intrinsic faceforward
86/// (vXf32 N) (vXf32 I) (vXf32 Ng)))
87///
88/// This only works for Vulkan shader targets.
89///
91 if (!STI.isShader())
92 return false;
93
94 // Match overall select pattern.
95 Register CondReg, TrueReg, FalseReg;
96 if (!mi_match(MI.getOperand(0).getReg(), MRI,
97 m_GISelect(m_Reg(CondReg), m_Reg(TrueReg), m_Reg(FalseReg))))
98 return false;
99
100 // Match the FCMP condition.
101 Register DotReg, CondZeroReg;
103 if (!mi_match(CondReg, MRI,
104 m_GFCmp(m_Pred(Pred), m_Reg(DotReg), m_Reg(CondZeroReg))))
105 return false;
106 if (Pred == CmpInst::FCMP_OGT || Pred == CmpInst::FCMP_UGT)
107 std::swap(DotReg, CondZeroReg);
108 else if (!(Pred == CmpInst::FCMP_OLT || Pred == CmpInst::FCMP_ULT))
109 return false;
110
111 // Check if FCMP is a comparison between a dot product and 0.
113 Register DotOperand1, DotOperand2;
114 // Check for scalar dot product.
115 if (!mi_match(DotReg, MRI,
116 m_GFMul(m_Reg(DotOperand1), m_Reg(DotOperand2))) ||
117 !MRI.getType(DotOperand1).isScalar() ||
118 !MRI.getType(DotOperand2).isScalar())
119 return false;
120 }
121
122 const ConstantFP *ZeroVal;
123 if (!mi_match(CondZeroReg, MRI, m_GFCst(ZeroVal)) || !ZeroVal->isZero())
124 return false;
125
126 // Check if select's false operand is the negation of the true operand.
127 auto AreNegatedConstantsOrSplats = [&](Register TrueReg, Register FalseReg) {
128 std::optional<FPValueAndVReg> TrueVal, FalseVal;
129 if (!mi_match(TrueReg, MRI, m_GFCstOrSplat(TrueVal)) ||
130 !mi_match(FalseReg, MRI, m_GFCstOrSplat(FalseVal)))
131 return false;
132 APFloat TrueValNegated = TrueVal->Value;
133 TrueValNegated.changeSign();
134 return FalseVal->Value.compare(TrueValNegated) == APFloat::cmpEqual;
135 };
136
137 if (!mi_match(TrueReg, MRI, m_GFNeg(m_SpecificReg(FalseReg))) &&
138 !mi_match(FalseReg, MRI, m_GFNeg(m_SpecificReg(TrueReg)))) {
139 std::optional<FPValueAndVReg> MulConstant;
140 GBuildVector *TrueInstr, *FalseInstr;
141 if (mi_match(TrueReg, MRI, m_GBuildVector(TrueInstr)) &&
142 mi_match(FalseReg, MRI, m_GBuildVector(FalseInstr)) &&
143 TrueInstr->getNumOperands() == FalseInstr->getNumOperands()) {
144 for (unsigned I = 1; I < TrueInstr->getNumOperands(); ++I)
145 if (!AreNegatedConstantsOrSplats(TrueInstr->getOperand(I).getReg(),
146 FalseInstr->getOperand(I).getReg()))
147 return false;
148 } else if (mi_match(TrueReg, MRI,
149 m_GFMul(m_SpecificReg(FalseReg),
150 m_GFCstOrSplat(MulConstant))) ||
151 mi_match(FalseReg, MRI,
152 m_GFMul(m_SpecificReg(TrueReg),
153 m_GFCstOrSplat(MulConstant))) ||
154 mi_match(TrueReg, MRI,
155 m_GFMul(m_GFCstOrSplat(MulConstant),
156 m_SpecificReg(FalseReg))) ||
157 mi_match(FalseReg, MRI,
158 m_GFMul(m_GFCstOrSplat(MulConstant),
159 m_SpecificReg(TrueReg)))) {
160 if (!MulConstant || !MulConstant->Value.isMinusOne())
161 return false;
162 } else if (!AreNegatedConstantsOrSplats(TrueReg, FalseReg))
163 return false;
164 }
165
166 return true;
167}
168
170 // Extract the operands for N, I, and Ng from the match criteria.
171 Register CondReg = MI.getOperand(1).getReg();
172 MachineInstr *CondInstr = MRI.getVRegDef(CondReg);
173 Register DotReg = CondInstr->getOperand(2).getReg();
174 CmpInst::Predicate Pred = cast<GFCmp>(CondInstr)->getCond();
175 if (Pred == CmpInst::FCMP_OGT || Pred == CmpInst::FCMP_UGT)
176 DotReg = CondInstr->getOperand(3).getReg();
177 MachineInstr *DotInstr = MRI.getVRegDef(DotReg);
178 Register DotOperand1, DotOperand2;
179 if (DotInstr->getOpcode() == TargetOpcode::G_FMUL) {
180 DotOperand1 = DotInstr->getOperand(1).getReg();
181 DotOperand2 = DotInstr->getOperand(2).getReg();
182 } else {
183 DotOperand1 = DotInstr->getOperand(2).getReg();
184 DotOperand2 = DotInstr->getOperand(3).getReg();
185 }
186 Register TrueReg = MI.getOperand(2).getReg();
187 Register FalseReg = MI.getOperand(3).getReg();
188 MachineInstr *TrueInstr = MRI.getVRegDef(TrueReg);
189 if (TrueInstr->getOpcode() == TargetOpcode::G_FNEG ||
190 TrueInstr->getOpcode() == TargetOpcode::G_FMUL)
191 std::swap(TrueReg, FalseReg);
192
193 Register ResultReg = MI.getOperand(0).getReg();
194 Builder.setInstrAndDebugLoc(MI);
195 Builder.buildIntrinsic(Intrinsic::spv_faceforward, ResultReg)
196 .addUse(TrueReg) // N
197 .addUse(DotOperand1) // I
198 .addUse(DotOperand2); // Ng
199
200 MI.eraseFromParent();
201}
202
204 Register ResReg = MI.getOperand(0).getReg();
205 Register InReg = MI.getOperand(2).getReg();
206 uint32_t Rows = MI.getOperand(3).getImm();
207 uint32_t Cols = MI.getOperand(4).getImm();
208
209 Builder.setInstrAndDebugLoc(MI);
210
211 // A 1xN or Nx1 transpose is a pure reshape.
212 if (Rows == 1 || Cols == 1) {
213 Builder.buildCopy(ResReg, InReg);
214 MI.eraseFromParent();
215 return;
216 }
217
219 for (uint32_t K = 0; K < Rows * Cols; ++K) {
220 uint32_t R = K / Cols;
221 uint32_t C = K % Cols;
222 Mask.push_back(C * Rows + R);
223 }
224
225 Builder.buildShuffleVector(ResReg, InReg, InReg, Mask);
226 MI.eraseFromParent();
227}
228
230SPIRVCombinerHelper::extractColumns(Register MatrixReg, uint32_t NumberOfCols,
231 SPIRVTypeInst SpvColType,
232 SPIRVGlobalRegistry *GR) const {
233 // If the matrix is a single colunm, return that single column.
234 if (NumberOfCols == 1)
235 return {MatrixReg};
236
238 LLT ColTy = GR->getRegType(SpvColType);
239 for (uint32_t J = 0; J < NumberOfCols; ++J)
241 Builder.buildUnmerge(Cols, MatrixReg);
242 for (Register R : Cols) {
243 setRegClassType(R, SpvColType, GR, &MRI, Builder.getMF());
244 }
245 return Cols;
246}
247
249SPIRVCombinerHelper::extractRows(Register MatrixReg, uint32_t NumRows,
250 uint32_t NumCols, SPIRVTypeInst SpvRowType,
251 SPIRVGlobalRegistry *GR) const {
253 LLT VecTy = GR->getRegType(SpvRowType);
254
255 // If there is only one column, then each row is a scalar that needs
256 // to be extracted.
257 if (NumCols == 1) {
258 assert(!isVectorType(SpvRowType));
259 for (uint32_t I = 0; I < NumRows; ++I)
260 Rows.push_back(MRI.createGenericVirtualRegister(VecTy));
261 Builder.buildUnmerge(Rows, MatrixReg);
262 for (Register R : Rows) {
263 setRegClassType(R, SpvRowType, GR, &MRI, Builder.getMF());
264 }
265 return Rows;
266 }
267
268 // If the matrix is a single row return that row.
269 if (NumRows == 1) {
270 return {MatrixReg};
271 }
272
273 for (uint32_t I = 0; I < NumRows; ++I) {
274 SmallVector<int, 4> Mask;
275 for (uint32_t k = 0; k < NumCols; ++k)
276 Mask.push_back(k * NumRows + I);
277 Rows.push_back(Builder.buildShuffleVector(VecTy, MatrixReg, MatrixReg, Mask)
278 .getReg(0));
279 }
280 for (Register R : Rows) {
281 setRegClassType(R, SpvRowType, GR, &MRI, Builder.getMF());
282 }
283 return Rows;
284}
285
286Register SPIRVCombinerHelper::computeDotProduct(Register RowA, Register ColB,
287 SPIRVTypeInst SpvVecType,
288 SPIRVGlobalRegistry *GR) const {
289 SPIRVTypeInst SpvScalarType = GR->getScalarOrVectorComponentType(SpvVecType);
290 bool IsFloatOp = SpvScalarType->getOpcode() == SPIRV::OpTypeFloat;
291 LLT VecTy = GR->getRegType(SpvVecType);
292
293 Register DotRes;
294 if (isVectorType(SpvVecType)) {
295 LLT ScalarTy = VecTy.getElementType();
296 Intrinsic::SPVIntrinsics DotIntrinsic =
297 (IsFloatOp ? Intrinsic::spv_fdot : Intrinsic::spv_udot);
298 DotRes = Builder.buildIntrinsic(DotIntrinsic, {ScalarTy})
299 .addUse(RowA)
300 .addUse(ColB)
301 .getReg(0);
302 } else {
303 if (IsFloatOp)
304 DotRes = Builder.buildFMul(VecTy, RowA, ColB).getReg(0);
305 else
306 DotRes = Builder.buildMul(VecTy, RowA, ColB).getReg(0);
307 }
308 setRegClassType(DotRes, SpvScalarType, GR, &MRI, Builder.getMF());
309 return DotRes;
310}
311
312SmallVector<Register, 16> SPIRVCombinerHelper::computeDotProducts(
314 SPIRVTypeInst SpvVecType, SPIRVGlobalRegistry *GR) const {
315 SmallVector<Register, 16> ResultScalars;
316 for (uint32_t J = 0; J < ColsB.size(); ++J) {
317 for (uint32_t I = 0; I < RowsA.size(); ++I) {
318 ResultScalars.push_back(
319 computeDotProduct(RowsA[I], ColsB[J], SpvVecType, GR));
320 }
321 }
322 return ResultScalars;
323}
324
326SPIRVCombinerHelper::getDotProductVectorType(Register ResReg, uint32_t K,
327 SPIRVGlobalRegistry *GR) const {
328 // Loop over all non debug uses of ResReg
329 Type *ScalarResType = nullptr;
330 for (auto &UseMI : MRI.use_instructions(ResReg)) {
331 if (UseMI.getOpcode() != TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS)
332 continue;
333
334 if (!isSpvIntrinsic(UseMI, Intrinsic::spv_assign_type))
335 continue;
336
337 Type *Ty = getMDOperandAsType(UseMI.getOperand(2).getMetadata(), 0);
338 if (Ty->isVectorTy())
339 ScalarResType = cast<VectorType>(Ty)->getElementType();
340 else
341 ScalarResType = Ty;
342 assert(ScalarResType->isIntegerTy() || ScalarResType->isFloatingPointTy());
343 break;
344 }
345 if (!ScalarResType)
346 llvm_unreachable("Could not determine scalar result type");
347 Type *VecType =
348 (K > 1 ? FixedVectorType::get(ScalarResType, K) : ScalarResType);
349 return GR->getOrCreateSPIRVType(VecType, Builder,
350 SPIRV::AccessQualifier::None, false);
351}
352
354 Register ResReg = MI.getOperand(0).getReg();
355 Register AReg = MI.getOperand(2).getReg();
356 Register BReg = MI.getOperand(3).getReg();
357 uint32_t NumRowsA = MI.getOperand(4).getImm();
358 uint32_t NumColsA = MI.getOperand(5).getImm();
359 uint32_t NumColsB = MI.getOperand(6).getImm();
360
361 Builder.setInstrAndDebugLoc(MI);
362
364 MI.getMF()->getSubtarget<SPIRVSubtarget>().getSPIRVGlobalRegistry();
365
366 SPIRVTypeInst SpvVecType = getDotProductVectorType(ResReg, NumColsA, GR);
368 extractColumns(BReg, NumColsB, SpvVecType, GR);
370 extractRows(AReg, NumRowsA, NumColsA, SpvVecType, GR);
371 SmallVector<Register, 16> ResultScalars =
372 computeDotProducts(RowsA, ColsB, SpvVecType, GR);
373
374 if (ResultScalars.size() == 1)
375 Builder.buildCopy(ResReg, ResultScalars[0]);
376 else
377 Builder.buildBuildVector(ResReg, ResultScalars);
378 MI.eraseFromParent();
379}
MachineInstrBuilder & UseMI
assert(UImm &&(UImm !=~static_cast< T >(0)) &&"Invalid immediate!")
static GCRegistry::Add< ShadowStackGC > C("shadow-stack", "Very portable GC for uncooperative code generators")
static GCRegistry::Add< OcamlGC > B("ocaml", "ocaml 3.10-compatible GC")
Declares convenience wrapper classes for interpreting MachineInstr instances as specific generic oper...
IRTranslator LLVM IR MI
static std::pair< Value *, APInt > getMask(Value *WideMask, unsigned Factor, ElementCount LeafValueEC)
#define I(x, y, z)
Definition MD5.cpp:57
Contains matchers for matching SSA Machine Instructions.
Promote Memory to Register
Definition Mem2Reg.cpp:110
void changeSign()
Definition APFloat.h:1401
Represent a constant reference to an array (0 or more elements consecutively in memory),...
Definition ArrayRef.h:40
size_t size() const
Get the array size.
Definition ArrayRef.h:141
Predicate
This enumeration lists the possible predicates for CmpInst subclasses.
Definition InstrTypes.h:740
@ FCMP_OLT
0 1 0 0 True if ordered and less than
Definition InstrTypes.h:746
@ FCMP_OGT
0 0 1 0 True if ordered and greater than
Definition InstrTypes.h:744
@ FCMP_ULT
1 1 0 0 True if unordered or less than
Definition InstrTypes.h:754
@ FCMP_UGT
1 0 1 0 True if unordered or greater than
Definition InstrTypes.h:752
MachineRegisterInfo & MRI
const LegalizerInfo * LI
MachineDominatorTree * MDT
GISelValueTracking * VT
GISelChangeObserver & Observer
MachineIRBuilder & Builder
ConstantFP - Floating Point Values [float, double].
Definition Constants.h:420
bool isZero() const
Return true if the value is positive or negative zero.
Definition Constants.h:467
static LLVM_ABI FixedVectorType * get(Type *ElementType, unsigned NumElts)
Definition Type.cpp:843
Represents a G_BUILD_VECTOR.
Abstract class that contains various methods for clients to notify about changes.
LLT getElementType() const
Returns the vector's element type. Only valid for vector types.
DominatorTree Class - Concrete subclass of DominatorTreeBase that is used to compute a normal dominat...
Helper class to build MachineInstr.
Representation of each machine instruction.
unsigned getOpcode() const
Returns the opcode of this MachineInstr.
unsigned getNumOperands() const
Retuns the total number of operands.
const MachineOperand & getOperand(unsigned i) const
Register getReg() const
getReg - Returns the register number.
const MachineFunction & getMF() const
LLVM_ABI Register createGenericVirtualRegister(LLT Ty, StringRef Name="")
Create and return a new generic virtual register with low-level type Ty.
Wrapper class representing virtual and physical registers.
Definition Register.h:20
void applyMatrixMultiply(MachineInstr &MI) const
bool matchSelectToFaceForward(MachineInstr &MI) const
This match is part of a combine that rewrites select(fcmp(dot(I, Ng), 0), N, -N) to faceforward(N,...
void applyMatrixTranspose(MachineInstr &MI) const
bool matchFDivToNormalize(MachineInstr &MI) const
This match is part of a combine that rewrites X / length(X) to normalize(X) (vXf32 (g_fdiv (vXf32 X) ...
LLVM_ABI CombinerHelper(GISelChangeObserver &Observer, MachineIRBuilder &B, bool IsPreLegalize, GISelValueTracking *VT=nullptr, MachineDominatorTree *MDT=nullptr, const LegalizerInfo *LI=nullptr)
void applySPIRVFaceForward(MachineInstr &MI) const
SPIRVCombinerHelper(GISelChangeObserver &Observer, MachineIRBuilder &B, bool IsPreLegalize, GISelValueTracking *VT, MachineDominatorTree *MDT, const LegalizerInfo *LI, const SPIRVSubtarget &STI)
void applySPIRVNormalize(MachineInstr &MI) const
const SPIRVSubtarget & STI
LLT getRegType(SPIRVTypeInst SpvType) const
SPIRVTypeInst getScalarOrVectorComponentType(SPIRVTypeInst Type) const
SPIRVTypeInst getOrCreateSPIRVType(const Type *Type, MachineInstr &I, SPIRV::AccessQualifier::AccessQualifier AQ, bool EmitIR)
void push_back(const T &Elt)
This is a 'vector' (really, a variable-sized array), optimized for the case when the array is small.
bool isVectorTy() const
True if this is an instance of VectorType.
Definition Type.h:283
bool isFloatingPointTy() const
Return true if this is one of the floating-point types.
Definition Type.h:186
bool isIntegerTy() const
True if this is an instance of IntegerType.
Definition Type.h:252
#define llvm_unreachable(msg)
Marks that the current location is not supposed to be reachable.
constexpr std::underlying_type_t< E > Mask()
Get a bitmask with 1s in all places up to the high-order bit of E's largest value.
operand_type_match m_Reg()
GInstrBind< GBuildVector > m_GBuildVector(GBuildVector *&Inst)
operand_type_match m_Pred()
SpecificConstantMatch m_ZeroInt()
Convenience matchers for specific integer values.
TernaryOp_match< Src0Ty, Src1Ty, Src2Ty, TargetOpcode::G_SELECT > m_GISelect(const Src0Ty &Src0, const Src1Ty &Src1, const Src2Ty &Src2)
bool mi_match(Reg R, const MachineRegisterInfo &MRI, Pattern &&P)
SpecificRegisterMatch m_SpecificReg(Register RequestedReg)
Matches a register only if it is equal to RequestedReg.
UnaryOp_match< SrcTy, TargetOpcode::G_FNEG > m_GFNeg(const SrcTy &Src)
GFCstAndRegMatch m_GFCst(std::optional< FPValueAndVReg > &FPValReg)
GFCstOrSplatGFCstMatch m_GFCstOrSplat(std::optional< FPValueAndVReg > &FPValReg)
BinaryOp_match< LHS, RHS, TargetOpcode::G_FMUL, true > m_GFMul(const LHS &L, const RHS &R)
GInstrBind< GIntrinsic > m_GIntrinsic(GIntrinsic *&Inst)
Binds the defining instruction of Reg if it is a GIntrinsic (any of the four G_INTRINSIC* opcodes).
CompareOp_match< Pred, LHS, RHS, TargetOpcode::G_FCMP > m_GFCmp(const Pred &P, const LHS &L, const RHS &R)
This is an optimization pass for GlobalISel generic memory operations.
bool all_of(R &&range, UnaryPredicate P)
Provide wrappers to std::all_of which take ranges instead of having to pass begin/end explicitly.
Definition STLExtras.h:1755
bool isVectorType(SPIRVTypeInst SPVTy)
void setRegClassType(Register Reg, SPIRVTypeInst SpvType, SPIRVGlobalRegistry *GR, MachineRegisterInfo *MRI, const MachineFunction &MF, bool Force)
class LLVM_GSL_OWNER SmallVector
Forward declaration of SmallVector so that calculateSmallVectorDefaultInlinedElements can reference s...
decltype(auto) cast(const From &Val)
cast<X> - Return the argument parameter cast to the specified type.
Definition Casting.h:559
Type * getMDOperandAsType(const MDNode *N, unsigned I)
bool isSpvIntrinsic(const MachineInstr &MI, Intrinsic::ID IntrinsicID)
void swap(llvm::BitVector &LHS, llvm::BitVector &RHS)
Implement std::swap in terms of BitVector swap.
Definition BitVector.h:880