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