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