1//===- SPIRVLegalizerInfo.cpp --- SPIR-V Legalization Rules ------*- 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// This file implements the targeting of the Machinelegalizer class for SPIR-V.
10//
11//===----------------------------------------------------------------------===//
12
13#include "SPIRVLegalizerInfo.h"
14#include "SPIRV.h"
15#include "SPIRVGlobalRegistry.h"
16#include "SPIRVSubtarget.h"
17#include "SPIRVUtils.h"
18#include "llvm/CodeGen/GlobalISel/GenericMachineInstrs.h"
19#include "llvm/CodeGen/GlobalISel/LegalizerHelper.h"
20#include "llvm/CodeGen/GlobalISel/MachineIRBuilder.h"
21#include "llvm/CodeGen/MachineInstr.h"
22#include "llvm/CodeGen/MachineRegisterInfo.h"
23#include "llvm/CodeGen/TargetOpcodes.h"
24#include "llvm/IR/IntrinsicsSPIRV.h"
25#include "llvm/Support/Debug.h"
26#include "llvm/Support/MathExtras.h"
27
28using namespace llvm;
29using namespace llvm::LegalizeActions;
30using namespace llvm::LegalityPredicates;
31
32#define DEBUG_TYPE "spirv-legalizer"
33
34LegalityPredicate typeOfExtendedScalars(unsigned TypeIdx, bool IsExtendedInts) {
35 return [IsExtendedInts, TypeIdx](const LegalityQuery &Query) {
36 const LLT Ty = Query.Types[TypeIdx];
37 return IsExtendedInts && Ty.isValid() && Ty.isScalar();
38 };
39}
40
41LegalityPredicate typeOfLongVectors(unsigned TypeIdx, bool IsLongVecs) {
42 return [TypeIdx, IsLongVecs](const LegalityQuery &Query) {
43 const LLT Ty = Query.Types[TypeIdx];
44 return IsLongVecs && Ty.isValid() && Ty.isVector();
45 };
46}
47
48SPIRVLegalizerInfo::SPIRVLegalizerInfo(const SPIRVSubtarget &ST) {
49 using namespace TargetOpcode;
50
51 this->ST = &ST;
52 GR = ST.getSPIRVGlobalRegistry();
53
54 const LLT s1 = LLT::scalar(SizeInBits: 1);
55 const LLT s8 = LLT::scalar(SizeInBits: 8);
56 const LLT s16 = LLT::scalar(SizeInBits: 16);
57 const LLT s32 = LLT::scalar(SizeInBits: 32);
58 const LLT s64 = LLT::scalar(SizeInBits: 64);
59 const LLT s128 = LLT::scalar(SizeInBits: 128);
60
61 const LLT v16s64 = LLT::fixed_vector(NumElements: 16, ScalarSizeInBits: 64);
62 const LLT v16s32 = LLT::fixed_vector(NumElements: 16, ScalarSizeInBits: 32);
63 const LLT v16s16 = LLT::fixed_vector(NumElements: 16, ScalarSizeInBits: 16);
64 const LLT v16s8 = LLT::fixed_vector(NumElements: 16, ScalarSizeInBits: 8);
65 const LLT v16s1 = LLT::fixed_vector(NumElements: 16, ScalarSizeInBits: 1);
66
67 const LLT v8s64 = LLT::fixed_vector(NumElements: 8, ScalarSizeInBits: 64);
68 const LLT v8s32 = LLT::fixed_vector(NumElements: 8, ScalarSizeInBits: 32);
69 const LLT v8s16 = LLT::fixed_vector(NumElements: 8, ScalarSizeInBits: 16);
70 const LLT v8s8 = LLT::fixed_vector(NumElements: 8, ScalarSizeInBits: 8);
71 const LLT v8s1 = LLT::fixed_vector(NumElements: 8, ScalarSizeInBits: 1);
72
73 const LLT v4s64 = LLT::fixed_vector(NumElements: 4, ScalarSizeInBits: 64);
74 const LLT v4s32 = LLT::fixed_vector(NumElements: 4, ScalarSizeInBits: 32);
75 const LLT v4s16 = LLT::fixed_vector(NumElements: 4, ScalarSizeInBits: 16);
76 const LLT v4s8 = LLT::fixed_vector(NumElements: 4, ScalarSizeInBits: 8);
77 const LLT v4s1 = LLT::fixed_vector(NumElements: 4, ScalarSizeInBits: 1);
78
79 const LLT v3s64 = LLT::fixed_vector(NumElements: 3, ScalarSizeInBits: 64);
80 const LLT v3s32 = LLT::fixed_vector(NumElements: 3, ScalarSizeInBits: 32);
81 const LLT v3s16 = LLT::fixed_vector(NumElements: 3, ScalarSizeInBits: 16);
82 const LLT v3s8 = LLT::fixed_vector(NumElements: 3, ScalarSizeInBits: 8);
83 const LLT v3s1 = LLT::fixed_vector(NumElements: 3, ScalarSizeInBits: 1);
84
85 const LLT v2s64 = LLT::fixed_vector(NumElements: 2, ScalarSizeInBits: 64);
86 const LLT v2s32 = LLT::fixed_vector(NumElements: 2, ScalarSizeInBits: 32);
87 const LLT v2s16 = LLT::fixed_vector(NumElements: 2, ScalarSizeInBits: 16);
88 const LLT v2s8 = LLT::fixed_vector(NumElements: 2, ScalarSizeInBits: 8);
89 const LLT v2s1 = LLT::fixed_vector(NumElements: 2, ScalarSizeInBits: 1);
90
91 const unsigned PSize = ST.getPointerSize();
92 const LLT p0 = LLT::pointer(AddressSpace: 0, SizeInBits: PSize); // Function
93 const LLT p1 = LLT::pointer(AddressSpace: 1, SizeInBits: PSize); // CrossWorkgroup
94 const LLT p2 = LLT::pointer(AddressSpace: 2, SizeInBits: PSize); // UniformConstant
95 const LLT p3 = LLT::pointer(AddressSpace: 3, SizeInBits: PSize); // Workgroup
96 const LLT p4 = LLT::pointer(AddressSpace: 4, SizeInBits: PSize); // Generic
97 const LLT p5 =
98 LLT::pointer(AddressSpace: 5, SizeInBits: PSize); // Input, SPV_INTEL_usm_storage_classes (Device)
99 const LLT p6 = LLT::pointer(AddressSpace: 6, SizeInBits: PSize); // SPV_INTEL_usm_storage_classes (Host)
100 const LLT p7 = LLT::pointer(AddressSpace: 7, SizeInBits: PSize); // Input
101 const LLT p8 = LLT::pointer(AddressSpace: 8, SizeInBits: PSize); // Output
102 const LLT p9 =
103 LLT::pointer(AddressSpace: 9, SizeInBits: PSize); // CodeSectionINTEL, SPV_INTEL_function_pointers
104 const LLT p10 = LLT::pointer(AddressSpace: 10, SizeInBits: PSize); // Private
105 const LLT p11 = LLT::pointer(AddressSpace: 11, SizeInBits: PSize); // StorageBuffer
106 const LLT p12 = LLT::pointer(AddressSpace: 12, SizeInBits: PSize); // Uniform
107 const LLT p13 = LLT::pointer(AddressSpace: 13, SizeInBits: PSize); // PushConstant
108
109 // TODO: remove copy-pasting here by using concatenation in some way.
110 auto allPtrsScalarsAndVectors = {
111 p0, p1, p2, p3, p4, p5, p6, p7, p8,
112 p9, p10, p11, p12, p13, s1, s8, s16, s32,
113 s64, s128, v2s1, v2s8, v2s16, v2s32, v2s64, v3s1, v3s8,
114 v3s16, v3s32, v3s64, v4s1, v4s8, v4s16, v4s32, v4s64, v8s1,
115 v8s8, v8s16, v8s32, v8s64, v16s1, v16s8, v16s16, v16s32, v16s64};
116
117 auto allVectors = {v2s1, v2s8, v2s16, v2s32, v2s64, v3s1, v3s8,
118 v3s16, v3s32, v3s64, v4s1, v4s8, v4s16, v4s32,
119 v4s64, v8s1, v8s8, v8s16, v8s32, v8s64, v16s1,
120 v16s8, v16s16, v16s32, v16s64};
121
122 auto allShaderVectors = {v2s1, v2s8, v2s16, v2s32, v2s64,
123 v3s1, v3s8, v3s16, v3s32, v3s64,
124 v4s1, v4s8, v4s16, v4s32, v4s64};
125
126 auto allScalars = {s1, s8, s16, s32, s64};
127
128 auto allScalarsAndVectors = {
129 s1, s8, s16, s32, s64, s128, v2s1, v2s8,
130 v2s16, v2s32, v2s64, v3s1, v3s8, v3s16, v3s32, v3s64,
131 v4s1, v4s8, v4s16, v4s32, v4s64, v8s1, v8s8, v8s16,
132 v8s32, v8s64, v16s1, v16s8, v16s16, v16s32, v16s64};
133
134 auto allShaderScalarsAndVectors = {
135 s1, s8, s16, s32, s64, s128, v2s1, v2s8, v2s16, v2s32, v2s64,
136 v3s1, v3s8, v3s16, v3s32, v3s64, v4s1, v4s8, v4s16, v4s32, v4s64};
137
138 auto &allowedScalarsAndVectors =
139 ST.isShader() ? allShaderScalarsAndVectors : allScalarsAndVectors;
140
141 auto allIntScalarsAndVectors = {
142 s8, s16, s32, s64, s128, v2s8, v2s16, v2s32, v2s64,
143 v3s8, v3s16, v3s32, v3s64, v4s8, v4s16, v4s32, v4s64, v8s8,
144 v8s16, v8s32, v8s64, v16s8, v16s16, v16s32, v16s64};
145
146 auto allBoolScalarsAndVectors = {s1, v2s1, v3s1, v4s1, v8s1, v16s1};
147 auto allBoolVectors = {v2s1, v3s1, v4s1, v8s1, v16s1};
148
149 auto allIntScalars = {s8, s16, s32, s64, s128};
150
151 auto allShaderIntVectors = {v2s8, v2s16, v2s32, v2s64, v3s8, v3s16,
152 v3s32, v3s64, v4s8, v4s16, v4s32, v4s64};
153
154 auto allIntVectors = {v2s8, v2s16, v2s32, v2s64, v3s8, v3s16, v3s32,
155 v3s64, v4s8, v4s16, v4s32, v4s64, v8s8, v8s16,
156 v8s32, v8s64, v16s8, v16s16, v16s32, v16s64};
157
158 auto &allowedIntVectorTypes =
159 ST.isShader() ? allShaderIntVectors : allIntVectors;
160
161 auto allFloatScalarsAndF16Vector2AndVector4s = {s16, s32, s64, v2s16, v4s16};
162
163 auto allFloatScalars = {s16, s32, s64};
164
165 auto allFloatScalarsAndVectors = {
166 s16, s32, s64, v2s16, v2s32, v2s64, v3s16, v3s32, v3s64,
167 v4s16, v4s32, v4s64, v8s16, v8s32, v8s64, v16s16, v16s32, v16s64};
168
169 auto allShaderFloatVectors = {v2s16, v2s32, v2s64, v3s16, v3s32,
170 v3s64, v4s16, v4s32, v4s64};
171
172 auto allFloatVectors = {v2s16, v2s32, v2s64, v3s16, v3s32,
173 v3s64, v4s16, v4s32, v4s64, v8s16,
174 v8s32, v8s64, v16s16, v16s32, v16s64};
175
176 auto &allowedFloatVectorTypes =
177 ST.isShader() ? allShaderFloatVectors : allFloatVectors;
178
179 auto allFloatAndIntScalarsAndPtrs = {s8, s16, s32, s64, p0, p1,
180 p2, p3, p4, p5, p6, p7,
181 p8, p9, p10, p11, p12, p13};
182
183 auto allPtrs = {p0, p1, p2, p3, p4, p5, p6, p7, p8, p9, p10, p11, p12, p13};
184
185 auto &allowedVectorTypes = ST.isShader() ? allShaderVectors : allVectors;
186
187 bool HasArbitraryPrecisionInts = ST.canUseExtension(
188 E: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_integers);
189 bool IsExtendedInts =
190 HasArbitraryPrecisionInts ||
191 ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_bit_instructions) ||
192 ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_int4);
193 bool IsLongVecs = ST.canUseExtension(E: SPIRV::Extension::SPV_EXT_long_vector);
194 auto ExtendedIntScalarsAndVectors =
195 [IsExtendedInts](const LegalityQuery &Query) {
196 const LLT Ty = Query.Types[0];
197 return IsExtendedInts && Ty.isValid() &&
198 !Ty.isPointerOrPointerVector() && Ty.getScalarSizeInBits() > 1;
199 };
200 auto ExtendedScalarsAndVectorsProduct = [IsExtendedInts](
201 const LegalityQuery &Query) {
202 const LLT Ty1 = Query.Types[0], Ty2 = Query.Types[1];
203 return IsExtendedInts && Ty1.isValid() && Ty2.isValid() &&
204 !Ty1.isPointerOrPointerVector() && !Ty2.isPointerOrPointerVector();
205 };
206 auto ExtendedPtrsScalarsAndVectors =
207 [IsExtendedInts](const LegalityQuery &Query) {
208 const LLT Ty = Query.Types[0];
209 return IsExtendedInts && Ty.isValid();
210 };
211
212 // The universal validation rules in the SPIR-V specification state that
213 // vector sizes are typically limited to 2, 3, or 4. However, larger vector
214 // sizes (8 and 16) are enabled when the Kernel capability is present. For
215 // shader execution models, vector sizes are strictly limited to 4. In
216 // non-shader contexts, vector sizes of 8 and 16 are also permitted, but
217 // arbitrary sizes (e.g., 6 or 11) are not.
218 uint32_t MaxVectorSize = ST.isShader() ? 4 : 16;
219 LLVM_DEBUG(dbgs() << "MaxVectorSize: " << MaxVectorSize << "\n");
220
221 for (auto Opc : getTypeFoldingSupportedOpcodes()) {
222 switch (Opc) {
223 case G_EXTRACT_VECTOR_ELT:
224 case G_UREM:
225 case G_SREM:
226 case G_UDIV:
227 case G_SDIV:
228 case G_FREM:
229 case G_SELECT:
230 break;
231 default:
232 getActionDefinitionsBuilder(Opcode: Opc)
233 .customFor(Types: allScalars)
234 .customFor(Types: allowedVectorTypes)
235 .customIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
236 .moreElementsToNextPow2(TypeIdx: 0)
237 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
238 Mutation: LegalizeMutations::changeElementCountTo(
239 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
240 .custom();
241 break;
242 }
243 }
244
245 getActionDefinitionsBuilder(Opcodes: {G_UREM, G_SREM, G_SDIV, G_UDIV, G_FREM})
246 .customFor(Types: allScalars)
247 .customFor(Types: allowedVectorTypes)
248 .customIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
249 .scalarizeIf(Predicate: numElementsNotPow2(TypeIdx: 0), TypeIdx: 0)
250 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
251 Mutation: LegalizeMutations::changeElementCountTo(
252 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
253 .custom();
254
255 getActionDefinitionsBuilder(Opcode: G_SELECT)
256 .customFor(Types: allScalars)
257 .customFor(Types: allowedVectorTypes)
258 .customIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
259 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
260 Mutation: LegalizeMutations::changeElementCountTo(
261 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
262 .custom();
263
264 getActionDefinitionsBuilder(Opcodes: {G_FMA, G_STRICT_FMA})
265 .legalFor(Types: allScalars)
266 .legalFor(Types: allowedVectorTypes)
267 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
268 .moreElementsToNextPow2(TypeIdx: 0)
269 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
270 Mutation: LegalizeMutations::changeElementCountTo(
271 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
272 .alwaysLegal();
273
274 getActionDefinitionsBuilder(Opcode: G_INTRINSIC_W_SIDE_EFFECTS).custom();
275
276 getActionDefinitionsBuilder(Opcode: G_SHUFFLE_VECTOR)
277 .legalForCartesianProduct(Types0: allowedVectorTypes, Types1: allowedVectorTypes)
278 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
279 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 1, IsLongVecs))
280 .moreElementsToNextPow2(TypeIdx: 0)
281 .lowerIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize))
282 .moreElementsToNextPow2(TypeIdx: 1)
283 .lowerIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 1, Size: MaxVectorSize));
284
285 getActionDefinitionsBuilder(Opcode: G_EXTRACT_VECTOR_ELT)
286 .customIf(Predicate: typeOfLongVectors(TypeIdx: 1, IsLongVecs))
287 .moreElementsToNextPow2(TypeIdx: 1)
288 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 1, Size: MaxVectorSize),
289 Mutation: LegalizeMutations::changeElementCountTo(
290 TypeIdx: 1, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
291 .custom();
292
293 getActionDefinitionsBuilder(Opcode: G_INSERT_VECTOR_ELT)
294 .customIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
295 .moreElementsToNextPow2(TypeIdx: 0)
296 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
297 Mutation: LegalizeMutations::changeElementCountTo(
298 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
299 .custom();
300
301 // Illegal G_UNMERGE_VALUES instructions should be handled
302 // during the combine phase.
303 getActionDefinitionsBuilder(Opcode: G_BUILD_VECTOR)
304 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
305 .legalIf(Predicate: vectorElementCountIsLessThanOrEqualTo(TypeIdx: 0, Size: MaxVectorSize))
306 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
307 Mutation: LegalizeMutations::changeElementCountTo(
308 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)));
309
310 // When entering the legalizer, there should be no G_BITCAST instructions.
311 // They should all be calls to the `spv_bitcast` intrinsic. The call to
312 // the intrinsic will be converted to a G_BITCAST during legalization if
313 // the vectors are not legal. After using the rules to legalize a G_BITCAST,
314 // we turn it back into a call to the intrinsic with a custom rule to avoid
315 // potential machine verifier failures.
316 getActionDefinitionsBuilder(Opcode: G_BITCAST)
317 .customIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
318 .moreElementsToNextPow2(TypeIdx: 0)
319 .moreElementsToNextPow2(TypeIdx: 1)
320 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
321 Mutation: LegalizeMutations::changeElementCountTo(
322 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
323 .lowerIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 1, Size: MaxVectorSize))
324 .custom();
325
326 // If the result is still illegal, the combiner should be able to remove it.
327 getActionDefinitionsBuilder(Opcode: G_CONCAT_VECTORS)
328 .legalForCartesianProduct(Types0: allowedVectorTypes, Types1: allowedVectorTypes)
329 .legalIf(Predicate: LegalityPredicates::any(P0: typeOfLongVectors(TypeIdx: 0, IsLongVecs),
330 P1: typeOfLongVectors(TypeIdx: 1, IsLongVecs)));
331
332 getActionDefinitionsBuilder(Opcode: G_SPLAT_VECTOR)
333 .legalFor(Types: allowedVectorTypes)
334 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
335 .moreElementsToNextPow2(TypeIdx: 0)
336 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
337 Mutation: LegalizeMutations::changeElementSizeTo(TypeIdx: 0, FromTypeIdx: MaxVectorSize))
338 .alwaysLegal();
339
340 // Vector Reduction Operations
341 getActionDefinitionsBuilder(
342 Opcodes: {G_VECREDUCE_SMIN, G_VECREDUCE_SMAX, G_VECREDUCE_UMIN, G_VECREDUCE_UMAX,
343 G_VECREDUCE_ADD, G_VECREDUCE_MUL, G_VECREDUCE_FMUL, G_VECREDUCE_FMIN,
344 G_VECREDUCE_FMAX, G_VECREDUCE_FMINIMUM, G_VECREDUCE_FMAXIMUM,
345 G_VECREDUCE_OR, G_VECREDUCE_AND, G_VECREDUCE_XOR})
346 .legalFor(Types: allowedVectorTypes)
347 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
348 .scalarize(TypeIdx: 1)
349 .lower();
350
351 getActionDefinitionsBuilder(Opcodes: {G_VECREDUCE_SEQ_FADD, G_VECREDUCE_SEQ_FMUL})
352 .scalarize(TypeIdx: 2)
353 .lower();
354
355 // Illegal G_UNMERGE_VALUES instructions should be handled
356 // during the combine phase.
357 getActionDefinitionsBuilder(Opcode: G_UNMERGE_VALUES)
358 .legalIf(Predicate: LegalityPredicates::any(P0: typeOfLongVectors(TypeIdx: 0, IsLongVecs),
359 P1: typeOfLongVectors(TypeIdx: 1, IsLongVecs)))
360 .legalIf(Predicate: vectorElementCountIsLessThanOrEqualTo(TypeIdx: 1, Size: MaxVectorSize));
361
362 getActionDefinitionsBuilder(Opcodes: {G_MEMCPY, G_MEMCPY_INLINE, G_MEMMOVE})
363 .unsupportedIf(Predicate: LegalityPredicates::any(P0: typeIs(TypeIdx: 0, TypesInit: p9), P1: typeIs(TypeIdx: 1, TypesInit: p9)))
364 .legalIf(Predicate: all(P0: typeInSet(TypeIdx: 0, TypesInit: allPtrs), P1: typeInSet(TypeIdx: 1, TypesInit: allPtrs)));
365
366 getActionDefinitionsBuilder(Opcodes: {G_MEMSET, G_MEMSET_INLINE})
367 .unsupportedIf(Predicate: typeIs(TypeIdx: 0, TypesInit: p9))
368 .legalIf(Predicate: all(P0: typeInSet(TypeIdx: 0, TypesInit: allPtrs), P1: typeInSet(TypeIdx: 1, TypesInit: allIntScalars)));
369
370 getActionDefinitionsBuilder(Opcode: G_ADDRSPACE_CAST)
371 .legalForCartesianProduct(Types0: allPtrs, Types1: allPtrs);
372
373 // Should we be legalizing bad scalar sizes like s5 here instead
374 // of handling them in the instruction selector?
375 getActionDefinitionsBuilder(Opcodes: {G_LOAD, G_STORE})
376 .unsupportedIf(Predicate: typeIs(TypeIdx: 1, TypesInit: p9))
377 .legalForCartesianProduct(Types0: allowedVectorTypes, Types1: allPtrs)
378 .legalForCartesianProduct(Types0: allPtrs, Types1: allPtrs)
379 .legalIf(Predicate: isScalar(TypeIdx: 0))
380 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
381 .custom();
382
383 getActionDefinitionsBuilder(Opcodes: {G_SMIN, G_SMAX, G_UMIN, G_UMAX, G_ABS,
384 G_BITREVERSE, G_SADDSAT, G_UADDSAT, G_SSUBSAT,
385 G_USUBSAT, G_SCMP, G_UCMP})
386 .legalFor(Types: allIntScalars)
387 .legalFor(Types: allowedIntVectorTypes)
388 .legalIf(Predicate: ExtendedIntScalarsAndVectors)
389 // LLVM i1 maps to OpTypeBool, not OpTypeInt.
390 .scalarizeIf(Predicate: typeInSet(TypeIdx: 0, TypesInit: allBoolVectors), TypeIdx: 0)
391 .minScalar(TypeIdx: 0, Ty: s32)
392 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
393 Mutation: LegalizeMutations::changeElementCountTo(
394 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
395 .moreElementsToNextPow2(TypeIdx: 0);
396
397 getActionDefinitionsBuilder(Opcodes: {G_SSHLSAT, G_USHLSAT}).lower();
398
399 getActionDefinitionsBuilder(Opcodes: {G_FLDEXP, G_STRICT_FLDEXP})
400 .legalForCartesianProduct(Types0: allFloatScalarsAndVectors, Types1: allIntScalars);
401
402 getActionDefinitionsBuilder(Opcodes: {G_FPTOSI, G_FPTOUI})
403 .legalForCartesianProduct(Types0: allIntScalarsAndVectors,
404 Types1: allFloatScalarsAndVectors);
405
406 getActionDefinitionsBuilder(Opcodes: {G_FPTOSI_SAT, G_FPTOUI_SAT})
407 .legalForCartesianProduct(Types0: allIntScalarsAndVectors,
408 Types1: allFloatScalarsAndVectors);
409
410 getActionDefinitionsBuilder(Opcodes: {G_SITOFP, G_UITOFP})
411 .legalForCartesianProduct(Types0: allFloatScalarsAndVectors,
412 Types1: allScalarsAndVectors);
413
414 getActionDefinitionsBuilder(Opcode: G_CTPOP)
415 .legalForCartesianProduct(Types: allIntScalarsAndVectors)
416 .legalIf(Predicate: ExtendedScalarsAndVectorsProduct)
417 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs));
418
419 getActionDefinitionsBuilder(Opcodes: {G_TRUNC, G_ZEXT, G_SEXT, G_ANYEXT})
420 .legalForCartesianProduct(Types: allowedScalarsAndVectors)
421 .legalIf(Predicate: ExtendedScalarsAndVectorsProduct)
422 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
423 .moreElementsToNextPow2(TypeIdx: 0)
424 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
425 Mutation: LegalizeMutations::changeElementCountTo(
426 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)));
427
428 getActionDefinitionsBuilder(Opcode: G_SEXT_INREG)
429 .lowerIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
430 .moreElementsToNextPow2(TypeIdx: 0)
431 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
432 Mutation: LegalizeMutations::changeElementCountTo(
433 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
434 .lower();
435
436 getActionDefinitionsBuilder(Opcode: G_PHI)
437 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
438 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
439 Mutation: LegalizeMutations::changeElementCountTo(
440 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
441 .legalFor(Types: allPtrsScalarsAndVectors)
442 .legalIf(Predicate: ExtendedPtrsScalarsAndVectors)
443 .moreElementsToNextPow2(TypeIdx: 0);
444
445 getActionDefinitionsBuilder(Opcode: G_BITCAST).legalIf(
446 Predicate: all(P0: typeInSet(TypeIdx: 0, TypesInit: allPtrsScalarsAndVectors),
447 P1: typeInSet(TypeIdx: 1, TypesInit: allPtrsScalarsAndVectors)));
448
449 getActionDefinitionsBuilder(Opcodes: {G_IMPLICIT_DEF, G_FREEZE})
450 .legalFor(Types: {s1, s128})
451 .legalFor(Types: allFloatAndIntScalarsAndPtrs)
452 .legalFor(Types: allowedVectorTypes)
453 .legalIf(Predicate: [](const LegalityQuery &Query) {
454 return Query.Types[0].isPointerVector();
455 })
456 .legalIf(Predicate: typeOfLongVectors(TypeIdx: 0, IsLongVecs))
457 .moreElementsToNextPow2(TypeIdx: 0)
458 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
459 Mutation: LegalizeMutations::changeElementCountTo(
460 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)));
461
462 getActionDefinitionsBuilder(Opcodes: {G_STACKSAVE, G_STACKRESTORE}).alwaysLegal();
463
464 getActionDefinitionsBuilder(Opcode: G_INTTOPTR)
465 .legalForCartesianProduct(Types0: allPtrs, Types1: allIntScalars)
466 .legalIf(
467 Predicate: all(P0: typeInSet(TypeIdx: 0, TypesInit: allPtrs), P1: typeOfExtendedScalars(TypeIdx: 1, IsExtendedInts)))
468 .legalIf(Predicate: [](const LegalityQuery &Query) {
469 const LLT DstTy = Query.Types[0];
470 const LLT SrcTy = Query.Types[1];
471 return DstTy.isPointerVector() && SrcTy.isVector() &&
472 !SrcTy.isPointer() &&
473 DstTy.getNumElements() == SrcTy.getNumElements();
474 });
475 getActionDefinitionsBuilder(Opcode: G_PTRTOINT)
476 .legalForCartesianProduct(Types0: allIntScalars, Types1: allPtrs)
477 .legalIf(
478 Predicate: all(P0: typeOfExtendedScalars(TypeIdx: 0, IsExtendedInts), P1: typeInSet(TypeIdx: 1, TypesInit: allPtrs)))
479 .legalIf(Predicate: [](const LegalityQuery &Query) {
480 const LLT DstTy = Query.Types[0];
481 const LLT SrcTy = Query.Types[1];
482 return SrcTy.isPointerVector() && DstTy.isVector() &&
483 !DstTy.isPointer() &&
484 DstTy.getNumElements() == SrcTy.getNumElements();
485 });
486 getActionDefinitionsBuilder(Opcode: G_PTR_ADD)
487 .legalForCartesianProduct(Types0: allPtrs, Types1: allIntScalars)
488 .legalIf(
489 Predicate: all(P0: typeInSet(TypeIdx: 0, TypesInit: allPtrs), P1: typeOfExtendedScalars(TypeIdx: 1, IsExtendedInts)));
490
491 getActionDefinitionsBuilder(Opcode: G_PTRMASK)
492 .legalForCartesianProduct(Types0: allPtrs, Types1: allIntScalars)
493 .legalIf(
494 Predicate: all(P0: typeInSet(TypeIdx: 0, TypesInit: allPtrs), P1: typeOfExtendedScalars(TypeIdx: 1, IsExtendedInts)))
495 .legalIf(Predicate: [](const LegalityQuery &Query) {
496 const LLT PtrTy = Query.Types[0];
497 const LLT MaskTy = Query.Types[1];
498 return PtrTy.isPointerVector() && MaskTy.isVector() &&
499 !MaskTy.isPointer() &&
500 PtrTy.getNumElements() == MaskTy.getNumElements();
501 });
502
503 // ST.canDirectlyComparePointers() for pointer args is supported in
504 // legalizeCustom().
505 getActionDefinitionsBuilder(Opcode: G_ICMP)
506 .unsupportedIf(Predicate: LegalityPredicates::any(
507 P0: all(P0: typeIs(TypeIdx: 0, TypesInit: p9), P1: typeInSet(TypeIdx: 1, TypesInit: allPtrs), args: typeIsNot(TypeIdx: 1, Type: p9)),
508 P1: all(P0: typeInSet(TypeIdx: 0, TypesInit: allPtrs), P1: typeIsNot(TypeIdx: 0, Type: p9), args: typeIs(TypeIdx: 1, TypesInit: p9))))
509 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 1, Size: MaxVectorSize),
510 Mutation: LegalizeMutations::changeElementCountTo(
511 TypeIdx: 1, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
512 .legalIf(Predicate: [IsExtendedInts](const LegalityQuery &Query) {
513 const LLT Ty = Query.Types[1];
514 return IsExtendedInts && Ty.isValid() && !Ty.isPointerOrPointerVector();
515 })
516 .customIf(Predicate: all(P0: typeInSet(TypeIdx: 0, TypesInit: allBoolScalarsAndVectors),
517 P1: typeInSet(TypeIdx: 1, TypesInit: allPtrsScalarsAndVectors)));
518
519 getActionDefinitionsBuilder(Opcode: G_FCMP)
520 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 1, Size: MaxVectorSize),
521 Mutation: LegalizeMutations::changeElementCountTo(
522 TypeIdx: 1, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
523 .legalIf(Predicate: all(P0: typeInSet(TypeIdx: 0, TypesInit: allBoolScalarsAndVectors),
524 P1: typeInSet(TypeIdx: 1, TypesInit: allFloatScalarsAndVectors)));
525
526 getActionDefinitionsBuilder(Opcodes: {G_ATOMICRMW_OR, G_ATOMICRMW_ADD, G_ATOMICRMW_AND,
527 G_ATOMICRMW_MAX, G_ATOMICRMW_MIN,
528 G_ATOMICRMW_SUB, G_ATOMICRMW_XOR,
529 G_ATOMICRMW_UMAX, G_ATOMICRMW_UMIN})
530 .legalForCartesianProduct(Types0: allIntScalars, Types1: allPtrs);
531
532 getActionDefinitionsBuilder(
533 Opcodes: {G_ATOMICRMW_FADD, G_ATOMICRMW_FSUB, G_ATOMICRMW_FMIN, G_ATOMICRMW_FMAX})
534 .legalForCartesianProduct(Types0: allFloatScalarsAndF16Vector2AndVector4s,
535 Types1: allPtrs);
536
537 getActionDefinitionsBuilder(Opcode: G_ATOMICRMW_XCHG)
538 .legalForCartesianProduct(Types0: allFloatAndIntScalarsAndPtrs, Types1: allPtrs);
539
540 getActionDefinitionsBuilder(Opcode: G_ATOMIC_CMPXCHG_WITH_SUCCESS).lower();
541 // TODO: add proper legalization rules.
542 getActionDefinitionsBuilder(Opcode: G_ATOMIC_CMPXCHG).alwaysLegal();
543 getActionDefinitionsBuilder(Opcode: G_PREFETCH).alwaysLegal();
544
545 getActionDefinitionsBuilder(Opcodes: {G_UADDO, G_USUBO, G_UMULO, G_SMULO})
546 .alwaysLegal();
547
548 getActionDefinitionsBuilder(Opcodes: {G_SADDO, G_SSUBO}).lower();
549
550 // Lowering widens s64 to s128, which needs
551 // SPV_ALTERA_arbitrary_precision_integers. Mark s64 unsupported otherwise.
552 auto &MulFix = getActionDefinitionsBuilder(Opcodes: {G_SMULFIX, G_UMULFIX});
553 if (!HasArbitraryPrecisionInts)
554 MulFix.unsupportedFor(Types: {s64});
555 MulFix.lower();
556
557 getActionDefinitionsBuilder(Opcodes: {G_LROUND, G_LLROUND})
558 .legalForCartesianProduct(Types0: allIntScalarsAndVectors,
559 Types1: allFloatScalarsAndVectors);
560
561 // FP conversions.
562 getActionDefinitionsBuilder(Opcodes: {G_FPTRUNC, G_FPEXT})
563 .legalForCartesianProduct(Types: allFloatScalarsAndVectors);
564
565 // Pointer-handling.
566 getActionDefinitionsBuilder(Opcode: G_FRAME_INDEX).legalFor(Types: {p0});
567
568 getActionDefinitionsBuilder(Opcode: G_GLOBAL_VALUE).legalFor(Types: allPtrs);
569
570 // Control-flow. In some cases (e.g. constants) s1 may be promoted to s32.
571 getActionDefinitionsBuilder(Opcode: G_BR).alwaysLegal();
572 getActionDefinitionsBuilder(Opcode: G_BRCOND).legalFor(Types: {s1, s32});
573
574 getActionDefinitionsBuilder(Opcode: G_FFREXP).legalForCartesianProduct(
575 Types0: allFloatScalarsAndVectors, Types1: {s32, v2s32, v3s32, v4s32, v8s32, v16s32});
576
577 // TODO: Review the target OpenCL and GLSL Extended Instruction Set specs to
578 // tighten these requirements. Many of these math functions are only legal on
579 // specific bitwidths, so they are not selectable for
580 // allFloatScalarsAndVectors.
581 // clang-format off
582 getActionDefinitionsBuilder(Opcodes: {G_STRICT_FSQRT,
583 G_FPOW,
584 G_FEXP,
585 G_FMODF,
586 G_FSINCOS,
587 G_FEXP2,
588 G_FEXP10,
589 G_FLOG,
590 G_FLOG2,
591 G_FLOG10,
592 G_FABS,
593 G_FMINNUM,
594 G_FMAXNUM,
595 G_FCEIL,
596 G_FCOS,
597 G_FSIN,
598 G_FTAN,
599 G_FACOS,
600 G_FASIN,
601 G_FATAN,
602 G_FATAN2,
603 G_FCOSH,
604 G_FSINH,
605 G_FTANH,
606 G_FSQRT,
607 G_FFLOOR,
608 G_FRINT,
609 G_FNEARBYINT,
610 G_INTRINSIC_ROUND,
611 G_INTRINSIC_TRUNC,
612 G_FMINIMUM,
613 G_FMAXIMUM,
614 G_INTRINSIC_ROUNDEVEN})
615 .legalFor(Types: allFloatScalars)
616 .legalFor(Types: allowedFloatVectorTypes)
617 .fewerElementsIf(Predicate: vectorElementCountIsGreaterThan(TypeIdx: 0, Size: MaxVectorSize),
618 Mutation: LegalizeMutations::changeElementCountTo(
619 TypeIdx: 0, EC: ElementCount::getFixed(MinVal: MaxVectorSize)))
620 .moreElementsToNextPow2(TypeIdx: 0);
621 // clang-format on
622
623 getActionDefinitionsBuilder(Opcode: G_FCOPYSIGN)
624 .legalForCartesianProduct(Types0: allFloatScalarsAndVectors,
625 Types1: allFloatScalarsAndVectors);
626
627 getActionDefinitionsBuilder(Opcode: G_FPOWI).legalForCartesianProduct(
628 Types0: allFloatScalarsAndVectors, Types1: allIntScalarsAndVectors);
629
630 if (ST.canUseExtInstSet(E: SPIRV::InstructionSet::OpenCL_std)) {
631 getActionDefinitionsBuilder(
632 Opcodes: {G_CTTZ, G_CTTZ_ZERO_POISON, G_CTLZ, G_CTLZ_ZERO_POISON})
633 .legalForCartesianProduct(Types0: allIntScalarsAndVectors,
634 Types1: allIntScalarsAndVectors);
635
636 // Struct return types become a single scalar, so cannot easily legalize.
637 getActionDefinitionsBuilder(Opcodes: {G_SMULH, G_UMULH}).alwaysLegal();
638 }
639
640 getActionDefinitionsBuilder(Opcode: G_IS_FPCLASS).custom();
641
642 getActionDefinitionsBuilder(Opcodes: {G_INTRINSIC, G_INTRINSIC_CONVERGENT,
643 G_INTRINSIC_CONVERGENT_W_SIDE_EFFECTS})
644 .alwaysLegal();
645 getActionDefinitionsBuilder(Opcode: G_FENCE).alwaysLegal();
646 getActionDefinitionsBuilder(Opcodes: {G_TRAP, G_DEBUGTRAP, G_UBSANTRAP}).alwaysLegal();
647
648 verify(MII: *ST.getInstrInfo());
649}
650
651static bool legalizeExtractVectorElt(LegalizerHelper &Helper,
652 MachineInstr &MI) {
653 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
654 Register DstReg = MI.getOperand(i: 0).getReg();
655 Register SrcReg = MI.getOperand(i: 1).getReg();
656 Register IdxReg = MI.getOperand(i: 2).getReg();
657
658 MIRBuilder
659 .buildIntrinsic(ID: Intrinsic::spv_extractelt, Res: ArrayRef<Register>{DstReg})
660 .addUse(RegNo: SrcReg)
661 .addUse(RegNo: IdxReg);
662 MI.eraseFromParent();
663 return true;
664}
665
666static bool legalizeInsertVectorElt(LegalizerHelper &Helper, MachineInstr &MI) {
667 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
668 Register DstReg = MI.getOperand(i: 0).getReg();
669 Register SrcReg = MI.getOperand(i: 1).getReg();
670 Register ValReg = MI.getOperand(i: 2).getReg();
671 Register IdxReg = MI.getOperand(i: 3).getReg();
672
673 MIRBuilder
674 .buildIntrinsic(ID: Intrinsic::spv_insertelt, Res: ArrayRef<Register>{DstReg})
675 .addUse(RegNo: SrcReg)
676 .addUse(RegNo: ValReg)
677 .addUse(RegNo: IdxReg);
678 MI.eraseFromParent();
679 return true;
680}
681
682static Register convertPtrToInt(Register Reg, LLT ConvTy, SPIRVTypeInst SpvType,
683 LegalizerHelper &Helper,
684 MachineRegisterInfo &MRI,
685 SPIRVGlobalRegistry *GR) {
686 Register ConvReg = MRI.createGenericVirtualRegister(Ty: ConvTy);
687 MRI.setRegClass(Reg: ConvReg, RC: GR->getRegClass(SpvType));
688 GR->assignSPIRVTypeToVReg(Type: SpvType, VReg: ConvReg, MF: Helper.MIRBuilder.getMF());
689 Helper.MIRBuilder.buildInstr(Opcode: TargetOpcode::G_PTRTOINT)
690 .addDef(RegNo: ConvReg)
691 .addUse(RegNo: Reg);
692 return ConvReg;
693}
694
695static bool needsVectorLegalization(const LLT &Ty, const SPIRVSubtarget &ST) {
696 if (!Ty.isVector() ||
697 ST.canUseExtension(E: SPIRV::Extension::SPV_EXT_long_vector))
698 return false;
699 unsigned NumElements = Ty.getNumElements();
700 unsigned MaxVectorSize = ST.isShader() ? 4 : 16;
701 return (NumElements > 4 && !isPowerOf2_32(Value: NumElements)) ||
702 NumElements > MaxVectorSize;
703}
704
705static bool legalizeLoad(LegalizerHelper &Helper, MachineInstr &MI,
706 SPIRVGlobalRegistry *GR) {
707 MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
708 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
709 Register DstReg = MI.getOperand(i: 0).getReg();
710 Register PtrReg = MI.getOperand(i: 1).getReg();
711 LLT DstTy = MRI.getType(Reg: DstReg);
712
713 if (!DstTy.isVector())
714 return true;
715
716 const SPIRVSubtarget &ST = MI.getMF()->getSubtarget<SPIRVSubtarget>();
717 if (!needsVectorLegalization(Ty: DstTy, ST))
718 return true;
719
720 SmallVector<Register, 8> SplitRegs;
721 LLT EltTy = DstTy.getElementType();
722 unsigned NumElts = DstTy.getNumElements();
723
724 LLT PtrTy = MRI.getType(Reg: PtrReg);
725 auto Zero = MIRBuilder.buildConstant(Res: LLT::scalar(SizeInBits: 32), Val: 0);
726
727 for (unsigned i = 0; i < NumElts; ++i) {
728 auto Idx = MIRBuilder.buildConstant(Res: LLT::scalar(SizeInBits: 32), Val: i);
729 Register EltPtr = MRI.createGenericVirtualRegister(Ty: PtrTy);
730
731 MIRBuilder.buildIntrinsic(ID: Intrinsic::spv_gep, Res: ArrayRef<Register>{EltPtr})
732 .addImm(Val: 1) // InBounds
733 .addUse(RegNo: PtrReg)
734 .addUse(RegNo: Zero.getReg(Idx: 0))
735 .addUse(RegNo: Idx.getReg(Idx: 0));
736
737 MachinePointerInfo EltPtrInfo;
738 Align EltAlign = Align(1);
739 if (!MI.memoperands_empty()) {
740 MachineMemOperand *MMO = *MI.memoperands_begin();
741 EltPtrInfo =
742 MMO->getPointerInfo().getWithOffset(O: i * EltTy.getSizeInBytes());
743 EltAlign = commonAlignment(A: MMO->getAlign(), Offset: i * EltTy.getSizeInBytes());
744 }
745
746 Register EltReg = MRI.createGenericVirtualRegister(Ty: EltTy);
747 MIRBuilder.buildLoad(Res: EltReg, Addr: EltPtr, PtrInfo: EltPtrInfo, Alignment: EltAlign);
748 SplitRegs.push_back(Elt: EltReg);
749 }
750
751 MIRBuilder.buildBuildVector(Res: DstReg, Ops: SplitRegs);
752 MI.eraseFromParent();
753 return true;
754}
755
756static bool legalizeStore(LegalizerHelper &Helper, MachineInstr &MI,
757 SPIRVGlobalRegistry *GR) {
758 MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
759 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
760 Register ValReg = MI.getOperand(i: 0).getReg();
761 Register PtrReg = MI.getOperand(i: 1).getReg();
762 LLT ValTy = MRI.getType(Reg: ValReg);
763
764 assert(ValTy.isVector() && "Expected vector store");
765
766 SmallVector<Register, 8> SplitRegs;
767 LLT EltTy = ValTy.getElementType();
768 unsigned NumElts = ValTy.getNumElements();
769
770 for (unsigned i = 0; i < NumElts; ++i)
771 SplitRegs.push_back(Elt: MRI.createGenericVirtualRegister(Ty: EltTy));
772
773 MIRBuilder.buildUnmerge(Res: SplitRegs, Op: ValReg);
774
775 LLT PtrTy = MRI.getType(Reg: PtrReg);
776 auto Zero = MIRBuilder.buildConstant(Res: LLT::scalar(SizeInBits: 32), Val: 0);
777
778 for (unsigned i = 0; i < NumElts; ++i) {
779 auto Idx = MIRBuilder.buildConstant(Res: LLT::scalar(SizeInBits: 32), Val: i);
780 Register EltPtr = MRI.createGenericVirtualRegister(Ty: PtrTy);
781
782 MIRBuilder.buildIntrinsic(ID: Intrinsic::spv_gep, Res: ArrayRef<Register>{EltPtr})
783 .addImm(Val: 1) // InBounds
784 .addUse(RegNo: PtrReg)
785 .addUse(RegNo: Zero.getReg(Idx: 0))
786 .addUse(RegNo: Idx.getReg(Idx: 0));
787
788 MachinePointerInfo EltPtrInfo;
789 Align EltAlign = Align(1);
790 if (!MI.memoperands_empty()) {
791 MachineMemOperand *MMO = *MI.memoperands_begin();
792 EltPtrInfo =
793 MMO->getPointerInfo().getWithOffset(O: i * EltTy.getSizeInBytes());
794 EltAlign = commonAlignment(A: MMO->getAlign(), Offset: i * EltTy.getSizeInBytes());
795 }
796
797 MIRBuilder.buildStore(Val: SplitRegs[i], Addr: EltPtr, PtrInfo: EltPtrInfo, Alignment: EltAlign);
798 }
799
800 MI.eraseFromParent();
801 return true;
802}
803
804bool SPIRVLegalizerInfo::legalizeCustom(
805 LegalizerHelper &Helper, MachineInstr &MI,
806 LostDebugLocObserver &LocObserver) const {
807 MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
808 switch (MI.getOpcode()) {
809 default:
810 // TODO: implement legalization for other opcodes.
811 return true;
812 case TargetOpcode::G_BITCAST:
813 return legalizeBitcast(Helper, MI);
814 case TargetOpcode::G_EXTRACT_VECTOR_ELT:
815 return legalizeExtractVectorElt(Helper, MI);
816 case TargetOpcode::G_INSERT_VECTOR_ELT:
817 return legalizeInsertVectorElt(Helper, MI);
818 case TargetOpcode::G_INTRINSIC:
819 case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
820 return legalizeIntrinsic(Helper, MI);
821 case TargetOpcode::G_IS_FPCLASS:
822 return legalizeIsFPClass(Helper, MI, LocObserver);
823 case TargetOpcode::G_ICMP: {
824 auto &Op0 = MI.getOperand(i: 2);
825 auto &Op1 = MI.getOperand(i: 3);
826 Register Reg0 = Op0.getReg();
827 Register Reg1 = Op1.getReg();
828 CmpInst::Predicate Cond =
829 static_cast<CmpInst::Predicate>(MI.getOperand(i: 1).getPredicate());
830 if ((!ST->canDirectlyComparePointers() ||
831 (Cond != CmpInst::ICMP_EQ && Cond != CmpInst::ICMP_NE)) &&
832 MRI.getType(Reg: Reg0).isPointer() && MRI.getType(Reg: Reg1).isPointer()) {
833 LLT ConvT = LLT::scalar(SizeInBits: ST->getPointerSize());
834 Type *LLVMTy = IntegerType::get(C&: MI.getMF()->getFunction().getContext(),
835 NumBits: ST->getPointerSize());
836 SPIRVTypeInst SpirvTy = GR->getOrCreateSPIRVType(
837 Type: LLVMTy, MIRBuilder&: Helper.MIRBuilder, AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
838 Op0.setReg(convertPtrToInt(Reg: Reg0, ConvTy: ConvT, SpvType: SpirvTy, Helper, MRI, GR));
839 Op1.setReg(convertPtrToInt(Reg: Reg1, ConvTy: ConvT, SpvType: SpirvTy, Helper, MRI, GR));
840 }
841 return true;
842 }
843 case TargetOpcode::G_LOAD:
844 return legalizeLoad(Helper, MI, GR);
845 case TargetOpcode::G_STORE:
846 return legalizeStore(Helper, MI, GR);
847 }
848}
849
850static MachineInstrBuilder
851createStackTemporaryForVector(LegalizerHelper &Helper, SPIRVGlobalRegistry *GR,
852 Register SrcReg, LLT SrcTy,
853 MachinePointerInfo &PtrInfo, Align &VecAlign) {
854 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
855 MachineRegisterInfo &MRI = *MIRBuilder.getMRI();
856
857 VecAlign = Helper.getStackTemporaryAlignment(Type: SrcTy);
858 auto StackTemp = Helper.createStackTemporary(
859 Bytes: TypeSize::getFixed(ExactSize: SrcTy.getSizeInBytes()), Alignment: VecAlign, PtrInfo);
860
861 // Set the type of StackTemp to a pointer to an array of the element type.
862 SPIRVTypeInst SpvSrcTy = GR->getSPIRVTypeForVReg(VReg: SrcReg);
863 SPIRVTypeInst EltSpvTy = GR->getScalarOrVectorComponentType(Type: SpvSrcTy);
864 const Type *LLVMEltTy = GR->getTypeForSPIRVType(Ty: EltSpvTy);
865 const Type *LLVMArrTy =
866 ArrayType::get(ElementType: const_cast<Type *>(LLVMEltTy), NumElements: SrcTy.getNumElements());
867 SPIRVTypeInst ArrSpvTy = GR->getOrCreateSPIRVType(
868 Type: LLVMArrTy, MIRBuilder, AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
869 SPIRVTypeInst PtrToArrSpvTy = GR->getOrCreateSPIRVPointerType(
870 BaseType: ArrSpvTy, MIRBuilder, SC: SPIRV::StorageClass::Function);
871
872 Register StackReg = StackTemp.getReg(Idx: 0);
873 MRI.setRegClass(Reg: StackReg, RC: GR->getRegClass(SpvType: PtrToArrSpvTy));
874 GR->assignSPIRVTypeToVReg(Type: PtrToArrSpvTy, VReg: StackReg, MF: MIRBuilder.getMF());
875
876 return StackTemp;
877}
878
879static bool legalizeSpvBitcast(LegalizerHelper &Helper, MachineInstr &MI,
880 SPIRVGlobalRegistry *GR) {
881 LLVM_DEBUG(dbgs() << "Found a bitcast instruction\n");
882 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
883 MachineRegisterInfo &MRI = *MIRBuilder.getMRI();
884 const SPIRVSubtarget &ST = MI.getMF()->getSubtarget<SPIRVSubtarget>();
885
886 Register DstReg = MI.getOperand(i: 0).getReg();
887 Register SrcReg = MI.getOperand(i: 2).getReg();
888 LLT DstTy = MRI.getType(Reg: DstReg);
889 LLT SrcTy = MRI.getType(Reg: SrcReg);
890
891 // If an spv_bitcast needs to be legalized, we convert it to G_BITCAST to
892 // allow using the generic legalization rules.
893 if (needsVectorLegalization(Ty: DstTy, ST) ||
894 needsVectorLegalization(Ty: SrcTy, ST)) {
895 LLVM_DEBUG(dbgs() << "Replacing with a G_BITCAST\n");
896 MIRBuilder.buildBitcast(Dst: DstReg, Src: SrcReg);
897 MI.eraseFromParent();
898 }
899 return true;
900}
901
902static bool legalizeSpvInsertElt(LegalizerHelper &Helper, MachineInstr &MI,
903 SPIRVGlobalRegistry *GR) {
904 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
905 MachineRegisterInfo &MRI = *MIRBuilder.getMRI();
906 const SPIRVSubtarget &ST = MI.getMF()->getSubtarget<SPIRVSubtarget>();
907
908 Register DstReg = MI.getOperand(i: 0).getReg();
909 LLT DstTy = MRI.getType(Reg: DstReg);
910
911 if (needsVectorLegalization(Ty: DstTy, ST)) {
912 Register SrcReg = MI.getOperand(i: 2).getReg();
913 Register ValReg = MI.getOperand(i: 3).getReg();
914 LLT SrcTy = MRI.getType(Reg: SrcReg);
915 MachineOperand &IdxOperand = MI.getOperand(i: 4);
916
917 if (getImm(MO: IdxOperand, MRI: &MRI)) {
918 uint64_t IdxVal = foldImm(MO: IdxOperand, MRI: &MRI);
919 if (IdxVal < SrcTy.getNumElements()) {
920 SmallVector<Register, 8> Regs;
921 SPIRVTypeInst ElementType =
922 GR->getScalarOrVectorComponentType(Type: GR->getSPIRVTypeForVReg(VReg: DstReg));
923 LLT ElementLLTTy = GR->getRegType(SpvType: ElementType);
924 for (unsigned I = 0, E = SrcTy.getNumElements(); I < E; ++I) {
925 Register Reg = MRI.createGenericVirtualRegister(Ty: ElementLLTTy);
926 MRI.setRegClass(Reg, RC: GR->getRegClass(SpvType: ElementType));
927 GR->assignSPIRVTypeToVReg(Type: ElementType, VReg: Reg, MF: *MI.getMF());
928 Regs.push_back(Elt: Reg);
929 }
930 MIRBuilder.buildUnmerge(Res: Regs, Op: SrcReg);
931 Regs[IdxVal] = ValReg;
932 MIRBuilder.buildBuildVector(Res: DstReg, Ops: Regs);
933 MI.eraseFromParent();
934 return true;
935 }
936 }
937
938 LLT EltTy = SrcTy.getElementType();
939 Align VecAlign;
940 MachinePointerInfo PtrInfo;
941 auto StackTemp = createStackTemporaryForVector(Helper, GR, SrcReg, SrcTy,
942 PtrInfo, VecAlign);
943
944 MIRBuilder.buildStore(Val: SrcReg, Addr: StackTemp, PtrInfo, Alignment: VecAlign);
945
946 Register IdxReg = IdxOperand.getReg();
947 LLT PtrTy = MRI.getType(Reg: StackTemp.getReg(Idx: 0));
948 Register EltPtr = MRI.createGenericVirtualRegister(Ty: PtrTy);
949 auto Zero = MIRBuilder.buildConstant(Res: LLT::scalar(SizeInBits: 32), Val: 0);
950
951 MIRBuilder.buildIntrinsic(ID: Intrinsic::spv_gep, Res: ArrayRef<Register>{EltPtr})
952 .addImm(Val: 1) // InBounds
953 .addUse(RegNo: StackTemp.getReg(Idx: 0))
954 .addUse(RegNo: Zero.getReg(Idx: 0))
955 .addUse(RegNo: IdxReg);
956
957 MachinePointerInfo EltPtrInfo = MachinePointerInfo(PtrTy.getAddressSpace());
958 Align EltAlign = Helper.getStackTemporaryAlignment(Type: EltTy);
959 MIRBuilder.buildStore(Val: ValReg, Addr: EltPtr, PtrInfo: EltPtrInfo, Alignment: EltAlign);
960
961 MIRBuilder.buildLoad(Res: DstReg, Addr: StackTemp, PtrInfo, Alignment: VecAlign);
962 MI.eraseFromParent();
963 return true;
964 }
965 return true;
966}
967
968static bool legalizeSpvExtractElt(LegalizerHelper &Helper, MachineInstr &MI,
969 SPIRVGlobalRegistry *GR) {
970 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
971 MachineRegisterInfo &MRI = *MIRBuilder.getMRI();
972 const SPIRVSubtarget &ST = MI.getMF()->getSubtarget<SPIRVSubtarget>();
973
974 Register SrcReg = MI.getOperand(i: 2).getReg();
975 LLT SrcTy = MRI.getType(Reg: SrcReg);
976
977 if (needsVectorLegalization(Ty: SrcTy, ST)) {
978 Register DstReg = MI.getOperand(i: 0).getReg();
979 MachineOperand &IdxOperand = MI.getOperand(i: 3);
980
981 if (getImm(MO: IdxOperand, MRI: &MRI)) {
982 uint64_t IdxVal = foldImm(MO: IdxOperand, MRI: &MRI);
983 if (IdxVal < SrcTy.getNumElements()) {
984 LLT DstTy = MRI.getType(Reg: DstReg);
985 SmallVector<Register, 8> Regs;
986 SPIRVTypeInst DstSpvTy = GR->getSPIRVTypeForVReg(VReg: DstReg);
987 for (unsigned I = 0, E = SrcTy.getNumElements(); I < E; ++I) {
988 if (I == IdxVal) {
989 Regs.push_back(Elt: DstReg);
990 } else {
991 Register Reg = MRI.createGenericVirtualRegister(Ty: DstTy);
992 MRI.setRegClass(Reg, RC: GR->getRegClass(SpvType: DstSpvTy));
993 GR->assignSPIRVTypeToVReg(Type: DstSpvTy, VReg: Reg, MF: *MI.getMF());
994 Regs.push_back(Elt: Reg);
995 }
996 }
997 MIRBuilder.buildUnmerge(Res: Regs, Op: SrcReg);
998 MI.eraseFromParent();
999 return true;
1000 }
1001 }
1002
1003 LLT EltTy = SrcTy.getElementType();
1004 Align VecAlign;
1005 MachinePointerInfo PtrInfo;
1006 auto StackTemp = createStackTemporaryForVector(Helper, GR, SrcReg, SrcTy,
1007 PtrInfo, VecAlign);
1008
1009 MIRBuilder.buildStore(Val: SrcReg, Addr: StackTemp, PtrInfo, Alignment: VecAlign);
1010
1011 Register IdxReg = IdxOperand.getReg();
1012 LLT PtrTy = MRI.getType(Reg: StackTemp.getReg(Idx: 0));
1013 Register EltPtr = MRI.createGenericVirtualRegister(Ty: PtrTy);
1014 auto Zero = MIRBuilder.buildConstant(Res: LLT::scalar(SizeInBits: 32), Val: 0);
1015
1016 MIRBuilder.buildIntrinsic(ID: Intrinsic::spv_gep, Res: ArrayRef<Register>{EltPtr})
1017 .addImm(Val: 1) // InBounds
1018 .addUse(RegNo: StackTemp.getReg(Idx: 0))
1019 .addUse(RegNo: Zero.getReg(Idx: 0))
1020 .addUse(RegNo: IdxReg);
1021
1022 MachinePointerInfo EltPtrInfo = MachinePointerInfo(PtrTy.getAddressSpace());
1023 Align EltAlign = Helper.getStackTemporaryAlignment(Type: EltTy);
1024 MIRBuilder.buildLoad(Res: DstReg, Addr: EltPtr, PtrInfo: EltPtrInfo, Alignment: EltAlign);
1025
1026 MI.eraseFromParent();
1027 return true;
1028 }
1029 return true;
1030}
1031
1032static bool legalizeSpvConstComposite(LegalizerHelper &Helper, MachineInstr &MI,
1033 SPIRVGlobalRegistry *GR) {
1034 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
1035 MachineRegisterInfo &MRI = *MIRBuilder.getMRI();
1036 const SPIRVSubtarget &ST = MI.getMF()->getSubtarget<SPIRVSubtarget>();
1037
1038 Register DstReg = MI.getOperand(i: 0).getReg();
1039 LLT DstTy = MRI.getType(Reg: DstReg);
1040
1041 if (!needsVectorLegalization(Ty: DstTy, ST))
1042 return true;
1043
1044 SmallVector<Register, 8> SrcRegs;
1045 if (MI.getNumOperands() == 2) {
1046 // The "null" case: no values are attached.
1047 LLT EltTy = DstTy.getElementType();
1048 auto Zero = MIRBuilder.buildConstant(Res: EltTy, Val: 0);
1049 SPIRVTypeInst SpvDstTy = GR->getSPIRVTypeForVReg(VReg: DstReg);
1050 SPIRVTypeInst SpvEltTy = GR->getScalarOrVectorComponentType(Type: SpvDstTy);
1051 GR->assignSPIRVTypeToVReg(Type: SpvEltTy, VReg: Zero.getReg(Idx: 0), MF: MIRBuilder.getMF());
1052 for (unsigned i = 0; i < DstTy.getNumElements(); ++i)
1053 SrcRegs.push_back(Elt: Zero.getReg(Idx: 0));
1054 } else {
1055 for (unsigned i = 2; i < MI.getNumOperands(); ++i) {
1056 SrcRegs.push_back(Elt: MI.getOperand(i).getReg());
1057 }
1058 }
1059 MIRBuilder.buildBuildVector(Res: DstReg, Ops: SrcRegs);
1060 MI.eraseFromParent();
1061 return true;
1062}
1063
1064static SmallVector<Register, 16> unmergeToScalars(Register Reg,
1065 MachineIRBuilder &MIRBuilder,
1066 SPIRVGlobalRegistry *GR) {
1067 LLT Ty = MIRBuilder.getMRI()->getType(Reg);
1068 if (!Ty.isVector())
1069 return {Reg};
1070 SPIRVTypeInst EltSpvTy =
1071 GR->getScalarOrVectorComponentType(Type: GR->getSPIRVTypeForVReg(VReg: Reg));
1072 unsigned NumElts = Ty.getNumElements();
1073 SmallVector<Register, 16> Elts;
1074 for (unsigned I = 0; I < NumElts; ++I)
1075 Elts.push_back(Elt: createVirtualRegister(SpvType: EltSpvTy, GR, MIRBuilder));
1076 MIRBuilder.buildUnmerge(Res: Elts, Op: Reg);
1077 return Elts;
1078}
1079
1080static Register buildVectorPart(ArrayRef<Register> Elts,
1081 MachineIRBuilder &MIRBuilder,
1082 SPIRVGlobalRegistry *GR) {
1083 if (Elts.size() == 1)
1084 return Elts[0];
1085 SPIRVTypeInst PartSpvTy =
1086 GR->getOrCreateSPIRVVectorType(BaseType: GR->getSPIRVTypeForVReg(VReg: Elts[0]),
1087 NumElements: Elts.size(), MIRBuilder, /*EmitIR=*/true);
1088 Register Part = createVirtualRegister(SpvType: PartSpvTy, GR, MIRBuilder);
1089 MIRBuilder.buildBuildVector(Res: Part, Ops: Elts);
1090 return Part;
1091}
1092
1093// Split an elementwise intrinsic with an illegal vector width into intrinsics
1094// on legal vector widths.
1095static bool legalizeElementwiseIntrinsic(LegalizerHelper &Helper,
1096 GIntrinsic &MI,
1097 SPIRVGlobalRegistry *GR) {
1098 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
1099 MachineRegisterInfo &MRI = *MIRBuilder.getMRI();
1100 const SPIRVSubtarget &ST = MI.getMF()->getSubtarget<SPIRVSubtarget>();
1101
1102 if (!Intrinsic::isTriviallyScalarizable(id: MI.getIntrinsicID()))
1103 return true;
1104 Register DstReg = MI.getReg(Idx: 0);
1105 LLT DstTy = MRI.getType(Reg: DstReg);
1106 if (!needsVectorLegalization(Ty: DstTy, ST))
1107 return true;
1108
1109 unsigned NumElts = DstTy.getNumElements();
1110 unsigned MaxVectorSize = ST.isShader() ? 4 : 16;
1111 unsigned PartSize = NumElts > MaxVectorSize ? MaxVectorSize : 4;
1112
1113 SmallDenseMap<Register, SmallVector<Register, 16>, 4> OpElts;
1114 for (const MachineOperand &MO : drop_begin(RangeOrContainer: MI.explicit_uses())) {
1115 if (!MO.isReg() || !MRI.getType(Reg: MO.getReg()).isVector())
1116 continue;
1117 auto [It, Inserted] = OpElts.try_emplace(Key: MO.getReg());
1118 if (Inserted)
1119 It->second = unmergeToScalars(Reg: MO.getReg(), MIRBuilder, GR);
1120 }
1121
1122 SPIRVTypeInst DstEltSpvTy =
1123 GR->getScalarOrVectorComponentType(Type: GR->getSPIRVTypeForVReg(VReg: DstReg));
1124 SmallVector<Register, 16> DstElts;
1125 for (unsigned Offset = 0; Offset < NumElts; Offset += PartSize) {
1126 unsigned Size = std::min(a: PartSize, b: NumElts - Offset);
1127 SPIRVTypeInst PartSpvTy =
1128 Size == 1 ? DstEltSpvTy
1129 : GR->getOrCreateSPIRVVectorType(BaseType: DstEltSpvTy, NumElements: Size,
1130 MIRBuilder, /*EmitIR=*/true);
1131 SmallDenseMap<Register, Register, 4> PartRegs;
1132 SmallVector<MachineOperand> PartOps;
1133 for (const MachineOperand &MO : drop_begin(RangeOrContainer: MI.explicit_uses())) {
1134 auto EltsIt = MO.isReg() ? OpElts.find(Val: MO.getReg()) : OpElts.end();
1135 if (EltsIt == OpElts.end()) {
1136 PartOps.push_back(Elt: MO);
1137 continue;
1138 }
1139 auto [It, Inserted] = PartRegs.try_emplace(Key: MO.getReg());
1140 if (Inserted)
1141 It->second = buildVectorPart(
1142 Elts: ArrayRef(EltsIt->second).slice(N: Offset, M: Size), MIRBuilder, GR);
1143 PartOps.push_back(Elt: MachineOperand::CreateReg(Reg: It->second, /*isDef=*/false));
1144 }
1145 Register PartDst = createVirtualRegister(SpvType: PartSpvTy, GR, MIRBuilder);
1146 auto Part = MIRBuilder.buildIntrinsic(
1147 ID: MI.getIntrinsicID(), Res: ArrayRef<Register>{PartDst}, HasSideEffects: MI.hasSideEffects(),
1148 isConvergent: MI.isConvergent());
1149 for (const MachineOperand &MO : PartOps)
1150 Part.add(MO);
1151 Part->setFlags(MI.getFlags());
1152 append_range(C&: DstElts, R: unmergeToScalars(Reg: PartDst, MIRBuilder, GR));
1153 }
1154
1155 MIRBuilder.buildBuildVector(Res: DstReg, Ops: DstElts);
1156 MI.eraseFromParent();
1157 return true;
1158}
1159
1160bool SPIRVLegalizerInfo::legalizeIntrinsic(LegalizerHelper &Helper,
1161 MachineInstr &MI) const {
1162 LLVM_DEBUG(dbgs() << "legalizeIntrinsic: " << MI);
1163 auto IntrinsicID = cast<GIntrinsic>(Val&: MI).getIntrinsicID();
1164 switch (IntrinsicID) {
1165 case Intrinsic::spv_bitcast:
1166 return legalizeSpvBitcast(Helper, MI, GR);
1167 case Intrinsic::spv_insertelt:
1168 return legalizeSpvInsertElt(Helper, MI, GR);
1169 case Intrinsic::spv_extractelt:
1170 return legalizeSpvExtractElt(Helper, MI, GR);
1171 case Intrinsic::spv_const_composite:
1172 return legalizeSpvConstComposite(Helper, MI, GR);
1173 }
1174 return legalizeElementwiseIntrinsic(Helper, MI&: cast<GIntrinsic>(Val&: MI), GR);
1175}
1176
1177bool SPIRVLegalizerInfo::legalizeBitcast(LegalizerHelper &Helper,
1178 MachineInstr &MI) const {
1179 // Once the G_BITCAST is using vectors that are allowed, we turn it back into
1180 // an spv_bitcast to avoid verifier problems when the register types are the
1181 // same for the source and the result. Note that the SPIR-V types associated
1182 // with the bitcast can be different even if the register types are the same.
1183 MachineIRBuilder &MIRBuilder = Helper.MIRBuilder;
1184 Register DstReg = MI.getOperand(i: 0).getReg();
1185 Register SrcReg = MI.getOperand(i: 1).getReg();
1186 SmallVector<Register, 1> DstRegs = {DstReg};
1187 MIRBuilder.buildIntrinsic(ID: Intrinsic::spv_bitcast, Res: DstRegs).addUse(RegNo: SrcReg);
1188 MI.eraseFromParent();
1189 return true;
1190}
1191
1192// Note this code was copied from LegalizerHelper::lowerISFPCLASS and adjusted
1193// to ensure that all instructions created during the lowering have SPIR-V types
1194// assigned to them.
1195bool SPIRVLegalizerInfo::legalizeIsFPClass(
1196 LegalizerHelper &Helper, MachineInstr &MI,
1197 LostDebugLocObserver &LocObserver) const {
1198 auto [DstReg, DstTy, SrcReg, SrcTy] = MI.getFirst2RegLLTs();
1199 FPClassTest Mask = static_cast<FPClassTest>(MI.getOperand(i: 2).getImm());
1200
1201 auto &MIRBuilder = Helper.MIRBuilder;
1202 auto &MF = MIRBuilder.getMF();
1203 MachineRegisterInfo &MRI = MF.getRegInfo();
1204
1205 Type *LLVMDstTy =
1206 IntegerType::get(C&: MIRBuilder.getContext(), NumBits: DstTy.getScalarSizeInBits());
1207 if (DstTy.isVector())
1208 LLVMDstTy = VectorType::get(ElementType: LLVMDstTy, EC: DstTy.getElementCount());
1209 SPIRVTypeInst SPIRVDstTy = GR->getOrCreateSPIRVType(
1210 Type: LLVMDstTy, MIRBuilder, AQ: SPIRV::AccessQualifier::ReadWrite,
1211 /*EmitIR*/ true);
1212
1213 unsigned BitSize = SrcTy.getScalarSizeInBits();
1214 const fltSemantics &Semantics = getFltSemanticForLLT(Ty: SrcTy.getScalarType());
1215
1216 LLT IntTy = LLT::scalar(SizeInBits: BitSize);
1217 Type *LLVMIntTy = IntegerType::get(C&: MIRBuilder.getContext(), NumBits: BitSize);
1218 if (SrcTy.isVector()) {
1219 IntTy = LLT::vector(EC: SrcTy.getElementCount(), ScalarTy: IntTy);
1220 LLVMIntTy = VectorType::get(ElementType: LLVMIntTy, EC: SrcTy.getElementCount());
1221 }
1222 SPIRVTypeInst SPIRVIntTy = GR->getOrCreateSPIRVType(
1223 Type: LLVMIntTy, MIRBuilder, AQ: SPIRV::AccessQualifier::ReadWrite,
1224 /*EmitIR*/ true);
1225
1226 // Clang doesn't support capture of structured bindings:
1227 LLT DstTyCopy = DstTy;
1228 const auto assignSPIRVTy = [&](MachineInstrBuilder &&MI) {
1229 // Assign this MI's (assumed only) destination to one of the two types we
1230 // expect: either the G_IS_FPCLASS's destination type, or the integer type
1231 // bitcast from the source type.
1232 LLT MITy = MRI.getType(Reg: MI.getReg(Idx: 0));
1233 assert((MITy == IntTy || MITy == DstTyCopy) &&
1234 "Unexpected LLT type while lowering G_IS_FPCLASS");
1235 SPIRVTypeInst SPVTy = MITy == IntTy ? SPIRVIntTy : SPIRVDstTy;
1236 GR->assignSPIRVTypeToVReg(Type: SPVTy, VReg: MI.getReg(Idx: 0), MF);
1237 return MI;
1238 };
1239
1240 // Helper to build and assign a constant in one go
1241 const auto buildSPIRVConstant = [&](LLT Ty, auto &&C) -> MachineInstrBuilder {
1242 if (!Ty.isFixedVector())
1243 return assignSPIRVTy(MIRBuilder.buildConstant(Ty, C));
1244 auto ScalarC = MIRBuilder.buildConstant(Ty.getScalarType(), C);
1245 assert((Ty == IntTy || Ty == DstTyCopy) &&
1246 "Unexpected LLT type while lowering constant for G_IS_FPCLASS");
1247 SPIRVTypeInst VecEltTy = GR->getOrCreateSPIRVType(
1248 Type: (Ty == IntTy ? LLVMIntTy : LLVMDstTy)->getScalarType(), MIRBuilder,
1249 AQ: SPIRV::AccessQualifier::ReadWrite,
1250 /*EmitIR*/ true);
1251 GR->assignSPIRVTypeToVReg(Type: VecEltTy, VReg: ScalarC.getReg(0), MF);
1252 return assignSPIRVTy(MIRBuilder.buildSplatBuildVector(Res: Ty, Src: ScalarC));
1253 };
1254
1255 if (Mask == fcNone) {
1256 MIRBuilder.buildCopy(Res: DstReg, Op: buildSPIRVConstant(DstTy, 0));
1257 MI.eraseFromParent();
1258 return true;
1259 }
1260 if (Mask == fcAllFlags) {
1261 MIRBuilder.buildCopy(Res: DstReg, Op: buildSPIRVConstant(DstTy, 1));
1262 MI.eraseFromParent();
1263 return true;
1264 }
1265
1266 // Note that rather than creating a COPY here (between a floating-point and
1267 // integer type of the same size) we create a SPIR-V bitcast immediately. We
1268 // can't create a G_BITCAST because the LLTs are the same, and we can't seem
1269 // to correctly lower COPYs to SPIR-V bitcasts at this moment.
1270 Register ResVReg = MRI.createGenericVirtualRegister(Ty: IntTy);
1271 MRI.setRegClass(Reg: ResVReg, RC: GR->getRegClass(SpvType: SPIRVIntTy));
1272 GR->assignSPIRVTypeToVReg(Type: SPIRVIntTy, VReg: ResVReg, MF: Helper.MIRBuilder.getMF());
1273 auto AsInt = MIRBuilder.buildInstr(Opcode: SPIRV::OpBitcast)
1274 .addDef(RegNo: ResVReg)
1275 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: SPIRVIntTy))
1276 .addUse(RegNo: SrcReg);
1277 AsInt = assignSPIRVTy(std::move(AsInt));
1278
1279 // Various masks.
1280 APInt SignBit = APInt::getSignMask(BitWidth: BitSize);
1281 APInt ValueMask = APInt::getSignedMaxValue(numBits: BitSize); // All bits but sign.
1282 APInt Inf = APFloat::getInf(Sem: Semantics).bitcastToAPInt(); // Exp and int bit.
1283 APInt ExpMask = Inf;
1284 APInt AllOneMantissa = APFloat::getLargest(Sem: Semantics).bitcastToAPInt() & ~Inf;
1285 APInt QNaNBitMask =
1286 APInt::getOneBitSet(numBits: BitSize, BitNo: AllOneMantissa.getActiveBits() - 1);
1287 APInt InversionMask = APInt::getAllOnes(numBits: DstTy.getScalarSizeInBits());
1288
1289 auto SignBitC = buildSPIRVConstant(IntTy, SignBit);
1290 auto ValueMaskC = buildSPIRVConstant(IntTy, ValueMask);
1291 auto InfC = buildSPIRVConstant(IntTy, Inf);
1292 auto ExpMaskC = buildSPIRVConstant(IntTy, ExpMask);
1293 auto ZeroC = buildSPIRVConstant(IntTy, 0);
1294
1295 auto Abs = assignSPIRVTy(MIRBuilder.buildAnd(Dst: IntTy, Src0: AsInt, Src1: ValueMaskC));
1296 auto Sign = assignSPIRVTy(
1297 MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_NE, Res: DstTy, Op0: AsInt, Op1: Abs));
1298
1299 auto Res = buildSPIRVConstant(DstTy, 0);
1300
1301 const auto appendToRes = [&](MachineInstrBuilder &&ToAppend) {
1302 Res = assignSPIRVTy(
1303 MIRBuilder.buildOr(Dst: DstTyCopy, Src0: Res, Src1: assignSPIRVTy(std::move(ToAppend))));
1304 };
1305
1306 // Tests that involve more than one class should be processed first.
1307 if ((Mask & fcFinite) == fcFinite) {
1308 // finite(V) ==> abs(V) u< exp_mask
1309 appendToRes(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_ULT, Res: DstTy, Op0: Abs,
1310 Op1: ExpMaskC));
1311 Mask &= ~fcFinite;
1312 } else if ((Mask & fcFinite) == fcPosFinite) {
1313 // finite(V) && V > 0 ==> V u< exp_mask
1314 appendToRes(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_ULT, Res: DstTy, Op0: AsInt,
1315 Op1: ExpMaskC));
1316 Mask &= ~fcPosFinite;
1317 } else if ((Mask & fcFinite) == fcNegFinite) {
1318 // finite(V) && V < 0 ==> abs(V) u< exp_mask && signbit == 1
1319 auto Cmp = assignSPIRVTy(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_ULT,
1320 Res: DstTy, Op0: Abs, Op1: ExpMaskC));
1321 appendToRes(MIRBuilder.buildAnd(Dst: DstTy, Src0: Cmp, Src1: Sign));
1322 Mask &= ~fcNegFinite;
1323 }
1324
1325 if (FPClassTest PartialCheck = Mask & (fcZero | fcSubnormal)) {
1326 // fcZero | fcSubnormal => test all exponent bits are 0
1327 // TODO: Handle sign bit specific cases
1328 // TODO: Handle inverted case
1329 if (PartialCheck == (fcZero | fcSubnormal)) {
1330 auto ExpBits = assignSPIRVTy(MIRBuilder.buildAnd(Dst: IntTy, Src0: AsInt, Src1: ExpMaskC));
1331 appendToRes(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: DstTy,
1332 Op0: ExpBits, Op1: ZeroC));
1333 Mask &= ~PartialCheck;
1334 }
1335 }
1336
1337 // Check for individual classes.
1338 if (FPClassTest PartialCheck = Mask & fcZero) {
1339 if (PartialCheck == fcPosZero)
1340 appendToRes(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: DstTy,
1341 Op0: AsInt, Op1: ZeroC));
1342 else if (PartialCheck == fcZero)
1343 appendToRes(
1344 MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: DstTy, Op0: Abs, Op1: ZeroC));
1345 else // fcNegZero
1346 appendToRes(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: DstTy,
1347 Op0: AsInt, Op1: SignBitC));
1348 }
1349
1350 if (FPClassTest PartialCheck = Mask & fcSubnormal) {
1351 // issubnormal(V) ==> unsigned(abs(V) - 1) u< (all mantissa bits set)
1352 // issubnormal(V) && V>0 ==> unsigned(V - 1) u< (all mantissa bits set)
1353 auto V = (PartialCheck == fcPosSubnormal) ? AsInt : Abs;
1354 auto OneC = buildSPIRVConstant(IntTy, 1);
1355 auto VMinusOne = MIRBuilder.buildSub(Dst: IntTy, Src0: V, Src1: OneC);
1356 auto SubnormalRes = assignSPIRVTy(
1357 MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_ULT, Res: DstTy, Op0: VMinusOne,
1358 Op1: buildSPIRVConstant(IntTy, AllOneMantissa)));
1359 if (PartialCheck == fcNegSubnormal)
1360 SubnormalRes = MIRBuilder.buildAnd(Dst: DstTy, Src0: SubnormalRes, Src1: Sign);
1361 appendToRes(std::move(SubnormalRes));
1362 }
1363
1364 if (FPClassTest PartialCheck = Mask & fcInf) {
1365 if (PartialCheck == fcPosInf)
1366 appendToRes(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: DstTy,
1367 Op0: AsInt, Op1: InfC));
1368 else if (PartialCheck == fcInf)
1369 appendToRes(
1370 MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: DstTy, Op0: Abs, Op1: InfC));
1371 else { // fcNegInf
1372 APInt NegInf = APFloat::getInf(Sem: Semantics, Negative: true).bitcastToAPInt();
1373 auto NegInfC = buildSPIRVConstant(IntTy, NegInf);
1374 appendToRes(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: DstTy,
1375 Op0: AsInt, Op1: NegInfC));
1376 }
1377 }
1378
1379 if (FPClassTest PartialCheck = Mask & fcNan) {
1380 auto InfWithQnanBitC =
1381 buildSPIRVConstant(IntTy, std::move(Inf) | QNaNBitMask);
1382 if (PartialCheck == fcNan) {
1383 // isnan(V) ==> abs(V) u> int(inf)
1384 appendToRes(
1385 MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_UGT, Res: DstTy, Op0: Abs, Op1: InfC));
1386 } else if (PartialCheck == fcQNan) {
1387 // isquiet(V) ==> abs(V) u>= (unsigned(Inf) | quiet_bit)
1388 appendToRes(MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_UGE, Res: DstTy, Op0: Abs,
1389 Op1: InfWithQnanBitC));
1390 } else { // fcSNan
1391 // issignaling(V) ==> abs(V) u> unsigned(Inf) &&
1392 // abs(V) u< (unsigned(Inf) | quiet_bit)
1393 auto IsNan = assignSPIRVTy(
1394 MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_UGT, Res: DstTy, Op0: Abs, Op1: InfC));
1395 auto IsNotQnan = assignSPIRVTy(MIRBuilder.buildICmp(
1396 Pred: CmpInst::Predicate::ICMP_ULT, Res: DstTy, Op0: Abs, Op1: InfWithQnanBitC));
1397 appendToRes(MIRBuilder.buildAnd(Dst: DstTy, Src0: IsNan, Src1: IsNotQnan));
1398 }
1399 }
1400
1401 if (FPClassTest PartialCheck = Mask & fcNormal) {
1402 // isnormal(V) ==> (0 u< exp u< max_exp) ==> (unsigned(exp-1) u<
1403 // (max_exp-1))
1404 APInt ExpLSB = ExpMask & ~(ExpMask.shl(shiftAmt: 1));
1405 auto ExpMinusOne = assignSPIRVTy(
1406 MIRBuilder.buildSub(Dst: IntTy, Src0: Abs, Src1: buildSPIRVConstant(IntTy, ExpLSB)));
1407 APInt MaxExpMinusOne = std::move(ExpMask) - ExpLSB;
1408 auto NormalRes = assignSPIRVTy(
1409 MIRBuilder.buildICmp(Pred: CmpInst::Predicate::ICMP_ULT, Res: DstTy, Op0: ExpMinusOne,
1410 Op1: buildSPIRVConstant(IntTy, MaxExpMinusOne)));
1411 if (PartialCheck == fcNegNormal)
1412 NormalRes = MIRBuilder.buildAnd(Dst: DstTy, Src0: NormalRes, Src1: Sign);
1413 else if (PartialCheck == fcPosNormal) {
1414 auto PosSign = assignSPIRVTy(MIRBuilder.buildXor(
1415 Dst: DstTy, Src0: Sign, Src1: buildSPIRVConstant(DstTy, InversionMask)));
1416 NormalRes = MIRBuilder.buildAnd(Dst: DstTy, Src0: NormalRes, Src1: PosSign);
1417 }
1418 appendToRes(std::move(NormalRes));
1419 }
1420
1421 MIRBuilder.buildCopy(Res: DstReg, Op: Res);
1422 MI.eraseFromParent();
1423 return true;
1424}
1425