1//===-- SPIRVPreLegalizer.cpp - prepare IR for legalization -----*- C++ -*-===//
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// The pass prepares IR for legalization: it assigns SPIR-V types to registers
10// and removes intrinsics which holded these types during IR translation.
11// Also it processes constants and registers them in GR to avoid duplication.
12//
13//===----------------------------------------------------------------------===//
14
15#include "SPIRV.h"
16#include "SPIRVSubtarget.h"
17#include "SPIRVUtils.h"
18#include "llvm/ADT/PostOrderIterator.h"
19#include "llvm/CodeGen/GlobalISel/CSEInfo.h"
20#include "llvm/CodeGen/GlobalISel/GISelValueTracking.h"
21#include "llvm/CodeGen/GlobalISel/MIPatternMatch.h"
22#include "llvm/CodeGen/MachineFunctionAnalysisManager.h"
23#include "llvm/CodeGen/MachinePassManager.h"
24#include "llvm/IR/Analysis.h"
25#include "llvm/IR/Attributes.h"
26#include "llvm/IR/Constants.h"
27#include "llvm/IR/InstrTypes.h"
28#include "llvm/IR/IntrinsicsSPIRV.h"
29#include "llvm/Support/MathExtras.h"
30
31#define DEBUG_TYPE "spirv-prelegalizer"
32
33using namespace llvm;
34using namespace llvm::MIPatternMatch;
35
36namespace {
37class SPIRVPreLegalizerLegacy : public MachineFunctionPass {
38public:
39 static char ID;
40 SPIRVPreLegalizerLegacy() : MachineFunctionPass(ID) {}
41 bool runOnMachineFunction(MachineFunction &MF) override;
42 void getAnalysisUsage(AnalysisUsage &AU) const override;
43};
44} // namespace
45
46void SPIRVPreLegalizerLegacy::getAnalysisUsage(AnalysisUsage &AU) const {
47 AU.addPreserved<GISelValueTrackingAnalysisLegacy>();
48 MachineFunctionPass::getAnalysisUsage(AU);
49}
50
51static inline void invalidateAndEraseMI(SPIRVGlobalRegistry *GR,
52 MachineInstr *MI) {
53 GR->invalidateMachineInstr(MI);
54 MI->eraseFromParent();
55}
56
57static void
58addConstantsToTrack(MachineFunction &MF, SPIRVGlobalRegistry *GR,
59 const SPIRVSubtarget &STI,
60 DenseMap<MachineInstr *, Type *> &TargetExtConstTypes) {
61 MachineRegisterInfo &MRI = MF.getRegInfo();
62 DenseMap<MachineInstr *, Register> RegsAlreadyAddedToDT;
63 SmallVector<MachineInstr *, 10> ToErase, ToEraseComposites;
64 for (MachineBasicBlock &MBB : MF) {
65 for (MachineInstr &MI : MBB) {
66 if (!isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_track_constant))
67 continue;
68 ToErase.push_back(Elt: &MI);
69 Register SrcReg = MI.getOperand(i: 2).getReg();
70 auto *Const =
71 cast<Constant>(Val: cast<ConstantAsMetadata>(
72 Val: MI.getOperand(i: 3).getMetadata()->getOperand(I: 0))
73 ->getValue());
74 if (auto *GV = dyn_cast<GlobalValue>(Val: Const)) {
75 Register Reg = GR->find(V: GV, MF: &MF);
76 if (!Reg.isValid()) {
77 GR->add(V: GV, MI: MRI.getVRegDef(Reg: SrcReg));
78 GR->addGlobalObject(V: GV, MF: &MF, R: SrcReg);
79 } else
80 RegsAlreadyAddedToDT[&MI] = Reg;
81 } else {
82 Register Reg = GR->find(V: Const, MF: &MF);
83 if (!Reg.isValid()) {
84 if (auto *ConstVec = dyn_cast<ConstantDataVector>(Val: Const)) {
85 auto *BuildVec = MRI.getVRegDef(Reg: SrcReg);
86 assert(BuildVec &&
87 BuildVec->getOpcode() == TargetOpcode::G_BUILD_VECTOR);
88 GR->add(V: Const, MI: BuildVec);
89 for (unsigned i = 0; i < ConstVec->getNumElements(); ++i) {
90 // Ensure that OpConstantComposite reuses a constant when it's
91 // already created and available in the same machine function.
92 Constant *ElemConst = ConstVec->getElementAsConstant(i);
93 Register ElemReg = GR->find(V: ElemConst, MF: &MF);
94 if (!ElemReg.isValid())
95 GR->add(V: ElemConst,
96 MI: MRI.getVRegDef(Reg: BuildVec->getOperand(i: 1 + i).getReg()));
97 else
98 BuildVec->getOperand(i: 1 + i).setReg(ElemReg);
99 }
100 }
101 if (Const->getType()->isTargetExtTy()) {
102 // remember association so that we can restore it when assign types
103 MachineInstr *SrcMI = MRI.getVRegDef(Reg: SrcReg);
104 if (SrcMI)
105 GR->add(V: Const, MI: SrcMI);
106 if (SrcMI && (SrcMI->getOpcode() == TargetOpcode::G_CONSTANT ||
107 SrcMI->getOpcode() == TargetOpcode::G_IMPLICIT_DEF))
108 TargetExtConstTypes[SrcMI] = Const->getType();
109 if (Const->isNullValue()) {
110 MachineBasicBlock &DepMBB = MF.front();
111 MachineIRBuilder MIB(DepMBB, DepMBB.getFirstNonPHI());
112 SPIRVTypeInst ExtType = GR->getOrCreateSPIRVType(
113 Type: Const->getType(), MIRBuilder&: MIB, AQ: SPIRV::AccessQualifier::ReadWrite,
114 EmitIR: true);
115 assert(SrcMI && "Expected source instruction to be valid");
116 SrcMI->setDesc(STI.getInstrInfo()->get(Opcode: SPIRV::OpConstantNull));
117 SrcMI->addOperand(Op: MachineOperand::CreateReg(
118 Reg: GR->getSPIRVTypeID(SpirvType: ExtType), isDef: false));
119 }
120 }
121 } else {
122 RegsAlreadyAddedToDT[&MI] = Reg;
123 // This MI is unused and will be removed. If the MI uses
124 // const_composite, it will be unused and should be removed too.
125 assert(MI.getOperand(2).isReg() && "Reg operand is expected");
126 MachineInstr *SrcMI = MRI.getVRegDef(Reg: MI.getOperand(i: 2).getReg());
127 if (SrcMI && isSpvIntrinsic(MI: *SrcMI, IntrinsicID: Intrinsic::spv_const_composite))
128 ToEraseComposites.push_back(Elt: SrcMI);
129 }
130 }
131 }
132 }
133 for (MachineInstr *MI : ToErase) {
134 Register Reg = MI->getOperand(i: 2).getReg();
135 auto It = RegsAlreadyAddedToDT.find(Val: MI);
136 if (It != RegsAlreadyAddedToDT.end())
137 Reg = It->second;
138 auto *RC = MRI.getRegClassOrNull(Reg: MI->getOperand(i: 0).getReg());
139 if (!MRI.getRegClassOrNull(Reg) && RC)
140 MRI.setRegClass(Reg, RC);
141 MRI.replaceRegWith(FromReg: MI->getOperand(i: 0).getReg(), ToReg: Reg);
142 invalidateAndEraseMI(GR, MI);
143 }
144 for (MachineInstr *MI : ToEraseComposites)
145 invalidateAndEraseMI(GR, MI);
146}
147
148static void foldConstantsIntoIntrinsics(MachineFunction &MF,
149 SPIRVGlobalRegistry *GR,
150 MachineIRBuilder MIB) {
151 SmallVector<MachineInstr *, 64> ToErase;
152 for (MachineBasicBlock &MBB : MF) {
153 for (MachineInstr &MI : MBB) {
154 if (!isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_assign_name))
155 continue;
156 const MDNode *MD = MI.getOperand(i: 2).getMetadata();
157 StringRef ValueName = cast<MDString>(Val: MD->getOperand(I: 0))->getString();
158 if (ValueName.size() > 0) {
159 MIB.setInsertPt(MBB&: *MI.getParent(), II: MI);
160 buildOpName(Target: MI.getOperand(i: 1).getReg(), Name: ValueName, MIRBuilder&: MIB);
161 }
162 ToErase.push_back(Elt: &MI);
163 }
164 for (MachineInstr *MI : ToErase)
165 invalidateAndEraseMI(GR, MI);
166 ToErase.clear();
167 }
168}
169
170static MachineInstr *findAssignTypeInstr(Register Reg,
171 MachineRegisterInfo *MRI) {
172 for (MachineRegisterInfo::use_instr_iterator I = MRI->use_instr_begin(RegNo: Reg),
173 IE = MRI->use_instr_end();
174 I != IE; ++I) {
175 MachineInstr *UseMI = &*I;
176 if ((isSpvIntrinsic(MI: *UseMI, IntrinsicID: Intrinsic::spv_assign_ptr_type) ||
177 isSpvIntrinsic(MI: *UseMI, IntrinsicID: Intrinsic::spv_assign_type)) &&
178 UseMI->getOperand(i: 1).getReg() == Reg)
179 return UseMI;
180 }
181 return nullptr;
182}
183
184static void buildOpBitcast(SPIRVGlobalRegistry *GR, MachineIRBuilder &MIB,
185 Register ResVReg, Register OpReg) {
186 SPIRVTypeInst ResType = GR->getSPIRVTypeForVReg(VReg: ResVReg);
187 SPIRVTypeInst OpType = GR->getSPIRVTypeForVReg(VReg: OpReg);
188 assert(ResType && OpType && "Operand types are expected");
189 if (!GR->isBitcastCompatible(Type1: ResType, Type2: OpType))
190 report_fatal_error(reason: "incompatible result and operand types in a bitcast");
191 MachineRegisterInfo *MRI = MIB.getMRI();
192 if (!MRI->getRegClassOrNull(Reg: ResVReg))
193 MRI->setRegClass(Reg: ResVReg, RC: GR->getRegClass(SpvType: ResType));
194 if (ResType == OpType)
195 MIB.buildInstr(Opcode: TargetOpcode::COPY).addDef(RegNo: ResVReg).addUse(RegNo: OpReg);
196 else
197 MIB.buildInstr(Opcode: SPIRV::OpBitcast)
198 .addDef(RegNo: ResVReg)
199 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: ResType))
200 .addUse(RegNo: OpReg);
201}
202
203// We lower G_BITCAST to OpBitcast here to avoid a MachineVerifier error.
204// The verifier checks if the source and destination LLTs of a G_BITCAST are
205// different, but this check is too strict for SPIR-V's typed pointers, which
206// may have the same LLT but different SPIRV type (e.g. pointers to different
207// pointee types). By lowering to OpBitcast here, we bypass the verifier's
208// check. See discussion in https://github.com/llvm/llvm-project/pull/110270
209// for more context.
210//
211// We also handle the llvm.spv.bitcast intrinsic here. If the source and
212// destination SPIR-V types are the same, we lower it to a COPY to enable
213// further optimizations like copy propagation.
214static void lowerBitcasts(MachineFunction &MF, SPIRVGlobalRegistry *GR,
215 MachineIRBuilder MIB) {
216 SmallVector<MachineInstr *, 16> ToErase;
217 for (MachineBasicBlock &MBB : MF) {
218 for (MachineInstr &MI : MBB) {
219 if (isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_bitcast)) {
220 Register DstReg = MI.getOperand(i: 0).getReg();
221 Register SrcReg = MI.getOperand(i: 2).getReg();
222 SPIRVTypeInst DstType = GR->getSPIRVTypeForVReg(VReg: DstReg);
223 assert(
224 DstType &&
225 "Expected destination SPIR-V type to have been assigned already.");
226 SPIRVTypeInst SrcType = GR->getSPIRVTypeForVReg(VReg: SrcReg);
227 assert(SrcType &&
228 "Expected source SPIR-V type to have been assigned already.");
229 if (DstType == SrcType) {
230 MIB.setInsertPt(MBB&: *MI.getParent(), II: MI);
231 MIB.buildCopy(Res: DstReg, Op: SrcReg);
232 ToErase.push_back(Elt: &MI);
233 continue;
234 }
235 }
236
237 if (MI.getOpcode() != TargetOpcode::G_BITCAST)
238 continue;
239
240 MIB.setInsertPt(MBB&: *MI.getParent(), II: MI);
241 buildOpBitcast(GR, MIB, ResVReg: MI.getOperand(i: 0).getReg(),
242 OpReg: MI.getOperand(i: 1).getReg());
243 ToErase.push_back(Elt: &MI);
244 }
245 }
246 for (MachineInstr *MI : ToErase)
247 invalidateAndEraseMI(GR, MI);
248}
249
250static void insertBitcasts(MachineFunction &MF, SPIRVGlobalRegistry *GR,
251 MachineIRBuilder MIB) {
252 // Get access to information about available extensions
253 const SPIRVSubtarget *ST =
254 static_cast<const SPIRVSubtarget *>(&MIB.getMF().getSubtarget());
255 SmallVector<MachineInstr *, 10> ToErase;
256 for (MachineBasicBlock &MBB : MF) {
257 for (MachineInstr &MI : MBB) {
258 if (!isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_ptrcast))
259 continue;
260 assert(MI.getOperand(2).isReg());
261 MIB.setInsertPt(MBB&: *MI.getParent(), II: MI);
262 ToErase.push_back(Elt: &MI);
263 Register Def = MI.getOperand(i: 0).getReg();
264 Register Source = MI.getOperand(i: 2).getReg();
265 Type *ElemTy = getMDOperandAsType(N: MI.getOperand(i: 3).getMetadata(), I: 0);
266 auto SC =
267 isa<FunctionType>(Val: ElemTy) &&
268 ST->canUseExtension(
269 E: SPIRV::Extension::SPV_INTEL_function_pointers)
270 ? SPIRV::StorageClass::CodeSectionINTEL
271 : addressSpaceToStorageClass(AddrSpace: MI.getOperand(i: 4).getImm(), STI: *ST);
272 SPIRVTypeInst AssignedPtrType =
273 GR->getOrCreateSPIRVPointerType(BaseType: ElemTy, I&: MI, SC);
274
275 // If the ptrcast would be redundant, replace all uses with the source
276 // register.
277 MachineRegisterInfo *MRI = MIB.getMRI();
278 // For untyped pointers the SPIR-V pointer type does not encode the
279 // pointee, so two pointers with different element types share the same
280 // pointer type. The element type still matters because it selects the
281 // Base Type operand of OpUntyped*AccessChainKHR. Treat the cast as
282 // redundant only when the source already carries the same element type.
283 // Otherwise keep a distinct register so the element type is preserved.
284 bool Redundant =
285 AssignedPtrType->getOpcode() == SPIRV::OpTypeUntypedPointerKHR
286 ? GR->getUntypedPtrElementType(Reg: Source) ==
287 GR->getOrCreateSPIRVType(Type: ElemTy, MIRBuilder&: MIB,
288 AQ: SPIRV::AccessQualifier::ReadWrite,
289 /*EmitIR=*/true)
290 : GR->getSPIRVTypeForVReg(VReg: Source) == AssignedPtrType;
291 if (Redundant) {
292 // Erase Def's assign type instruction if we are going to replace Def.
293 if (MachineInstr *AssignMI = findAssignTypeInstr(Reg: Def, MRI))
294 ToErase.push_back(Elt: AssignMI);
295 MRI->replaceRegWith(FromReg: Def, ToReg: Source);
296 } else {
297 if (!GR->getSPIRVTypeForVReg(VReg: Def, MF: &MF))
298 GR->assignSPIRVTypeToVReg(Type: AssignedPtrType, VReg: Def, MF);
299 MIB.buildBitcast(Dst: Def, Src: Source);
300 }
301 }
302 }
303 for (MachineInstr *MI : ToErase)
304 invalidateAndEraseMI(GR, MI);
305}
306
307// Translating GV, IRTranslator sometimes generates following IR:
308// %1 = G_GLOBAL_VALUE
309// %2 = COPY %1
310// %3 = G_ADDRSPACE_CAST %2
311//
312// or
313//
314// %1 = G_ZEXT %2
315// G_MEMCPY ... %2 ...
316//
317// New registers have no SPIRV type and no register class info.
318//
319// Set SPIRV type for GV, propagate it from GV to other instructions,
320// also set register classes.
321static SPIRVTypeInst propagateSPIRVType(MachineInstr *MI,
322 SPIRVGlobalRegistry *GR,
323 MachineRegisterInfo &MRI,
324 MachineIRBuilder &MIB) {
325 SPIRVTypeInst SpvType = nullptr;
326 assert(MI && "Machine instr is expected");
327 if (MI->getOperand(i: 0).isReg()) {
328 Register Reg = MI->getOperand(i: 0).getReg();
329 SpvType = GR->getSPIRVTypeForVReg(VReg: Reg);
330 if (!SpvType) {
331 switch (MI->getOpcode()) {
332 case TargetOpcode::G_FCONSTANT:
333 case TargetOpcode::G_CONSTANT: {
334 MIB.setInsertPt(MBB&: *MI->getParent(), II: MI);
335 Type *Ty = MI->getOperand(i: 1).getCImm()->getType();
336 SpvType = GR->getOrCreateSPIRVType(
337 Type: Ty, MIRBuilder&: MIB, AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
338 break;
339 }
340 case TargetOpcode::G_GLOBAL_VALUE: {
341 MIB.setInsertPt(MBB&: *MI->getParent(), II: MI);
342 const GlobalValue *Global = MI->getOperand(i: 1).getGlobal();
343 Type *ElementTy = toTypedPointer(Ty: GR->getDeducedGlobalValueType(Global));
344 unsigned AddrSpace = Global->getType()->getAddressSpace();
345 // Function pointers use CodeSectionINTEL storage class in SPIR-V when
346 // the SPV_INTEL_function_pointers extension is enabled.
347 const SPIRVSubtarget &ST = MIB.getMF().getSubtarget<SPIRVSubtarget>();
348 if (isa<Function>(Val: Global) &&
349 ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_function_pointers))
350 AddrSpace =
351 storageClassToAddressSpace(SC: SPIRV::StorageClass::CodeSectionINTEL);
352 auto *Ty = TypedPointerType::get(ElementType: ElementTy, AddressSpace: AddrSpace);
353 SpvType = GR->getOrCreateSPIRVType(
354 Type: Ty, MIRBuilder&: MIB, AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
355 break;
356 }
357 case TargetOpcode::G_ANYEXT:
358 case TargetOpcode::G_SEXT:
359 case TargetOpcode::G_ZEXT: {
360 if (MI->getOperand(i: 1).isReg()) {
361 if (MachineInstr *DefInstr =
362 MRI.getVRegDef(Reg: MI->getOperand(i: 1).getReg())) {
363 if (SPIRVTypeInst Def =
364 propagateSPIRVType(MI: DefInstr, GR, MRI, MIB)) {
365 unsigned CurrentBW = GR->getScalarOrVectorBitWidth(Type: Def);
366 unsigned ExpectedBW =
367 std::max(a: MRI.getType(Reg).getScalarSizeInBits(), b: CurrentBW);
368 unsigned NumElements = GR->getScalarOrVectorComponentCount(Type: Def);
369 SpvType = GR->getOrCreateSPIRVIntegerType(BitWidth: ExpectedBW, MIRBuilder&: MIB);
370 if (NumElements > 1)
371 SpvType = GR->getOrCreateSPIRVVectorType(BaseType: SpvType, NumElements,
372 MIRBuilder&: MIB, EmitIR: true);
373 }
374 }
375 }
376 break;
377 }
378 case TargetOpcode::G_PTRTOINT:
379 SpvType = GR->getOrCreateSPIRVIntegerType(
380 BitWidth: MRI.getType(Reg).getScalarSizeInBits(), MIRBuilder&: MIB);
381 break;
382 case TargetOpcode::G_TRUNC:
383 case TargetOpcode::G_ADDRSPACE_CAST:
384 case TargetOpcode::G_PTR_ADD:
385 case TargetOpcode::COPY: {
386 MachineOperand &Op = MI->getOperand(i: 1);
387 MachineInstr *Def = Op.isReg() ? MRI.getVRegDef(Reg: Op.getReg()) : nullptr;
388 if (Def)
389 SpvType = propagateSPIRVType(MI: Def, GR, MRI, MIB);
390 break;
391 }
392 default:
393 break;
394 }
395 if (SpvType) {
396 // check if the address space needs correction
397 LLT RegType = MRI.getType(Reg);
398 if (SpvType.isPointer() && RegType.isPointer() &&
399 storageClassToAddressSpace(SC: GR->getPointerStorageClass(Type: SpvType)) !=
400 RegType.getAddressSpace()) {
401 // Don't correct CodeSectionINTEL back to Function for function
402 // pointer G_GLOBAL_VALUE - the LLVM register has address space 0
403 // but the SPIR-V type was intentionally set to CodeSectionINTEL.
404 bool SkipCorrection =
405 MI->getOpcode() == TargetOpcode::G_GLOBAL_VALUE &&
406 GR->getPointerStorageClass(Type: SpvType) ==
407 SPIRV::StorageClass::CodeSectionINTEL;
408 if (!SkipCorrection) {
409 const SPIRVSubtarget &ST =
410 MI->getParent()->getParent()->getSubtarget<SPIRVSubtarget>();
411 auto TSC =
412 addressSpaceToStorageClass(AddrSpace: RegType.getAddressSpace(), STI: ST);
413 SpvType = GR->changePointerStorageClass(PtrType: SpvType, SC: TSC, I&: *MI);
414 }
415 }
416 GR->assignSPIRVTypeToVReg(Type: SpvType, VReg: Reg, MF: MIB.getMF());
417 }
418 if (!MRI.getRegClassOrNull(Reg))
419 MRI.setRegClass(Reg, RC: SpvType ? GR->getRegClass(SpvType)
420 : &SPIRV::iIDRegClass);
421 }
422 }
423 return SpvType;
424}
425
426// To support current approach and limitations wrt. bit width here we widen a
427// scalar register with a bit width greater than 1 to valid sizes and cap it to
428// 128 width.
429static unsigned widenBitWidthToNextPow2(unsigned BitWidth) {
430 if (BitWidth == 1)
431 return 1; // No need to widen 1-bit values
432 return std::min(a: std::max<unsigned>(a: PowerOf2Ceil(A: BitWidth), b: 8u), b: 128u);
433}
434
435static std::optional<unsigned>
436getNarrowScalarWidth(Register Reg, const MachineRegisterInfo &MRI) {
437 LLT Ty = MRI.getType(Reg);
438 if (!Ty.isScalar())
439 return std::nullopt;
440 unsigned W = Ty.getScalarSizeInBits();
441 // <= and not == because widenBitWidthToNextPow2 caps at 128.
442 if (widenBitWidthToNextPow2(BitWidth: W) <= W)
443 return std::nullopt;
444 return W;
445}
446
447static void widenScalarType(Register Reg, MachineRegisterInfo &MRI) {
448 LLT RegType = MRI.getType(Reg);
449 if (!RegType.isScalar())
450 return;
451 unsigned CurrentWidth = RegType.getScalarSizeInBits();
452 unsigned NewWidth = widenBitWidthToNextPow2(BitWidth: CurrentWidth);
453 if (NewWidth != CurrentWidth)
454 MRI.setType(VReg: Reg, Ty: LLT::scalar(SizeInBits: NewWidth));
455}
456
457static void widenCImmType(MachineOperand &MOP) {
458 const ConstantInt *CImmVal = MOP.getCImm();
459 unsigned CurrentWidth = CImmVal->getBitWidth();
460 unsigned NewWidth = widenBitWidthToNextPow2(BitWidth: CurrentWidth);
461 if (NewWidth != CurrentWidth) {
462 // Replace the immediate value with the widened version
463 MOP.setCImm(ConstantInt::get(Context&: CImmVal->getType()->getContext(),
464 V: CImmVal->getValue().zextOrTrunc(width: NewWidth)));
465 }
466}
467
468static void setInsertPtAfterDef(MachineIRBuilder &MIB, MachineInstr *Def) {
469 MachineBasicBlock &MBB = *Def->getParent();
470 MachineBasicBlock::iterator DefIt =
471 Def->getNextNode() ? Def->getNextNode()->getIterator() : MBB.end();
472 // Skip all the PHI and debug instructions.
473 while (DefIt != MBB.end() &&
474 (DefIt->isPHI() || DefIt->isDebugOrPseudoInstr()))
475 DefIt = std::next(x: DefIt);
476 MIB.setInsertPt(MBB, II: DefIt);
477}
478
479namespace llvm {
480void updateRegType(Register Reg, Type *Ty, SPIRVTypeInst SpvType,
481 SPIRVGlobalRegistry *GR, MachineIRBuilder &MIB,
482 MachineRegisterInfo &MRI) {
483 assert((Ty || SpvType) && "Either LLVM or SPIRV type is expected.");
484 MachineInstr *Def = MRI.getVRegDef(Reg);
485 setInsertPtAfterDef(MIB, Def);
486 if (!SpvType)
487 SpvType = GR->getOrCreateSPIRVType(Type: Ty, MIRBuilder&: MIB,
488 AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
489 if (!MRI.getRegClassOrNull(Reg))
490 MRI.setRegClass(Reg, RC: GR->getRegClass(SpvType));
491 if (!MRI.getType(Reg).isValid())
492 MRI.setType(VReg: Reg, Ty: GR->getRegType(SpvType));
493 GR->assignSPIRVTypeToVReg(Type: SpvType, VReg: Reg, MF: MIB.getMF());
494}
495
496void processInstr(MachineInstr &MI, MachineIRBuilder &MIB,
497 MachineRegisterInfo &MRI, SPIRVGlobalRegistry *GR,
498 SPIRVTypeInst KnownResType) {
499 MIB.setInsertPt(MBB&: *MI.getParent(), II: MI.getIterator());
500 for (auto &Op : MI.operands()) {
501 if (!Op.isReg() || Op.isDef())
502 continue;
503 Register OpReg = Op.getReg();
504 SPIRVTypeInst SpvType = GR->getSPIRVTypeForVReg(VReg: OpReg);
505 if (!SpvType && KnownResType) {
506 SpvType = KnownResType;
507 GR->assignSPIRVTypeToVReg(Type: KnownResType, VReg: OpReg, MF: *MI.getMF());
508 }
509 assert(SpvType);
510 if (!MRI.getRegClassOrNull(Reg: OpReg))
511 MRI.setRegClass(Reg: OpReg, RC: GR->getRegClass(SpvType));
512 if (!MRI.getType(Reg: OpReg).isValid())
513 MRI.setType(VReg: OpReg, Ty: GR->getRegType(SpvType));
514 }
515}
516} // namespace llvm
517
518// Sign-sensitive integer ops: their result depends on the value of the input
519// sign bit at position (width-1). On sub-pow2 widths the general widening
520// loop is a pure LLT relabel, which leaves the sign bit at the *original*
521// position instead of the widened MSB. These ops therefore need an explicit
522// G_SEXT_INREG on each value operand to move the sign bit up.
523//
524// Signed-vs-unsigned G_ICMP is distinguished by its predicate operand.
525//
526// TODO: follow-up PRs will add the remaining sign-sensitive opcodes
527// (e.g. G_SADDSAT/G_SSUBSAT, signed overflow ops).
528static bool isSignSensitiveOp(const MachineInstr &MI) {
529 switch (MI.getOpcode()) {
530 case TargetOpcode::G_ASHR:
531 case TargetOpcode::G_SDIV:
532 case TargetOpcode::G_SREM:
533 case TargetOpcode::G_SITOFP:
534 case TargetOpcode::G_SMIN:
535 case TargetOpcode::G_SMAX:
536 return true;
537 case TargetOpcode::G_ICMP:
538 return CmpInst::isSigned(
539 Pred: static_cast<CmpInst::Predicate>(MI.getOperand(i: 1).getPredicate()));
540 default:
541 return false;
542 }
543}
544
545struct NarrowWideningInfo {
546 // Width before widening of each sign-sensitive value-operand vreg (one entry
547 // per vreg).
548 DenseMap<Register, unsigned> OrigWidth;
549 // Sign-sensitive ops whose value operand(s) need replacing, ordered for
550 // reproducible vreg numbering.
551 SmallVector<MachineInstr *> SignSensitiveWorklist;
552 // Keyed by instruction, not vreg: G_TRUNC handling can replace the source.
553 SmallVector<std::pair<MachineInstr *, unsigned>> BitCountWorklist;
554};
555
556// G_CTTZ_ZERO_POISON is absent because its low bits are known non-zero, G_CTLS
557// because the backend does not select it.
558static bool isWidthSensitiveBitCountOp(unsigned Opcode) {
559 switch (Opcode) {
560 case TargetOpcode::G_CTLZ:
561 case TargetOpcode::G_CTLZ_ZERO_POISON:
562 case TargetOpcode::G_CTTZ:
563 case TargetOpcode::G_CTPOP:
564 return true;
565 default:
566 return false;
567 }
568}
569
570// Collect ops whose semantics depend on the operand width along with their
571// pre-widening widths, before later passes retype those vregs to pow2 LLTs
572// and the original width is no longer recoverable.
573static NarrowWideningInfo
574recordNarrowOperandWidths(MachineFunction &MF, const MachineRegisterInfo &MRI) {
575 NarrowWideningInfo Info;
576 auto RecordIfNarrow = [&](Register Reg) {
577 std::optional<unsigned> W = getNarrowScalarWidth(Reg, MRI);
578 if (!W)
579 return false;
580 Info.OrigWidth.try_emplace(Key: Reg, Args&: *W);
581 return true;
582 };
583 for (MachineBasicBlock &MBB : MF) {
584 for (MachineInstr &MI : MBB) {
585 if (isWidthSensitiveBitCountOp(Opcode: MI.getOpcode())) {
586 if (std::optional<unsigned> W =
587 getNarrowScalarWidth(Reg: MI.getOperand(i: 1).getReg(), MRI))
588 Info.BitCountWorklist.emplace_back(Args: &MI, Args&: *W);
589 continue;
590 }
591 if (!isSignSensitiveOp(MI))
592 continue;
593 bool NeedsRewrite = false;
594 for (const MachineOperand &MO : MI.all_uses())
595 NeedsRewrite = RecordIfNarrow(MO.getReg()) || NeedsRewrite;
596 if (NeedsRewrite) {
597 Info.SignSensitiveWorklist.push_back(Elt: &MI);
598 // Record the result too so it gets masked.
599 RecordIfNarrow(MI.getOperand(i: 0).getReg());
600 }
601 }
602 }
603 return Info;
604}
605
606// For every recorded sign-sensitive op, insert G_SEXT_INREG on each value
607// operand whose original width was narrower than the widened pow2 width and
608// retype the operand's vreg LLT in place to the widened width.
609//
610// Info must have been populated by recordNarrowOperandWidths before
611// other passes retyped the vregs; otherwise the narrow widths needed here
612// are lost.
613//
614// TODO: handle vector operands.
615static void widenSignSensitiveOps(MachineFunction &MF, SPIRVGlobalRegistry *GR,
616 MachineIRBuilder &MIB,
617 MachineRegisterInfo &MRI,
618 const NarrowWideningInfo &Info) {
619 // Emit G_SEXT_INREG from Reg's recorded narrow width; retypes Reg to the
620 // widened width and returns the sign-extended vreg.
621 auto SignExtendReg = [&](Register Reg, unsigned OldW,
622 MachineInstr &MI) -> Register {
623 unsigned NewW = widenBitWidthToNextPow2(BitWidth: OldW);
624 LLT NewLLT = LLT::scalar(SizeInBits: NewW);
625 MIB.setInsertPt(MBB&: *MI.getParent(), II: MI.getIterator());
626 SPIRVTypeInst SpvTy = GR->getOrCreateSPIRVIntegerType(BitWidth: NewW, MIRBuilder&: MIB);
627 Register SExted = MRI.createGenericVirtualRegister(Ty: NewLLT);
628 GR->assignSPIRVTypeToVReg(Type: SpvTy, VReg: SExted, MF);
629 MRI.setRegClass(Reg: SExted, RC: GR->getRegClass(SpvType: SpvTy));
630 MRI.setType(VReg: Reg, Ty: NewLLT);
631 MIB.buildSExtInReg(Res: SExted, Op: Reg, ImmOp: OldW);
632 return SExted;
633 };
634
635 // The wide op yields a sign-extended value, mask it back to OldW bits because
636 // later users assume the upper bits are zero.
637 auto MaskResult = [&](MachineInstr &MI, unsigned OldW) {
638 Register DstReg = MI.getOperand(i: 0).getReg();
639 widenScalarType(Reg: DstReg, MRI);
640 MIB.setInstrAndDebugLoc(MI);
641 SPIRVTypeInst SpvTy =
642 GR->getOrCreateSPIRVIntegerType(BitWidth: widenBitWidthToNextPow2(BitWidth: OldW), MIRBuilder&: MIB);
643 Register Result = createVirtualRegister(SpvType: SpvTy, GR, MIRBuilder&: MIB);
644 MI.getOperand(i: 0).setReg(Result);
645 setInsertPtAfterDef(MIB, Def: &MI);
646 MIB.buildZExtInReg(Res: DstReg, Op: Result, ImmOp: OldW);
647 };
648
649 // TODO: when the same narrow vreg feeds multiple sign-sensitive ops (e.g.
650 // sdiv %x, %y and srem %x, %y), emit one shared G_SEXT_INREG instead of one
651 // per use.
652 const TargetRegisterInfo &TRI = *MRI.getTargetRegisterInfo();
653 for (MachineInstr *MI : Info.SignSensitiveWorklist) {
654 for (const MachineOperand &MO : MI->all_uses()) {
655 Register Reg = MO.getReg();
656 auto It = Info.OrigWidth.find(Val: Reg);
657 if (It == Info.OrigWidth.end())
658 continue;
659 // substituteRegister fills every slot holding Reg at once: SignExtendReg
660 // retypes Reg in place, so a second sext would read the widened width.
661 MI->substituteRegister(FromReg: Reg, ToReg: SignExtendReg(Reg, It->second, *MI),
662 /*SubIdx=*/0, RegInfo: TRI);
663 }
664 Register Dst = MI->getOperand(i: 0).getReg();
665 if (auto It = Info.OrigWidth.find(Val: Dst); It != Info.OrigWidth.end())
666 MaskResult(*MI, It->second);
667 }
668}
669
670// LegalizerHelper::widenScalar has the same cases but cannot be reached: the
671// relabel retypes every narrow scalar to a pow2 LLT, so no illegal narrow type
672// ever reaches the legalizer.
673//
674// TODO: handle vector operands.
675static void widenBitCountOps(SPIRVGlobalRegistry *GR, MachineIRBuilder &MIB,
676 MachineRegisterInfo &MRI,
677 const NarrowWideningInfo &Info) {
678 for (auto [MI, OldWidth] : Info.BitCountWorklist) {
679 Register SrcReg = MI->getOperand(i: 1).getReg();
680 unsigned NewWidth = widenBitWidthToNextPow2(BitWidth: OldWidth);
681 LLT NewTy = LLT::scalar(SizeInBits: NewWidth);
682 widenScalarType(Reg: SrcReg, MRI);
683 MIB.setInstrAndDebugLoc(*MI);
684 SPIRVTypeInst SpvTy = GR->getOrCreateSPIRVIntegerType(BitWidth: NewWidth, MIRBuilder&: MIB);
685
686 // The G_TRUNC lowering masks its result to the narrow width, so a source
687 // coming from it needs no second mask.
688 APInt Cst;
689 bool HighBitsAlreadyZero =
690 mi_match(R: SrcReg, MRI, P: m_GAnd(L: m_Reg(), R: m_ICst(Cst))) &&
691 Cst.isSubsetOf(RHS: APInt::getLowBitsSet(numBits: Cst.getBitWidth(), loBitsSet: OldWidth));
692 auto ClearHighBits = [&](unsigned Width) -> Register {
693 if (HighBitsAlreadyZero)
694 return SrcReg;
695 Register Masked = createVirtualRegister(SpvType: SpvTy, GR, MIRBuilder&: MIB);
696 MIB.buildZExtInReg(Res: Masked, Op: SrcReg, ImmOp: Width);
697 return Masked;
698 };
699
700 Register Input;
701 switch (MI->getOpcode()) {
702 case TargetOpcode::G_CTLZ_ZERO_POISON: {
703 // Shifting up to the widened MSB moves the poison out too, so no
704 // adjustment.
705 Input = createVirtualRegister(SpvType: SpvTy, GR, MIRBuilder&: MIB);
706 auto Diff = MIB.buildConstant(Res: NewTy, Val: NewWidth - OldWidth);
707 MIB.buildShl(Dst: Input, Src0: SrcReg, Src1: Diff);
708 break;
709 }
710 case TargetOpcode::G_CTTZ: {
711 // Keeps an all-zero narrow value counting exactly OldWidth zeros.
712 Input = createVirtualRegister(SpvType: SpvTy, GR, MIRBuilder&: MIB);
713 auto TopBit =
714 MIB.buildConstant(Res: NewTy, Val: APInt::getOneBitSet(numBits: NewWidth, BitNo: OldWidth));
715 MIB.buildOr(Dst: Input, Src0: SrcReg, Src1: TopBit);
716 break;
717 }
718 case TargetOpcode::G_CTPOP:
719 Input = ClearHighBits(OldWidth);
720 break;
721 case TargetOpcode::G_CTLZ: {
722 // Clearing the extra bits adds leading zeros the count has to drop.
723 Input = ClearHighBits(OldWidth);
724 Register DstReg = MI->getOperand(i: 0).getReg();
725 widenScalarType(Reg: DstReg, MRI);
726 Register Count = createVirtualRegister(SpvType: SpvTy, GR, MIRBuilder&: MIB);
727 MI->getOperand(i: 0).setReg(Count);
728 setInsertPtAfterDef(MIB, Def: MI);
729 auto Diff = MIB.buildConstant(Res: NewTy, Val: NewWidth - OldWidth);
730 MIB.buildSub(Dst: DstReg, Src0: Count, Src1: Diff);
731 break;
732 }
733 default:
734 llvm_unreachable("unexpected width-sensitive bit-count opcode");
735 }
736 MI->getOperand(i: 1).setReg(Input);
737 }
738}
739
740static void
741generateAssignInstrs(MachineFunction &MF, SPIRVGlobalRegistry *GR,
742 MachineIRBuilder MIB,
743 DenseMap<MachineInstr *, Type *> &TargetExtConstTypes) {
744 // Get access to information about available extensions
745 const SPIRVSubtarget *ST =
746 static_cast<const SPIRVSubtarget *>(&MIB.getMF().getSubtarget());
747
748 MachineRegisterInfo &MRI = MF.getRegInfo();
749 SmallVector<MachineInstr *, 10> ToErase;
750 DenseMap<MachineInstr *, Register> RegsAlreadyAddedToDT;
751
752 bool IsExtendedInts =
753 ST->canUseExtension(
754 E: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_integers) ||
755 ST->canUseExtension(E: SPIRV::Extension::SPV_KHR_bit_instructions) ||
756 ST->canUseExtension(E: SPIRV::Extension::SPV_INTEL_int4);
757
758 if (!IsExtendedInts) {
759 // Without arbitrary precision integer extensions, SPIR-V only supports
760 // integer widths of 8, 16, 32, 64. Non-standard widths (e.g., i24, i40)
761 // must be widened to the next power of two.
762 //
763 // Record the original widths of width-sensitive operands before either
764 // the G_TRUNC handling or the general widening loop retypes vregs, then
765 // rewrite those ops after G_TRUNC processing using the recorded widths.
766 NarrowWideningInfo WideningInfo = recordNarrowOperandWidths(MF, MRI);
767
768 // G_TRUNC requires special handling because its semantics depend on the
769 // original destination width. For example:
770 // %dst:s24 = G_TRUNC %src:s64
771 // After widening s24 to s32, we cannot simply do:
772 // %dst:s32 = G_TRUNC %src:s64
773 // because this would keep 32 bits instead of 24. Instead, we insert a
774 // G_AND to mask the value to the original width:
775 // %mask:s64 = G_CONSTANT 0xFFFFFF ; 24-bit mask
776 // %masked:s64 = G_AND %src:s64, %mask
777 // %dst:s32 = G_TRUNC %masked:s64
778 // If src and dst widen to the same size, G_TRUNC is replaced entirely:
779 // %mask:s64 = G_CONSTANT 0xFFFFFFFFFF ; 40-bit mask
780 // %dst:s64 = G_AND %src:s64, %mask
781 SmallVector<MachineInstr *, 8> TruncToRemove;
782 for (MachineBasicBlock &MBB : MF) {
783 for (MachineInstr &MI : MBB) {
784 unsigned MIOp = MI.getOpcode();
785 if (MIOp != TargetOpcode::G_TRUNC)
786 continue;
787 assert(MI.getNumOperands() == 2);
788 assert(MI.getOperand(0).isReg());
789 assert(MI.getOperand(1).isReg());
790
791 Register DstReg = MI.getOperand(i: 0).getReg();
792 Register SrcReg = MI.getOperand(i: 1).getReg();
793
794 LLT DstTy = MRI.getType(Reg: DstReg);
795 LLT SrcTy = MRI.getType(Reg: SrcReg);
796 assert((DstTy.isScalar() || DstTy.isVector()) &&
797 (SrcTy.isScalar() || SrcTy.isVector()) &&
798 "Expected scalar or vector G_TRUNC types");
799 assert(DstTy.isVector() == SrcTy.isVector() &&
800 "Expected matching scalar/vector G_TRUNC types");
801 assert((!DstTy.isVector() ||
802 DstTy.getElementCount() == SrcTy.getElementCount()) &&
803 "Expected equal vector element counts");
804
805 unsigned OriginalDstWidth = DstTy.getScalarSizeInBits();
806 unsigned OriginalSrcWidth = SrcTy.getScalarSizeInBits();
807
808 unsigned NewDstWidth = widenBitWidthToNextPow2(BitWidth: OriginalDstWidth);
809 unsigned NewSrcWidth = widenBitWidthToNextPow2(BitWidth: OriginalSrcWidth);
810 LLT NewDstTy = DstTy.changeElementSize(NewEltSize: NewDstWidth);
811 LLT NewSrcTy = SrcTy.changeElementSize(NewEltSize: NewSrcWidth);
812
813 // No Dst width change means no truncation semantics change, but the
814 // source still needs a legal type.
815 if (OriginalDstWidth == NewDstWidth) {
816 MRI.setType(VReg: SrcReg, Ty: NewSrcTy);
817 continue;
818 }
819
820 MRI.setType(VReg: SrcReg, Ty: NewSrcTy);
821 MRI.setType(VReg: DstReg, Ty: NewDstTy);
822
823 MIB.setInsertPt(MBB, II: MI.getIterator());
824 APInt Mask = APInt::getLowBitsSet(numBits: NewSrcWidth, loBitsSet: OriginalDstWidth);
825 MachineInstrBuilder MaskReg =
826 DstTy.isVector()
827 ? MIB.buildBuildVectorConstant(
828 Res: NewSrcTy,
829 Ops: SmallVector<APInt, 4>(DstTy.getNumElements(), Mask))
830 : MIB.buildConstant(Res: NewSrcTy, Val: Mask);
831 Register MaskedReg = MRI.createGenericVirtualRegister(Ty: NewSrcTy);
832 MIB.buildAnd(Dst: MaskedReg, Src0: SrcReg, Src1: MaskReg);
833
834 if (NewSrcWidth == NewDstWidth) {
835 // Rekey OrigWidth from DstReg to MaskedReg so widenSignSensitiveOps
836 // still sees the narrow original width after replaceRegWith.
837 if (auto It = WideningInfo.OrigWidth.find(Val: DstReg);
838 It != WideningInfo.OrigWidth.end()) {
839 unsigned W = It->second;
840 WideningInfo.OrigWidth.erase(I: It);
841 WideningInfo.OrigWidth.try_emplace(Key: MaskedReg, Args&: W);
842 }
843 MRI.replaceRegWith(FromReg: DstReg, ToReg: MaskedReg);
844 TruncToRemove.push_back(Elt: &MI);
845 } else {
846 MI.getOperand(i: 1).setReg(MaskedReg);
847 }
848 }
849 }
850 for (MachineInstr *MI : TruncToRemove)
851 MI->eraseFromParent();
852
853 widenSignSensitiveOps(MF, GR, MIB, MRI, Info: WideningInfo);
854 widenBitCountOps(GR, MIB, MRI, Info: WideningInfo);
855 }
856
857 for (MachineBasicBlock *MBB : post_order(G: &MF)) {
858 if (MBB->empty())
859 continue;
860
861 bool ReachedBegin = false;
862 for (auto MII = std::prev(x: MBB->end()), Begin = MBB->begin();
863 !ReachedBegin;) {
864 MachineInstr &MI = *MII;
865 unsigned MIOp = MI.getOpcode();
866
867 if (!IsExtendedInts) {
868 // validate bit width of scalar registers and constant immediates
869 for (auto &MOP : MI.operands()) {
870 if (MOP.isReg())
871 widenScalarType(Reg: MOP.getReg(), MRI);
872 else if (MOP.isCImm())
873 widenCImmType(MOP);
874 }
875 }
876
877 if (isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_assign_ptr_type)) {
878 Register Reg = MI.getOperand(i: 1).getReg();
879 MIB.setInsertPt(MBB&: *MI.getParent(), II: MI.getIterator());
880 Type *ElementTy = getMDOperandAsType(N: MI.getOperand(i: 2).getMetadata(), I: 0);
881 auto SC = addressSpaceToStorageClass(AddrSpace: MI.getOperand(i: 3).getImm(), STI: *ST);
882 if (SC == SPIRV::StorageClass::Function &&
883 isa<FunctionType>(Val: ElementTy) &&
884 ST->canUseExtension(E: SPIRV::Extension::SPV_INTEL_function_pointers))
885 SC = SPIRV::StorageClass::CodeSectionINTEL;
886 SPIRVTypeInst AssignedPtrType =
887 GR->getOrCreateSPIRVPointerType(BaseType: ElementTy, I&: MI, SC);
888
889 // For untyped pointers, store the element type for later use.
890 if (ST->canUseExtension(E: SPIRV::Extension::SPV_KHR_untyped_pointers) &&
891 !ST->isShader()) {
892 SPIRVTypeInst ElemSpvType = GR->getOrCreateSPIRVType(
893 Type: ElementTy, MIRBuilder&: MIB, AQ: SPIRV::AccessQualifier::ReadWrite,
894 /*EmitIR=*/true);
895 GR->setUntypedPtrElementType(Reg, ElemType: ElemSpvType);
896 }
897
898 // The intrinsic also carries vector-of-pointer values produced by
899 // scalarized vector GEPs; wrap the pointer in OpTypeVector to match
900 // the vreg's LLT.
901 LLT RegTy = MRI.getType(Reg);
902 if (RegTy.isValid() && RegTy.isVector())
903 AssignedPtrType = GR->getOrCreateSPIRVVectorType(
904 BaseType: AssignedPtrType, NumElements: RegTy.getNumElements(), MIRBuilder&: MIB,
905 /*EmitIR=*/true);
906 MachineInstr *Def = MRI.getVRegDef(Reg);
907 assert(Def && "Expecting an instruction that defines the register");
908 // G_GLOBAL_VALUE already has type info.
909 if (Def->getOpcode() != TargetOpcode::G_GLOBAL_VALUE)
910 updateRegType(Reg, Ty: nullptr, SpvType: AssignedPtrType, GR, MIB,
911 MRI&: MF.getRegInfo());
912 ToErase.push_back(Elt: &MI);
913 } else if (isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_assign_type)) {
914 Register Reg = MI.getOperand(i: 1).getReg();
915 Type *Ty = getMDOperandAsType(N: MI.getOperand(i: 2).getMetadata(), I: 0);
916 MachineInstr *Def = MRI.getVRegDef(Reg);
917 assert(Def && "Expecting an instruction that defines the register");
918 // G_GLOBAL_VALUE already has type info.
919 if (Def->getOpcode() != TargetOpcode::G_GLOBAL_VALUE)
920 updateRegType(Reg, Ty, SpvType: nullptr, GR, MIB, MRI&: MF.getRegInfo());
921 if (Def->getOpcode() == TargetOpcode::COPY && isVector1(Ty))
922 updateRegType(Reg: passCopy(Def, MRI: &MF.getRegInfo())->getOperand(i: 0).getReg(),
923 Ty, SpvType: nullptr, GR, MIB, MRI&: MF.getRegInfo());
924 ToErase.push_back(Elt: &MI);
925 } else if (MIOp == TargetOpcode::FAKE_USE && MI.getNumOperands() > 0) {
926 MachineInstr *MdMI = MI.getPrevNode();
927 if (MdMI && isSpvIntrinsic(MI: *MdMI, IntrinsicID: Intrinsic::spv_value_md)) {
928 // It's an internal service info from before IRTranslator passes.
929 MachineInstr *Def = getVRegDef(MRI, Reg: MI.getOperand(i: 0).getReg());
930 for (unsigned I = 1, E = MI.getNumOperands(); I != E && Def; ++I)
931 if (getVRegDef(MRI, Reg: MI.getOperand(i: I).getReg()) != Def)
932 Def = nullptr;
933 if (Def) {
934 const MDNode *MD = MdMI->getOperand(i: 1).getMetadata();
935 StringRef ValueName =
936 cast<MDString>(Val: MD->getOperand(I: 1))->getString();
937 const MDNode *TypeMD = cast<MDNode>(Val: MD->getOperand(I: 0));
938 Type *ValueTy = getMDOperandAsType(N: TypeMD, I: 0);
939 GR->addValueAttrs(Key: Def, Val: std::make_pair(x&: ValueTy, y: ValueName.str()));
940 }
941 ToErase.push_back(Elt: MdMI);
942 }
943 ToErase.push_back(Elt: &MI);
944 } else if (MIOp == TargetOpcode::G_CONSTANT ||
945 MIOp == TargetOpcode::G_FCONSTANT ||
946 MIOp == TargetOpcode::G_BUILD_VECTOR) {
947 // %rc = G_CONSTANT ty Val
948 // Ensure %rc has a valid SPIR-V type assigned in the Global Registry.
949 Register Reg = MI.getOperand(i: 0).getReg();
950 bool NeedAssignType = !GR->getSPIRVTypeForVReg(VReg: Reg);
951 Type *Ty = nullptr;
952 if (MIOp == TargetOpcode::G_CONSTANT) {
953 auto TargetExtIt = TargetExtConstTypes.find(Val: &MI);
954 Ty = TargetExtIt == TargetExtConstTypes.end()
955 ? MI.getOperand(i: 1).getCImm()->getType()
956 : TargetExtIt->second;
957 const ConstantInt *OpCI = MI.getOperand(i: 1).getCImm();
958 // TODO: we may wish to analyze here if OpCI is zero and LLT RegType =
959 // MRI.getType(Reg); RegType.isPointer() is true, so that we observe
960 // at this point not i64/i32 constant but null pointer in the
961 // corresponding address space of RegType.getAddressSpace(). This may
962 // help to successfully validate the case when a OpConstantComposite's
963 // constituent has type that does not match Result Type of
964 // OpConstantComposite (see, for example,
965 // pointers/PtrCast-null-in-OpSpecConstantOp.ll).
966 Register PrimaryReg = GR->find(V: OpCI, MF: &MF);
967 if (!PrimaryReg.isValid()) {
968 GR->add(V: OpCI, MI: &MI);
969 } else if (PrimaryReg != Reg &&
970 MRI.getType(Reg) == MRI.getType(Reg: PrimaryReg)) {
971 auto *RCReg = MRI.getRegClassOrNull(Reg);
972 auto *RCPrimary = MRI.getRegClassOrNull(Reg: PrimaryReg);
973 if (!RCReg || RCPrimary == RCReg) {
974 RegsAlreadyAddedToDT[&MI] = PrimaryReg;
975 ToErase.push_back(Elt: &MI);
976 NeedAssignType = false;
977 }
978 }
979 } else if (MIOp == TargetOpcode::G_FCONSTANT) {
980 Ty = MI.getOperand(i: 1).getFPImm()->getType();
981 } else {
982 assert(MIOp == TargetOpcode::G_BUILD_VECTOR);
983 Type *ElemTy = nullptr;
984 MachineInstr *ElemMI = MRI.getVRegDef(Reg: MI.getOperand(i: 1).getReg());
985 assert(ElemMI);
986
987 if (ElemMI->getOpcode() == TargetOpcode::G_CONSTANT) {
988 ElemTy = ElemMI->getOperand(i: 1).getCImm()->getType();
989 } else if (ElemMI->getOpcode() == TargetOpcode::G_FCONSTANT) {
990 ElemTy = ElemMI->getOperand(i: 1).getFPImm()->getType();
991 } else {
992 if (SPIRVTypeInst ElemSpvType =
993 GR->getSPIRVTypeForVReg(VReg: MI.getOperand(i: 1).getReg(), MF: &MF))
994 ElemTy = const_cast<Type *>(GR->getTypeForSPIRVType(Ty: ElemSpvType));
995 }
996 if (ElemTy)
997 Ty = VectorType::get(
998 ElementType: ElemTy, NumElements: MI.getNumExplicitOperands() - MI.getNumExplicitDefs(),
999 Scalable: false);
1000 else
1001 NeedAssignType = false;
1002 }
1003 if (NeedAssignType)
1004 updateRegType(Reg, Ty, SpvType: nullptr, GR, MIB, MRI);
1005 } else if (MIOp == TargetOpcode::G_GLOBAL_VALUE) {
1006 propagateSPIRVType(MI: &MI, GR, MRI, MIB);
1007 }
1008
1009 if (MII == Begin)
1010 ReachedBegin = true;
1011 else
1012 --MII;
1013 }
1014 }
1015 for (MachineInstr *MI : ToErase) {
1016 auto It = RegsAlreadyAddedToDT.find(Val: MI);
1017 if (It != RegsAlreadyAddedToDT.end())
1018 MRI.replaceRegWith(FromReg: MI->getOperand(i: 0).getReg(), ToReg: It->second);
1019 invalidateAndEraseMI(GR, MI);
1020 }
1021
1022 // Address the case when IRTranslator introduces instructions with new
1023 // registers without associated SPIRV type.
1024 for (MachineBasicBlock &MBB : MF) {
1025 for (MachineInstr &MI : MBB) {
1026 switch (MI.getOpcode()) {
1027 case TargetOpcode::G_TRUNC:
1028 case TargetOpcode::G_ANYEXT:
1029 case TargetOpcode::G_SEXT:
1030 case TargetOpcode::G_ZEXT:
1031 case TargetOpcode::G_PTRTOINT:
1032 case TargetOpcode::COPY:
1033 case TargetOpcode::G_ADDRSPACE_CAST:
1034 propagateSPIRVType(MI: &MI, GR, MRI, MIB);
1035 break;
1036 }
1037 }
1038 }
1039}
1040
1041static void processInstrsWithTypeFolding(MachineFunction &MF,
1042 SPIRVGlobalRegistry *GR,
1043 MachineIRBuilder MIB) {
1044 MachineRegisterInfo &MRI = MF.getRegInfo();
1045 for (MachineBasicBlock &MBB : MF)
1046 for (MachineInstr &MI : MBB)
1047 if (isTypeFoldingSupported(Opcode: MI.getOpcode()))
1048 processInstr(MI, MIB, MRI, GR, KnownResType: nullptr);
1049}
1050
1051static Register
1052collectInlineAsmInstrOperands(MachineInstr *MI,
1053 SmallVector<unsigned, 4> *Ops = nullptr) {
1054 Register DefReg;
1055 unsigned StartOp = InlineAsm::MIOp_FirstOperand,
1056 AsmDescOp = InlineAsm::MIOp_FirstOperand;
1057 for (unsigned Idx = StartOp, MISz = MI->getNumOperands(); Idx != MISz;
1058 ++Idx) {
1059 const MachineOperand &MO = MI->getOperand(i: Idx);
1060 if (MO.isMetadata())
1061 continue;
1062 if (Idx == AsmDescOp && MO.isImm()) {
1063 // compute the index of the next operand descriptor
1064 const InlineAsm::Flag F(MO.getImm());
1065 AsmDescOp += 1 + F.getNumOperandRegisters();
1066 continue;
1067 }
1068 if (MO.isReg() && MO.isDef()) {
1069 if (!Ops)
1070 return MO.getReg();
1071 DefReg = MO.getReg();
1072 } else if (Ops) {
1073 Ops->push_back(Elt: Idx);
1074 }
1075 }
1076 return DefReg;
1077}
1078
1079static void
1080insertInlineAsmProcess(MachineFunction &MF, SPIRVGlobalRegistry *GR,
1081 const SPIRVSubtarget &ST, MachineIRBuilder MIRBuilder,
1082 const SmallVector<MachineInstr *> &ToProcess) {
1083 MachineRegisterInfo &MRI = MF.getRegInfo();
1084 Register AsmTargetReg;
1085 for (unsigned i = 0, Sz = ToProcess.size(); i + 1 < Sz; i += 2) {
1086 MachineInstr *I1 = ToProcess[i], *I2 = ToProcess[i + 1];
1087 assert(isSpvIntrinsic(*I1, Intrinsic::spv_inline_asm) && I2->isInlineAsm());
1088 MIRBuilder.setInsertPt(MBB&: *I2->getParent(), II: *I2);
1089
1090 if (!AsmTargetReg.isValid()) {
1091 // define vendor specific assembly target or dialect
1092 AsmTargetReg = MRI.createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: 32));
1093 MRI.setRegClass(Reg: AsmTargetReg, RC: &SPIRV::iIDRegClass);
1094 auto AsmTargetMIB =
1095 MIRBuilder.buildInstr(Opcode: SPIRV::OpAsmTargetINTEL).addDef(RegNo: AsmTargetReg);
1096 addStringImm(Str: ST.getTargetTripleAsStr(), MIB&: AsmTargetMIB);
1097 GR->add(Obj: AsmTargetMIB.getInstr(), MI: AsmTargetMIB);
1098 }
1099
1100 // create types
1101 const MDNode *IAMD = I1->getOperand(i: 1).getMetadata();
1102 FunctionType *FTy = cast<FunctionType>(Val: getMDOperandAsType(N: IAMD, I: 0));
1103 SmallVector<SPIRVTypeInst, 4> ArgTypes;
1104 for (const auto &ArgTy : FTy->params())
1105 ArgTypes.push_back(Elt: GR->getOrCreateSPIRVType(
1106 Type: ArgTy, MIRBuilder, AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true));
1107 SPIRVTypeInst RetType =
1108 GR->getOrCreateSPIRVType(Type: FTy->getReturnType(), MIRBuilder,
1109 AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
1110 SPIRVTypeInst FuncType = GR->getOrCreateOpTypeFunctionWithArgs(
1111 Ty: FTy, RetType, ArgTypes, MIRBuilder);
1112
1113 // define vendor specific assembly instructions string
1114 Register AsmReg = MRI.createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: 32));
1115 MRI.setRegClass(Reg: AsmReg, RC: &SPIRV::iIDRegClass);
1116 auto AsmMIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpAsmINTEL)
1117 .addDef(RegNo: AsmReg)
1118 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: RetType))
1119 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: FuncType))
1120 .addUse(RegNo: AsmTargetReg);
1121 // inline asm string:
1122 addStringImm(Str: I2->getOperand(i: InlineAsm::MIOp_AsmString).getSymbolName(),
1123 MIB&: AsmMIB);
1124 // inline asm constraint string:
1125 addStringImm(Str: cast<MDString>(Val: I1->getOperand(i: 2).getMetadata()->getOperand(I: 0))
1126 ->getString(),
1127 MIB&: AsmMIB);
1128 GR->add(Obj: AsmMIB.getInstr(), MI: AsmMIB);
1129
1130 // calls the inline assembly instruction
1131 unsigned ExtraInfo = I2->getOperand(i: InlineAsm::MIOp_ExtraInfo).getImm();
1132 if (ExtraInfo & InlineAsm::Extra_HasSideEffects)
1133 MIRBuilder.buildInstr(Opcode: SPIRV::OpDecorate)
1134 .addUse(RegNo: AsmReg)
1135 .addImm(Val: static_cast<uint32_t>(SPIRV::Decoration::SideEffectsINTEL));
1136
1137 Register DefReg = collectInlineAsmInstrOperands(MI: I2);
1138 if (!DefReg.isValid()) {
1139 DefReg = MRI.createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: 32));
1140 MRI.setRegClass(Reg: DefReg, RC: &SPIRV::iIDRegClass);
1141 SPIRVTypeInst VoidType = GR->getOrCreateSPIRVType(
1142 Type: Type::getVoidTy(C&: MF.getFunction().getContext()), MIRBuilder,
1143 AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
1144 GR->assignSPIRVTypeToVReg(Type: VoidType, VReg: DefReg, MF);
1145 }
1146
1147 auto AsmCall = MIRBuilder.buildInstr(Opcode: SPIRV::OpAsmCallINTEL)
1148 .addDef(RegNo: DefReg)
1149 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: RetType))
1150 .addUse(RegNo: AsmReg);
1151 for (unsigned IntrIdx = 3; IntrIdx < I1->getNumOperands(); ++IntrIdx)
1152 AsmCall.addUse(RegNo: I1->getOperand(i: IntrIdx).getReg());
1153
1154 // IRTranslator gets a bit confused when lowering inline ASM with outputs
1155 // and inserts a spurious COPY & TRUNC as registers are assumed to be i64;
1156 // we have to clean that up here to prevent erroneous trunc casts either on
1157 // a struct (for multiple outputs) or same width integers to get lowered
1158 // into SPIR-V
1159 if (MRI.hasOneUse(RegNo: DefReg)) {
1160 MachineInstr &CopyMI = *MRI.use_instr_begin(RegNo: DefReg);
1161 if (CopyMI.getOpcode() == TargetOpcode::COPY) {
1162 Register CopyDst = CopyMI.getOperand(i: 0).getReg();
1163 if (MRI.hasOneUse(RegNo: CopyDst)) {
1164 MachineInstr &TruncMI = *MRI.use_instr_begin(RegNo: CopyDst);
1165 if (TruncMI.getOpcode() == TargetOpcode::G_TRUNC) {
1166 MRI.setType(VReg: DefReg, Ty: GR->getRegType(SpvType: RetType));
1167 Register TruncReg = TruncMI.defs().begin()->getReg();
1168 MRI.replaceRegWith(FromReg: TruncReg, ToReg: DefReg);
1169 invalidateAndEraseMI(GR, MI: &TruncMI);
1170 invalidateAndEraseMI(GR, MI: &CopyMI);
1171 }
1172 }
1173 }
1174 }
1175 }
1176 for (MachineInstr *MI : ToProcess)
1177 invalidateAndEraseMI(GR, MI);
1178}
1179
1180static void insertInlineAsm(MachineFunction &MF, SPIRVGlobalRegistry *GR,
1181 const SPIRVSubtarget &ST,
1182 MachineIRBuilder MIRBuilder) {
1183 SmallVector<MachineInstr *> ToProcess;
1184 for (MachineBasicBlock &MBB : MF) {
1185 for (MachineInstr &MI : MBB) {
1186 if (isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_inline_asm) ||
1187 MI.getOpcode() == TargetOpcode::INLINEASM)
1188 ToProcess.push_back(Elt: &MI);
1189 }
1190 }
1191 if (ToProcess.size() == 0)
1192 return;
1193
1194 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_inline_assembly))
1195 report_fatal_error(reason: "Inline assembly instructions require the "
1196 "following SPIR-V extension: SPV_INTEL_inline_assembly",
1197 gen_crash_diag: false);
1198
1199 insertInlineAsmProcess(MF, GR, ST, MIRBuilder, ToProcess);
1200}
1201
1202static void insertSpirvDecorations(MachineFunction &MF, SPIRVGlobalRegistry *GR,
1203 MachineIRBuilder MIB) {
1204 const SPIRVSubtarget &ST = cast<SPIRVSubtarget>(Val: MIB.getMF().getSubtarget());
1205 SmallVector<MachineInstr *, 10> ToErase;
1206 for (MachineBasicBlock &MBB : MF) {
1207 for (MachineInstr &MI : MBB) {
1208 if (!isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_assign_decoration) &&
1209 !isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_assign_aliasing_decoration) &&
1210 !isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_assign_fpmaxerror_decoration))
1211 continue;
1212 MIB.setInsertPt(MBB&: *MI.getParent(), II: MI.getNextNode());
1213 if (isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_assign_decoration)) {
1214 buildOpSpirvDecorations(Reg: MI.getOperand(i: 1).getReg(), MIRBuilder&: MIB,
1215 GVarMD: MI.getOperand(i: 2).getMetadata(), ST);
1216 } else if (isSpvIntrinsic(MI,
1217 IntrinsicID: Intrinsic::spv_assign_fpmaxerror_decoration)) {
1218 ConstantFP *OpV = mdconst::dyn_extract<ConstantFP>(
1219 MD: MI.getOperand(i: 2).getMetadata()->getOperand(I: 0));
1220 uint32_t OpValue = OpV->getValueAPF().bitcastToAPInt().getZExtValue();
1221
1222 buildOpDecorate(Reg: MI.getOperand(i: 1).getReg(), MIRBuilder&: MIB,
1223 Dec: SPIRV::Decoration::FPMaxErrorDecorationINTEL,
1224 DecArgs: {OpValue});
1225 } else {
1226 GR->buildMemAliasingOpDecorate(Reg: MI.getOperand(i: 1).getReg(), MIRBuilder&: MIB,
1227 Dec: MI.getOperand(i: 2).getImm(),
1228 GVarMD: MI.getOperand(i: 3).getMetadata());
1229 }
1230
1231 ToErase.push_back(Elt: &MI);
1232 }
1233 }
1234 for (MachineInstr *MI : ToErase)
1235 invalidateAndEraseMI(GR, MI);
1236}
1237
1238// Returns the value of the switch case operand in Reg. The case value stays a
1239// G_CONSTANT until the module emits a SPIR-V constant for the same value, at
1240// which point the case register is replaced with the one defining that
1241// constant, which keeps its value in literal operands rather than in a CImm.
1242static const ConstantInt *getSwitchCaseValue(Register Reg,
1243 const MachineRegisterInfo &MRI,
1244 LLVMContext &Ctx) {
1245 APInt Val;
1246 if (mi_match(R: Reg, MRI, P: m_ICst(Cst&: Val)))
1247 return ConstantInt::get(Context&: Ctx, V: Val);
1248
1249 const MachineInstr *Def = nullptr;
1250 if (!mi_match(R: Reg, MRI, P: m_MInstr(MI&: Def)))
1251 llvm_unreachable("Switch case operand has no definition");
1252
1253 LLT Ty = MRI.getType(Reg);
1254 assert(Ty.isValid() && "Expected a typed switch case value");
1255 Val = APInt(Ty.getScalarSizeInBits(), 0);
1256
1257 switch (Def->getOpcode()) {
1258 case SPIRV::OpConstantNull:
1259 case SPIRV::OpConstantI:
1260 // The operands after the type are 32-bit literal words, least significant
1261 // first, as written by addNumImm(). OpConstantNull carries none, so it
1262 // decodes to zero without a case of its own.
1263 for (unsigned I = 2, E = Def->getNumExplicitOperands(); I != E; ++I) {
1264 uint32_t Word = static_cast<uint32_t>(Def->getOperand(i: I).getImm());
1265 Val |= APInt(Val.getBitWidth(), Word).shl(shiftAmt: (I - 2) * 32);
1266 }
1267 break;
1268 default:
1269 llvm_unreachable("Unexpected definition of a switch case value");
1270 }
1271 return ConstantInt::get(Context&: Ctx, V: Val);
1272}
1273
1274// LLVM allows the switches to use registers as cases, while SPIR-V required
1275// those to be immediate values. This function replaces such operands with the
1276// equivalent immediate constant.
1277static void processSwitchesConstants(MachineFunction &MF,
1278 SPIRVGlobalRegistry *GR,
1279 MachineIRBuilder MIB) {
1280 MachineRegisterInfo &MRI = MF.getRegInfo();
1281 LLVMContext &Ctx = MF.getFunction().getContext();
1282 for (MachineBasicBlock &MBB : MF) {
1283 for (MachineInstr &MI : MBB) {
1284 if (!isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_switch))
1285 continue;
1286
1287 SmallVector<MachineOperand, 8> NewOperands;
1288 NewOperands.push_back(Elt: MI.getOperand(i: 0)); // Opcode
1289 NewOperands.push_back(Elt: MI.getOperand(i: 1)); // Condition
1290 NewOperands.push_back(Elt: MI.getOperand(i: 2)); // Default
1291 for (unsigned i = 3; i < MI.getNumOperands(); i += 2) {
1292 Register Reg = MI.getOperand(i).getReg();
1293 NewOperands.push_back(
1294 Elt: MachineOperand::CreateCImm(CI: getSwitchCaseValue(Reg, MRI, Ctx)));
1295
1296 NewOperands.push_back(Elt: MI.getOperand(i: i + 1));
1297 }
1298
1299 assert(MI.getNumOperands() == NewOperands.size());
1300 while (MI.getNumOperands() > 0)
1301 MI.removeOperand(OpNo: 0);
1302 for (auto &MO : NewOperands)
1303 MI.addOperand(Op: MO);
1304 }
1305 }
1306}
1307
1308// Some instructions are used during CodeGen but should never be emitted.
1309// Cleaning up those.
1310static void cleanupHelperInstructions(MachineFunction &MF,
1311 SPIRVGlobalRegistry *GR) {
1312 SmallVector<MachineInstr *, 8> ToEraseMI;
1313 for (MachineBasicBlock &MBB : MF) {
1314 for (MachineInstr &MI : MBB) {
1315 if (isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_track_constant) ||
1316 MI.getOpcode() == TargetOpcode::G_BRINDIRECT)
1317 ToEraseMI.push_back(Elt: &MI);
1318 }
1319 }
1320
1321 for (MachineInstr *MI : ToEraseMI)
1322 invalidateAndEraseMI(GR, MI);
1323}
1324
1325// Find all usages of G_BLOCK_ADDR in our intrinsics and replace those
1326// operands/registers by the actual MBB it references.
1327static void processBlockAddr(MachineFunction &MF, SPIRVGlobalRegistry *GR,
1328 MachineIRBuilder MIB) {
1329 // Gather the reverse-mapping BB -> MBB.
1330 DenseMap<const BasicBlock *, MachineBasicBlock *> BB2MBB;
1331 for (MachineBasicBlock &MBB : MF)
1332 BB2MBB[MBB.getBasicBlock()] = &MBB;
1333
1334 // Gather instructions requiring patching. For now, only those can use
1335 // G_BLOCK_ADDR.
1336 SmallVector<MachineInstr *, 8> InstructionsToPatch;
1337 for (MachineBasicBlock &MBB : MF) {
1338 for (MachineInstr &MI : MBB) {
1339 if (isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_switch) ||
1340 isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_loop_merge) ||
1341 isSpvIntrinsic(MI, IntrinsicID: Intrinsic::spv_selection_merge))
1342 InstructionsToPatch.push_back(Elt: &MI);
1343 }
1344 }
1345
1346 // For each instruction to fix, we replace all the G_BLOCK_ADDR operands by
1347 // the actual MBB it references. Once those references have been updated, we
1348 // can cleanup remaining G_BLOCK_ADDR references.
1349 SmallPtrSet<MachineBasicBlock *, 8> ClearAddressTaken;
1350 SmallPtrSet<MachineInstr *, 8> ToEraseMI;
1351 MachineRegisterInfo &MRI = MF.getRegInfo();
1352 for (MachineInstr *MI : InstructionsToPatch) {
1353 SmallVector<MachineOperand, 8> NewOps;
1354 for (unsigned i = 0; i < MI->getNumOperands(); ++i) {
1355 // The operand is not a register, keep as-is.
1356 if (!MI->getOperand(i).isReg()) {
1357 NewOps.push_back(Elt: MI->getOperand(i));
1358 continue;
1359 }
1360
1361 Register Reg = MI->getOperand(i).getReg();
1362 MachineInstr *BuildMBB = MRI.getVRegDef(Reg);
1363 // The register is not the result of G_BLOCK_ADDR, keep as-is.
1364 if (!BuildMBB || BuildMBB->getOpcode() != TargetOpcode::G_BLOCK_ADDR) {
1365 NewOps.push_back(Elt: MI->getOperand(i));
1366 continue;
1367 }
1368
1369 assert(BuildMBB && BuildMBB->getOpcode() == TargetOpcode::G_BLOCK_ADDR &&
1370 BuildMBB->getOperand(1).isBlockAddress() &&
1371 BuildMBB->getOperand(1).getBlockAddress());
1372 BasicBlock *BB =
1373 BuildMBB->getOperand(i: 1).getBlockAddress()->getBasicBlock();
1374 auto It = BB2MBB.find(Val: BB);
1375 if (It == BB2MBB.end())
1376 report_fatal_error(reason: "cannot find a machine basic block by a basic block "
1377 "in a switch statement");
1378 MachineBasicBlock *ReferencedBlock = It->second;
1379 NewOps.push_back(Elt: MachineOperand::CreateMBB(MBB: ReferencedBlock));
1380
1381 ClearAddressTaken.insert(Ptr: ReferencedBlock);
1382 ToEraseMI.insert(Ptr: BuildMBB);
1383 }
1384
1385 // Replace the operands.
1386 assert(MI->getNumOperands() == NewOps.size());
1387 while (MI->getNumOperands() > 0)
1388 MI->removeOperand(OpNo: 0);
1389 for (auto &MO : NewOps)
1390 MI->addOperand(Op: MO);
1391
1392 if (MachineInstr *Next = MI->getNextNode()) {
1393 if (isSpvIntrinsic(MI: *Next, IntrinsicID: Intrinsic::spv_track_constant)) {
1394 ToEraseMI.insert(Ptr: Next);
1395 Next = MI->getNextNode();
1396 }
1397 if (Next && Next->getOpcode() == TargetOpcode::G_BRINDIRECT)
1398 ToEraseMI.insert(Ptr: Next);
1399 }
1400 }
1401
1402 // BlockAddress operands were used to keep information between passes,
1403 // let's undo the "address taken" status to reflect that Succ doesn't
1404 // actually correspond to an IR-level basic block.
1405 for (MachineBasicBlock *Succ : ClearAddressTaken)
1406 Succ->setAddressTakenIRBlock(nullptr);
1407
1408 // If we just delete G_BLOCK_ADDR instructions with BlockAddress operands,
1409 // this leaves their BasicBlock counterparts in a "address taken" status. This
1410 // would make AsmPrinter to generate a series of unneeded labels of a "Address
1411 // of block that was removed by CodeGen" kind. Let's first ensure that we
1412 // don't have a dangling BlockAddress constants by zapping the BlockAddress
1413 // nodes, and only after that proceed with erasing G_BLOCK_ADDR instructions.
1414 Constant *Replacement =
1415 ConstantInt::get(Ty: Type::getInt32Ty(C&: MF.getFunction().getContext()), V: 1);
1416 for (MachineInstr *BlockAddrI : ToEraseMI) {
1417 if (BlockAddrI->getOpcode() == TargetOpcode::G_BLOCK_ADDR) {
1418 BlockAddress *BA = const_cast<BlockAddress *>(
1419 BlockAddrI->getOperand(i: 1).getBlockAddress());
1420 BA->replaceAllUsesWith(
1421 V: ConstantExpr::getIntToPtr(C: Replacement, Ty: BA->getType()));
1422 BA->destroyConstant();
1423 }
1424 invalidateAndEraseMI(GR, MI: BlockAddrI);
1425 }
1426}
1427
1428static bool isImplicitFallthrough(MachineBasicBlock &MBB) {
1429 if (MBB.empty())
1430 return MBB.getNextNode() != nullptr;
1431
1432 // Branching SPIR-V intrinsics are not detected by this generic method.
1433 // Thus, we can only trust negative result.
1434 if (!MBB.canFallThrough())
1435 return false;
1436
1437 // Otherwise, we must manually check if we have a SPIR-V intrinsic which
1438 // prevent an implicit fallthrough.
1439 for (MachineBasicBlock::reverse_iterator It = MBB.rbegin(), E = MBB.rend();
1440 It != E; ++It) {
1441 if (isSpvIntrinsic(MI: *It, IntrinsicID: Intrinsic::spv_switch))
1442 return false;
1443 }
1444 return true;
1445}
1446
1447static void removeImplicitFallthroughs(MachineFunction &MF,
1448 MachineIRBuilder MIB) {
1449 // It is valid for MachineBasicBlocks to not finish with a branch instruction.
1450 // In such cases, they will simply fallthrough their immediate successor.
1451 for (MachineBasicBlock &MBB : MF) {
1452 if (!isImplicitFallthrough(MBB))
1453 continue;
1454
1455 assert(MBB.succ_size() == 1);
1456 MIB.setInsertPt(MBB, II: MBB.end());
1457 MIB.buildBr(Dest&: **MBB.successors().begin());
1458 }
1459}
1460
1461static bool runPreLegalizer(MachineFunction &MF) {
1462 // Initialize the type registry.
1463 const SPIRVSubtarget &ST = MF.getSubtarget<SPIRVSubtarget>();
1464 SPIRVGlobalRegistry *GR = ST.getSPIRVGlobalRegistry();
1465 GR->setCurrentFunc(MF);
1466 MachineIRBuilder MIB(MF);
1467 // a registry of target extension constants
1468 DenseMap<MachineInstr *, Type *> TargetExtConstTypes;
1469 // to keep record of tracked constants
1470 addConstantsToTrack(MF, GR, STI: ST, TargetExtConstTypes);
1471 foldConstantsIntoIntrinsics(MF, GR, MIB);
1472 insertBitcasts(MF, GR, MIB);
1473 generateAssignInstrs(MF, GR, MIB, TargetExtConstTypes);
1474
1475 processSwitchesConstants(MF, GR, MIB);
1476 processBlockAddr(MF, GR, MIB);
1477 cleanupHelperInstructions(MF, GR);
1478
1479 processInstrsWithTypeFolding(MF, GR, MIB);
1480 removeImplicitFallthroughs(MF, MIB);
1481 insertSpirvDecorations(MF, GR, MIB);
1482 insertInlineAsm(MF, GR, ST, MIRBuilder: MIB);
1483 lowerBitcasts(MF, GR, MIB);
1484
1485 return true;
1486}
1487
1488INITIALIZE_PASS(SPIRVPreLegalizerLegacy, DEBUG_TYPE, "SPIRV pre legalizer",
1489 false, false)
1490
1491char SPIRVPreLegalizerLegacy::ID = 0;
1492
1493FunctionPass *llvm::createSPIRVPreLegalizerLegacyPass() {
1494 return new SPIRVPreLegalizerLegacy();
1495}
1496
1497bool SPIRVPreLegalizerLegacy::runOnMachineFunction(MachineFunction &MF) {
1498 return runPreLegalizer(MF);
1499}
1500
1501PreservedAnalyses
1502SPIRVPreLegalizerPass::run(MachineFunction &MF,
1503 MachineFunctionAnalysisManager &MFAM) {
1504 bool Changed = runPreLegalizer(MF);
1505 if (!Changed)
1506 return PreservedAnalyses::all();
1507
1508 return getMachineFunctionPassPreservedAnalyses()
1509 .preserve<GISelValueTrackingAnalysis>();
1510}
1511