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
9#include "SPIRVCombinerHelper.h"
10#include "SPIRVGlobalRegistry.h"
11#include "SPIRVUtils.h"
12#include "llvm/CodeGen/GlobalISel/GenericMachineInstrs.h"
13#include "llvm/CodeGen/GlobalISel/MIPatternMatch.h"
14#include "llvm/IR/DerivedTypes.h"
15#include "llvm/IR/IntrinsicsSPIRV.h"
16#include "llvm/Target/TargetMachine.h"
17
18using namespace llvm;
19using namespace MIPatternMatch;
20
21SPIRVCombinerHelper::SPIRVCombinerHelper(
22 GISelChangeObserver &Observer, MachineIRBuilder &B, bool IsPreLegalize,
23 GISelValueTracking *VT, MachineDominatorTree *MDT, const LegalizerInfo *LI,
24 const SPIRVSubtarget &STI)
25 : CombinerHelper(Observer, B, IsPreLegalize, VT, MDT, LI), STI(STI) {}
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///
36bool SPIRVCombinerHelper::matchFDivToNormalize(MachineInstr &MI) const {
37 Register NumeratorReg = MI.getOperand(i: 1).getReg();
38 Register DivisorReg = MI.getOperand(i: 2).getReg();
39
40 // Match the divisor as a splat of length, inserted into lane 0.
41 MachineInstr *ShuffleInstr = MRI.getVRegDef(Reg: DivisorReg);
42 if (ShuffleInstr->getOpcode() != TargetOpcode::G_SHUFFLE_VECTOR)
43 return false;
44 if (!all_of(Range: cast<GShuffleVector>(Val: ShuffleInstr)->getMask(),
45 P: [](int M) { return M == 0; }))
46 return false;
47
48 MachineInstr *InsertInstr =
49 MRI.getVRegDef(Reg: ShuffleInstr->getOperand(i: 1).getReg());
50 if (!isSpvIntrinsic(MI: *InsertInstr, IntrinsicID: Intrinsic::spv_insertelt))
51 return false;
52 if (!mi_match(R: InsertInstr->getOperand(i: 4).getReg(), MRI, P: m_ZeroInt()))
53 return false;
54
55 MachineInstr *LengthInstr =
56 MRI.getVRegDef(Reg: InsertInstr->getOperand(i: 3).getReg());
57 if (!isSpvIntrinsic(MI: *LengthInstr, IntrinsicID: Intrinsic::spv_length))
58 return false;
59
60 // Check that length's argument is the same as the numerator.
61 return LengthInstr->getOperand(i: 2).getReg() == NumeratorReg;
62}
63
64void SPIRVCombinerHelper::applySPIRVNormalize(MachineInstr &MI) const {
65 // Extract the operand for X from the match criteria.
66 Register NumeratorReg = MI.getOperand(i: 1).getReg();
67 Register ResultReg = MI.getOperand(i: 0).getReg();
68
69 Builder.setInstrAndDebugLoc(MI);
70 Builder.buildIntrinsic(ID: Intrinsic::spv_normalize, Res: ResultReg)
71 .addUse(RegNo: 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///
90bool SPIRVCombinerHelper::matchSelectToFaceForward(MachineInstr &MI) const {
91 if (!STI.isShader())
92 return false;
93
94 // Match overall select pattern.
95 Register CondReg, TrueReg, FalseReg;
96 if (!mi_match(R: MI.getOperand(i: 0).getReg(), MRI,
97 P: m_GISelect(Src0: m_Reg(R&: CondReg), Src1: m_Reg(R&: TrueReg), Src2: m_Reg(R&: FalseReg))))
98 return false;
99
100 // Match the FCMP condition.
101 Register DotReg, CondZeroReg;
102 CmpInst::Predicate Pred;
103 if (!mi_match(R: CondReg, MRI,
104 P: m_GFCmp(P: m_Pred(P&: Pred), L: m_Reg(R&: DotReg), R: m_Reg(R&: CondZeroReg))))
105 return false;
106 if (Pred == CmpInst::FCMP_OGT || Pred == CmpInst::FCMP_UGT)
107 std::swap(a&: DotReg, b&: 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.
112 if (!mi_match(R: DotReg, MRI, P: m_GIntrinsic<Intrinsic::spv_fdot>())) {
113 Register DotOperand1, DotOperand2;
114 // Check for scalar dot product.
115 if (!mi_match(R: DotReg, MRI,
116 P: m_GFMul(L: m_Reg(R&: DotOperand1), R: m_Reg(R&: DotOperand2))) ||
117 !MRI.getType(Reg: DotOperand1).isScalar() ||
118 !MRI.getType(Reg: DotOperand2).isScalar())
119 return false;
120 }
121
122 const ConstantFP *ZeroVal;
123 if (!mi_match(R: CondZeroReg, MRI, P: m_GFCst(C&: 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(R: TrueReg, MRI, P: m_GFCstOrSplat(FPValReg&: TrueVal)) ||
130 !mi_match(R: FalseReg, MRI, P: m_GFCstOrSplat(FPValReg&: FalseVal)))
131 return false;
132 APFloat TrueValNegated = TrueVal->Value;
133 TrueValNegated.changeSign();
134 return FalseVal->Value.compare(RHS: TrueValNegated) == APFloat::cmpEqual;
135 };
136
137 if (!mi_match(R: TrueReg, MRI, P: m_GFNeg(Src: m_SpecificReg(RequestedReg: FalseReg))) &&
138 !mi_match(R: FalseReg, MRI, P: m_GFNeg(Src: m_SpecificReg(RequestedReg: TrueReg)))) {
139 std::optional<FPValueAndVReg> MulConstant;
140 GBuildVector *TrueInstr, *FalseInstr;
141 if (mi_match(R: TrueReg, MRI, P: m_GBuildVector(Inst&: TrueInstr)) &&
142 mi_match(R: FalseReg, MRI, P: m_GBuildVector(Inst&: FalseInstr)) &&
143 TrueInstr->getNumOperands() == FalseInstr->getNumOperands()) {
144 for (unsigned I = 1; I < TrueInstr->getNumOperands(); ++I)
145 if (!AreNegatedConstantsOrSplats(TrueInstr->getOperand(i: I).getReg(),
146 FalseInstr->getOperand(i: I).getReg()))
147 return false;
148 } else if (mi_match(R: TrueReg, MRI,
149 P: m_GFMul(L: m_SpecificReg(RequestedReg: FalseReg),
150 R: m_GFCstOrSplat(FPValReg&: MulConstant))) ||
151 mi_match(R: FalseReg, MRI,
152 P: m_GFMul(L: m_SpecificReg(RequestedReg: TrueReg),
153 R: m_GFCstOrSplat(FPValReg&: MulConstant))) ||
154 mi_match(R: TrueReg, MRI,
155 P: m_GFMul(L: m_GFCstOrSplat(FPValReg&: MulConstant),
156 R: m_SpecificReg(RequestedReg: FalseReg))) ||
157 mi_match(R: FalseReg, MRI,
158 P: m_GFMul(L: m_GFCstOrSplat(FPValReg&: MulConstant),
159 R: m_SpecificReg(RequestedReg: 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
169void SPIRVCombinerHelper::applySPIRVFaceForward(MachineInstr &MI) const {
170 // Extract the operands for N, I, and Ng from the match criteria.
171 Register CondReg = MI.getOperand(i: 1).getReg();
172 MachineInstr *CondInstr = MRI.getVRegDef(Reg: CondReg);
173 Register DotReg = CondInstr->getOperand(i: 2).getReg();
174 CmpInst::Predicate Pred = cast<GFCmp>(Val: CondInstr)->getCond();
175 if (Pred == CmpInst::FCMP_OGT || Pred == CmpInst::FCMP_UGT)
176 DotReg = CondInstr->getOperand(i: 3).getReg();
177 MachineInstr *DotInstr = MRI.getVRegDef(Reg: DotReg);
178 Register DotOperand1, DotOperand2;
179 if (DotInstr->getOpcode() == TargetOpcode::G_FMUL) {
180 DotOperand1 = DotInstr->getOperand(i: 1).getReg();
181 DotOperand2 = DotInstr->getOperand(i: 2).getReg();
182 } else {
183 DotOperand1 = DotInstr->getOperand(i: 2).getReg();
184 DotOperand2 = DotInstr->getOperand(i: 3).getReg();
185 }
186 Register TrueReg = MI.getOperand(i: 2).getReg();
187 Register FalseReg = MI.getOperand(i: 3).getReg();
188 MachineInstr *TrueInstr = MRI.getVRegDef(Reg: TrueReg);
189 if (TrueInstr->getOpcode() == TargetOpcode::G_FNEG ||
190 TrueInstr->getOpcode() == TargetOpcode::G_FMUL)
191 std::swap(a&: TrueReg, b&: FalseReg);
192
193 Register ResultReg = MI.getOperand(i: 0).getReg();
194 Builder.setInstrAndDebugLoc(MI);
195 Builder.buildIntrinsic(ID: Intrinsic::spv_faceforward, Res: ResultReg)
196 .addUse(RegNo: TrueReg) // N
197 .addUse(RegNo: DotOperand1) // I
198 .addUse(RegNo: DotOperand2); // Ng
199
200 MI.eraseFromParent();
201}
202
203void SPIRVCombinerHelper::applyMatrixTranspose(MachineInstr &MI) const {
204 Register ResReg = MI.getOperand(i: 0).getReg();
205 Register InReg = MI.getOperand(i: 2).getReg();
206 uint32_t Rows = MI.getOperand(i: 3).getImm();
207 uint32_t Cols = MI.getOperand(i: 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(Res: ResReg, Op: InReg);
214 MI.eraseFromParent();
215 return;
216 }
217
218 SmallVector<int, 16> Mask;
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(Elt: C * Rows + R);
223 }
224
225 Builder.buildShuffleVector(Res: ResReg, Src1: InReg, Src2: InReg, Mask);
226 MI.eraseFromParent();
227}
228
229SmallVector<Register, 4>
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
237 SmallVector<Register, 4> Cols;
238 LLT ColTy = GR->getRegType(SpvType: SpvColType);
239 for (uint32_t J = 0; J < NumberOfCols; ++J)
240 Cols.push_back(Elt: MRI.createGenericVirtualRegister(Ty: ColTy));
241 Builder.buildUnmerge(Res: Cols, Op: MatrixReg);
242 for (Register R : Cols) {
243 setRegClassType(Reg: R, SpvType: SpvColType, GR, MRI: &MRI, MF: Builder.getMF());
244 }
245 return Cols;
246}
247
248SmallVector<Register, 4>
249SPIRVCombinerHelper::extractRows(Register MatrixReg, uint32_t NumRows,
250 uint32_t NumCols, SPIRVTypeInst SpvRowType,
251 SPIRVGlobalRegistry *GR) const {
252 SmallVector<Register, 4> Rows;
253 LLT VecTy = GR->getRegType(SpvType: 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(Elt: MRI.createGenericVirtualRegister(Ty: VecTy));
261 Builder.buildUnmerge(Res: Rows, Op: MatrixReg);
262 for (Register R : Rows) {
263 setRegClassType(Reg: R, SpvType: SpvRowType, GR, MRI: &MRI, MF: 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(Elt: k * NumRows + I);
277 Rows.push_back(Elt: Builder.buildShuffleVector(Res: VecTy, Src1: MatrixReg, Src2: MatrixReg, Mask)
278 .getReg(Idx: 0));
279 }
280 for (Register R : Rows) {
281 setRegClassType(Reg: R, SpvType: SpvRowType, GR, MRI: &MRI, MF: 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(Type: SpvVecType);
290 bool IsFloatOp = SpvScalarType->getOpcode() == SPIRV::OpTypeFloat;
291 LLT VecTy = GR->getRegType(SpvType: SpvVecType);
292
293 Register DotRes;
294 if (isVectorType(SPVTy: SpvVecType)) {
295 LLT ScalarTy = VecTy.getElementType();
296 Intrinsic::SPVIntrinsics DotIntrinsic =
297 (IsFloatOp ? Intrinsic::spv_fdot : Intrinsic::spv_udot);
298 DotRes = Builder.buildIntrinsic(ID: DotIntrinsic, Res: {ScalarTy})
299 .addUse(RegNo: RowA)
300 .addUse(RegNo: ColB)
301 .getReg(Idx: 0);
302 } else {
303 if (IsFloatOp)
304 DotRes = Builder.buildFMul(Dst: VecTy, Src0: RowA, Src1: ColB).getReg(Idx: 0);
305 else
306 DotRes = Builder.buildMul(Dst: VecTy, Src0: RowA, Src1: ColB).getReg(Idx: 0);
307 }
308 setRegClassType(Reg: DotRes, SpvType: SpvScalarType, GR, MRI: &MRI, MF: Builder.getMF());
309 return DotRes;
310}
311
312SmallVector<Register, 16> SPIRVCombinerHelper::computeDotProducts(
313 ArrayRef<Register> RowsA, ArrayRef<Register> ColsB,
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 Elt: computeDotProduct(RowA: RowsA[I], ColB: ColsB[J], SpvVecType, GR));
320 }
321 }
322 return ResultScalars;
323}
324
325SPIRVTypeInst
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(Reg: ResReg)) {
331 if (UseMI.getOpcode() != TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS)
332 continue;
333
334 if (!isSpvIntrinsic(MI: UseMI, IntrinsicID: Intrinsic::spv_assign_type))
335 continue;
336
337 Type *Ty = getMDOperandAsType(N: UseMI.getOperand(i: 2).getMetadata(), I: 0);
338 if (Ty->isVectorTy())
339 ScalarResType = cast<VectorType>(Val: 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(ElementType: ScalarResType, NumElts: K) : ScalarResType);
349 return GR->getOrCreateSPIRVType(Type: VecType, MIRBuilder&: Builder,
350 AQ: SPIRV::AccessQualifier::None, EmitIR: false);
351}
352
353void SPIRVCombinerHelper::applyMatrixMultiply(MachineInstr &MI) const {
354 Register ResReg = MI.getOperand(i: 0).getReg();
355 Register AReg = MI.getOperand(i: 2).getReg();
356 Register BReg = MI.getOperand(i: 3).getReg();
357 uint32_t NumRowsA = MI.getOperand(i: 4).getImm();
358 uint32_t NumColsA = MI.getOperand(i: 5).getImm();
359 uint32_t NumColsB = MI.getOperand(i: 6).getImm();
360
361 Builder.setInstrAndDebugLoc(MI);
362
363 SPIRVGlobalRegistry *GR =
364 MI.getMF()->getSubtarget<SPIRVSubtarget>().getSPIRVGlobalRegistry();
365
366 SPIRVTypeInst SpvVecType = getDotProductVectorType(ResReg, K: NumColsA, GR);
367 SmallVector<Register, 4> ColsB =
368 extractColumns(MatrixReg: BReg, NumberOfCols: NumColsB, SpvColType: SpvVecType, GR);
369 SmallVector<Register, 4> RowsA =
370 extractRows(MatrixReg: AReg, NumRows: NumRowsA, NumCols: NumColsA, SpvRowType: SpvVecType, GR);
371 SmallVector<Register, 16> ResultScalars =
372 computeDotProducts(RowsA, ColsB, SpvVecType, GR);
373
374 if (ResultScalars.size() == 1)
375 Builder.buildCopy(Res: ResReg, Op: ResultScalars[0]);
376 else
377 Builder.buildBuildVector(Res: ResReg, Ops: ResultScalars);
378 MI.eraseFromParent();
379}
380