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"
16#include "llvm/IR/LLVMContext.h" // Explicitly include for LLVMContext
18
19using namespace llvm;
20using namespace MIPatternMatch;
21
27
28/// This match is part of a combine that
29/// rewrites length(X - Y) to distance(X, Y)
30/// (f32 (g_intrinsic length
31/// (g_fsub (vXf32 X) (vXf32 Y))))
32/// ->
33/// (f32 (g_intrinsic distance
34/// (vXf32 X) (vXf32 Y)))
35///
37 if (MI.getOpcode() != TargetOpcode::G_INTRINSIC ||
38 cast<GIntrinsic>(MI).getIntrinsicID() != Intrinsic::spv_length)
39 return false;
40
41 // First operand of MI is `G_INTRINSIC` so start at operand 2.
42 Register SubReg = MI.getOperand(2).getReg();
43 MachineInstr *SubInstr = MRI.getVRegDef(SubReg);
44 if (SubInstr->getOpcode() != TargetOpcode::G_FSUB)
45 return false;
46
47 return true;
48}
49
51 // Extract the operands for X and Y from the match criteria.
52 Register SubDestReg = MI.getOperand(2).getReg();
53 MachineInstr *SubInstr = MRI.getVRegDef(SubDestReg);
54 Register SubOperand1 = SubInstr->getOperand(1).getReg();
55 Register SubOperand2 = SubInstr->getOperand(2).getReg();
56 Register ResultReg = MI.getOperand(0).getReg();
57
58 Builder.setInstrAndDebugLoc(MI);
59 Builder.buildIntrinsic(Intrinsic::spv_distance, ResultReg)
60 .addUse(SubOperand1)
61 .addUse(SubOperand2);
62
63 MI.eraseFromParent();
64}
65
66/// This match is part of a combine that
67/// rewrites select(fcmp(dot(I, Ng), 0), N, -N) to faceforward(N, I, Ng)
68/// (vXf32 (g_select
69/// (g_fcmp
70/// (g_intrinsic dot(vXf32 I) (vXf32 Ng)
71/// 0)
72/// (vXf32 N)
73/// (vXf32 g_fneg (vXf32 N))))
74/// ->
75/// (vXf32 (g_intrinsic faceforward
76/// (vXf32 N) (vXf32 I) (vXf32 Ng)))
77///
78/// This only works for Vulkan shader targets.
79///
81 if (!STI.isShader())
82 return false;
83
84 // Match overall select pattern.
85 Register CondReg, TrueReg, FalseReg;
86 if (!mi_match(MI.getOperand(0).getReg(), MRI,
87 m_GISelect(m_Reg(CondReg), m_Reg(TrueReg), m_Reg(FalseReg))))
88 return false;
89
90 // Match the FCMP condition.
91 Register DotReg, CondZeroReg;
93 if (!mi_match(CondReg, MRI,
94 m_GFCmp(m_Pred(Pred), m_Reg(DotReg), m_Reg(CondZeroReg))))
95 return false;
96 if (Pred == CmpInst::FCMP_OGT || Pred == CmpInst::FCMP_UGT)
97 std::swap(DotReg, CondZeroReg);
98 else if (!(Pred == CmpInst::FCMP_OLT || Pred == CmpInst::FCMP_ULT))
99 return false;
100
101 // Check if FCMP is a comparison between a dot product and 0.
102 MachineInstr *DotInstr = MRI.getVRegDef(DotReg);
103 if (DotInstr->getOpcode() != TargetOpcode::G_INTRINSIC ||
104 cast<GIntrinsic>(DotInstr)->getIntrinsicID() != Intrinsic::spv_fdot) {
105 Register DotOperand1, DotOperand2;
106 // Check for scalar dot product.
107 if (!mi_match(DotReg, MRI,
108 m_GFMul(m_Reg(DotOperand1), m_Reg(DotOperand2))) ||
109 !MRI.getType(DotOperand1).isScalar() ||
110 !MRI.getType(DotOperand2).isScalar())
111 return false;
112 }
113
114 const ConstantFP *ZeroVal;
115 if (!mi_match(CondZeroReg, MRI, m_GFCst(ZeroVal)) || !ZeroVal->isZero())
116 return false;
117
118 // Check if select's false operand is the negation of the true operand.
119 auto AreNegatedConstantsOrSplats = [&](Register TrueReg, Register FalseReg) {
120 std::optional<FPValueAndVReg> TrueVal, FalseVal;
121 if (!mi_match(TrueReg, MRI, m_GFCstOrSplat(TrueVal)) ||
122 !mi_match(FalseReg, MRI, m_GFCstOrSplat(FalseVal)))
123 return false;
124 APFloat TrueValNegated = TrueVal->Value;
125 TrueValNegated.changeSign();
126 return FalseVal->Value.compare(TrueValNegated) == APFloat::cmpEqual;
127 };
128
129 if (!mi_match(TrueReg, MRI, m_GFNeg(m_SpecificReg(FalseReg))) &&
130 !mi_match(FalseReg, MRI, m_GFNeg(m_SpecificReg(TrueReg)))) {
131 std::optional<FPValueAndVReg> MulConstant;
132 MachineInstr *TrueInstr = MRI.getVRegDef(TrueReg);
133 MachineInstr *FalseInstr = MRI.getVRegDef(FalseReg);
134 if (TrueInstr->getOpcode() == TargetOpcode::G_BUILD_VECTOR &&
135 FalseInstr->getOpcode() == TargetOpcode::G_BUILD_VECTOR &&
136 TrueInstr->getNumOperands() == FalseInstr->getNumOperands()) {
137 for (unsigned I = 1; I < TrueInstr->getNumOperands(); ++I)
138 if (!AreNegatedConstantsOrSplats(TrueInstr->getOperand(I).getReg(),
139 FalseInstr->getOperand(I).getReg()))
140 return false;
141 } else if (mi_match(TrueReg, MRI,
142 m_GFMul(m_SpecificReg(FalseReg),
143 m_GFCstOrSplat(MulConstant))) ||
144 mi_match(FalseReg, MRI,
145 m_GFMul(m_SpecificReg(TrueReg),
146 m_GFCstOrSplat(MulConstant))) ||
147 mi_match(TrueReg, MRI,
148 m_GFMul(m_GFCstOrSplat(MulConstant),
149 m_SpecificReg(FalseReg))) ||
150 mi_match(FalseReg, MRI,
151 m_GFMul(m_GFCstOrSplat(MulConstant),
152 m_SpecificReg(TrueReg)))) {
153 if (!MulConstant || !MulConstant->Value.isMinusOne())
154 return false;
155 } else if (!AreNegatedConstantsOrSplats(TrueReg, FalseReg))
156 return false;
157 }
158
159 return true;
160}
161
163 // Extract the operands for N, I, and Ng from the match criteria.
164 Register CondReg = MI.getOperand(1).getReg();
165 MachineInstr *CondInstr = MRI.getVRegDef(CondReg);
166 Register DotReg = CondInstr->getOperand(2).getReg();
167 CmpInst::Predicate Pred = cast<GFCmp>(CondInstr)->getCond();
168 if (Pred == CmpInst::FCMP_OGT || Pred == CmpInst::FCMP_UGT)
169 DotReg = CondInstr->getOperand(3).getReg();
170 MachineInstr *DotInstr = MRI.getVRegDef(DotReg);
171 Register DotOperand1, DotOperand2;
172 if (DotInstr->getOpcode() == TargetOpcode::G_FMUL) {
173 DotOperand1 = DotInstr->getOperand(1).getReg();
174 DotOperand2 = DotInstr->getOperand(2).getReg();
175 } else {
176 DotOperand1 = DotInstr->getOperand(2).getReg();
177 DotOperand2 = DotInstr->getOperand(3).getReg();
178 }
179 Register TrueReg = MI.getOperand(2).getReg();
180 Register FalseReg = MI.getOperand(3).getReg();
181 MachineInstr *TrueInstr = MRI.getVRegDef(TrueReg);
182 if (TrueInstr->getOpcode() == TargetOpcode::G_FNEG ||
183 TrueInstr->getOpcode() == TargetOpcode::G_FMUL)
184 std::swap(TrueReg, FalseReg);
185
186 Register ResultReg = MI.getOperand(0).getReg();
187 Builder.setInstrAndDebugLoc(MI);
188 Builder.buildIntrinsic(Intrinsic::spv_faceforward, ResultReg)
189 .addUse(TrueReg) // N
190 .addUse(DotOperand1) // I
191 .addUse(DotOperand2); // Ng
192
193 MI.eraseFromParent();
194}
195
197 return MI.getOpcode() == TargetOpcode::G_INTRINSIC &&
198 cast<GIntrinsic>(MI).getIntrinsicID() == Intrinsic::matrix_transpose;
199}
200
202 Register ResReg = MI.getOperand(0).getReg();
203 Register InReg = MI.getOperand(2).getReg();
204 uint32_t Rows = MI.getOperand(3).getImm();
205 uint32_t Cols = MI.getOperand(4).getImm();
206
207 Builder.setInstrAndDebugLoc(MI);
208
209 // A 1xN or Nx1 transpose is a pure reshape.
210 if (Rows == 1 || Cols == 1) {
211 Builder.buildCopy(ResReg, InReg);
212 MI.eraseFromParent();
213 return;
214 }
215
217 for (uint32_t K = 0; K < Rows * Cols; ++K) {
218 uint32_t R = K / Cols;
219 uint32_t C = K % Cols;
220 Mask.push_back(C * Rows + R);
221 }
222
223 Builder.buildShuffleVector(ResReg, InReg, InReg, Mask);
224 MI.eraseFromParent();
225}
226
228 return MI.getOpcode() == TargetOpcode::G_INTRINSIC &&
229 cast<GIntrinsic>(MI).getIntrinsicID() == Intrinsic::matrix_multiply;
230}
231
233SPIRVCombinerHelper::extractColumns(Register MatrixReg, uint32_t NumberOfCols,
234 SPIRVTypeInst SpvColType,
235 SPIRVGlobalRegistry *GR) const {
236 // If the matrix is a single colunm, return that single column.
237 if (NumberOfCols == 1)
238 return {MatrixReg};
239
241 LLT ColTy = GR->getRegType(SpvColType);
242 for (uint32_t J = 0; J < NumberOfCols; ++J)
244 Builder.buildUnmerge(Cols, MatrixReg);
245 for (Register R : Cols) {
246 setRegClassType(R, SpvColType, GR, &MRI, Builder.getMF());
247 }
248 return Cols;
249}
250
252SPIRVCombinerHelper::extractRows(Register MatrixReg, uint32_t NumRows,
253 uint32_t NumCols, SPIRVTypeInst SpvRowType,
254 SPIRVGlobalRegistry *GR) const {
256 LLT VecTy = GR->getRegType(SpvRowType);
257
258 // If there is only one column, then each row is a scalar that needs
259 // to be extracted.
260 if (NumCols == 1) {
261 assert(SpvRowType->getOpcode() != SPIRV::OpTypeVector);
262 for (uint32_t I = 0; I < NumRows; ++I)
263 Rows.push_back(MRI.createGenericVirtualRegister(VecTy));
264 Builder.buildUnmerge(Rows, MatrixReg);
265 for (Register R : Rows) {
266 setRegClassType(R, SpvRowType, GR, &MRI, Builder.getMF());
267 }
268 return Rows;
269 }
270
271 // If the matrix is a single row return that row.
272 if (NumRows == 1) {
273 return {MatrixReg};
274 }
275
276 for (uint32_t I = 0; I < NumRows; ++I) {
277 SmallVector<int, 4> Mask;
278 for (uint32_t k = 0; k < NumCols; ++k)
279 Mask.push_back(k * NumRows + I);
280 Rows.push_back(Builder.buildShuffleVector(VecTy, MatrixReg, MatrixReg, Mask)
281 .getReg(0));
282 }
283 for (Register R : Rows) {
284 setRegClassType(R, SpvRowType, GR, &MRI, Builder.getMF());
285 }
286 return Rows;
287}
288
289Register SPIRVCombinerHelper::computeDotProduct(Register RowA, Register ColB,
290 SPIRVTypeInst SpvVecType,
291 SPIRVGlobalRegistry *GR) const {
292 bool IsVectorOp = SpvVecType->getOpcode() == SPIRV::OpTypeVector;
293 SPIRVTypeInst SpvScalarType = GR->getScalarOrVectorComponentType(SpvVecType);
294 bool IsFloatOp = SpvScalarType->getOpcode() == SPIRV::OpTypeFloat;
295 LLT VecTy = GR->getRegType(SpvVecType);
296
297 Register DotRes;
298 if (IsVectorOp) {
299 LLT ScalarTy = VecTy.getElementType();
300 Intrinsic::SPVIntrinsics DotIntrinsic =
301 (IsFloatOp ? Intrinsic::spv_fdot : Intrinsic::spv_udot);
302 DotRes = Builder.buildIntrinsic(DotIntrinsic, {ScalarTy})
303 .addUse(RowA)
304 .addUse(ColB)
305 .getReg(0);
306 } else {
307 if (IsFloatOp)
308 DotRes = Builder.buildFMul(VecTy, RowA, ColB).getReg(0);
309 else
310 DotRes = Builder.buildMul(VecTy, RowA, ColB).getReg(0);
311 }
312 setRegClassType(DotRes, SpvScalarType, GR, &MRI, Builder.getMF());
313 return DotRes;
314}
315
316SmallVector<Register, 16> SPIRVCombinerHelper::computeDotProducts(
318 SPIRVTypeInst SpvVecType, SPIRVGlobalRegistry *GR) const {
319 SmallVector<Register, 16> ResultScalars;
320 for (uint32_t J = 0; J < ColsB.size(); ++J) {
321 for (uint32_t I = 0; I < RowsA.size(); ++I) {
322 ResultScalars.push_back(
323 computeDotProduct(RowsA[I], ColsB[J], SpvVecType, GR));
324 }
325 }
326 return ResultScalars;
327}
328
330SPIRVCombinerHelper::getDotProductVectorType(Register ResReg, uint32_t K,
331 SPIRVGlobalRegistry *GR) const {
332 // Loop over all non debug uses of ResReg
333 Type *ScalarResType = nullptr;
334 for (auto &UseMI : MRI.use_instructions(ResReg)) {
335 if (UseMI.getOpcode() != TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS)
336 continue;
337
338 if (!isSpvIntrinsic(UseMI, Intrinsic::spv_assign_type))
339 continue;
340
341 Type *Ty = getMDOperandAsType(UseMI.getOperand(2).getMetadata(), 0);
342 if (Ty->isVectorTy())
343 ScalarResType = cast<VectorType>(Ty)->getElementType();
344 else
345 ScalarResType = Ty;
346 assert(ScalarResType->isIntegerTy() || ScalarResType->isFloatingPointTy());
347 break;
348 }
349 if (!ScalarResType)
350 llvm_unreachable("Could not determine scalar result type");
351 Type *VecType =
352 (K > 1 ? FixedVectorType::get(ScalarResType, K) : ScalarResType);
353 return GR->getOrCreateSPIRVType(VecType, Builder,
354 SPIRV::AccessQualifier::None, false);
355}
356
358 Register ResReg = MI.getOperand(0).getReg();
359 Register AReg = MI.getOperand(2).getReg();
360 Register BReg = MI.getOperand(3).getReg();
361 uint32_t NumRowsA = MI.getOperand(4).getImm();
362 uint32_t NumColsA = MI.getOperand(5).getImm();
363 uint32_t NumColsB = MI.getOperand(6).getImm();
364
365 Builder.setInstrAndDebugLoc(MI);
366
368 MI.getMF()->getSubtarget<SPIRVSubtarget>().getSPIRVGlobalRegistry();
369
370 SPIRVTypeInst SpvVecType = getDotProductVectorType(ResReg, NumColsA, GR);
372 extractColumns(BReg, NumColsB, SpvVecType, GR);
374 extractRows(AReg, NumRowsA, NumColsA, SpvVecType, GR);
375 SmallVector<Register, 16> ResultScalars =
376 computeDotProducts(RowsA, ColsB, SpvVecType, GR);
377
378 if (ResultScalars.size() == 1)
379 Builder.buildCopy(ResReg, ResultScalars[0]);
380 else
381 Builder.buildBuildVector(ResReg, ResultScalars);
382 MI.eraseFromParent();
383}
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
#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:1393
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:867
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 matchMatrixTranspose(MachineInstr &MI) const
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)
bool matchMatrixMultiply(MachineInstr &MI) const
const SPIRVSubtarget & STI
void applySPIRVDistance(MachineInstr &MI) const
bool matchLengthToDistance(MachineInstr &MI) const
This match is part of a combine that rewrites length(X - Y) to distance(X, Y) (f32 (g_intrinsic lengt...
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:288
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:257
#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()
operand_type_match m_Pred()
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)
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.
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