1//===-- NVPTXISelLowering.cpp - NVPTX DAG Lowering Implementation ---------===//
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 defines the interfaces that NVPTX uses to lower LLVM code into a
10// selection DAG.
11//
12//===----------------------------------------------------------------------===//
13
14#include "NVPTXISelLowering.h"
15#include "MCTargetDesc/NVPTXBaseInfo.h"
16#include "NVPTX.h"
17#include "NVPTXMachineFunctionInfo.h"
18#include "NVPTXSelectionDAGInfo.h"
19#include "NVPTXSubtarget.h"
20#include "NVPTXTargetMachine.h"
21#include "NVPTXTargetObjectFile.h"
22#include "NVPTXUtilities.h"
23#include "NVVMProperties.h"
24#include "llvm/ADT/APFloat.h"
25#include "llvm/ADT/APInt.h"
26#include "llvm/ADT/STLExtras.h"
27#include "llvm/ADT/SmallVector.h"
28#include "llvm/ADT/StringRef.h"
29#include "llvm/CodeGen/Analysis.h"
30#include "llvm/CodeGen/ISDOpcodes.h"
31#include "llvm/CodeGen/MachineFrameInfo.h"
32#include "llvm/CodeGen/MachineFunction.h"
33#include "llvm/CodeGen/MachineJumpTableInfo.h"
34#include "llvm/CodeGen/MachineMemOperand.h"
35#include "llvm/CodeGen/SDPatternMatch.h"
36#include "llvm/CodeGen/SelectionDAG.h"
37#include "llvm/CodeGen/SelectionDAGNodes.h"
38#include "llvm/CodeGen/TargetCallingConv.h"
39#include "llvm/CodeGen/TargetLowering.h"
40#include "llvm/CodeGen/ValueTypes.h"
41#include "llvm/CodeGenTypes/MachineValueType.h"
42#include "llvm/IR/Argument.h"
43#include "llvm/IR/Attributes.h"
44#include "llvm/IR/Constants.h"
45#include "llvm/IR/DataLayout.h"
46#include "llvm/IR/DerivedTypes.h"
47#include "llvm/IR/DiagnosticInfo.h"
48#include "llvm/IR/FPEnv.h"
49#include "llvm/IR/Function.h"
50#include "llvm/IR/GlobalValue.h"
51#include "llvm/IR/IRBuilder.h"
52#include "llvm/IR/Instruction.h"
53#include "llvm/IR/Instructions.h"
54#include "llvm/IR/IntrinsicsNVPTX.h"
55#include "llvm/IR/LLVMContext.h"
56#include "llvm/IR/Module.h"
57#include "llvm/IR/NVVMIntrinsicUtils.h"
58#include "llvm/IR/Type.h"
59#include "llvm/IR/Value.h"
60#include "llvm/MC/MCContext.h"
61#include "llvm/MC/MCSymbol.h"
62#include "llvm/Support/Alignment.h"
63#include "llvm/Support/AtomicOrdering.h"
64#include "llvm/Support/Casting.h"
65#include "llvm/Support/CodeGen.h"
66#include "llvm/Support/CommandLine.h"
67#include "llvm/Support/ErrorHandling.h"
68#include "llvm/Support/KnownBits.h"
69#include "llvm/Support/NVPTXAddrSpace.h"
70#include "llvm/Target/TargetMachine.h"
71#include "llvm/Target/TargetOptions.h"
72#include <algorithm>
73#include <cassert>
74#include <cmath>
75#include <cstdint>
76#include <iterator>
77#include <optional>
78#include <tuple>
79#include <utility>
80#include <vector>
81
82#define DEBUG_TYPE "nvptx-lower"
83
84using namespace llvm;
85
86static cl::opt<bool> sched4reg(
87 "nvptx-sched4reg",
88 cl::desc("NVPTX Specific: schedule for register pressue"), cl::init(Val: false));
89
90static cl::opt<unsigned> FMAContractLevelOpt(
91 "nvptx-fma-level", cl::Hidden,
92 cl::desc("NVPTX Specific: FMA contraction (0: don't do it"
93 " 1: do it 2: do it aggressively"),
94 cl::init(Val: 2));
95
96static cl::opt<NVPTX::DivPrecisionLevel> UsePrecDivF32(
97 "nvptx-prec-divf32", cl::Hidden,
98 cl::desc(
99 "NVPTX Specific: Override the precision of the lowering for f32 fdiv"),
100 cl::values(
101 clEnumValN(NVPTX::DivPrecisionLevel::Approx, "0", "Use div.approx"),
102 clEnumValN(NVPTX::DivPrecisionLevel::Full, "1", "Use div.full"),
103 clEnumValN(NVPTX::DivPrecisionLevel::IEEE754, "2",
104 "Use IEEE Compliant F32 div.rnd if available (default)"),
105 clEnumValN(NVPTX::DivPrecisionLevel::IEEE754_NoFTZ, "3",
106 "Use IEEE Compliant F32 div.rnd if available, no FTZ")),
107 cl::init(Val: NVPTX::DivPrecisionLevel::IEEE754));
108
109static cl::opt<bool> UsePrecSqrtF32(
110 "nvptx-prec-sqrtf32", cl::Hidden,
111 cl::desc("NVPTX Specific: 0 use sqrt.approx, 1 use sqrt.rn."),
112 cl::init(Val: true));
113
114// PTX atom.add.f32 has fixed FTZ behavior that may not match the function's
115// (see shouldExpandAtomicRMWInIR), so we'd normally fall back to a CAS loop
116// when they disagree. This option (enabled by default) allows using atom.add
117// anyway, trading correct denormal handling for the speed of the native
118// instruction.
119static cl::opt<bool> AllowFTZAtomics(
120 "nvptx-allow-ftz-atomics", cl::Hidden,
121 cl::desc("NVPTX Specific: Lower atomicrmw fadd to atom.add even when its "
122 "FTZ behavior does not match the function's denormal mode."),
123 cl::init(Val: true));
124
125/// Whereas CUDA's implementation (see libdevice) uses ex2.approx for exp2(), it
126/// does NOT use lg2.approx for log2, so this is disabled by default.
127static cl::opt<bool> UseApproxLog2F32(
128 "nvptx-approx-log2f32",
129 cl::desc("NVPTX Specific: whether to use lg2.approx for log2"),
130 cl::init(Val: false));
131
132NVPTX::DivPrecisionLevel
133NVPTXTargetLowering::getDivF32Level(const MachineFunction &MF,
134 const SDNode &N) const {
135 // If nvptx-prec-div32=N is used on the command-line, always honor it
136 if (UsePrecDivF32.getNumOccurrences() > 0)
137 return UsePrecDivF32;
138
139 const SDNodeFlags Flags = N.getFlags();
140 if (Flags.hasApproximateFuncs())
141 return NVPTX::DivPrecisionLevel::Approx;
142
143 return NVPTX::DivPrecisionLevel::IEEE754;
144}
145
146bool NVPTXTargetLowering::usePrecSqrtF32(const SDNode *N) const {
147 // If nvptx-prec-sqrtf32 is used on the command-line, always honor it
148 if (UsePrecSqrtF32.getNumOccurrences() > 0)
149 return UsePrecSqrtF32;
150
151 if (N) {
152 const SDNodeFlags Flags = N->getFlags();
153 if (Flags.hasApproximateFuncs())
154 return false;
155 }
156
157 return true;
158}
159
160bool NVPTXTargetLowering::useF32FTZ(const MachineFunction &MF) const {
161 return MF.getDenormalMode(FPType: APFloat::IEEEsingle()).Output ==
162 DenormalMode::PreserveSign;
163}
164
165static bool IsPTXVectorType(MVT VT) {
166 switch (VT.SimpleTy) {
167 default:
168 return false;
169 case MVT::v2i1:
170 case MVT::v4i1:
171 case MVT::v2i8:
172 case MVT::v4i8:
173 case MVT::v8i8: // <2 x i8x4>
174 case MVT::v16i8: // <4 x i8x4>
175 case MVT::v2i16:
176 case MVT::v4i16:
177 case MVT::v8i16: // <4 x i16x2>
178 case MVT::v2i32:
179 case MVT::v4i32:
180 case MVT::v2i64:
181 case MVT::v2f16:
182 case MVT::v4f16:
183 case MVT::v8f16: // <4 x f16x2>
184 case MVT::v2bf16:
185 case MVT::v4bf16:
186 case MVT::v8bf16: // <4 x bf16x2>
187 case MVT::v2f32:
188 case MVT::v4f32:
189 case MVT::v2f64:
190 case MVT::v4i64:
191 case MVT::v4f64:
192 case MVT::v8i32:
193 case MVT::v8f32:
194 case MVT::v16f16: // <8 x f16x2>
195 case MVT::v16bf16: // <8 x bf16x2>
196 case MVT::v16i16: // <8 x i16x2>
197 case MVT::v32i8: // <8 x i8x4>
198 return true;
199 }
200}
201
202// When legalizing vector loads/stores, this function is called, which does two
203// things:
204// 1. Determines Whether the vector is something we want to custom lower,
205// std::nullopt is returned if we do not want to custom lower it.
206// 2. If we do want to handle it, returns two parameters:
207// - unsigned int NumElts - The number of elements in the final vector
208// - EVT EltVT - The type of the elements in the final vector
209static std::optional<std::pair<unsigned int, MVT>>
210getVectorLoweringShape(EVT VectorEVT, const NVPTXSubtarget &STI,
211 unsigned AddressSpace) {
212 const bool CanLowerTo256Bit = STI.has256BitVectorLoadStore(AS: AddressSpace);
213
214 if (CanLowerTo256Bit && VectorEVT.isScalarInteger() &&
215 VectorEVT.getSizeInBits() == 256)
216 return {{4, MVT::i64}};
217
218 if (!VectorEVT.isSimple())
219 return std::nullopt;
220 const MVT VectorVT = VectorEVT.getSimpleVT();
221
222 if (!VectorVT.isVector()) {
223 if (VectorVT == MVT::i128 || VectorVT == MVT::f128)
224 return {{2, MVT::i64}};
225 return std::nullopt;
226 }
227
228 const MVT EltVT = VectorVT.getVectorElementType();
229 const unsigned NumElts = VectorVT.getVectorNumElements();
230
231 // The size of the PTX virtual register that holds a packed type.
232 unsigned PackRegSize;
233
234 // We only handle "native" vector sizes for now, e.g. <4 x double> is not
235 // legal. We can (and should) split that into 2 stores of <2 x double> here
236 // but I'm leaving that as a TODO for now.
237 switch (VectorVT.SimpleTy) {
238 default:
239 return std::nullopt;
240
241 case MVT::v4i64:
242 case MVT::v4f64:
243 // This is a "native" vector type iff the address space is global and the
244 // target supports 256-bit loads/stores
245 if (!CanLowerTo256Bit)
246 return std::nullopt;
247 [[fallthrough]];
248 case MVT::v2i8:
249 case MVT::v2i64:
250 case MVT::v2f64:
251 // This is a "native" vector type
252 return std::pair(NumElts, EltVT);
253
254 case MVT::v16f16: // <8 x f16x2>
255 case MVT::v16bf16: // <8 x bf16x2>
256 case MVT::v16i16: // <8 x i16x2>
257 case MVT::v32i8: // <8 x i8x4>
258 // This can be upsized into a "native" vector type iff the address space is
259 // global and the target supports 256-bit loads/stores.
260 if (!CanLowerTo256Bit)
261 return std::nullopt;
262 [[fallthrough]];
263 case MVT::v2i16: // <1 x i16x2>
264 case MVT::v2f16: // <1 x f16x2>
265 case MVT::v2bf16: // <1 x bf16x2>
266 case MVT::v4i8: // <1 x i8x4>
267 case MVT::v4i16: // <2 x i16x2>
268 case MVT::v4f16: // <2 x f16x2>
269 case MVT::v4bf16: // <2 x bf16x2>
270 case MVT::v8i8: // <2 x i8x4>
271 case MVT::v8f16: // <4 x f16x2>
272 case MVT::v8bf16: // <4 x bf16x2>
273 case MVT::v8i16: // <4 x i16x2>
274 case MVT::v16i8: // <4 x i8x4>
275 PackRegSize = 32;
276 break;
277
278 case MVT::v8f32: // <4 x f32x2>
279 case MVT::v8i32: // <4 x i32x2>
280 // This is a "native" vector type iff the address space is global and the
281 // target supports 256-bit loads/stores
282 if (!CanLowerTo256Bit)
283 return std::nullopt;
284 [[fallthrough]];
285 case MVT::v2f32: // <1 x f32x2>
286 case MVT::v4f32: // <2 x f32x2>
287 case MVT::v2i32: // <1 x i32x2>
288 case MVT::v4i32: // <2 x i32x2>
289 if (!STI.hasF32x2Instructions())
290 return std::pair(NumElts, EltVT);
291 PackRegSize = 64;
292 break;
293 }
294
295 // If we reach here, then we can pack 2 or more elements into a single 32-bit
296 // or 64-bit PTX register and treat the vector as a new vector containing
297 // packed elements.
298
299 // Number of elements to pack in one word.
300 const unsigned NPerReg = PackRegSize / EltVT.getSizeInBits();
301
302 return std::pair(NumElts / NPerReg, MVT::getVectorVT(VT: EltVT, NumElements: NPerReg));
303}
304
305/// ComputePTXValueVTs - For the given Type \p Ty, returns the set of primitive
306/// legal-ish MVTs that compose it. Unlike ComputeValueVTs, this will legalize
307/// the types as required by the calling convention (with special handling for
308/// i8s).
309/// NOTE: This is a band-aid for code that expects ComputeValueVTs to return the
310/// same number of types as the Ins/Outs arrays in LowerFormalArguments,
311/// LowerCall, and LowerReturn.
312static void ComputePTXValueVTs(const TargetLowering &TLI, const DataLayout &DL,
313 LLVMContext &Ctx, CallingConv::ID CallConv,
314 Type *Ty, SmallVectorImpl<EVT> &ValueVTs,
315 SmallVectorImpl<uint64_t> &Offsets,
316 uint64_t StartingOffset = 0) {
317 SmallVector<EVT, 16> TempVTs;
318 SmallVector<uint64_t, 16> TempOffsets;
319 ComputeValueVTs(TLI, DL, Ty, ValueVTs&: TempVTs, /*MemVTs=*/nullptr, FixedOffsets: &TempOffsets,
320 StartingOffset);
321
322 for (const auto [VT, Off] : zip(t&: TempVTs, u&: TempOffsets)) {
323 MVT RegisterVT = TLI.getRegisterTypeForCallingConv(Context&: Ctx, CC: CallConv, VT);
324 unsigned NumRegs = TLI.getNumRegistersForCallingConv(Context&: Ctx, CC: CallConv, VT);
325
326 // Since we actually can load/store b8, we need to ensure that we'll use
327 // the original sized type for any i8s or i8 vectors.
328 if (VT.getScalarType() == MVT::i8) {
329 if (RegisterVT == MVT::i16)
330 RegisterVT = MVT::i8;
331 else if (RegisterVT == MVT::v2i16)
332 RegisterVT = MVT::v2i8;
333 else
334 assert(RegisterVT == MVT::v4i8 &&
335 "Expected v4i8, v2i16, or i16 for i8 RegisterVT");
336 }
337
338 // TODO: This is horribly incorrect for cases where the vector elements are
339 // not a multiple of bytes (ex i1) and legal or i8. However, this problem
340 // has existed for as long as NVPTX has and no one has complained, so we'll
341 // leave it for now.
342 for (unsigned I : seq(Size: NumRegs)) {
343 ValueVTs.push_back(Elt: RegisterVT);
344 Offsets.push_back(Elt: Off + I * RegisterVT.getStoreSize());
345 }
346 }
347}
348
349// We return an EVT that can hold N VTs
350// If the VT is a vector, the resulting EVT is a flat vector with the same
351// element type as VT's element type.
352static EVT getVectorizedVT(EVT VT, unsigned N, LLVMContext &C) {
353 if (N == 1)
354 return VT;
355
356 return VT.isVector() ? EVT::getVectorVT(Context&: C, VT: VT.getScalarType(),
357 NumElements: VT.getVectorNumElements() * N)
358 : EVT::getVectorVT(Context&: C, VT, NumElements: N);
359}
360
361static SDValue getExtractVectorizedValue(SDValue V, unsigned I, EVT VT,
362 const SDLoc &dl, SelectionDAG &DAG) {
363 if (V.getValueType() == VT) {
364 assert(I == 0 && "Index must be 0 for scalar value");
365 return V;
366 }
367
368 if (!VT.isVector())
369 return DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL: dl, VT, N1: V,
370 N2: DAG.getVectorIdxConstant(Val: I, DL: dl));
371
372 return DAG.getNode(
373 Opcode: ISD::EXTRACT_SUBVECTOR, DL: dl, VT, N1: V,
374 N2: DAG.getVectorIdxConstant(Val: I * VT.getVectorNumElements(), DL: dl));
375}
376
377template <typename T>
378static inline SDValue getBuildVectorizedValue(unsigned N, const SDLoc &dl,
379 SelectionDAG &DAG, T GetElement) {
380 if (N == 1)
381 return GetElement(0);
382
383 SmallVector<SDValue, 8> Values;
384 for (const unsigned I : llvm::seq(Size: N)) {
385 SDValue Val = GetElement(I);
386 if (Val.getValueType().isVector())
387 DAG.ExtractVectorElements(Op: Val, Args&: Values);
388 else
389 Values.push_back(Elt: Val);
390 }
391
392 EVT VT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: Values[0].getValueType(),
393 NumElements: Values.size());
394 return DAG.getBuildVector(VT, DL: dl, Ops: Values);
395}
396
397/// PromoteScalarIntegerPTX
398/// Used to make sure the arguments/returns are suitable for passing
399/// and promote them to a larger size if they're not.
400///
401/// The promoted type is placed in \p PromoteVT if the function returns true.
402static EVT promoteScalarIntegerPTX(const EVT VT) {
403 if (VT.isScalarInteger()) {
404 switch (PowerOf2Ceil(A: VT.getFixedSizeInBits())) {
405 default:
406 llvm_unreachable(
407 "Promotion is not suitable for scalars of size larger than 64-bits");
408 case 1:
409 return MVT::i1;
410 case 2:
411 case 4:
412 case 8:
413 return MVT::i8;
414 case 16:
415 return MVT::i16;
416 case 32:
417 return MVT::i32;
418 case 64:
419 return MVT::i64;
420 }
421 }
422 return VT;
423}
424
425// Check whether we can merge loads/stores of some of the pieces of a
426// flattened function parameter or return value into a single vector
427// load/store.
428//
429// The flattened parameter is represented as a list of EVTs and
430// offsets, and the whole structure is aligned to ParamAlignment. This
431// function determines whether we can load/store pieces of the
432// parameter starting at index Idx using a single vectorized op of
433// size AccessSize. If so, it returns the number of param pieces
434// covered by the vector op. Otherwise, it returns 1.
435template <typename T>
436static unsigned canMergeParamLoadStoresStartingAt(
437 unsigned Idx, uint32_t AccessSize, const SmallVectorImpl<EVT> &ValueVTs,
438 const SmallVectorImpl<T> &Offsets, Align ParamAlignment) {
439
440 // Can't vectorize if param alignment is not sufficient.
441 if (ParamAlignment < AccessSize)
442 return 1;
443 // Can't vectorize if offset is not aligned.
444 if (Offsets[Idx] & (AccessSize - 1))
445 return 1;
446
447 EVT EltVT = ValueVTs[Idx];
448 unsigned EltSize = EltVT.getStoreSize();
449
450 // Element is too large to vectorize.
451 if (EltSize >= AccessSize)
452 return 1;
453
454 unsigned NumElts = AccessSize / EltSize;
455 // Can't vectorize if AccessBytes if not a multiple of EltSize.
456 if (AccessSize != EltSize * NumElts)
457 return 1;
458
459 // We don't have enough elements to vectorize.
460 if (Idx + NumElts > ValueVTs.size())
461 return 1;
462
463 // PTX ISA can only deal with 2- and 4-element vector ops.
464 if (NumElts != 4 && NumElts != 2)
465 return 1;
466
467 for (unsigned j = Idx + 1; j < Idx + NumElts; ++j) {
468 // Types do not match.
469 if (ValueVTs[j] != EltVT)
470 return 1;
471
472 // Elements are not contiguous.
473 if (Offsets[j] - Offsets[j - 1] != EltSize)
474 return 1;
475 }
476 // OK. We can vectorize ValueVTs[i..i+NumElts)
477 return NumElts;
478}
479
480// Computes whether and how we can vectorize the loads/stores of a
481// flattened function parameter or return value.
482//
483// The flattened parameter is represented as the list of ValueVTs and
484// Offsets, and is aligned to ParamAlignment bytes. We return a vector
485// of the same size as ValueVTs indicating how each piece should be
486// loaded/stored (i.e. as a scalar, or as part of a vector
487// load/store).
488template <typename T>
489static SmallVector<unsigned, 16>
490VectorizePTXValueVTs(const SmallVectorImpl<EVT> &ValueVTs,
491 const SmallVectorImpl<T> &Offsets, Align ParamAlignment,
492 bool IsVAArg = false) {
493 // Set vector size to match ValueVTs and mark all elements as
494 // scalars by default.
495
496 if (IsVAArg)
497 return SmallVector<unsigned>(ValueVTs.size(), 1);
498
499 SmallVector<unsigned, 16> VectorInfo;
500
501 const auto GetNumElts = [&](unsigned I) -> unsigned {
502 for (const unsigned AccessSize : {16, 8, 4, 2}) {
503 const unsigned NumElts = canMergeParamLoadStoresStartingAt(
504 I, AccessSize, ValueVTs, Offsets, ParamAlignment);
505 assert((NumElts == 1 || NumElts == 2 || NumElts == 4) &&
506 "Unexpected vectorization size");
507 if (NumElts != 1)
508 return NumElts;
509 }
510 return 1;
511 };
512
513 // Check what we can vectorize using 128/64/32-bit accesses.
514 for (unsigned I = 0, E = ValueVTs.size(); I != E;) {
515 const unsigned NumElts = GetNumElts(I);
516 VectorInfo.push_back(Elt: NumElts);
517 I += NumElts;
518 }
519 assert(std::accumulate(VectorInfo.begin(), VectorInfo.end(), 0u) ==
520 ValueVTs.size());
521 return VectorInfo;
522}
523
524// NVPTXTargetLowering Constructor.
525NVPTXTargetLowering::NVPTXTargetLowering(const NVPTXTargetMachine &TM,
526 const NVPTXSubtarget &STI)
527 : TargetLowering(TM, STI), STI(STI), GlobalUniqueCallSite(0) {
528 // always lower memset, memcpy, and memmove intrinsics to load/store
529 // instructions, rather
530 // then generating calls to memset, mempcy or memmove.
531 MaxStoresPerMemset = MaxStoresPerMemsetOptSize = (unsigned)0xFFFFFFFF;
532 MaxStoresPerMemcpy = MaxStoresPerMemcpyOptSize = (unsigned) 0xFFFFFFFF;
533 MaxStoresPerMemmove = MaxStoresPerMemmoveOptSize = (unsigned) 0xFFFFFFFF;
534
535 setBooleanContents(ZeroOrNegativeOneBooleanContent);
536 setBooleanVectorContents(ZeroOrNegativeOneBooleanContent);
537
538 // Jump is Expensive. Don't create extra control flow for 'and', 'or'
539 // condition branches.
540 setJumpIsExpensive(true);
541
542 // Wide divides are _very_ slow. Try to reduce the width of the divide if
543 // possible.
544 addBypassSlowDiv(SlowBitWidth: 64, FastBitWidth: 32);
545
546 // By default, use the Source scheduling
547 if (sched4reg)
548 setSchedulingPreference(Sched::RegPressure);
549 else
550 setSchedulingPreference(Sched::Source);
551
552 auto setFP16OperationAction = [&](unsigned Op, MVT VT, LegalizeAction Action,
553 LegalizeAction NoF16Action) {
554 bool IsOpSupported = STI.allowFP16Math();
555 switch (Op) {
556 // Several FP16 instructions are available on sm_80 only.
557 case ISD::FMINNUM:
558 case ISD::FMAXNUM:
559 case ISD::FMAXNUM_IEEE:
560 case ISD::FMINNUM_IEEE:
561 case ISD::FMAXIMUM:
562 case ISD::FMINIMUM:
563 case ISD::FMAXIMUMNUM:
564 case ISD::FMINIMUMNUM:
565 IsOpSupported &= STI.hasFeature(Feature: NVPTX::SM80);
566 break;
567 case ISD::FEXP2:
568 case ISD::FTANH:
569 IsOpSupported &=
570 STI.hasFeature(Feature: NVPTX::SM75) && STI.hasFeature(Feature: NVPTX::PTX70);
571 break;
572 }
573 setOperationAction(Op, VT, Action: IsOpSupported ? Action : NoF16Action);
574 };
575
576 auto setBF16OperationAction = [&](unsigned Op, MVT VT, LegalizeAction Action,
577 LegalizeAction NoBF16Action) {
578 bool IsOpSupported = STI.hasNativeBF16Support(Opcode: Op);
579 setOperationAction(
580 Op, VT, Action: IsOpSupported ? Action : NoBF16Action);
581 };
582
583 auto setI16x2OperationAction = [&](unsigned Op, MVT VT, LegalizeAction Action,
584 LegalizeAction NoI16x2Action) {
585 bool IsOpSupported = false;
586 // instructions are available on sm_90 only
587 switch (Op) {
588 case ISD::ADD:
589 case ISD::SMAX:
590 case ISD::SMIN:
591 case ISD::UMIN:
592 case ISD::UMAX:
593 IsOpSupported =
594 STI.hasFeature(Feature: NVPTX::SM90) && STI.hasFeature(Feature: NVPTX::PTX80);
595 break;
596 }
597 setOperationAction(Op, VT, Action: IsOpSupported ? Action : NoI16x2Action);
598 };
599
600 addRegisterClass(VT: MVT::i1, RC: &NVPTX::B1RegClass);
601 addRegisterClass(VT: MVT::i16, RC: &NVPTX::B16RegClass);
602 addRegisterClass(VT: MVT::v2i16, RC: &NVPTX::B32RegClass);
603 addRegisterClass(VT: MVT::v4i8, RC: &NVPTX::B32RegClass);
604 addRegisterClass(VT: MVT::i32, RC: &NVPTX::B32RegClass);
605 addRegisterClass(VT: MVT::i64, RC: &NVPTX::B64RegClass);
606 addRegisterClass(VT: MVT::f32, RC: &NVPTX::B32RegClass);
607 addRegisterClass(VT: MVT::f64, RC: &NVPTX::B64RegClass);
608 addRegisterClass(VT: MVT::f16, RC: &NVPTX::B16RegClass);
609 addRegisterClass(VT: MVT::v2f16, RC: &NVPTX::B32RegClass);
610 addRegisterClass(VT: MVT::bf16, RC: &NVPTX::B16RegClass);
611 addRegisterClass(VT: MVT::v2bf16, RC: &NVPTX::B32RegClass);
612
613 if (STI.hasF32x2Instructions()) {
614 addRegisterClass(VT: MVT::v2f32, RC: &NVPTX::B64RegClass);
615 addRegisterClass(VT: MVT::v2i32, RC: &NVPTX::B64RegClass);
616 }
617
618 // Conversion to/from FP16/FP16x2 is always legal.
619 setOperationAction(Op: ISD::BUILD_VECTOR, VT: MVT::v2f16, Action: Custom);
620 setOperationAction(Op: ISD::EXTRACT_VECTOR_ELT, VT: MVT::v2f16, Action: Custom);
621 setOperationAction(Op: ISD::INSERT_VECTOR_ELT, VT: MVT::v2f16, Action: Expand);
622 setOperationAction(Op: ISD::VECTOR_SHUFFLE, VT: MVT::v2f16, Action: Expand);
623
624 setOperationAction(Op: ISD::READCYCLECOUNTER, VT: MVT::i64, Action: Legal);
625 if (STI.hasFeature(Feature: NVPTX::SM30))
626 setOperationAction(Op: ISD::READSTEADYCOUNTER, VT: MVT::i64, Action: Legal);
627
628 setFP16OperationAction(ISD::SETCC, MVT::f16, Legal, Promote);
629 setFP16OperationAction(ISD::SETCC, MVT::v2f16, Legal, Expand);
630
631 // Conversion to/from BFP16/BFP16x2 is always legal.
632 setOperationAction(Op: ISD::BUILD_VECTOR, VT: MVT::v2bf16, Action: Custom);
633 setOperationAction(Op: ISD::EXTRACT_VECTOR_ELT, VT: MVT::v2bf16, Action: Custom);
634 setOperationAction(Op: ISD::INSERT_VECTOR_ELT, VT: MVT::v2bf16, Action: Expand);
635 setOperationAction(Op: ISD::VECTOR_SHUFFLE, VT: MVT::v2bf16, Action: Expand);
636
637 setBF16OperationAction(ISD::SETCC, MVT::v2bf16, Legal, Expand);
638 setBF16OperationAction(ISD::SETCC, MVT::bf16, Legal, Promote);
639 if (getOperationAction(Op: ISD::SETCC, VT: MVT::bf16) == Promote)
640 AddPromotedToType(Opc: ISD::SETCC, OrigVT: MVT::bf16, DestVT: MVT::f32);
641
642 // Conversion to/from i16/i16x2 is always legal.
643 setOperationAction(Op: ISD::BUILD_VECTOR, VT: MVT::v2i16, Action: Custom);
644 setOperationAction(Op: ISD::EXTRACT_VECTOR_ELT, VT: MVT::v2i16, Action: Custom);
645 setOperationAction(Op: ISD::INSERT_VECTOR_ELT, VT: MVT::v2i16, Action: Expand);
646 setOperationAction(Op: ISD::VECTOR_SHUFFLE, VT: MVT::v2i16, Action: Expand);
647
648 setOperationAction(Op: ISD::BUILD_VECTOR, VT: MVT::v4i8, Action: Custom);
649 setOperationAction(Op: ISD::EXTRACT_VECTOR_ELT, VT: MVT::v4i8, Action: Custom);
650 setOperationAction(Op: ISD::INSERT_VECTOR_ELT, VT: MVT::v4i8, Action: Custom);
651 setOperationAction(Op: ISD::VECTOR_SHUFFLE, VT: MVT::v4i8, Action: Custom);
652
653 // No support for these operations with v2f32/v2i32
654 setOperationAction(Ops: ISD::INSERT_VECTOR_ELT, VTs: {MVT::v2f32, MVT::v2i32}, Action: Expand);
655 setOperationAction(Ops: ISD::VECTOR_SHUFFLE, VTs: {MVT::v2f32, MVT::v2i32}, Action: Expand);
656
657 setOperationAction(Op: ISD::TRUNCATE, VT: MVT::v2i16, Action: Expand);
658 setOperationAction(Ops: {ISD::ANY_EXTEND, ISD::ZERO_EXTEND, ISD::SIGN_EXTEND},
659 VT: MVT::v2i32, Action: Expand);
660
661 // Need custom lowering in case the index is dynamic.
662 if (STI.hasF32x2Instructions())
663 setOperationAction(Ops: ISD::EXTRACT_VECTOR_ELT, VTs: {MVT::v2f32, MVT::v2i32},
664 Action: Custom);
665
666 // Custom conversions to/from v2i8.
667 setOperationAction(Op: ISD::BITCAST, VT: MVT::v2i8, Action: Custom);
668
669 // Only logical ops can be done on v4i8/v2i32 directly, others must be done
670 // elementwise.
671 setOperationAction(
672 Ops: {ISD::ABS, ISD::ADD, ISD::ADDC, ISD::ADDE,
673 ISD::BITREVERSE, ISD::CTLZ, ISD::CTPOP, ISD::CTTZ,
674 ISD::FP_TO_SINT, ISD::FP_TO_UINT, ISD::FSHL, ISD::FSHR,
675 ISD::MUL, ISD::MULHS, ISD::MULHU, ISD::PARITY,
676 ISD::ROTL, ISD::ROTR, ISD::SADDO, ISD::SADDO_CARRY,
677 ISD::SADDSAT, ISD::SDIV, ISD::SDIVREM, ISD::SELECT_CC,
678 ISD::SETCC, ISD::SHL, ISD::SINT_TO_FP, ISD::SMAX,
679 ISD::SMIN, ISD::SMULO, ISD::SMUL_LOHI, ISD::SRA,
680 ISD::SREM, ISD::SRL, ISD::SSHLSAT, ISD::SSUBO,
681 ISD::SSUBO_CARRY, ISD::SSUBSAT, ISD::SUB, ISD::SUBC,
682 ISD::SUBE, ISD::UADDO, ISD::UADDO_CARRY, ISD::UADDSAT,
683 ISD::UDIV, ISD::UDIVREM, ISD::UINT_TO_FP, ISD::UMAX,
684 ISD::UMIN, ISD::UMULO, ISD::UMUL_LOHI, ISD::UREM,
685 ISD::USHLSAT, ISD::USUBO, ISD::USUBO_CARRY, ISD::VSELECT,
686 ISD::USUBSAT},
687 VTs: {MVT::v4i8, MVT::v2i32}, Action: Expand);
688
689 // Operations not directly supported by NVPTX.
690 for (MVT VT : {MVT::bf16, MVT::f16, MVT::v2bf16, MVT::v2f16, MVT::f32,
691 MVT::v2f32, MVT::f64, MVT::i1, MVT::i8, MVT::i16, MVT::v2i16,
692 MVT::v4i8, MVT::i32, MVT::v2i32, MVT::i64}) {
693 setOperationAction(Op: ISD::SELECT_CC, VT, Action: Expand);
694 setOperationAction(Op: ISD::BR_CC, VT, Action: Expand);
695 }
696
697 setOperationAction(Ops: ISD::SDIVREM, VTs: {MVT::i32, MVT::i64}, Action: Expand);
698 setOperationAction(Ops: ISD::UDIVREM, VTs: {MVT::i32, MVT::i64}, Action: Expand);
699
700 // We don't want ops like FMINIMUM or UMAX to be lowered to SETCC+VSELECT.
701 setOperationAction(Ops: ISD::VSELECT, VTs: {MVT::v2f32, MVT::v2i32}, Action: Expand);
702
703 // Some SIGN_EXTEND_INREG can be done using cvt instruction.
704 // For others we will expand to a SHL/SRA pair.
705 setOperationAction(Op: ISD::SIGN_EXTEND_INREG, VT: MVT::i64, Action: Legal);
706 setOperationAction(Op: ISD::SIGN_EXTEND_INREG, VT: MVT::i32, Action: Legal);
707 setOperationAction(Op: ISD::SIGN_EXTEND_INREG, VT: MVT::i16, Action: Legal);
708 setOperationAction(Op: ISD::SIGN_EXTEND_INREG, VT: MVT::i8 , Action: Legal);
709 setOperationAction(Op: ISD::SIGN_EXTEND_INREG, VT: MVT::i1, Action: Expand);
710 setOperationAction(Ops: ISD::SIGN_EXTEND_INREG, VTs: {MVT::v2i16, MVT::v2i32}, Action: Expand);
711
712 setOperationAction(Op: ISD::SHL_PARTS, VT: MVT::i32 , Action: Custom);
713 setOperationAction(Op: ISD::SRA_PARTS, VT: MVT::i32 , Action: Custom);
714 setOperationAction(Op: ISD::SRL_PARTS, VT: MVT::i32 , Action: Custom);
715 setOperationAction(Op: ISD::SHL_PARTS, VT: MVT::i64 , Action: Custom);
716 setOperationAction(Op: ISD::SRA_PARTS, VT: MVT::i64 , Action: Custom);
717 setOperationAction(Op: ISD::SRL_PARTS, VT: MVT::i64 , Action: Custom);
718
719 if (STI.hasCLMAD())
720 setOperationAction(Ops: {ISD::CLMUL, ISD::CLMULH}, VT: MVT::i64, Action: Legal);
721 setOperationAction(Op: ISD::BITREVERSE, VT: MVT::i32, Action: Legal);
722 setOperationAction(Op: ISD::BITREVERSE, VT: MVT::i64, Action: Legal);
723
724 setOperationAction(Ops: {ISD::ROTL, ISD::ROTR},
725 VTs: {MVT::i8, MVT::i16, MVT::v2i16, MVT::i32, MVT::i64},
726 Action: Expand);
727
728 if (STI.hasHWROT32()) {
729 setOperationAction(Ops: {ISD::FSHL, ISD::FSHR}, VT: MVT::i32, Action: Legal);
730 setOperationAction(Ops: {ISD::ROTL, ISD::ROTR, ISD::FSHL, ISD::FSHR}, VT: MVT::i64,
731 Action: Custom);
732 }
733
734 setOperationAction(Op: ISD::BR_JT, VT: MVT::Other, Action: STI.hasBrx() ? Legal : Expand);
735 setOperationAction(Op: ISD::BRIND, VT: MVT::Other, Action: Expand);
736
737 // We want to legalize constant related memmove and memcopy
738 // intrinsics.
739 setOperationAction(Op: ISD::INTRINSIC_W_CHAIN, VT: MVT::Other, Action: Custom);
740
741 // FP extload/truncstore is not legal in PTX. We need to expand all these.
742 for (auto FloatVTs :
743 {MVT::fp_valuetypes(), MVT::fp_fixedlen_vector_valuetypes()}) {
744 for (MVT ValVT : FloatVTs) {
745 for (MVT MemVT : FloatVTs) {
746 setLoadExtAction(ExtType: ISD::EXTLOAD, ValVT, MemVT, Action: Expand);
747 setTruncStoreAction(ValVT, MemVT, Action: Expand);
748 }
749 }
750 }
751
752 // To improve CodeGen we'll legalize any-extend loads to zext loads. This is
753 // how they'll be lowered in ISel anyway, and by doing this a little earlier
754 // we allow for more DAG combine opportunities.
755 for (auto IntVTs :
756 {MVT::integer_valuetypes(), MVT::integer_fixedlen_vector_valuetypes()})
757 for (MVT ValVT : IntVTs)
758 for (MVT MemVT : IntVTs)
759 if (isTypeLegal(VT: ValVT))
760 setLoadExtAction(ExtType: ISD::EXTLOAD, ValVT, MemVT, Action: Custom);
761
762 // PTX does not support load / store predicate registers
763 setOperationAction(Ops: {ISD::LOAD, ISD::STORE}, VT: MVT::i1, Action: Custom);
764 for (MVT VT : MVT::integer_valuetypes()) {
765 setLoadExtAction(ExtTypes: {ISD::SEXTLOAD, ISD::ZEXTLOAD, ISD::EXTLOAD}, ValVT: VT, MemVT: MVT::i1,
766 Action: Promote);
767 setTruncStoreAction(ValVT: VT, MemVT: MVT::i1, Action: Expand);
768 }
769
770 // Disable generations of extload/truncstore for v2i32/v2i16/v2i8. The generic
771 // expansion for these nodes when they are unaligned is incorrect if the
772 // type is a vector.
773 //
774 // TODO: Fix the generic expansion for these nodes found in
775 // TargetLowering::expandUnalignedLoad/Store.
776 setLoadExtAction(ExtTypes: {ISD::EXTLOAD, ISD::SEXTLOAD, ISD::ZEXTLOAD}, ValVT: MVT::v2i16,
777 MemVT: MVT::v2i8, Action: Expand);
778 setLoadExtAction(ExtTypes: {ISD::EXTLOAD, ISD::SEXTLOAD, ISD::ZEXTLOAD}, ValVT: MVT::v2i32,
779 MemVTs: {MVT::v2i8, MVT::v2i16}, Action: Expand);
780 setTruncStoreAction(ValVT: MVT::v2i16, MemVT: MVT::v2i8, Action: Expand);
781 setTruncStoreAction(ValVT: MVT::v2i32, MemVT: MVT::v2i16, Action: Expand);
782 setTruncStoreAction(ValVT: MVT::v2i32, MemVT: MVT::v2i8, Action: Expand);
783
784 // Register custom handling for illegal type loads/stores. We'll try to custom
785 // lower almost all illegal types and logic in the lowering will discard cases
786 // we can't handle.
787 setOperationAction(Ops: {ISD::LOAD, ISD::STORE}, VTs: {MVT::i128, MVT::i256, MVT::f128},
788 Action: Custom);
789 for (MVT VT : MVT::fixedlen_vector_valuetypes())
790 if (!isTypeLegal(VT) && VT.getStoreSizeInBits() <= 256)
791 setOperationAction(Ops: {ISD::STORE, ISD::LOAD, ISD::MSTORE, ISD::MLOAD}, VT,
792 Action: Custom);
793
794 // Custom legalization for LDU intrinsics.
795 // TODO: The logic to lower these is not very robust and we should rewrite it.
796 // Perhaps LDU should not be represented as an intrinsic at all.
797 setOperationAction(Op: ISD::INTRINSIC_W_CHAIN, VT: MVT::i8, Action: Custom);
798 for (MVT VT : MVT::fixedlen_vector_valuetypes())
799 if (IsPTXVectorType(VT))
800 setOperationAction(Op: ISD::INTRINSIC_W_CHAIN, VT, Action: Custom);
801
802 setCondCodeAction(CCs: {ISD::SETNE, ISD::SETEQ, ISD::SETUGE, ISD::SETULE,
803 ISD::SETUGT, ISD::SETULT, ISD::SETGT, ISD::SETLT,
804 ISD::SETGE, ISD::SETLE},
805 VT: MVT::i1, Action: Expand);
806
807 // This is legal in NVPTX
808 setOperationAction(Op: ISD::ConstantFP, VT: MVT::f64, Action: Legal);
809 setOperationAction(Op: ISD::ConstantFP, VT: MVT::f32, Action: Legal);
810 setOperationAction(Op: ISD::ConstantFP, VT: MVT::f16, Action: Legal);
811 setOperationAction(Op: ISD::ConstantFP, VT: MVT::bf16, Action: Legal);
812
813 setOperationAction(Ops: ISD::DYNAMIC_STACKALLOC, VTs: {MVT::i32, MVT::i64}, Action: Custom);
814 setOperationAction(Ops: {ISD::STACKRESTORE, ISD::STACKSAVE}, VT: MVT::Other, Action: Custom);
815
816 // TRAP can be lowered to PTX trap
817 setOperationAction(Op: ISD::TRAP, VT: MVT::Other, Action: Legal);
818 // DEBUGTRAP can be lowered to PTX brkpt
819 setOperationAction(Op: ISD::DEBUGTRAP, VT: MVT::Other, Action: Legal);
820
821 // Support varargs.
822 setOperationAction(Op: ISD::VASTART, VT: MVT::Other, Action: Custom);
823 setOperationAction(Op: ISD::VAARG, VT: MVT::Other, Action: Custom);
824 setOperationAction(Op: ISD::VACOPY, VT: MVT::Other, Action: Expand);
825 setOperationAction(Op: ISD::VAEND, VT: MVT::Other, Action: Expand);
826
827 setOperationAction(Ops: {ISD::SMIN, ISD::SMAX, ISD::UMIN, ISD::UMAX},
828 VTs: {MVT::i16, MVT::i32, MVT::i64}, Action: Legal);
829 // PTX abs.s is undefined for INT_MIN, so ISD::ABS (which requires
830 // abs(INT_MIN) == INT_MIN) must be expanded. ABS_MIN_POISON matches
831 // PTX abs semantics since INT_MIN input is poison/undefined.
832 setOperationAction(Ops: ISD::ABS, VTs: {MVT::i16, MVT::i32, MVT::i64}, Action: Expand);
833 setOperationAction(Ops: ISD::ABS_MIN_POISON, VTs: {MVT::i16, MVT::i32, MVT::i64},
834 Action: Legal);
835
836 setOperationAction(Ops: {ISD::CTPOP, ISD::CTLZ, ISD::CTLZ_ZERO_POISON}, VT: MVT::i16,
837 Action: Promote);
838 setOperationAction(Ops: {ISD::CTPOP, ISD::CTLZ}, VT: MVT::i32, Action: Legal);
839 setOperationAction(Ops: {ISD::CTPOP, ISD::CTLZ}, VT: MVT::i64, Action: Custom);
840
841 setI16x2OperationAction(ISD::ABS_MIN_POISON, MVT::v2i16, Legal, Custom);
842 setI16x2OperationAction(ISD::SMIN, MVT::v2i16, Legal, Custom);
843 setI16x2OperationAction(ISD::SMAX, MVT::v2i16, Legal, Custom);
844 setI16x2OperationAction(ISD::UMIN, MVT::v2i16, Legal, Custom);
845 setI16x2OperationAction(ISD::UMAX, MVT::v2i16, Legal, Custom);
846 setI16x2OperationAction(ISD::CTPOP, MVT::v2i16, Legal, Expand);
847 setI16x2OperationAction(ISD::CTLZ, MVT::v2i16, Legal, Expand);
848
849 setI16x2OperationAction(ISD::ADD, MVT::v2i16, Legal, Custom);
850 setI16x2OperationAction(ISD::SUB, MVT::v2i16, Legal, Custom);
851 setI16x2OperationAction(ISD::MUL, MVT::v2i16, Legal, Custom);
852 setI16x2OperationAction(ISD::SHL, MVT::v2i16, Legal, Custom);
853 setI16x2OperationAction(ISD::SREM, MVT::v2i16, Legal, Custom);
854 setI16x2OperationAction(ISD::UREM, MVT::v2i16, Legal, Custom);
855
856 // Other arithmetic and logic ops are unsupported.
857 setOperationAction(Ops: {ISD::SDIV, ISD::UDIV, ISD::SRA, ISD::SRL, ISD::MULHS,
858 ISD::MULHU, ISD::FP_TO_SINT, ISD::FP_TO_UINT,
859 ISD::SINT_TO_FP, ISD::UINT_TO_FP, ISD::SETCC},
860 VTs: {MVT::v2i16, MVT::v2i32}, Action: Expand);
861
862 // v2i32 is not supported for any arithmetic operations
863 setOperationAction(Ops: {ISD::ABS, ISD::SMIN, ISD::SMAX, ISD::UMIN, ISD::UMAX,
864 ISD::CTPOP, ISD::CTLZ, ISD::ADD, ISD::SUB, ISD::MUL,
865 ISD::SHL, ISD::SRA, ISD::SRL, ISD::OR, ISD::AND, ISD::XOR,
866 ISD::SREM, ISD::UREM},
867 VT: MVT::v2i32, Action: Expand);
868
869 setOperationAction(Op: ISD::ADDC, VT: MVT::i32, Action: Legal);
870 setOperationAction(Op: ISD::ADDE, VT: MVT::i32, Action: Legal);
871 setOperationAction(Op: ISD::SUBC, VT: MVT::i32, Action: Legal);
872 setOperationAction(Op: ISD::SUBE, VT: MVT::i32, Action: Legal);
873 if (STI.hasFeature(Feature: NVPTX::PTX43)) {
874 setOperationAction(Op: ISD::ADDC, VT: MVT::i64, Action: Legal);
875 setOperationAction(Op: ISD::ADDE, VT: MVT::i64, Action: Legal);
876 setOperationAction(Op: ISD::SUBC, VT: MVT::i64, Action: Legal);
877 setOperationAction(Op: ISD::SUBE, VT: MVT::i64, Action: Legal);
878 }
879
880 setOperationAction(Op: ISD::CTTZ, VT: MVT::i16, Action: Expand);
881 setOperationAction(Ops: ISD::CTTZ, VTs: {MVT::v2i16, MVT::v2i32}, Action: Expand);
882 setOperationAction(Op: ISD::CTTZ, VT: MVT::i32, Action: Expand);
883 setOperationAction(Op: ISD::CTTZ, VT: MVT::i64, Action: Expand);
884
885 // PTX does not directly support SELP of i1, so promote to i32 first
886 setOperationAction(Op: ISD::SELECT, VT: MVT::i1, Action: Custom);
887
888 // PTX cannot multiply two i64s in a single instruction.
889 setOperationAction(Op: ISD::SMUL_LOHI, VT: MVT::i64, Action: Expand);
890 setOperationAction(Op: ISD::UMUL_LOHI, VT: MVT::i64, Action: Expand);
891
892 // We have some custom DAG combine patterns for these nodes
893 setTargetDAGCombine({ISD::ADD,
894 ISD::AND,
895 ISD::EXTRACT_VECTOR_ELT,
896 ISD::FADD,
897 ISD::FMAXNUM,
898 ISD::FMINNUM,
899 ISD::FMAXIMUM,
900 ISD::FMINIMUM,
901 ISD::FMAXIMUMNUM,
902 ISD::FMINIMUMNUM,
903 ISD::MUL,
904 ISD::SELECT,
905 ISD::SHL,
906 ISD::SREM,
907 ISD::UREM,
908 ISD::VSELECT,
909 ISD::BUILD_VECTOR,
910 ISD::ADDRSPACECAST,
911 ISD::LOAD,
912 ISD::STORE,
913 ISD::ZERO_EXTEND,
914 ISD::SIGN_EXTEND,
915 ISD::INTRINSIC_WO_CHAIN});
916
917 // If the vector operands require register coalescing, scalarize instead
918 if (STI.hasF32x2Instructions())
919 setTargetDAGCombine({ISD::FMA, ISD::FMUL, ISD::FSUB});
920
921 // setcc for f16x2 and bf16x2 needs special handling to prevent
922 // legalizer's attempt to scalarize it due to v2i1 not being legal.
923 if (STI.allowFP16Math() || STI.hasBF16Math())
924 setTargetDAGCombine(ISD::SETCC);
925
926 // Vector reduction operations. These may be turned into shuffle or tree
927 // reductions depending on what instructions are available for each type.
928 for (MVT VT : MVT::fixedlen_vector_valuetypes()) {
929 MVT EltVT = VT.getVectorElementType();
930 if (EltVT == MVT::f32 || EltVT == MVT::f64) {
931 setOperationAction(Ops: {ISD::VECREDUCE_FMAX, ISD::VECREDUCE_FMIN,
932 ISD::VECREDUCE_FMAXIMUM, ISD::VECREDUCE_FMINIMUM},
933 VT, Action: Custom);
934 }
935 }
936
937 // Promote fp16 arithmetic if fp16 hardware isn't available or the
938 // user passed --nvptx-no-fp16-math. The flag is useful because,
939 // although sm_53+ GPUs have some sort of FP16 support in
940 // hardware, only sm_53 and sm_60 have full implementation. Others
941 // only have token amount of hardware and are likely to run faster
942 // by using fp32 units instead.
943 for (const auto &Op : {ISD::FADD, ISD::FMUL, ISD::FSUB, ISD::FMA}) {
944 setFP16OperationAction(Op, MVT::f16, Legal, Promote);
945 setFP16OperationAction(Op, MVT::v2f16, Legal, Expand);
946 setBF16OperationAction(Op, MVT::v2bf16, Legal, Expand);
947 // bf16 must be promoted to f32.
948 setBF16OperationAction(Op, MVT::bf16, Legal, Promote);
949 if (getOperationAction(Op, VT: MVT::bf16) == Promote)
950 AddPromotedToType(Opc: Op, OrigVT: MVT::bf16, DestVT: MVT::f32);
951 setOperationAction(Op, VT: MVT::v2f32,
952 Action: STI.hasF32x2Instructions() ? Legal : Expand);
953 }
954
955 // On SM80, we select add/mul/sub as fma to avoid promotion to float
956 for (const auto &Op : {ISD::FADD, ISD::FMUL, ISD::FSUB}) {
957 for (const auto &VT : {MVT::bf16, MVT::v2bf16}) {
958 if (!STI.hasNativeBF16Support(Opcode: Op) && STI.hasNativeBF16Support(Opcode: ISD::FMA)) {
959 setOperationAction(Op, VT, Action: Custom);
960 }
961 }
962 }
963
964 // f16/f16x2 neg was introduced in PTX 60, SM_53.
965 const bool IsFP16FP16x2NegAvailable = STI.hasFeature(Feature: NVPTX::SM53) &&
966 STI.hasFeature(Feature: NVPTX::PTX60) &&
967 STI.allowFP16Math();
968 for (const auto &VT : {MVT::f16, MVT::v2f16})
969 setOperationAction(Op: ISD::FNEG, VT,
970 Action: IsFP16FP16x2NegAvailable ? Legal : Expand);
971
972 setBF16OperationAction(ISD::FNEG, MVT::bf16, Legal, Expand);
973 setBF16OperationAction(ISD::FNEG, MVT::v2bf16, Legal, Expand);
974 setOperationAction(Op: ISD::FNEG, VT: MVT::v2f32, Action: Expand);
975 // (would be) Library functions.
976
977 // These map to conversion instructions for scalar FP types.
978 for (const auto &Op : {ISD::FCEIL, ISD::FFLOOR, ISD::FNEARBYINT, ISD::FRINT,
979 ISD::FROUNDEVEN, ISD::FTRUNC}) {
980 setOperationAction(Op, VT: MVT::f16, Action: Legal);
981 setOperationAction(Op, VT: MVT::f32, Action: Legal);
982 setOperationAction(Op, VT: MVT::f64, Action: Legal);
983 setOperationAction(Op, VT: MVT::v2f16, Action: Expand);
984 setOperationAction(Op, VT: MVT::v2bf16, Action: Expand);
985 setOperationAction(Op, VT: MVT::v2f32, Action: Expand);
986 setBF16OperationAction(Op, MVT::bf16, Legal, Promote);
987 if (getOperationAction(Op, VT: MVT::bf16) == Promote)
988 AddPromotedToType(Opc: Op, OrigVT: MVT::bf16, DestVT: MVT::f32);
989 }
990
991 if (!STI.hasFeature(Feature: NVPTX::SM80) || !STI.hasFeature(Feature: NVPTX::PTX71)) {
992 setOperationAction(Op: ISD::BF16_TO_FP, VT: MVT::f32, Action: Expand);
993 }
994 if (!STI.hasFeature(Feature: NVPTX::SM90)) {
995 for (MVT VT : {MVT::bf16, MVT::f32, MVT::f64}) {
996 setOperationAction(Op: ISD::FP_EXTEND, VT, Action: Custom);
997 setOperationAction(Op: ISD::FP_ROUND, VT, Action: Custom);
998 }
999 }
1000
1001 // Expand nearest-even rounding and diagnose unsupported conversions.
1002 setOperationAction(Ops: ISD::FPTRUNC_ROUND, VTs: {MVT::f16, MVT::bf16, MVT::f32},
1003 Action: Custom);
1004
1005 // Expand v2f32 = fp_extend
1006 setOperationAction(Op: ISD::FP_EXTEND, VT: MVT::v2f32, Action: Expand);
1007 // Expand v2[b]f16 = fp_round v2f32
1008 setOperationAction(Ops: ISD::FP_ROUND, VTs: {MVT::v2bf16, MVT::v2f16}, Action: Expand);
1009
1010 // sm_80 only has conversions between f32 and bf16. Custom lower all other
1011 // bf16 conversions.
1012 if (!STI.hasFeature(Feature: NVPTX::SM90)) {
1013 for (MVT VT : {MVT::i1, MVT::i16, MVT::i32, MVT::i64}) {
1014 setOperationAction(
1015 Ops: {ISD::SINT_TO_FP, ISD::UINT_TO_FP, ISD::FP_TO_SINT, ISD::FP_TO_UINT},
1016 VT, Action: Custom);
1017 }
1018 setOperationAction(
1019 Ops: {ISD::SINT_TO_FP, ISD::UINT_TO_FP, ISD::FP_TO_SINT, ISD::FP_TO_UINT},
1020 VT: MVT::bf16, Action: Custom);
1021 }
1022
1023 setOperationAction(Ops: {ISD::FP_TO_SINT, ISD::FP_TO_UINT}, VT: MVT::i1, Action: Custom);
1024 setOperationAction(Op: ISD::FROUND, VT: MVT::f16, Action: Promote);
1025 setOperationAction(Op: ISD::FROUND, VT: MVT::v2f16, Action: Expand);
1026 setOperationAction(Op: ISD::FROUND, VT: MVT::v2bf16, Action: Expand);
1027 setOperationAction(Op: ISD::FROUND, VT: MVT::f32, Action: Custom);
1028 setOperationAction(Op: ISD::FROUND, VT: MVT::f64, Action: Custom);
1029 setOperationAction(Op: ISD::FROUND, VT: MVT::bf16, Action: Promote);
1030 AddPromotedToType(Opc: ISD::FROUND, OrigVT: MVT::bf16, DestVT: MVT::f32);
1031
1032 setOperationAction(Ops: {ISD::LROUND, ISD::LLROUND}, VTs: {MVT::f32, MVT::f64}, Action: Expand);
1033
1034 // 'Expand' implements FCOPYSIGN without calling an external library.
1035 setOperationAction(Op: ISD::FCOPYSIGN, VT: MVT::f16, Action: Expand);
1036 setOperationAction(Op: ISD::FCOPYSIGN, VT: MVT::v2f16, Action: Expand);
1037 setOperationAction(Op: ISD::FCOPYSIGN, VT: MVT::bf16, Action: Expand);
1038 setOperationAction(Op: ISD::FCOPYSIGN, VT: MVT::v2bf16, Action: Expand);
1039 setOperationAction(Op: ISD::FCOPYSIGN, VT: MVT::f32, Action: Custom);
1040 setOperationAction(Op: ISD::FCOPYSIGN, VT: MVT::f64, Action: Custom);
1041
1042 // These map to corresponding instructions for f32/f64. f16 must be
1043 // promoted to f32. v2f16 is expanded to f16, which is then promoted
1044 // to f32.
1045 for (const auto &Op : {ISD::FDIV, ISD::FSQRT, ISD::FSIN, ISD::FCOS}) {
1046 setOperationAction(Op, VT: MVT::f16, Action: Promote);
1047 setOperationAction(Op, VT: MVT::f32, Action: Legal);
1048 // Only div/sqrt are legal for f64.
1049 if (Op == ISD::FDIV || Op == ISD::FSQRT) {
1050 setOperationAction(Op, VT: MVT::f64, Action: Legal);
1051 }
1052 setOperationAction(Ops: Op, VTs: {MVT::v2f16, MVT::v2bf16, MVT::v2f32}, Action: Expand);
1053 setOperationAction(Op, VT: MVT::bf16, Action: Promote);
1054 AddPromotedToType(Opc: Op, OrigVT: MVT::bf16, DestVT: MVT::f32);
1055 }
1056 // Expand remainders at the IR level using ExpandIRInsts.
1057 setOperationAction(Ops: ISD::FREM, VTs: {MVT::f16, MVT::bf16, MVT::f32, MVT::f64},
1058 Action: Expand);
1059
1060 // FTANH support:
1061 // - f32 (sm_75+, PTX 7.0+)
1062 // - f16/f16x2 (sm_75+, PTX 7.0+)
1063 // - bf16/bf16x2 (sm_90+, PTX 7.8+)
1064 // When f16/bf16 types aren't supported, they are promoted/expanded to f32.
1065 if (STI.hasFeature(Feature: NVPTX::SM75) && STI.hasFeature(Feature: NVPTX::PTX70))
1066 setOperationAction(Op: ISD::FTANH, VT: MVT::f32, Action: Legal);
1067 setOperationAction(Op: ISD::FTANH, VT: MVT::v2f32, Action: Expand);
1068
1069 // Scalar f16/bf16: promote to f32 when not natively supported.
1070 setFP16OperationAction(ISD::FTANH, MVT::f16, Legal, Promote);
1071 setBF16OperationAction(ISD::FTANH, MVT::bf16, Legal, Promote);
1072 if (getOperationAction(Op: ISD::FTANH, VT: MVT::bf16) == Promote)
1073 AddPromotedToType(Opc: ISD::FTANH, OrigVT: MVT::bf16, DestVT: MVT::f32);
1074
1075 // Vector v2f16/v2bf16: expand when not natively supported.
1076 setFP16OperationAction(ISD::FTANH, MVT::v2f16, Legal, Expand);
1077 setBF16OperationAction(ISD::FTANH, MVT::v2bf16, Legal, Expand);
1078
1079 setOperationAction(Ops: ISD::FABS, VTs: {MVT::f32, MVT::f64}, Action: Legal);
1080 setOperationAction(Op: ISD::FABS, VT: MVT::v2f32, Action: Expand);
1081 if (STI.hasFeature(Feature: NVPTX::PTX65)) {
1082 setFP16OperationAction(ISD::FABS, MVT::f16, Legal, Promote);
1083 setFP16OperationAction(ISD::FABS, MVT::v2f16, Legal, Expand);
1084 } else {
1085 setOperationAction(Op: ISD::FABS, VT: MVT::f16, Action: Promote);
1086 setOperationAction(Op: ISD::FABS, VT: MVT::v2f16, Action: Expand);
1087 }
1088 setBF16OperationAction(ISD::FABS, MVT::v2bf16, Legal, Expand);
1089 setBF16OperationAction(ISD::FABS, MVT::bf16, Legal, Promote);
1090 if (getOperationAction(Op: ISD::FABS, VT: MVT::bf16) == Promote)
1091 AddPromotedToType(Opc: ISD::FABS, OrigVT: MVT::bf16, DestVT: MVT::f32);
1092
1093 for (const auto &Op :
1094 {ISD::FMINNUM, ISD::FMAXNUM, ISD::FMINIMUMNUM, ISD::FMAXIMUMNUM}) {
1095 setOperationAction(Op, VT: MVT::f32, Action: Legal);
1096 setOperationAction(Op, VT: MVT::f64, Action: Legal);
1097 setFP16OperationAction(Op, MVT::f16, Legal, Promote);
1098 setFP16OperationAction(Op, MVT::v2f16, Legal, Expand);
1099 setBF16OperationAction(Op, MVT::v2bf16, Legal, Expand);
1100 setBF16OperationAction(Op, MVT::bf16, Legal, Promote);
1101 if (getOperationAction(Op, VT: MVT::bf16) == Promote)
1102 AddPromotedToType(Opc: Op, OrigVT: MVT::bf16, DestVT: MVT::f32);
1103 setOperationAction(Op, VT: MVT::v2f32, Action: Expand);
1104 }
1105 bool SupportsF32MinMaxNaN = STI.hasFeature(Feature: NVPTX::SM80);
1106 for (const auto &Op : {ISD::FMINIMUM, ISD::FMAXIMUM}) {
1107 setOperationAction(Op, VT: MVT::f32, Action: SupportsF32MinMaxNaN ? Legal : Expand);
1108 setFP16OperationAction(Op, MVT::f16, Legal, Expand);
1109 setFP16OperationAction(Op, MVT::v2f16, Legal, Expand);
1110 setBF16OperationAction(Op, MVT::bf16, Legal, Expand);
1111 setBF16OperationAction(Op, MVT::v2bf16, Legal, Expand);
1112 setOperationAction(Op, VT: MVT::v2f32, Action: Expand);
1113 }
1114
1115 // Custom lowering for inline asm with 128-bit operands
1116 setOperationAction(Op: ISD::CopyToReg, VT: MVT::i128, Action: Custom);
1117 setOperationAction(Op: ISD::CopyFromReg, VT: MVT::i128, Action: Custom);
1118
1119 // FEXP2 support:
1120 // - f32
1121 // - f16/f16x2 (sm_70+, PTX 7.0+)
1122 // - bf16/bf16x2 (sm_90+, PTX 7.8+)
1123 // When f16/bf16 types aren't supported, they are promoted/expanded to f32.
1124 setOperationAction(Op: ISD::FEXP2, VT: MVT::f32, Action: Legal);
1125 setOperationAction(Op: ISD::FEXP2, VT: MVT::v2f32, Action: Expand);
1126 setFP16OperationAction(ISD::FEXP2, MVT::f16, Legal, Promote);
1127 setFP16OperationAction(ISD::FEXP2, MVT::v2f16, Legal, Expand);
1128 setBF16OperationAction(ISD::FEXP2, MVT::bf16, Legal, Promote);
1129 setBF16OperationAction(ISD::FEXP2, MVT::v2bf16, Legal, Expand);
1130
1131 // FLOG2 supports f32 only
1132 // f16/bf16 types aren't supported, but they are promoted/expanded to f32.
1133 if (UseApproxLog2F32) {
1134 setOperationAction(Op: ISD::FLOG2, VT: MVT::f32, Action: Legal);
1135 setOperationPromotedToType(Opc: ISD::FLOG2, OrigVT: MVT::f16, DestVT: MVT::f32);
1136 setOperationPromotedToType(Opc: ISD::FLOG2, OrigVT: MVT::bf16, DestVT: MVT::f32);
1137 setOperationAction(Ops: ISD::FLOG2, VTs: {MVT::v2f16, MVT::v2bf16, MVT::v2f32},
1138 Action: Expand);
1139 }
1140
1141 setOperationAction(Ops: ISD::ADDRSPACECAST, VTs: {MVT::i32, MVT::i64}, Action: Custom);
1142
1143 // atom.b128 is legal in PTX but since we don't represent i128 as a legal
1144 // type, we need to custom lower it.
1145 setOperationAction(Ops: {ISD::ATOMIC_CMP_SWAP, ISD::ATOMIC_SWAP}, VT: MVT::i128,
1146 Action: Custom);
1147
1148 // Now deduce the information based on the above mentioned
1149 // actions
1150 computeRegisterProperties(TRI: STI.getRegisterInfo());
1151
1152 // PTX support for 16-bit CAS is emulated. Only use 32+
1153 setMinCmpXchgSizeInBits(STI.getMinCmpXchgSizeInBits());
1154 setMaxAtomicSizeInBitsSupported(STI.hasAtomSwap128() ? 128 : 64);
1155 setMaxDivRemBitWidthSupported(64);
1156 setMaxLargeFPConvertBitWidthSupported(64);
1157
1158 // Custom lowering for tcgen05.ld vector operands
1159 setOperationAction(Ops: ISD::INTRINSIC_W_CHAIN,
1160 VTs: {MVT::v1i32, MVT::v2i32, MVT::v4i32, MVT::v8i32,
1161 MVT::v16i32, MVT::v32i32, MVT::v64i32, MVT::v128i32,
1162 MVT::v2f32, MVT::v4f32, MVT::v8f32, MVT::v16f32,
1163 MVT::v32f32, MVT::v64f32, MVT::v128f32},
1164 Action: Custom);
1165
1166 // Custom lowering for tcgen05.st vector operands and the st.async
1167 // i128 (.b128) operand. MVT::i8 is needed for the st.async.{sys,gpu} b8
1168 // variant.
1169 setOperationAction(Ops: ISD::INTRINSIC_VOID,
1170 VTs: {MVT::i8, MVT::v1i32, MVT::v2i32, MVT::v4i32, MVT::v8i32,
1171 MVT::v16i32, MVT::v32i32, MVT::v64i32, MVT::v128i32,
1172 MVT::i128, MVT::Other},
1173 Action: Custom);
1174
1175 // Enable custom lowering for the following:
1176 // * MVT::i128 - clusterlaunchcontrol
1177 // * MVT::i32 - prmt
1178 // * MVT::v4f32 - cvt_rs fp{4/6/8}x4 intrinsics
1179 // * MVT::Other - internal.addrspace.wrap
1180 setOperationAction(Ops: ISD::INTRINSIC_WO_CHAIN,
1181 VTs: {MVT::i32, MVT::i128, MVT::v4f32, MVT::Other}, Action: Custom);
1182
1183 // Custom lowering for bswap
1184 setOperationAction(Ops: ISD::BSWAP, VTs: {MVT::i16, MVT::i32, MVT::i64, MVT::v2i16},
1185 Action: Custom);
1186}
1187
1188TargetLoweringBase::LegalizeTypeAction
1189NVPTXTargetLowering::getPreferredVectorAction(MVT VT) const {
1190 if (!VT.isScalableVector() && VT.getVectorNumElements() != 1 &&
1191 VT.getScalarType() == MVT::i1)
1192 return TypeSplitVector;
1193 return TargetLoweringBase::getPreferredVectorAction(VT);
1194}
1195
1196SDValue NVPTXTargetLowering::getSqrtEstimate(SDValue Operand, SelectionDAG &DAG,
1197 int Enabled, int &ExtraSteps,
1198 bool &UseOneConst,
1199 bool Reciprocal) const {
1200 if (!(Enabled == ReciprocalEstimate::Enabled ||
1201 (Enabled == ReciprocalEstimate::Unspecified && !usePrecSqrtF32())))
1202 return SDValue();
1203
1204 if (ExtraSteps == ReciprocalEstimate::Unspecified)
1205 ExtraSteps = 0;
1206
1207 SDLoc DL(Operand);
1208 EVT VT = Operand.getValueType();
1209 bool Ftz = useF32FTZ(MF: DAG.getMachineFunction());
1210
1211 auto MakeIntrinsicCall = [&](Intrinsic::ID IID) {
1212 return DAG.getNode(Opcode: ISD::INTRINSIC_WO_CHAIN, DL, VT,
1213 N1: DAG.getConstant(Val: IID, DL, VT: MVT::i32), N2: Operand);
1214 };
1215
1216 // The sqrt and rsqrt refinement processes assume we always start out with an
1217 // approximation of the rsqrt. Therefore, if we're going to do any refinement
1218 // (i.e. ExtraSteps > 0), we must return an rsqrt. But if we're *not* doing
1219 // any refinement, we must return a regular sqrt.
1220 if (Reciprocal || ExtraSteps > 0) {
1221 if (VT == MVT::f32)
1222 return MakeIntrinsicCall(Ftz ? Intrinsic::nvvm_rsqrt_approx_ftz_f
1223 : Intrinsic::nvvm_rsqrt_approx_f);
1224 else if (VT == MVT::f64)
1225 return MakeIntrinsicCall(Intrinsic::nvvm_rsqrt_approx_d);
1226 else
1227 return SDValue();
1228 } else {
1229 if (VT == MVT::f32)
1230 return MakeIntrinsicCall(Ftz ? Intrinsic::nvvm_sqrt_approx_ftz_f
1231 : Intrinsic::nvvm_sqrt_approx_f);
1232 else {
1233 // There's no sqrt.approx.f64 instruction, so we emit
1234 // reciprocal(rsqrt(x)). This is faster than
1235 // select(x == 0, 0, x * rsqrt(x)). (In fact, it's faster than plain
1236 // x * rsqrt(x).)
1237 return DAG.getNode(
1238 Opcode: ISD::INTRINSIC_WO_CHAIN, DL, VT,
1239 N1: DAG.getConstant(Val: Intrinsic::nvvm_rcp_approx_ftz_d, DL, VT: MVT::i32),
1240 N2: MakeIntrinsicCall(Intrinsic::nvvm_rsqrt_approx_d));
1241 }
1242 }
1243}
1244
1245static MachinePointerInfo refinePtrAS(SDValue &Ptr, SelectionDAG &DAG) {
1246 // Load directly from the source address space of a cast to generic.
1247 unsigned SrcAS = ADDRESS_SPACE_GENERIC;
1248 if (Ptr->getOpcode() == ISD::ADDRSPACECAST) {
1249 const auto *ASC = cast<AddrSpaceCastSDNode>(Val&: Ptr);
1250 if (ASC->getDestAddressSpace() == ADDRESS_SPACE_GENERIC) {
1251 Ptr = ASC->getOperand(Num: 0);
1252 SrcAS = ASC->getSrcAddressSpace();
1253 }
1254 }
1255
1256 // Preserve the alloca's address space through frame-index inference.
1257 if (const auto *FIN = dyn_cast<FrameIndexSDNode>(Val&: Ptr))
1258 if (const AllocaInst *AI =
1259 DAG.getMachineFunction().getFrameInfo().getObjectAllocation(
1260 ObjectIdx: FIN->getIndex()))
1261 return MachinePointerInfo(AI);
1262
1263 return MachinePointerInfo(SrcAS);
1264}
1265
1266static ISD::NodeType getExtOpcode(const ISD::ArgFlagsTy &Flags) {
1267 if (Flags.isSExt())
1268 return ISD::SIGN_EXTEND;
1269 if (Flags.isZExt())
1270 return ISD::ZERO_EXTEND;
1271 return ISD::ANY_EXTEND;
1272}
1273
1274static SDValue correctParamType(SDValue V, EVT ExpectedVT,
1275 ISD::ArgFlagsTy Flags, SelectionDAG &DAG,
1276 SDLoc dl) {
1277 const EVT ActualVT = V.getValueType();
1278 assert((ActualVT == ExpectedVT ||
1279 (ExpectedVT.isInteger() && ActualVT.isInteger())) &&
1280 "Non-integer argument type size mismatch");
1281 if (ExpectedVT.bitsGT(VT: ActualVT))
1282 return DAG.getNode(Opcode: getExtOpcode(Flags), DL: dl, VT: ExpectedVT, Operand: V);
1283 if (ExpectedVT.bitsLT(VT: ActualVT))
1284 return DAG.getNode(Opcode: ISD::TRUNCATE, DL: dl, VT: ExpectedVT, Operand: V);
1285
1286 return V;
1287}
1288
1289static SDValue getSymbolNode(SelectionDAG &DAG, MCSymbol *Sym, EVT T) {
1290 return DAG.getNode(Opcode: NVPTXISD::Symbol, DL: SDLoc(), VT: T, Operand: DAG.getMCSymbol(Sym, VT: T));
1291}
1292
1293static SDValue getSymbolNode(SelectionDAG &DAG, const Twine &Name, EVT T) {
1294 MCContext &Ctx = DAG.getMachineFunction().getContext();
1295 return getSymbolNode(DAG, Sym: Ctx.getOrCreateSymbol(Name), T);
1296}
1297
1298SDValue NVPTXTargetLowering::LowerCall(TargetLowering::CallLoweringInfo &CLI,
1299 SmallVectorImpl<SDValue> &InVals) const {
1300
1301 if (CLI.IsVarArg &&
1302 (!STI.hasFeature(Feature: NVPTX::PTX60) || !STI.hasFeature(Feature: NVPTX::SM30)))
1303 report_fatal_error(
1304 reason: "Support for variadic functions (unsized array parameter) introduced "
1305 "in PTX ISA version 6.0 and requires target sm_30.");
1306
1307 SelectionDAG &DAG = CLI.DAG;
1308 SDLoc dl = CLI.DL;
1309 const SmallVectorImpl<ISD::InputArg> &Ins = CLI.Ins;
1310 SDValue Callee = CLI.Callee;
1311 ArgListTy &Args = CLI.getArgs();
1312 Type *RetTy = CLI.RetTy;
1313 const CallBase *CB = CLI.CB;
1314 const DataLayout &DL = DAG.getDataLayout();
1315 LLVMContext &Ctx = *DAG.getContext();
1316
1317 const auto GetI32 = [&](const unsigned I) {
1318 return DAG.getConstant(Val: I, DL: dl, VT: MVT::i32);
1319 };
1320
1321 const unsigned UniqueCallSite = GlobalUniqueCallSite++;
1322 const SDValue CallChain = CLI.Chain;
1323 const SDValue StartChain =
1324 DAG.getCALLSEQ_START(Chain: CallChain, InSize: UniqueCallSite, OutSize: 0, DL: dl);
1325 SDValue DeclareGlue = StartChain.getValue(R: 1);
1326
1327 SmallVector<SDValue, 16> CallPrereqs{StartChain};
1328
1329 const auto MakeDeclareScalarParam = [&](SDValue Symbol, unsigned Size) {
1330 // PTX ABI requires integral types to be at least 32 bits in size. FP16 is
1331 // loaded/stored using i16, so it's handled here as well.
1332 const unsigned SizeBits = promoteScalarArgumentSize(Size: Size * 8);
1333 SDValue Declare =
1334 DAG.getNode(Opcode: NVPTXISD::DeclareScalarParam, DL: dl, ResultTys: {MVT::Other, MVT::Glue},
1335 Ops: {StartChain, Symbol, GetI32(SizeBits), DeclareGlue});
1336 CallPrereqs.push_back(Elt: Declare);
1337 DeclareGlue = Declare.getValue(R: 1);
1338 return Declare;
1339 };
1340
1341 const auto MakeDeclareArrayParam = [&](SDValue Symbol, Align Align,
1342 unsigned Size) {
1343 SDValue Declare = DAG.getNode(
1344 Opcode: NVPTXISD::DeclareArrayParam, DL: dl, ResultTys: {MVT::Other, MVT::Glue},
1345 Ops: {StartChain, Symbol, GetI32(Align.value()), GetI32(Size), DeclareGlue});
1346 CallPrereqs.push_back(Elt: Declare);
1347 DeclareGlue = Declare.getValue(R: 1);
1348 return Declare;
1349 };
1350
1351 // Variadic arguments.
1352 //
1353 // Normally, for each argument, we declare a param scalar or a param
1354 // byte array in the .param space, and store the argument value to that
1355 // param scalar or array starting at offset 0.
1356 //
1357 // In the case of the first variadic argument, we declare a vararg byte array
1358 // with size 0. The exact size of this array isn't known at this point, so
1359 // it'll be patched later. All the variadic arguments will be stored to this
1360 // array at a certain offset (which gets tracked by 'VAOffset'). The offset is
1361 // initially set to 0, so it can be used for non-variadic arguments (which use
1362 // 0 offset) to simplify the code.
1363 //
1364 // After all vararg is processed, 'VAOffset' holds the size of the
1365 // vararg byte array.
1366 assert((CLI.IsVarArg || CLI.Args.size() <= CLI.NumFixedArgs) &&
1367 "Non-VarArg function with extra arguments");
1368
1369 const unsigned FirstVAArg = CLI.NumFixedArgs; // position of first variadic
1370 unsigned VAOffset = 0; // current offset in the param array
1371
1372 const SDValue VADeclareParam =
1373 CLI.Args.size() > FirstVAArg
1374 ? MakeDeclareArrayParam(
1375 getCallParamSymbolNode(DAG, I: FirstVAArg, T: MVT::i32),
1376 Align(STI.getMaxRequiredAlignment()), 0)
1377 : SDValue();
1378
1379 // Args.size() and Outs.size() need not match.
1380 // Outs.size() will be larger
1381 // * if there is an aggregate argument with multiple fields (each field
1382 // showing up separately in Outs)
1383 // * if there is a vector argument with more than typical vector-length
1384 // elements (generally if more than 4) where each vector element is
1385 // individually present in Outs.
1386 // So a different index should be used for indexing into Outs/OutVals.
1387 // See similar issue in LowerFormalArguments.
1388 auto AllOuts = ArrayRef(CLI.Outs);
1389 auto AllOutVals = ArrayRef(CLI.OutVals);
1390 assert(AllOuts.size() == AllOutVals.size() &&
1391 "Outs and OutVals must be the same size");
1392 // Declare the .params or .reg need to pass values
1393 // to the function
1394 for (const auto E : llvm::enumerate(First&: Args)) {
1395 const auto ArgI = E.index();
1396 const auto Arg = E.value();
1397 const auto ArgOuts =
1398 AllOuts.take_while(Pred: [&](auto O) { return O.OrigArgIndex == ArgI; });
1399 const auto ArgOutVals = AllOutVals.take_front(N: ArgOuts.size());
1400 AllOuts = AllOuts.drop_front(N: ArgOuts.size());
1401 AllOutVals = AllOutVals.drop_front(N: ArgOuts.size());
1402
1403 const bool IsVAArg = (ArgI >= FirstVAArg);
1404 const bool IsByVal = Arg.IsByVal;
1405
1406 const SDValue ParamSymbol =
1407 getCallParamSymbolNode(DAG, I: IsVAArg ? FirstVAArg : ArgI, T: MVT::i32);
1408
1409 assert((!IsByVal || Arg.IndirectType) &&
1410 "byval arg must have indirect type");
1411 Type *ETy = (IsByVal ? Arg.IndirectType : Arg.Ty);
1412
1413 const Align ArgAlign = [&]() {
1414 const unsigned ParamIdx = ArgI + AttributeList::FirstArgIndex;
1415 if (IsByVal)
1416 return getDeviceByValParamAlign(CB, ArgTy: ETy, AttrIdx: ParamIdx, DL);
1417 return getPTXParamAlign(CB, Ty: Arg.Ty, AttrIdx: ParamIdx, DL);
1418 }();
1419
1420 const unsigned TySize = DL.getTypeAllocSize(Ty: ETy);
1421 assert((!IsByVal || TySize == ArgOuts[0].Flags.getByValSize()) &&
1422 "type size mismatch");
1423
1424 const SDValue ArgDeclare = [&]() {
1425 if (IsVAArg)
1426 return VADeclareParam;
1427
1428 if (IsByVal || shouldPassAsArray(Ty: Arg.Ty))
1429 return MakeDeclareArrayParam(ParamSymbol, ArgAlign, TySize);
1430
1431 assert(ArgOuts.size() == 1 && "We must pass only one value as non-array");
1432 assert((ArgOuts[0].VT.isInteger() || ArgOuts[0].VT.isFloatingPoint()) &&
1433 "Only int and float types are supported as non-array arguments");
1434
1435 return MakeDeclareScalarParam(ParamSymbol, TySize);
1436 }();
1437
1438 if (IsByVal) {
1439 assert(ArgOutVals.size() == 1 && "We must pass only one value as byval");
1440 SDValue SrcPtr = ArgOutVals[0];
1441 const MachinePointerInfo SrcPtrInfo = refinePtrAS(Ptr&: SrcPtr, DAG);
1442 // Don't use Flags.getNonZeroByValAlign as this includes the stackalign,
1443 // which does not apply to the source pointer.
1444 const Align BaseSrcAlign = [&]() {
1445 // The align attribute on a byval argument indicates the known alignment
1446 // of the pointer passed to the function.
1447 if (CB)
1448 if (const MaybeAlign A = CB->getParamAlign(ArgNo: ArgI))
1449 return *A;
1450 // Fall back to the default alignment for the type.
1451 // TODO: This might be too aggressive but we haven't had a problem with
1452 // it yet.
1453 return getPTXParamTypeAlign(ArgTy: ETy, DL);
1454 }();
1455
1456 if (IsVAArg)
1457 VAOffset = alignTo(Size: VAOffset, A: ArgAlign);
1458
1459 SmallVector<EVT, 4> ValueVTs, MemVTs;
1460 SmallVector<TypeSize, 4> Offsets;
1461 ComputeValueVTs(TLI: *this, DL, Ty: ETy, ValueVTs, MemVTs: &MemVTs, Offsets: &Offsets);
1462
1463 unsigned J = 0;
1464 const auto VI = VectorizePTXValueVTs(ValueVTs: MemVTs, Offsets, ParamAlignment: ArgAlign, IsVAArg);
1465 for (const unsigned NumElts : VI) {
1466 EVT LoadVT = getVectorizedVT(VT: MemVTs[J], N: NumElts, C&: Ctx);
1467 Align SrcAlign = commonAlignment(A: BaseSrcAlign, Offset: Offsets[J]);
1468 SDValue SrcAddr = DAG.getObjectPtrOffset(SL: dl, Ptr: SrcPtr, Offset: Offsets[J]);
1469 SDValue SrcLoad =
1470 DAG.getLoad(VT: LoadVT, dl, Chain: CallChain, Ptr: SrcAddr,
1471 PtrInfo: SrcPtrInfo.getWithOffset(O: Offsets[J]), Alignment: SrcAlign);
1472
1473 TypeSize ParamOffset = Offsets[J].getWithIncrement(RHS: VAOffset);
1474 Align ParamAlign = commonAlignment(A: ArgAlign, Offset: ParamOffset);
1475 SDValue ParamAddr =
1476 DAG.getObjectPtrOffset(SL: dl, Ptr: ParamSymbol, Offset: ParamOffset);
1477 SDValue StoreParam = DAG.getStore(
1478 Chain: ArgDeclare, dl, Val: SrcLoad, Ptr: ParamAddr,
1479 PtrInfo: MachinePointerInfo(NVPTX::AddressSpace::DeviceParam), Alignment: ParamAlign);
1480 CallPrereqs.push_back(Elt: StoreParam);
1481
1482 J += NumElts;
1483 }
1484 if (IsVAArg)
1485 VAOffset += TySize;
1486 } else {
1487 SmallVector<EVT, 16> VTs;
1488 SmallVector<uint64_t, 16> Offsets;
1489 ComputePTXValueVTs(TLI: *this, DL, Ctx, CallConv: CLI.CallConv, Ty: Arg.Ty, ValueVTs&: VTs, Offsets,
1490 StartingOffset: VAOffset);
1491 assert(VTs.size() == Offsets.size() && "Size mismatch");
1492 assert(VTs.size() == ArgOuts.size() && "Size mismatch");
1493
1494 // PTX Interoperability Guide 3.3(A): [Integer] Values shorter
1495 // than 32-bits are sign extended or zero extended, depending on
1496 // whether they are signed or unsigned types. This case applies
1497 // only to scalar parameters and not to aggregate values.
1498 const bool ExtendIntegerParam =
1499 Arg.Ty->isIntegerTy() && DL.getTypeAllocSizeInBits(Ty: Arg.Ty) < 32;
1500
1501 const auto GetStoredValue = [&](const unsigned I) {
1502 SDValue StVal = ArgOutVals[I];
1503 assert(promoteScalarIntegerPTX(StVal.getValueType()) ==
1504 StVal.getValueType() &&
1505 "OutVal type should always be legal");
1506
1507 const EVT VTI = promoteScalarIntegerPTX(VT: VTs[I]);
1508 const EVT StoreVT =
1509 ExtendIntegerParam ? MVT::i32 : (VTI == MVT::i1 ? MVT::i8 : VTI);
1510
1511 return correctParamType(V: StVal, ExpectedVT: StoreVT, Flags: ArgOuts[I].Flags, DAG, dl);
1512 };
1513
1514 unsigned J = 0;
1515 const auto VI = VectorizePTXValueVTs(ValueVTs: VTs, Offsets, ParamAlignment: ArgAlign, IsVAArg);
1516 for (const unsigned NumElts : VI) {
1517 const EVT EltVT = promoteScalarIntegerPTX(VT: VTs[J]);
1518
1519 unsigned Offset;
1520 if (IsVAArg) {
1521 // TODO: We may need to support vector types that can be passed
1522 // as scalars in variadic arguments.
1523 assert(NumElts == 1 &&
1524 "Vectorization should be disabled for vaargs.");
1525
1526 // Align each part of the variadic argument to their type.
1527 VAOffset = alignTo(Size: VAOffset, A: DAG.getEVTAlign(MemoryVT: EltVT));
1528 Offset = VAOffset;
1529
1530 const EVT TheStoreType = ExtendIntegerParam ? MVT::i32 : EltVT;
1531 VAOffset += DL.getTypeAllocSize(Ty: TheStoreType.getTypeForEVT(Context&: Ctx));
1532 } else {
1533 assert(VAOffset == 0 && "VAOffset must be 0 for non-VA args");
1534 Offset = Offsets[J];
1535 }
1536
1537 SDValue Ptr =
1538 DAG.getObjectPtrOffset(SL: dl, Ptr: ParamSymbol, Offset: TypeSize::getFixed(ExactSize: Offset));
1539
1540 const MaybeAlign CurrentAlign = ExtendIntegerParam
1541 ? MaybeAlign(std::nullopt)
1542 : commonAlignment(A: ArgAlign, Offset);
1543
1544 SDValue Val =
1545 getBuildVectorizedValue(N: NumElts, dl, DAG, GetElement: [&](unsigned K) {
1546 return GetStoredValue(J + K);
1547 });
1548
1549 SDValue StoreParam = DAG.getStore(
1550 Chain: ArgDeclare, dl, Val, Ptr,
1551 PtrInfo: MachinePointerInfo(NVPTX::AddressSpace::DeviceParam), Alignment: CurrentAlign);
1552 CallPrereqs.push_back(Elt: StoreParam);
1553
1554 J += NumElts;
1555 }
1556 }
1557 }
1558
1559 // Handle Result
1560 if (!Ins.empty()) {
1561 const SDValue RetSymbol = getSymbolNode(DAG, Name: "retval0", T: MVT::i32);
1562 const unsigned ResultSize = DL.getTypeAllocSize(Ty: RetTy);
1563 if (shouldPassAsArray(Ty: RetTy)) {
1564 const Align RetAlign =
1565 getPTXParamAlign(CB, Ty: RetTy, AttrIdx: AttributeList::ReturnIndex, DL);
1566 MakeDeclareArrayParam(RetSymbol, RetAlign, ResultSize);
1567 } else {
1568 MakeDeclareScalarParam(RetSymbol, ResultSize);
1569 }
1570 }
1571
1572 // Set the size of the vararg param byte array if the callee is a variadic
1573 // function and the variadic part is not empty.
1574 if (VADeclareParam) {
1575 SDValue DeclareParamOps[] = {VADeclareParam.getOperand(i: 0),
1576 VADeclareParam.getOperand(i: 1),
1577 VADeclareParam.getOperand(i: 2), GetI32(VAOffset),
1578 VADeclareParam.getOperand(i: 4)};
1579 DAG.MorphNodeTo(N: VADeclareParam.getNode(), Opc: VADeclareParam.getOpcode(),
1580 VTs: VADeclareParam->getVTList(), Ops: DeclareParamOps);
1581 }
1582
1583 const auto *Func = dyn_cast<GlobalAddressSDNode>(Val: Callee.getNode());
1584 const auto *CalleeF = Func ? dyn_cast<Function>(Val: Func->getGlobal()) : nullptr;
1585
1586 // If the type of the callsite does not match that of the function, convert
1587 // the callsite to an indirect call.
1588 const bool ConvertToIndirectCall =
1589 CalleeF && CB->getFunctionType() != CalleeF->getFunctionType();
1590
1591 // Both indirect calls and libcalls have nullptr Func. In order to distinguish
1592 // between them we must rely on the call site value which is valid for
1593 // indirect calls but is always null for libcalls.
1594 const bool IsIndirectCall = (!Func && CB) || ConvertToIndirectCall;
1595
1596 if (isa<ExternalSymbolSDNode>(Val: Callee)) {
1597 Function* CalleeFunc = nullptr;
1598
1599 // Try to find the callee in the current module.
1600 Callee = DAG.getSymbolFunctionGlobalAddress(Op: Callee, TargetFunction: &CalleeFunc);
1601 assert(CalleeFunc != nullptr && "Libcall callee must be set.");
1602
1603 // Set the "libcall callee" attribute to indicate that the function
1604 // must always have a declaration.
1605 CalleeFunc->addFnAttr(Kind: "nvptx-libcall-callee", Val: "true");
1606 }
1607
1608 // In the indirect function call case, PTX requires a prototype of the form:
1609 // proto_0 : .callprototype(.param .b32 _) _ (.param .b32 _);
1610 // Where the label is to be used as the last arg of the call instruction.
1611 // We record the call site here and emit all prototypes at the
1612 // start of the function in the AsmPrinter.
1613 SDValue Proto = GetI32(0);
1614 if (IsIndirectCall) {
1615 auto *ProtoSymbol = DAG.getMachineFunction()
1616 .getInfo<NVPTXMachineFunctionInfo>()
1617 ->addCallPrototype(CB, MF&: DAG.getMachineFunction());
1618 Proto = DAG.getMCSymbol(Sym: ProtoSymbol, VT: MVT::i32);
1619 }
1620
1621 const bool IsUnknownIntrinsic =
1622 CalleeF && CalleeF->isIntrinsic() &&
1623 CalleeF->getIntrinsicID() == Intrinsic::not_intrinsic;
1624 if (IsUnknownIntrinsic) {
1625 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
1626 DAG.getMachineFunction().getFunction(),
1627 "call to unknown intrinsic '" + CalleeF->getName() +
1628 "' cannot be lowered by the NVPTX backend",
1629 dl.getDebugLoc()));
1630 }
1631
1632 const unsigned NumArgs =
1633 std::min<unsigned>(a: CLI.NumFixedArgs + 1, b: Args.size());
1634 /// CALL(Chain, IsConvergent, IsIndirectCall/IsUniform, NumReturns,
1635 /// NumParams, Callee, Proto)
1636 const SDValue CallToken = DAG.getTokenFactor(DL: dl, Vals&: CallPrereqs);
1637 const SDValue Call = DAG.getNode(
1638 Opcode: NVPTXISD::CALL, DL: dl, VT: MVT::Other,
1639 Ops: {CallToken, GetI32(CLI.IsConvergent), GetI32(IsIndirectCall),
1640 GetI32(Ins.empty() ? 0 : 1), GetI32(NumArgs), Callee, Proto});
1641
1642 SmallVector<SDValue, 16> LoadChains{Call};
1643 SmallVector<SDValue, 16> ProxyRegOps;
1644 if (!Ins.empty()) {
1645 SmallVector<EVT, 16> VTs;
1646 SmallVector<uint64_t, 16> Offsets;
1647 ComputePTXValueVTs(TLI: *this, DL, Ctx, CallConv: CLI.CallConv, Ty: RetTy, ValueVTs&: VTs, Offsets);
1648 assert(VTs.size() == Ins.size() && "Bad value decomposition");
1649
1650 const Align RetAlign =
1651 getPTXParamAlign(CB, Ty: RetTy, AttrIdx: AttributeList::ReturnIndex, DL);
1652 const SDValue RetSymbol = getSymbolNode(DAG, Name: "retval0", T: MVT::i32);
1653
1654 // PTX Interoperability Guide 3.3(A): [Integer] Values shorter than
1655 // 32-bits are sign extended or zero extended, depending on whether
1656 // they are signed or unsigned types.
1657 const bool ExtendIntegerRetVal =
1658 RetTy->isIntegerTy() && DL.getTypeAllocSizeInBits(Ty: RetTy) < 32;
1659
1660 unsigned I = 0;
1661 const auto VI = VectorizePTXValueVTs(ValueVTs: VTs, Offsets, ParamAlignment: RetAlign);
1662 for (const unsigned NumElts : VI) {
1663 const MaybeAlign CurrentAlign =
1664 ExtendIntegerRetVal ? MaybeAlign(std::nullopt)
1665 : commonAlignment(A: RetAlign, Offset: Offsets[I]);
1666
1667 const EVT VTI = promoteScalarIntegerPTX(VT: VTs[I]);
1668 const EVT LoadVT =
1669 ExtendIntegerRetVal ? MVT::i32 : (VTI == MVT::i1 ? MVT::i8 : VTI);
1670 const EVT VecVT = getVectorizedVT(VT: LoadVT, N: NumElts, C&: Ctx);
1671 SDValue Ptr =
1672 DAG.getObjectPtrOffset(SL: dl, Ptr: RetSymbol, Offset: TypeSize::getFixed(ExactSize: Offsets[I]));
1673
1674 SDValue R = DAG.getLoad(
1675 VT: VecVT, dl, Chain: Call, Ptr,
1676 PtrInfo: MachinePointerInfo(NVPTX::AddressSpace::DeviceParam), Alignment: CurrentAlign);
1677
1678 LoadChains.push_back(Elt: R.getValue(R: 1));
1679 for (const unsigned J : llvm::seq(Size: NumElts))
1680 ProxyRegOps.push_back(Elt: getExtractVectorizedValue(V: R, I: J, VT: LoadVT, dl, DAG));
1681 I += NumElts;
1682 }
1683 }
1684
1685 const SDValue EndToken = DAG.getTokenFactor(DL: dl, Vals&: LoadChains);
1686 const SDValue CallEnd = DAG.getCALLSEQ_END(Chain: EndToken, Size1: UniqueCallSite,
1687 Size2: UniqueCallSite + 1, Glue: SDValue(), DL: dl);
1688
1689 // Append ProxyReg instructions to the chain to make sure that `callseq_end`
1690 // will not get lost. Otherwise, during libcalls expansion, the nodes can become
1691 // dangling.
1692 for (const auto [I, Reg] : llvm::enumerate(First&: ProxyRegOps)) {
1693 SDValue Proxy =
1694 DAG.getNode(Opcode: NVPTXISD::ProxyReg, DL: dl, VT: Reg.getValueType(), Ops: {CallEnd, Reg});
1695 SDValue Ret = correctParamType(V: Proxy, ExpectedVT: Ins[I].VT, Flags: Ins[I].Flags, DAG, dl);
1696 InVals.push_back(Elt: Ret);
1697 }
1698
1699 // set IsTailCall to false for now, until we figure out how to express
1700 // tail call optimization in PTX
1701 CLI.IsTailCall = false;
1702 return CallEnd;
1703}
1704
1705SDValue NVPTXTargetLowering::LowerDYNAMIC_STACKALLOC(SDValue Op,
1706 SelectionDAG &DAG) const {
1707
1708 if (!STI.hasFeature(Feature: NVPTX::PTX73) || !STI.hasFeature(Feature: NVPTX::SM52)) {
1709 const Function &Fn = DAG.getMachineFunction().getFunction();
1710
1711 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
1712 Fn,
1713 "Support for dynamic alloca introduced in PTX ISA version 7.3 and "
1714 "requires target sm_52.",
1715 SDLoc(Op).getDebugLoc()));
1716 auto Ops = {DAG.getConstant(Val: 0, DL: SDLoc(), VT: Op.getValueType()),
1717 Op.getOperand(i: 0)};
1718 return DAG.getMergeValues(Ops, dl: SDLoc());
1719 }
1720
1721 SDLoc DL(Op.getNode());
1722 SDValue Chain = Op.getOperand(i: 0);
1723 SDValue Size = Op.getOperand(i: 1);
1724 uint64_t Align = Op.getConstantOperandVal(i: 2);
1725
1726 // The alignment on a ISD::DYNAMIC_STACKALLOC node may be 0 to indicate that
1727 // the default stack alignment should be used.
1728 if (Align == 0)
1729 Align = DAG.getSubtarget().getFrameLowering()->getStackAlign().value();
1730
1731 // The size for ptx alloca instruction is 64-bit for m64 and 32-bit for m32.
1732 const MVT LocalVT = getPointerTy(DL: DAG.getDataLayout(), AS: ADDRESS_SPACE_LOCAL);
1733
1734 SDValue Alloc =
1735 DAG.getNode(Opcode: NVPTXISD::DYNAMIC_STACKALLOC, DL, ResultTys: {LocalVT, MVT::Other},
1736 Ops: {Chain, DAG.getZExtOrTrunc(Op: Size, DL, VT: LocalVT),
1737 DAG.getTargetConstant(Val: Align, DL, VT: MVT::i32)});
1738
1739 // NVPTXLowerAlloca puts allocas in the local address space, so a local
1740 // pointer is requested here; escapes are explicit addrspacecasts in the IR.
1741 assert(Op.getValueType() == LocalVT && "Unexpected alloca pointer size");
1742
1743 return DAG.getMergeValues(Ops: {Alloc, SDValue(Alloc.getNode(), 1)}, dl: DL);
1744}
1745
1746SDValue NVPTXTargetLowering::LowerSTACKRESTORE(SDValue Op,
1747 SelectionDAG &DAG) const {
1748 SDLoc DL(Op.getNode());
1749 if (!STI.hasFeature(Feature: NVPTX::PTX73) || !STI.hasFeature(Feature: NVPTX::SM52)) {
1750 const Function &Fn = DAG.getMachineFunction().getFunction();
1751
1752 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
1753 Fn,
1754 "Support for stackrestore requires PTX ISA version >= 7.3 and target "
1755 ">= sm_52.",
1756 DL.getDebugLoc()));
1757 return Op.getOperand(i: 0);
1758 }
1759
1760 const MVT LocalVT = getPointerTy(DL: DAG.getDataLayout(), AS: ADDRESS_SPACE_LOCAL);
1761 SDValue Chain = Op.getOperand(i: 0);
1762 SDValue Ptr = Op.getOperand(i: 1);
1763 SDValue ASC = DAG.getAddrSpaceCast(dl: DL, VT: LocalVT, Ptr, SrcAS: ADDRESS_SPACE_GENERIC,
1764 DestAS: ADDRESS_SPACE_LOCAL);
1765 return DAG.getNode(Opcode: NVPTXISD::STACKRESTORE, DL, VT: MVT::Other, Ops: {Chain, ASC});
1766}
1767
1768SDValue NVPTXTargetLowering::LowerSTACKSAVE(SDValue Op,
1769 SelectionDAG &DAG) const {
1770 SDLoc DL(Op.getNode());
1771 if (!STI.hasFeature(Feature: NVPTX::PTX73) || !STI.hasFeature(Feature: NVPTX::SM52)) {
1772 const Function &Fn = DAG.getMachineFunction().getFunction();
1773
1774 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
1775 Fn,
1776 "Support for stacksave requires PTX ISA version >= 7.3 and target >= "
1777 "sm_52.",
1778 DL.getDebugLoc()));
1779 auto Ops = {DAG.getConstant(Val: 0, DL, VT: Op.getValueType()), Op.getOperand(i: 0)};
1780 return DAG.getMergeValues(Ops, dl: DL);
1781 }
1782
1783 const MVT LocalVT = getPointerTy(DL: DAG.getDataLayout(), AS: ADDRESS_SPACE_LOCAL);
1784 SDValue Chain = Op.getOperand(i: 0);
1785 SDValue SS =
1786 DAG.getNode(Opcode: NVPTXISD::STACKSAVE, DL, ResultTys: {LocalVT, MVT::Other}, Ops: Chain);
1787 SDValue ASC = DAG.getAddrSpaceCast(
1788 dl: DL, VT: Op.getValueType(), Ptr: SS, SrcAS: ADDRESS_SPACE_LOCAL, DestAS: ADDRESS_SPACE_GENERIC);
1789 return DAG.getMergeValues(Ops: {ASC, SDValue(SS.getNode(), 1)}, dl: DL);
1790}
1791
1792// By default CONCAT_VECTORS is lowered by ExpandVectorBuildThroughStack()
1793// (see LegalizeDAG.cpp). This is slow and uses local memory.
1794// We use extract/insert/build vector just as what LegalizeOp() does in llvm 2.5
1795SDValue
1796NVPTXTargetLowering::LowerCONCAT_VECTORS(SDValue Op, SelectionDAG &DAG) const {
1797 SDNode *Node = Op.getNode();
1798 SDLoc dl(Node);
1799 SmallVector<SDValue, 8> Ops;
1800 unsigned NumOperands = Node->getNumOperands();
1801 for (unsigned i = 0; i < NumOperands; ++i) {
1802 SDValue SubOp = Node->getOperand(Num: i);
1803 EVT VVT = SubOp.getNode()->getValueType(ResNo: 0);
1804 EVT EltVT = VVT.getVectorElementType();
1805 unsigned NumSubElem = VVT.getVectorNumElements();
1806 for (unsigned j = 0; j < NumSubElem; ++j) {
1807 Ops.push_back(Elt: DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL: dl, VT: EltVT, N1: SubOp,
1808 N2: DAG.getIntPtrConstant(Val: j, DL: dl)));
1809 }
1810 }
1811 return DAG.getBuildVector(VT: Node->getValueType(ResNo: 0), DL: dl, Ops);
1812}
1813
1814static SDValue getPRMT(SDValue A, SDValue B, SDValue Selector, SDLoc DL,
1815 SelectionDAG &DAG,
1816 unsigned Mode = NVPTX::PTXPrmtMode::NONE) {
1817 assert(A.getValueType() == MVT::i32 && B.getValueType() == MVT::i32 &&
1818 Selector.getValueType() == MVT::i32 && "PRMT must have i32 operands");
1819 return DAG.getNode(Opcode: NVPTXISD::PRMT, DL, VT: MVT::i32,
1820 Ops: {A, B, Selector, DAG.getConstant(Val: Mode, DL, VT: MVT::i32)});
1821}
1822
1823static SDValue getPRMT(SDValue A, SDValue B, uint64_t Selector, SDLoc DL,
1824 SelectionDAG &DAG,
1825 unsigned Mode = NVPTX::PTXPrmtMode::NONE) {
1826 return getPRMT(A, B, Selector: DAG.getConstant(Val: Selector, DL, VT: MVT::i32), DL, DAG, Mode);
1827}
1828
1829/// Reduces the elements using the scalar operations provided. The operations
1830/// are sorted descending in number of inputs they take. The flags on the
1831/// original reduction operation will be propagated to each scalar operation.
1832/// Nearby elements are grouped in tree reduction, unlike the shuffle reduction
1833/// used in ExpandReductions and SelectionDAG.
1834static SDValue buildTreeReduction(
1835 const SmallVector<SDValue> &Elements, EVT EltTy,
1836 ArrayRef<std::pair<unsigned /*NodeType*/, unsigned /*NumInputs*/>> Ops,
1837 const SDLoc &DL, const SDNodeFlags Flags, SelectionDAG &DAG) {
1838 // Build the reduction tree at each level, starting with all the elements.
1839 SmallVector<SDValue> Level = Elements;
1840
1841 unsigned OpIdx = 0;
1842 while (Level.size() > 1) {
1843 // Try to reduce this level using the current operator.
1844 const auto [Op, NumInputs] = Ops[OpIdx];
1845
1846 // Build the next level by partially reducing all elements.
1847 SmallVector<SDValue> ReducedLevel;
1848 unsigned I = 0, E = Level.size();
1849 for (; I + NumInputs <= E; I += NumInputs) {
1850 // Reduce elements in groups of [NumInputs], as much as possible.
1851 ReducedLevel.push_back(Elt: DAG.getNode(
1852 Opcode: Op, DL, VT: EltTy, Ops: ArrayRef<SDValue>(Level).slice(N: I, M: NumInputs), Flags));
1853 }
1854
1855 if (I < E) {
1856 // Handle leftover elements.
1857
1858 if (ReducedLevel.empty()) {
1859 // We didn't reduce anything at this level. We need to pick a smaller
1860 // operator.
1861 ++OpIdx;
1862 assert(OpIdx < Ops.size() && "no smaller operators for reduction");
1863 continue;
1864 }
1865
1866 // We reduced some things but there's still more left, meaning the
1867 // operator's number of inputs doesn't evenly divide this level size. Move
1868 // these elements to the next level.
1869 for (; I < E; ++I)
1870 ReducedLevel.push_back(Elt: Level[I]);
1871 }
1872
1873 // Process the next level.
1874 Level = ReducedLevel;
1875 }
1876
1877 return *Level.begin();
1878}
1879
1880// Get scalar reduction opcode
1881static ISD::NodeType getScalarOpcodeForReduction(unsigned ReductionOpcode) {
1882 switch (ReductionOpcode) {
1883 case ISD::VECREDUCE_FMAX:
1884 return ISD::FMAXNUM;
1885 case ISD::VECREDUCE_FMIN:
1886 return ISD::FMINNUM;
1887 case ISD::VECREDUCE_FMAXIMUM:
1888 return ISD::FMAXIMUM;
1889 case ISD::VECREDUCE_FMINIMUM:
1890 return ISD::FMINIMUM;
1891 default:
1892 llvm_unreachable("unhandled reduction opcode");
1893 }
1894}
1895
1896/// Get 3-input scalar reduction opcode
1897static std::optional<unsigned>
1898getScalar3OpcodeForReduction(unsigned ReductionOpcode) {
1899 switch (ReductionOpcode) {
1900 case ISD::VECREDUCE_FMAX:
1901 return NVPTXISD::FMAXNUM3;
1902 case ISD::VECREDUCE_FMIN:
1903 return NVPTXISD::FMINNUM3;
1904 case ISD::VECREDUCE_FMAXIMUM:
1905 return NVPTXISD::FMAXIMUM3;
1906 case ISD::VECREDUCE_FMINIMUM:
1907 return NVPTXISD::FMINIMUM3;
1908 default:
1909 return std::nullopt;
1910 }
1911}
1912
1913/// Lower reductions to either a sequence of operations or a tree if
1914/// reassociations are allowed. This method will use larger operations like
1915/// max3/min3 when the target supports them.
1916SDValue NVPTXTargetLowering::LowerVECREDUCE(SDValue Op,
1917 SelectionDAG &DAG) const {
1918 SDLoc DL(Op);
1919 const SDNodeFlags Flags = Op->getFlags();
1920 SDValue Vector = Op.getOperand(i: 0);
1921
1922 const unsigned Opcode = Op->getOpcode();
1923 const EVT EltTy = Vector.getValueType().getVectorElementType();
1924
1925 // Whether we can use 3-input min/max when expanding the reduction.
1926 const bool CanUseMinMax3 =
1927 EltTy == MVT::f32 && STI.hasFeature(Feature: NVPTX::SM100) &&
1928 STI.hasFeature(Feature: NVPTX::PTX88) &&
1929 (Opcode == ISD::VECREDUCE_FMAX || Opcode == ISD::VECREDUCE_FMIN ||
1930 Opcode == ISD::VECREDUCE_FMAXIMUM || Opcode == ISD::VECREDUCE_FMINIMUM);
1931
1932 // A list of SDNode opcodes with equivalent semantics, sorted descending by
1933 // number of inputs they take.
1934 SmallVector<std::pair<unsigned /*Op*/, unsigned /*NumIn*/>, 2> ScalarOps;
1935
1936 if (auto Opcode3Elem = getScalar3OpcodeForReduction(ReductionOpcode: Opcode);
1937 CanUseMinMax3 && Opcode3Elem)
1938 ScalarOps.push_back(Elt: {*Opcode3Elem, 3});
1939 ScalarOps.push_back(Elt: {getScalarOpcodeForReduction(ReductionOpcode: Opcode), 2});
1940
1941 SmallVector<SDValue> Elements;
1942 DAG.ExtractVectorElements(Op: Vector, Args&: Elements);
1943
1944 return buildTreeReduction(Elements, EltTy, Ops: ScalarOps, DL, Flags, DAG);
1945}
1946
1947SDValue NVPTXTargetLowering::LowerBITCAST(SDValue Op, SelectionDAG &DAG) const {
1948 // Handle bitcasting from v2i8 without hitting the default promotion
1949 // strategy which goes through stack memory.
1950 EVT FromVT = Op->getOperand(Num: 0)->getValueType(ResNo: 0);
1951 if (FromVT != MVT::v2i8) {
1952 return Op;
1953 }
1954
1955 // Pack vector elements into i16 and bitcast to final type
1956 SDLoc DL(Op);
1957 SDValue Vec0 = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i8,
1958 N1: Op->getOperand(Num: 0), N2: DAG.getIntPtrConstant(Val: 0, DL));
1959 SDValue Vec1 = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i8,
1960 N1: Op->getOperand(Num: 0), N2: DAG.getIntPtrConstant(Val: 1, DL));
1961 SDValue Extend0 = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: MVT::i16, Operand: Vec0);
1962 SDValue Extend1 = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: MVT::i16, Operand: Vec1);
1963 SDValue Const8 = DAG.getConstant(Val: 8, DL, VT: MVT::i16);
1964 SDValue AsInt = DAG.getNode(
1965 Opcode: ISD::OR, DL, VT: MVT::i16,
1966 Ops: {Extend0, DAG.getNode(Opcode: ISD::SHL, DL, VT: MVT::i16, Ops: {Extend1, Const8})});
1967 EVT ToVT = Op->getValueType(ResNo: 0);
1968 return DAG.getBitcast(VT: ToVT, V: AsInt);
1969}
1970
1971// We can init constant f16x2/v2i16/v4i8 with a single .b32 move. Normally it
1972// would get lowered as two constant loads and vector-packing move.
1973// Instead we want just a constant move:
1974// mov.b32 %r2, 0x40003C00
1975SDValue NVPTXTargetLowering::LowerBUILD_VECTOR(SDValue Op,
1976 SelectionDAG &DAG) const {
1977 EVT VT = Op->getValueType(ResNo: 0);
1978 if (!(NVPTX::isPackedVectorTy(VT) && VT.is32BitVector()))
1979 return Op;
1980 SDLoc DL(Op);
1981
1982 if (!llvm::all_of(Range: Op->ops(), P: [](SDValue Operand) {
1983 return Operand->isUndef() || isa<ConstantSDNode>(Val: Operand) ||
1984 isa<ConstantFPSDNode>(Val: Operand);
1985 })) {
1986 if (VT != MVT::v4i8)
1987 return Op;
1988 // Lower non-const v4i8 vector as byte-wise constructed i32, which allows us
1989 // to optimize calculation of constant parts.
1990 auto GetPRMT = [&](const SDValue Left, const SDValue Right, bool Cast,
1991 uint64_t SelectionValue) -> SDValue {
1992 SDValue L = Left;
1993 SDValue R = Right;
1994 if (Cast) {
1995 L = DAG.getAnyExtOrTrunc(Op: L, DL, VT: MVT::i32);
1996 R = DAG.getAnyExtOrTrunc(Op: R, DL, VT: MVT::i32);
1997 }
1998 return getPRMT(A: L, B: R, Selector: SelectionValue, DL, DAG);
1999 };
2000 auto PRMT__10 = GetPRMT(Op->getOperand(Num: 0), Op->getOperand(Num: 1), true, 0x3340);
2001 auto PRMT__32 = GetPRMT(Op->getOperand(Num: 2), Op->getOperand(Num: 3), true, 0x3340);
2002 auto PRMT3210 = GetPRMT(PRMT__10, PRMT__32, false, 0x5410);
2003 return DAG.getBitcast(VT, V: PRMT3210);
2004 }
2005
2006 // Get value or the Nth operand as an APInt(32). Undef values treated as 0.
2007 auto GetOperand = [](SDValue Op, int N) -> APInt {
2008 const SDValue &Operand = Op->getOperand(Num: N);
2009 EVT VT = Op->getValueType(ResNo: 0);
2010 if (Operand->isUndef())
2011 return APInt(32, 0);
2012 APInt Value;
2013 if (VT == MVT::v2f16 || VT == MVT::v2bf16)
2014 Value = cast<ConstantFPSDNode>(Val: Operand)->getValueAPF().bitcastToAPInt();
2015 else if (VT == MVT::v2i16 || VT == MVT::v4i8)
2016 Value = Operand->getAsAPIntVal();
2017 else
2018 llvm_unreachable("Unsupported type");
2019 // i8 values are carried around as i16, so we need to zero out upper bits,
2020 // so they do not get in the way of combining individual byte values
2021 if (VT == MVT::v4i8)
2022 Value = Value.trunc(width: 8);
2023 return Value.zext(width: 32);
2024 };
2025
2026 // Construct a 32-bit constant by shifting into place smaller values
2027 // (elements of the vector type VT).
2028 // For example, if VT has 2 elements, then N == 2:
2029 // ShiftAmount = 32 / N = 16
2030 // Value |= Op0 (b16) << 0
2031 // Value |= Op1 (b16) << 16
2032 // If N == 4:
2033 // ShiftAmount = 32 / N = 8
2034 // Value |= Op0 (b8) << 0
2035 // Value |= Op1 (b8) << 8
2036 // Value |= Op2 (b8) << 16
2037 // Value |= Op3 (b8) << 24
2038 // ...etc
2039 APInt Value(32, 0);
2040 const unsigned NumElements = VT.getVectorNumElements();
2041 assert(32 % NumElements == 0 && "must evenly divide bit length");
2042 const unsigned ShiftAmount = 32 / NumElements;
2043 for (unsigned ElementNo : seq(Size: NumElements))
2044 Value |= GetOperand(Op, ElementNo).shl(shiftAmt: ElementNo * ShiftAmount);
2045 SDValue Const = DAG.getConstant(Val: Value, DL, VT: MVT::i32);
2046 return DAG.getNode(Opcode: ISD::BITCAST, DL, VT: Op->getValueType(ResNo: 0), Operand: Const);
2047}
2048
2049SDValue NVPTXTargetLowering::LowerEXTRACT_VECTOR_ELT(SDValue Op,
2050 SelectionDAG &DAG) const {
2051 SDValue Index = Op->getOperand(Num: 1);
2052 SDValue Vector = Op->getOperand(Num: 0);
2053 SDLoc DL(Op);
2054 EVT VectorVT = Vector.getValueType();
2055
2056 if (VectorVT == MVT::v4i8) {
2057 SDValue Selector = DAG.getNode(Opcode: ISD::OR, DL, VT: MVT::i32,
2058 N1: DAG.getZExtOrTrunc(Op: Index, DL, VT: MVT::i32),
2059 N2: DAG.getConstant(Val: 0x7770, DL, VT: MVT::i32));
2060 SDValue PRMT = getPRMT(A: DAG.getBitcast(VT: MVT::i32, V: Vector),
2061 B: DAG.getConstant(Val: 0, DL, VT: MVT::i32), Selector, DL, DAG);
2062 SDValue Ext = DAG.getAnyExtOrTrunc(Op: PRMT, DL, VT: Op->getValueType(ResNo: 0));
2063 SDNodeFlags Flags;
2064 Flags.setNoSignedWrap(Ext.getScalarValueSizeInBits() > 8);
2065 Flags.setNoUnsignedWrap(Ext.getScalarValueSizeInBits() >= 8);
2066 Ext->setFlags(Flags);
2067 return Ext;
2068 }
2069
2070 // Constant index will be matched by tablegen.
2071 if (isa<ConstantSDNode>(Val: Index.getNode()))
2072 return Op;
2073
2074 // Extract individual elements and select one of them.
2075 assert(NVPTX::isPackedVectorTy(VectorVT) &&
2076 VectorVT.getVectorNumElements() == 2 && "Unexpected vector type.");
2077 EVT EltVT = VectorVT.getVectorElementType();
2078
2079 SDLoc dl(Op.getNode());
2080 SDValue E0 = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL: dl, VT: EltVT, N1: Vector,
2081 N2: DAG.getIntPtrConstant(Val: 0, DL: dl));
2082 SDValue E1 = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL: dl, VT: EltVT, N1: Vector,
2083 N2: DAG.getIntPtrConstant(Val: 1, DL: dl));
2084 return DAG.getSelectCC(DL: dl, LHS: Index, RHS: DAG.getIntPtrConstant(Val: 0, DL: dl), True: E0, False: E1,
2085 Cond: ISD::CondCode::SETEQ);
2086}
2087
2088SDValue NVPTXTargetLowering::LowerINSERT_VECTOR_ELT(SDValue Op,
2089 SelectionDAG &DAG) const {
2090 SDValue Vector = Op->getOperand(Num: 0);
2091 EVT VectorVT = Vector.getValueType();
2092
2093 if (VectorVT != MVT::v4i8)
2094 return Op;
2095 SDLoc DL(Op);
2096 SDValue Value = Op->getOperand(Num: 1);
2097 if (Value->isUndef())
2098 return Vector;
2099
2100 SDValue Index = Op->getOperand(Num: 2);
2101
2102 SDValue BFI =
2103 DAG.getNode(Opcode: NVPTXISD::BFI, DL, VT: MVT::i32,
2104 Ops: {DAG.getZExtOrTrunc(Op: Value, DL, VT: MVT::i32),
2105 DAG.getBitcast(VT: MVT::i32, V: Vector),
2106 DAG.getNode(Opcode: ISD::MUL, DL, VT: MVT::i32,
2107 N1: DAG.getZExtOrTrunc(Op: Index, DL, VT: MVT::i32),
2108 N2: DAG.getConstant(Val: 8, DL, VT: MVT::i32)),
2109 DAG.getConstant(Val: 8, DL, VT: MVT::i32)});
2110 return DAG.getNode(Opcode: ISD::BITCAST, DL, VT: Op->getValueType(ResNo: 0), Operand: BFI);
2111}
2112
2113SDValue NVPTXTargetLowering::LowerVECTOR_SHUFFLE(SDValue Op,
2114 SelectionDAG &DAG) const {
2115 SDValue V1 = Op.getOperand(i: 0);
2116 EVT VectorVT = V1.getValueType();
2117 if (VectorVT != MVT::v4i8 || Op.getValueType() != MVT::v4i8)
2118 return Op;
2119
2120 // Lower shuffle to PRMT instruction.
2121 const ShuffleVectorSDNode *SVN = cast<ShuffleVectorSDNode>(Val: Op.getNode());
2122 SDValue V2 = Op.getOperand(i: 1);
2123 uint32_t Selector = 0;
2124 for (auto I : llvm::enumerate(First: SVN->getMask())) {
2125 if (I.value() != -1) // -1 is a placeholder for undef.
2126 Selector |= (I.value() << (I.index() * 4));
2127 }
2128
2129 SDLoc DL(Op);
2130 SDValue PRMT = getPRMT(A: DAG.getBitcast(VT: MVT::i32, V: V1),
2131 B: DAG.getBitcast(VT: MVT::i32, V: V2), Selector, DL, DAG);
2132 return DAG.getBitcast(VT: Op.getValueType(), V: PRMT);
2133}
2134/// LowerShiftRightParts - Lower SRL_PARTS, SRA_PARTS, which
2135/// 1) returns two i32 values and take a 2 x i32 value to shift plus a shift
2136/// amount, or
2137/// 2) returns two i64 values and take a 2 x i64 value to shift plus a shift
2138/// amount.
2139SDValue NVPTXTargetLowering::LowerShiftRightParts(SDValue Op,
2140 SelectionDAG &DAG) const {
2141 assert(Op.getNumOperands() == 3 && "Not a double-shift!");
2142 assert(Op.getOpcode() == ISD::SRA_PARTS || Op.getOpcode() == ISD::SRL_PARTS);
2143
2144 EVT VT = Op.getValueType();
2145 unsigned VTBits = VT.getSizeInBits();
2146 SDLoc dl(Op);
2147 SDValue ShOpLo = Op.getOperand(i: 0);
2148 SDValue ShOpHi = Op.getOperand(i: 1);
2149 SDValue ShAmt = Op.getOperand(i: 2);
2150 unsigned Opc = (Op.getOpcode() == ISD::SRA_PARTS) ? ISD::SRA : ISD::SRL;
2151
2152 if (VTBits == 32 && STI.hasFeature(Feature: NVPTX::SM35)) {
2153 // For 32bit and sm35, we can use the funnel shift 'shf' instruction.
2154 // {dHi, dLo} = {aHi, aLo} >> Amt
2155 // dHi = aHi >> Amt
2156 // dLo = shf.r.clamp aLo, aHi, Amt
2157
2158 SDValue Hi = DAG.getNode(Opcode: Opc, DL: dl, VT, N1: ShOpHi, N2: ShAmt);
2159 SDValue Lo =
2160 DAG.getNode(Opcode: NVPTXISD::FSHR_CLAMP, DL: dl, VT, N1: ShOpHi, N2: ShOpLo, N3: ShAmt);
2161
2162 SDValue Ops[2] = { Lo, Hi };
2163 return DAG.getMergeValues(Ops, dl);
2164 } else {
2165 // {dHi, dLo} = {aHi, aLo} >> Amt
2166 // - if (Amt>=size) then
2167 // dLo = aHi >> (Amt-size)
2168 // dHi = aHi >> Amt (this is either all 0 or all 1)
2169 // else
2170 // dLo = (aLo >>logic Amt) | (aHi << (size-Amt))
2171 // dHi = aHi >> Amt
2172
2173 SDValue RevShAmt = DAG.getNode(Opcode: ISD::SUB, DL: dl, VT: MVT::i32,
2174 N1: DAG.getConstant(Val: VTBits, DL: dl, VT: MVT::i32),
2175 N2: ShAmt);
2176 SDValue Tmp1 = DAG.getNode(Opcode: ISD::SRL, DL: dl, VT, N1: ShOpLo, N2: ShAmt);
2177 SDValue ExtraShAmt = DAG.getNode(Opcode: ISD::SUB, DL: dl, VT: MVT::i32, N1: ShAmt,
2178 N2: DAG.getConstant(Val: VTBits, DL: dl, VT: MVT::i32));
2179 SDValue Tmp2 = DAG.getNode(Opcode: ISD::SHL, DL: dl, VT, N1: ShOpHi, N2: RevShAmt);
2180 SDValue FalseVal = DAG.getNode(Opcode: ISD::OR, DL: dl, VT, N1: Tmp1, N2: Tmp2);
2181 SDValue TrueVal = DAG.getNode(Opcode: Opc, DL: dl, VT, N1: ShOpHi, N2: ExtraShAmt);
2182
2183 SDValue Cmp = DAG.getSetCC(DL: dl, VT: MVT::i1, LHS: ShAmt,
2184 RHS: DAG.getConstant(Val: VTBits, DL: dl, VT: MVT::i32),
2185 Cond: ISD::SETGE);
2186 SDValue Hi = DAG.getNode(Opcode: Opc, DL: dl, VT, N1: ShOpHi, N2: ShAmt);
2187 SDValue Lo = DAG.getNode(Opcode: ISD::SELECT, DL: dl, VT, N1: Cmp, N2: TrueVal, N3: FalseVal);
2188
2189 SDValue Ops[2] = { Lo, Hi };
2190 return DAG.getMergeValues(Ops, dl);
2191 }
2192}
2193
2194/// LowerShiftLeftParts - Lower SHL_PARTS, which
2195/// 1) returns two i32 values and take a 2 x i32 value to shift plus a shift
2196/// amount, or
2197/// 2) returns two i64 values and take a 2 x i64 value to shift plus a shift
2198/// amount.
2199SDValue NVPTXTargetLowering::LowerShiftLeftParts(SDValue Op,
2200 SelectionDAG &DAG) const {
2201 assert(Op.getNumOperands() == 3 && "Not a double-shift!");
2202 assert(Op.getOpcode() == ISD::SHL_PARTS);
2203
2204 EVT VT = Op.getValueType();
2205 unsigned VTBits = VT.getSizeInBits();
2206 SDLoc dl(Op);
2207 SDValue ShOpLo = Op.getOperand(i: 0);
2208 SDValue ShOpHi = Op.getOperand(i: 1);
2209 SDValue ShAmt = Op.getOperand(i: 2);
2210
2211 if (VTBits == 32 && STI.hasFeature(Feature: NVPTX::SM35)) {
2212 // For 32bit and sm35, we can use the funnel shift 'shf' instruction.
2213 // {dHi, dLo} = {aHi, aLo} << Amt
2214 // dHi = shf.l.clamp aLo, aHi, Amt
2215 // dLo = aLo << Amt
2216
2217 SDValue Hi =
2218 DAG.getNode(Opcode: NVPTXISD::FSHL_CLAMP, DL: dl, VT, N1: ShOpHi, N2: ShOpLo, N3: ShAmt);
2219 SDValue Lo = DAG.getNode(Opcode: ISD::SHL, DL: dl, VT, N1: ShOpLo, N2: ShAmt);
2220
2221 SDValue Ops[2] = { Lo, Hi };
2222 return DAG.getMergeValues(Ops, dl);
2223 } else {
2224 // {dHi, dLo} = {aHi, aLo} << Amt
2225 // - if (Amt>=size) then
2226 // dLo = aLo << Amt (all 0)
2227 // dLo = aLo << (Amt-size)
2228 // else
2229 // dLo = aLo << Amt
2230 // dHi = (aHi << Amt) | (aLo >> (size-Amt))
2231
2232 SDValue RevShAmt = DAG.getNode(Opcode: ISD::SUB, DL: dl, VT: MVT::i32,
2233 N1: DAG.getConstant(Val: VTBits, DL: dl, VT: MVT::i32),
2234 N2: ShAmt);
2235 SDValue Tmp1 = DAG.getNode(Opcode: ISD::SHL, DL: dl, VT, N1: ShOpHi, N2: ShAmt);
2236 SDValue ExtraShAmt = DAG.getNode(Opcode: ISD::SUB, DL: dl, VT: MVT::i32, N1: ShAmt,
2237 N2: DAG.getConstant(Val: VTBits, DL: dl, VT: MVT::i32));
2238 SDValue Tmp2 = DAG.getNode(Opcode: ISD::SRL, DL: dl, VT, N1: ShOpLo, N2: RevShAmt);
2239 SDValue FalseVal = DAG.getNode(Opcode: ISD::OR, DL: dl, VT, N1: Tmp1, N2: Tmp2);
2240 SDValue TrueVal = DAG.getNode(Opcode: ISD::SHL, DL: dl, VT, N1: ShOpLo, N2: ExtraShAmt);
2241
2242 SDValue Cmp = DAG.getSetCC(DL: dl, VT: MVT::i1, LHS: ShAmt,
2243 RHS: DAG.getConstant(Val: VTBits, DL: dl, VT: MVT::i32),
2244 Cond: ISD::SETGE);
2245 SDValue Lo = DAG.getNode(Opcode: ISD::SHL, DL: dl, VT, N1: ShOpLo, N2: ShAmt);
2246 SDValue Hi = DAG.getNode(Opcode: ISD::SELECT, DL: dl, VT, N1: Cmp, N2: TrueVal, N3: FalseVal);
2247
2248 SDValue Ops[2] = { Lo, Hi };
2249 return DAG.getMergeValues(Ops, dl);
2250 }
2251}
2252
2253/// If the types match, convert the generic copysign to the NVPTXISD version,
2254/// otherwise bail ensuring that mismatched cases are properly expaned.
2255SDValue NVPTXTargetLowering::LowerFCOPYSIGN(SDValue Op,
2256 SelectionDAG &DAG) const {
2257 EVT VT = Op.getValueType();
2258 SDLoc DL(Op);
2259
2260 SDValue In1 = Op.getOperand(i: 0);
2261 SDValue In2 = Op.getOperand(i: 1);
2262 EVT SrcVT = In2.getValueType();
2263
2264 if (!SrcVT.bitsEq(VT))
2265 return SDValue();
2266
2267 return DAG.getNode(Opcode: NVPTXISD::FCOPYSIGN, DL, VT, N1: In1, N2: In2);
2268}
2269
2270SDValue NVPTXTargetLowering::LowerFROUND(SDValue Op, SelectionDAG &DAG) const {
2271 EVT VT = Op.getValueType();
2272
2273 if (VT == MVT::f32)
2274 return LowerFROUND32(Op, DAG);
2275
2276 if (VT == MVT::f64)
2277 return LowerFROUND64(Op, DAG);
2278
2279 llvm_unreachable("unhandled type");
2280}
2281
2282// This is the the rounding method used in CUDA libdevice in C like code:
2283// float roundf(float A)
2284// {
2285// float RoundedA = (float) (int) ( A > 0 ? (A + 0.5f) : (A - 0.5f));
2286// RoundedA = abs(A) > 0x1.0p23 ? A : RoundedA;
2287// return abs(A) < 0.5 ? (float)(int)A : RoundedA;
2288// }
2289SDValue NVPTXTargetLowering::LowerFROUND32(SDValue Op,
2290 SelectionDAG &DAG) const {
2291 SDLoc SL(Op);
2292 SDValue A = Op.getOperand(i: 0);
2293 EVT VT = Op.getValueType();
2294
2295 SDValue AbsA = DAG.getNode(Opcode: ISD::FABS, DL: SL, VT, Operand: A);
2296
2297 // RoundedA = (float) (int) ( A > 0 ? (A + 0.5f) : (A - 0.5f))
2298 SDValue Bitcast = DAG.getNode(Opcode: ISD::BITCAST, DL: SL, VT: MVT::i32, Operand: A);
2299 const unsigned SignBitMask = 0x80000000;
2300 SDValue Sign = DAG.getNode(Opcode: ISD::AND, DL: SL, VT: MVT::i32, N1: Bitcast,
2301 N2: DAG.getConstant(Val: SignBitMask, DL: SL, VT: MVT::i32));
2302 const unsigned PointFiveInBits = 0x3F000000;
2303 SDValue PointFiveWithSignRaw =
2304 DAG.getNode(Opcode: ISD::OR, DL: SL, VT: MVT::i32, N1: Sign,
2305 N2: DAG.getConstant(Val: PointFiveInBits, DL: SL, VT: MVT::i32));
2306 SDValue PointFiveWithSign =
2307 DAG.getNode(Opcode: ISD::BITCAST, DL: SL, VT, Operand: PointFiveWithSignRaw);
2308 SDValue AdjustedA = DAG.getNode(Opcode: ISD::FADD, DL: SL, VT, N1: A, N2: PointFiveWithSign);
2309 SDValue RoundedA = DAG.getNode(Opcode: ISD::FTRUNC, DL: SL, VT, Operand: AdjustedA);
2310
2311 // RoundedA = abs(A) > 0x1.0p23 ? A : RoundedA;
2312 EVT SetCCVT = getSetCCResultType(DL: DAG.getDataLayout(), Ctx&: *DAG.getContext(), VT);
2313 SDValue IsLarge =
2314 DAG.getSetCC(DL: SL, VT: SetCCVT, LHS: AbsA, RHS: DAG.getConstantFP(Val: pow(x: 2.0, y: 23.0), DL: SL, VT),
2315 Cond: ISD::SETOGT);
2316 RoundedA = DAG.getNode(Opcode: ISD::SELECT, DL: SL, VT, N1: IsLarge, N2: A, N3: RoundedA);
2317
2318 // return abs(A) < 0.5 ? (float)(int)A : RoundedA;
2319 SDValue IsSmall =DAG.getSetCC(DL: SL, VT: SetCCVT, LHS: AbsA,
2320 RHS: DAG.getConstantFP(Val: 0.5, DL: SL, VT), Cond: ISD::SETOLT);
2321 SDValue RoundedAForSmallA = DAG.getNode(Opcode: ISD::FTRUNC, DL: SL, VT, Operand: A);
2322 return DAG.getNode(Opcode: ISD::SELECT, DL: SL, VT, N1: IsSmall, N2: RoundedAForSmallA, N3: RoundedA);
2323}
2324
2325// The implementation of round(double) is similar to that of round(float) in
2326// that they both separate the value range into three regions and use a method
2327// specific to the region to round the values. However, round(double) first
2328// calculates the round of the absolute value and then adds the sign back while
2329// round(float) directly rounds the value with sign.
2330SDValue NVPTXTargetLowering::LowerFROUND64(SDValue Op,
2331 SelectionDAG &DAG) const {
2332 SDLoc SL(Op);
2333 SDValue A = Op.getOperand(i: 0);
2334 EVT VT = Op.getValueType();
2335
2336 SDValue AbsA = DAG.getNode(Opcode: ISD::FABS, DL: SL, VT, Operand: A);
2337
2338 // double RoundedA = (double) (int) (abs(A) + 0.5f);
2339 SDValue AdjustedA = DAG.getNode(Opcode: ISD::FADD, DL: SL, VT, N1: AbsA,
2340 N2: DAG.getConstantFP(Val: 0.5, DL: SL, VT));
2341 SDValue RoundedA = DAG.getNode(Opcode: ISD::FTRUNC, DL: SL, VT, Operand: AdjustedA);
2342
2343 // RoundedA = abs(A) < 0.5 ? (double)0 : RoundedA;
2344 EVT SetCCVT = getSetCCResultType(DL: DAG.getDataLayout(), Ctx&: *DAG.getContext(), VT);
2345 SDValue IsSmall =DAG.getSetCC(DL: SL, VT: SetCCVT, LHS: AbsA,
2346 RHS: DAG.getConstantFP(Val: 0.5, DL: SL, VT), Cond: ISD::SETOLT);
2347 RoundedA = DAG.getNode(Opcode: ISD::SELECT, DL: SL, VT, N1: IsSmall,
2348 N2: DAG.getConstantFP(Val: 0, DL: SL, VT),
2349 N3: RoundedA);
2350
2351 // Add sign to rounded_A
2352 RoundedA = DAG.getNode(Opcode: ISD::FCOPYSIGN, DL: SL, VT, N1: RoundedA, N2: A);
2353 DAG.getNode(Opcode: ISD::FTRUNC, DL: SL, VT, Operand: A);
2354
2355 // RoundedA = abs(A) > 0x1.0p52 ? A : RoundedA;
2356 SDValue IsLarge =
2357 DAG.getSetCC(DL: SL, VT: SetCCVT, LHS: AbsA, RHS: DAG.getConstantFP(Val: pow(x: 2.0, y: 52.0), DL: SL, VT),
2358 Cond: ISD::SETOGT);
2359 return DAG.getNode(Opcode: ISD::SELECT, DL: SL, VT, N1: IsLarge, N2: A, N3: RoundedA);
2360}
2361
2362static SDValue PromoteBinOpToF32(SDNode *N, SelectionDAG &DAG) {
2363 EVT VT = N->getValueType(ResNo: 0);
2364 EVT NVT = MVT::f32;
2365 if (VT.isVector()) {
2366 NVT = EVT::getVectorVT(Context&: *DAG.getContext(), VT: NVT, EC: VT.getVectorElementCount());
2367 }
2368 SDLoc DL(N);
2369 SDValue Tmp0 = DAG.getFPExtendOrRound(Op: N->getOperand(Num: 0), DL, VT: NVT);
2370 SDValue Tmp1 = DAG.getFPExtendOrRound(Op: N->getOperand(Num: 1), DL, VT: NVT);
2371 SDValue Res = DAG.getNode(Opcode: N->getOpcode(), DL, VT: NVT, N1: Tmp0, N2: Tmp1, Flags: N->getFlags());
2372 return DAG.getFPExtendOrRound(Op: Res, DL, VT);
2373}
2374
2375SDValue NVPTXTargetLowering::PromoteBinOpIfF32FTZ(SDValue Op,
2376 SelectionDAG &DAG) const {
2377 if (useF32FTZ(MF: DAG.getMachineFunction())) {
2378 return PromoteBinOpToF32(N: Op.getNode(), DAG);
2379 }
2380 return Op;
2381}
2382
2383SDValue NVPTXTargetLowering::LowerINT_TO_FP(SDValue Op,
2384 SelectionDAG &DAG) const {
2385 assert(!STI.hasFeature(NVPTX::SM90));
2386
2387 if (Op.getValueType() == MVT::bf16) {
2388 SDLoc Loc(Op);
2389 return DAG.getNode(
2390 Opcode: ISD::FP_ROUND, DL: Loc, VT: MVT::bf16,
2391 N1: DAG.getNode(Opcode: Op.getOpcode(), DL: Loc, VT: MVT::f32, Operand: Op.getOperand(i: 0)),
2392 N2: DAG.getIntPtrConstant(Val: 0, DL: Loc, /*isTarget=*/true));
2393 }
2394
2395 // Everything else is considered legal.
2396 return Op;
2397}
2398
2399SDValue NVPTXTargetLowering::LowerFP_TO_INT(SDValue Op,
2400 SelectionDAG &DAG) const {
2401 assert(!STI.hasFeature(NVPTX::SM90));
2402
2403 if (Op.getOperand(i: 0).getValueType() == MVT::bf16) {
2404 SDLoc Loc(Op);
2405 return DAG.getNode(
2406 Opcode: Op.getOpcode(), DL: Loc, VT: Op.getValueType(),
2407 Operand: DAG.getNode(Opcode: ISD::FP_EXTEND, DL: Loc, VT: MVT::f32, Operand: Op.getOperand(i: 0)));
2408 }
2409
2410 // Everything else is considered legal.
2411 return Op;
2412}
2413
2414SDValue NVPTXTargetLowering::LowerFP_ROUND(SDValue Op,
2415 SelectionDAG &DAG) const {
2416 EVT NarrowVT = Op.getValueType();
2417 SDValue Wide = Op.getOperand(i: 0);
2418 EVT WideVT = Wide.getValueType();
2419 if (NarrowVT.getScalarType() == MVT::bf16) {
2420 const TargetLowering *TLI = STI.getTargetLowering();
2421 if (!STI.hasFeature(Feature: NVPTX::SM80)) {
2422 return TLI->expandFP_ROUND(Node: Op.getNode(), DAG);
2423 }
2424 if (!STI.hasFeature(Feature: NVPTX::SM90)) {
2425 // sm_80 was the first architecture to support f32 -> bf16.
2426 if (WideVT.getScalarType() == MVT::f32) {
2427 return Op;
2428 }
2429 if (WideVT.getScalarType() == MVT::f64) {
2430 SDLoc Loc(Op);
2431 // Round-inexact-to-odd f64 to f32, then do the final rounding using
2432 // the hardware f32 -> bf16 instruction.
2433 SDValue rod = TLI->expandRoundInexactToOdd(
2434 ResultVT: WideVT.changeElementType(Context&: *DAG.getContext(), EltVT: MVT::f32), Op: Wide, DL: Loc,
2435 DAG);
2436 return DAG.getFPExtendOrRound(Op: rod, DL: Loc, VT: NarrowVT);
2437 }
2438 return TLI->expandFP_ROUND(Node: Op.getNode(), DAG);
2439 }
2440 }
2441
2442 // Everything else is considered legal.
2443 return Op;
2444}
2445
2446static SDValue lowerFPTRUNC_ROUND(SDValue Op, SelectionDAG &DAG,
2447 const NVPTXSubtarget &STI) {
2448 EVT SrcVT = Op.getOperand(i: 0).getValueType();
2449 EVT DstVT = Op.getValueType();
2450 auto RM = static_cast<RoundingMode>(Op.getConstantOperandVal(i: 1));
2451 if (RM == RoundingMode::NearestTiesToEven) {
2452 // Reuse the native selection and fallback expansion for ordinary fptrunc.
2453 SDLoc DL(Op);
2454 return DAG.getNode(Opcode: ISD::FP_ROUND, DL, VT: DstVT, N1: Op.getOperand(i: 0),
2455 N2: DAG.getIntPtrConstant(Val: 0, DL, /*isTarget=*/true),
2456 Flags: Op.getNode()->getFlags());
2457 }
2458
2459 bool RoundToInfinity =
2460 RM == RoundingMode::TowardNegative || RM == RoundingMode::TowardPositive;
2461
2462 bool Supported = RM == RoundingMode::TowardZero || RoundToInfinity;
2463 if (DstVT == MVT::bf16) {
2464 if (SrcVT == MVT::f64 || RoundToInfinity)
2465 Supported &= STI.hasFeature(Feature: NVPTX::SM90);
2466 else
2467 Supported &= STI.hasFeature(Feature: NVPTX::SM80);
2468 }
2469
2470 if (Supported)
2471 return Op;
2472
2473 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
2474 DAG.getMachineFunction().getFunction(),
2475 "unsupported conversion or rounding mode for llvm.fptrunc.round",
2476 SDLoc(Op).getDebugLoc()));
2477 return DAG.getPOISON(VT: DstVT);
2478}
2479
2480SDValue NVPTXTargetLowering::LowerFP_EXTEND(SDValue Op,
2481 SelectionDAG &DAG) const {
2482 SDValue Narrow = Op.getOperand(i: 0);
2483 EVT NarrowVT = Narrow.getValueType();
2484 EVT WideVT = Op.getValueType();
2485 if (NarrowVT.getScalarType() == MVT::bf16) {
2486 if (WideVT.getScalarType() == MVT::f32 &&
2487 (!STI.hasFeature(Feature: NVPTX::SM80) || !STI.hasFeature(Feature: NVPTX::PTX71))) {
2488 SDLoc Loc(Op);
2489 return DAG.getNode(Opcode: ISD::BF16_TO_FP, DL: Loc, VT: WideVT, Operand: Narrow);
2490 }
2491 if (WideVT.getScalarType() == MVT::f64 && !STI.hasFeature(Feature: NVPTX::SM90)) {
2492 EVT F32 = NarrowVT.changeElementType(Context&: *DAG.getContext(), EltVT: MVT::f32);
2493 SDLoc Loc(Op);
2494 if (STI.hasFeature(Feature: NVPTX::SM80) && STI.hasFeature(Feature: NVPTX::PTX71)) {
2495 Op = DAG.getNode(Opcode: ISD::FP_EXTEND, DL: Loc, VT: F32, Operand: Narrow);
2496 } else {
2497 Op = DAG.getNode(Opcode: ISD::BF16_TO_FP, DL: Loc, VT: F32, Operand: Narrow);
2498 }
2499 return DAG.getNode(Opcode: ISD::FP_EXTEND, DL: Loc, VT: WideVT, Operand: Op);
2500 }
2501 }
2502
2503 // Everything else is considered legal.
2504 return Op;
2505}
2506
2507static SDValue LowerVectorArith(SDValue Op, SelectionDAG &DAG) {
2508 SDLoc DL(Op);
2509 if (Op.getValueType() != MVT::v2i16)
2510 return Op;
2511 EVT EltVT = Op.getValueType().getVectorElementType();
2512 SmallVector<SDValue> VecElements;
2513 for (int I = 0, E = Op.getValueType().getVectorNumElements(); I < E; I++) {
2514 SmallVector<SDValue> ScalarArgs;
2515 llvm::transform(Range: Op->ops(), d_first: std::back_inserter(x&: ScalarArgs),
2516 F: [&](const SDUse &O) {
2517 return DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT,
2518 N1: O.get(), N2: DAG.getIntPtrConstant(Val: I, DL));
2519 });
2520 VecElements.push_back(Elt: DAG.getNode(Opcode: Op.getOpcode(), DL, VT: EltVT, Ops: ScalarArgs));
2521 }
2522 SDValue V =
2523 DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL, VT: Op.getValueType(), Ops: VecElements);
2524 return V;
2525}
2526
2527static SDValue lowerTcgen05St(SDValue Op, SelectionDAG &DAG,
2528 bool hasOffset = false) {
2529 // skip lowering if the vector operand is already legalized
2530 if (!Op->getOperand(Num: hasOffset ? 4 : 3).getValueType().isVector())
2531 return Op;
2532
2533 SDNode *N = Op.getNode();
2534 SDLoc DL(N);
2535 SmallVector<SDValue, 32> Ops;
2536
2537 // split the vector argument
2538 for (size_t I = 0; I < N->getNumOperands(); I++) {
2539 SDValue Val = N->getOperand(Num: I);
2540 EVT ValVT = Val.getValueType();
2541 if (ValVT.isVector()) {
2542 EVT EltVT = ValVT.getVectorElementType();
2543 for (unsigned J = 0, NElts = ValVT.getVectorNumElements(); J < NElts; J++)
2544 Ops.push_back(Elt: DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT, N1: Val,
2545 N2: DAG.getIntPtrConstant(Val: J, DL)));
2546 } else
2547 Ops.push_back(Elt: Val);
2548 }
2549
2550 MemIntrinsicSDNode *MemSD = cast<MemIntrinsicSDNode>(Val: N);
2551 SDValue Tcgen05StNode =
2552 DAG.getMemIntrinsicNode(Opcode: ISD::INTRINSIC_VOID, dl: DL, VTList: N->getVTList(), Ops,
2553 MemVT: MemSD->getMemoryVT(), MMO: MemSD->getMemOperand());
2554
2555 return Tcgen05StNode;
2556}
2557
2558static SDValue lowerBSWAP(SDValue Op, SelectionDAG &DAG) {
2559 SDLoc DL(Op);
2560 SDValue Src = Op.getOperand(i: 0);
2561 EVT VT = Op.getValueType();
2562
2563 switch (VT.getSimpleVT().SimpleTy) {
2564 case MVT::i16: {
2565 SDValue Extended = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT: MVT::i32, Operand: Src);
2566 SDValue Swapped =
2567 getPRMT(A: Extended, B: DAG.getConstant(Val: 0, DL, VT: MVT::i32), Selector: 0x7701, DL, DAG);
2568 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: MVT::i16, Operand: Swapped);
2569 }
2570 case MVT::i32: {
2571 return getPRMT(A: Src, B: DAG.getConstant(Val: 0, DL, VT: MVT::i32), Selector: 0x0123, DL, DAG);
2572 }
2573 case MVT::v2i16: {
2574 SDValue Converted = DAG.getBitcast(VT: MVT::i32, V: Src);
2575 SDValue Swapped =
2576 getPRMT(A: Converted, B: DAG.getConstant(Val: 0, DL, VT: MVT::i32), Selector: 0x2301, DL, DAG);
2577 return DAG.getNode(Opcode: ISD::BITCAST, DL, VT: MVT::v2i16, Operand: Swapped);
2578 }
2579 case MVT::i64: {
2580 SDValue UnpackSrc =
2581 DAG.getNode(Opcode: NVPTXISD::UNPACK_VECTOR, DL, ResultTys: {MVT::i32, MVT::i32}, Ops: Src);
2582 SDValue SwappedLow =
2583 getPRMT(A: UnpackSrc.getValue(R: 0), B: DAG.getConstant(Val: 0, DL, VT: MVT::i32), Selector: 0x0123,
2584 DL, DAG);
2585 SDValue SwappedHigh =
2586 getPRMT(A: UnpackSrc.getValue(R: 1), B: DAG.getConstant(Val: 0, DL, VT: MVT::i32), Selector: 0x0123,
2587 DL, DAG);
2588 return DAG.getNode(Opcode: NVPTXISD::BUILD_VECTOR, DL, VT: MVT::i64,
2589 Ops: {SwappedHigh, SwappedLow});
2590 }
2591 default:
2592 llvm_unreachable("unsupported type for bswap");
2593 }
2594}
2595
2596static SDValue lowerStAsyncWithMbarrier(SDValue Op, SelectionDAG &DAG) {
2597 const Function &Fn = DAG.getMachineFunction().getFunction();
2598 SDNode *N = Op.getNode();
2599 SDLoc DL(N);
2600 Intrinsic::ID IntrinsicID = N->getConstantOperandVal(Num: 1);
2601 SDValue DestAddr = N->getOperand(Num: 2);
2602 SDValue Value = N->getOperand(Num: 3);
2603 SDValue MbarAddr = N->getOperand(Num: 4);
2604
2605 MVT ValueVT = Value.getSimpleValueType();
2606
2607 if (ValueVT == MVT::i32 || ValueVT == MVT::i64)
2608 return Op;
2609
2610 if (ValueVT == MVT::i128) {
2611 SDValue Cast = DAG.getNode(Opcode: ISD::BITCAST, DL, VT: MVT::v2i64, Operand: Value);
2612 SDValue ValueLo = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i64, N1: Cast,
2613 N2: DAG.getIntPtrConstant(Val: 0, DL));
2614 SDValue ValueHi = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i64, N1: Cast,
2615 N2: DAG.getIntPtrConstant(Val: 1, DL));
2616 SDValue Ops[] = {N->getOperand(Num: 0), DestAddr, ValueLo, ValueHi, MbarAddr};
2617 return DAG.getNode(Opcode: NVPTXISD::ST_ASYNC_MBARRIER_B128, DL, VT: MVT::Other, Ops);
2618 }
2619
2620 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
2621 Fn,
2622 Twine("unsupported argument type ") + llvm::EVT(ValueVT).getEVTString() +
2623 " for " + llvm::Intrinsic::getName(id: IntrinsicID) + " intrinsic",
2624 DiagnosticLocation(DL.getDebugLoc())));
2625 return Op.getOperand(i: 0); // Return only the chain
2626}
2627
2628static SDValue lowerStAsyncRelease(SDValue Op, SelectionDAG &DAG) {
2629 const Function &Fn = DAG.getMachineFunction().getFunction();
2630 SDNode *N = Op.getNode();
2631 SDLoc DL(N);
2632 Intrinsic::ID IntrinsicID = N->getConstantOperandVal(Num: 1);
2633 SDValue DestAddr = N->getOperand(Num: 2);
2634 SDValue Value = N->getOperand(Num: 3);
2635
2636 MVT ValueVT = Value.getSimpleValueType();
2637
2638 if (ValueVT == MVT::i16 || ValueVT == MVT::i32 || ValueVT == MVT::i64)
2639 return Op;
2640
2641 if (ValueVT == MVT::i8) {
2642 unsigned OpCode;
2643 switch (IntrinsicID) {
2644 case Intrinsic::nvvm_st_async_sys:
2645 OpCode = NVPTXISD::ST_ASYNC_SYS_B8;
2646 break;
2647 case Intrinsic::nvvm_st_async_gpu:
2648 OpCode = NVPTXISD::ST_ASYNC_GPU_B8;
2649 break;
2650 case Intrinsic::nvvm_st_async_mmio_sys:
2651 OpCode = NVPTXISD::ST_ASYNC_MMIO_SYS_B8;
2652 break;
2653 default:
2654 llvm_unreachable("unexpected intrinsic ID for st.async.release");
2655 }
2656
2657 Value = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: MVT::i16, Operand: Value);
2658
2659 // The `.mmio` variant has no multimem form and therefore no `isMultimem`
2660 // operand.
2661 if (IntrinsicID == Intrinsic::nvvm_st_async_mmio_sys) {
2662 SDValue Ops[] = {N->getOperand(Num: 0), DestAddr, Value};
2663 return DAG.getNode(Opcode: OpCode, DL, VT: MVT::Other, Ops);
2664 }
2665
2666 SDValue IsMultimem =
2667 DAG.getTargetConstant(Val: N->getConstantOperandVal(Num: 4), DL, VT: MVT::i1);
2668 SDValue Ops[] = {N->getOperand(Num: 0), DestAddr, Value, IsMultimem};
2669 return DAG.getNode(Opcode: OpCode, DL, VT: MVT::Other, Ops);
2670 }
2671
2672 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
2673 Fn,
2674 Twine("unsupported argument type ") + llvm::EVT(ValueVT).getEVTString() +
2675 " for " + llvm::Intrinsic::getName(id: IntrinsicID) + " intrinsic",
2676 DiagnosticLocation(DL.getDebugLoc())));
2677 return Op.getOperand(i: 0); // Return only the chain
2678}
2679
2680static unsigned getTcgen05MMADisableOutputLane(unsigned IID) {
2681 switch (IID) {
2682 case Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg1:
2683 return NVPTXISD::TCGEN05_MMA_SHARED_DISABLE_OUTPUT_LANE_CG1;
2684 case Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg2:
2685 return NVPTXISD::TCGEN05_MMA_SHARED_DISABLE_OUTPUT_LANE_CG2;
2686 case Intrinsic::nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg1:
2687 return NVPTXISD::TCGEN05_MMA_SHARED_SCALE_D_DISABLE_OUTPUT_LANE_CG1;
2688 case Intrinsic::nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg2:
2689 return NVPTXISD::TCGEN05_MMA_SHARED_SCALE_D_DISABLE_OUTPUT_LANE_CG2;
2690 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg1:
2691 return NVPTXISD::TCGEN05_MMA_TENSOR_DISABLE_OUTPUT_LANE_CG1;
2692 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg2:
2693 return NVPTXISD::TCGEN05_MMA_TENSOR_DISABLE_OUTPUT_LANE_CG2;
2694 case Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1:
2695 return NVPTXISD::TCGEN05_MMA_TENSOR_SCALE_D_DISABLE_OUTPUT_LANE_CG1;
2696 case Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2:
2697 return NVPTXISD::TCGEN05_MMA_TENSOR_SCALE_D_DISABLE_OUTPUT_LANE_CG2;
2698 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg1_ashift:
2699 return NVPTXISD::TCGEN05_MMA_TENSOR_DISABLE_OUTPUT_LANE_CG1_ASHIFT;
2700 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg2_ashift:
2701 return NVPTXISD::TCGEN05_MMA_TENSOR_DISABLE_OUTPUT_LANE_CG2_ASHIFT;
2702 case Intrinsic::
2703 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1_ashift:
2704 return NVPTXISD::TCGEN05_MMA_TENSOR_SCALE_D_DISABLE_OUTPUT_LANE_CG1_ASHIFT;
2705 case Intrinsic::
2706 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2_ashift:
2707 return NVPTXISD::TCGEN05_MMA_TENSOR_SCALE_D_DISABLE_OUTPUT_LANE_CG2_ASHIFT;
2708 case Intrinsic::nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg1:
2709 return NVPTXISD::TCGEN05_MMA_SP_SHARED_DISABLE_OUTPUT_LANE_CG1;
2710 case Intrinsic::nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg2:
2711 return NVPTXISD::TCGEN05_MMA_SP_SHARED_DISABLE_OUTPUT_LANE_CG2;
2712 case Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg1:
2713 return NVPTXISD::TCGEN05_MMA_SP_SHARED_SCALE_D_DISABLE_OUTPUT_LANE_CG1;
2714 case Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg2:
2715 return NVPTXISD::TCGEN05_MMA_SP_SHARED_SCALE_D_DISABLE_OUTPUT_LANE_CG2;
2716 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1:
2717 return NVPTXISD::TCGEN05_MMA_SP_TENSOR_DISABLE_OUTPUT_LANE_CG1;
2718 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2:
2719 return NVPTXISD::TCGEN05_MMA_SP_TENSOR_DISABLE_OUTPUT_LANE_CG2;
2720 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1_ashift:
2721 return NVPTXISD::TCGEN05_MMA_SP_TENSOR_DISABLE_OUTPUT_LANE_CG1_ASHIFT;
2722 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2_ashift:
2723 return NVPTXISD::TCGEN05_MMA_SP_TENSOR_DISABLE_OUTPUT_LANE_CG2_ASHIFT;
2724 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1:
2725 return NVPTXISD::TCGEN05_MMA_SP_TENSOR_SCALE_D_DISABLE_OUTPUT_LANE_CG1;
2726 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2:
2727 return NVPTXISD::TCGEN05_MMA_SP_TENSOR_SCALE_D_DISABLE_OUTPUT_LANE_CG2;
2728 case Intrinsic::
2729 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1_ashift:
2730 return NVPTXISD::
2731 TCGEN05_MMA_SP_TENSOR_SCALE_D_DISABLE_OUTPUT_LANE_CG1_ASHIFT;
2732 case Intrinsic::
2733 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2_ashift:
2734 return NVPTXISD::
2735 TCGEN05_MMA_SP_TENSOR_SCALE_D_DISABLE_OUTPUT_LANE_CG2_ASHIFT;
2736 case Intrinsic::
2737 nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg1_decompress_b:
2738 return NVPTXISD::TCGEN05_MMA_SHARED_DISABLE_OUTPUT_LANE_CG1_DECOMPRESS_B;
2739 case Intrinsic::
2740 nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg2_decompress_b:
2741 return NVPTXISD::TCGEN05_MMA_SHARED_DISABLE_OUTPUT_LANE_CG2_DECOMPRESS_B;
2742 case Intrinsic::
2743 nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg1_decompress_b:
2744 return NVPTXISD::TCGEN05_MMA_TENSOR_DISABLE_OUTPUT_LANE_CG1_DECOMPRESS_B;
2745 case Intrinsic::
2746 nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg2_decompress_b:
2747 return NVPTXISD::TCGEN05_MMA_TENSOR_DISABLE_OUTPUT_LANE_CG2_DECOMPRESS_B;
2748 };
2749 llvm_unreachable("unhandled tcgen05.mma.disable_output_lane intrinsic");
2750}
2751
2752static SDValue LowerTcgen05MMADisableOutputLane(SDValue Op, SelectionDAG &DAG) {
2753 SDNode *N = Op.getNode();
2754 SDLoc DL(N);
2755 unsigned IID = cast<ConstantSDNode>(Val: N->getOperand(Num: 1))->getZExtValue();
2756
2757 SmallVector<SDValue, 16> Ops;
2758 // split the vector argument
2759 for (size_t I = 0; I < N->getNumOperands(); I++) {
2760 if (I == 1)
2761 continue; // skip IID
2762 SDValue Val = N->getOperand(Num: I);
2763 EVT ValVT = Val.getValueType();
2764 if (ValVT.isVector()) {
2765 EVT EltVT = ValVT.getVectorElementType();
2766 for (unsigned J = 0, NElts = ValVT.getVectorNumElements(); J < NElts; J++)
2767 Ops.push_back(Elt: DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT, N1: Val,
2768 N2: DAG.getIntPtrConstant(Val: J, DL)));
2769 } else
2770 Ops.push_back(Elt: Val);
2771 }
2772
2773 MemIntrinsicSDNode *MemSD = cast<MemIntrinsicSDNode>(Val: N);
2774 SDValue Tcgen05MMANode = DAG.getMemIntrinsicNode(
2775 Opcode: getTcgen05MMADisableOutputLane(IID), dl: DL, VTList: N->getVTList(), Ops,
2776 MemVT: MemSD->getMemoryVT(), MMO: MemSD->getMemOperand());
2777
2778 return Tcgen05MMANode;
2779}
2780
2781// Lower vector return type of tcgen05.ld intrinsics
2782static std::optional<std::pair<SDValue, SDValue>>
2783lowerTcgen05Ld(SDNode *N, SelectionDAG &DAG, bool HasOffset = false) {
2784 SDLoc DL(N);
2785 EVT ResVT = N->getValueType(ResNo: 0);
2786 if (!ResVT.isVector())
2787 return {}; // already legalized.
2788
2789 const unsigned NumElts = ResVT.getVectorNumElements();
2790
2791 // Create the return type of the instructions
2792 SmallVector<EVT, 5> ListVTs;
2793 for (unsigned i = 0; i < NumElts; ++i)
2794 ListVTs.push_back(Elt: MVT::i32);
2795
2796 ListVTs.push_back(Elt: N->getValueType(ResNo: 1)); // Chain
2797
2798 SDVTList ResVTs = DAG.getVTList(VTs: ListVTs);
2799
2800 SmallVector<SDValue, 8> Ops{N->getOperand(Num: 0), N->getOperand(Num: 1),
2801 N->getOperand(Num: 2)};
2802
2803 if (HasOffset) {
2804 Ops.push_back(Elt: N->getOperand(Num: 3)); // offset
2805 Ops.push_back(Elt: N->getOperand(Num: 4)); // Pack flag
2806 } else
2807 Ops.push_back(Elt: N->getOperand(Num: 3)); // Pack flag
2808
2809 MemIntrinsicSDNode *MemSD = cast<MemIntrinsicSDNode>(Val: N);
2810 SDValue NewNode =
2811 DAG.getMemIntrinsicNode(Opcode: ISD::INTRINSIC_W_CHAIN, dl: DL, VTList: ResVTs, Ops,
2812 MemVT: MemSD->getMemoryVT(), MMO: MemSD->getMemOperand());
2813
2814 // split the vector result
2815 SmallVector<SDValue, 4> ScalarRes;
2816 for (unsigned i = 0; i < NumElts; ++i) {
2817 SDValue Res = NewNode.getValue(R: i);
2818 ScalarRes.push_back(Elt: Res);
2819 }
2820
2821 SDValue Chain = NewNode.getValue(R: NumElts);
2822 SDValue BuildVector = DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL, VT: ResVT, Ops: ScalarRes);
2823 return {{BuildVector, Chain}};
2824}
2825
2826static SDValue reportInvalidTensormapReplaceUsage(SDValue Op, SelectionDAG &DAG,
2827 unsigned Val) {
2828 SDNode *N = Op.getNode();
2829 SDLoc DL(N);
2830
2831 const Function &Fn = DAG.getMachineFunction().getFunction();
2832
2833 unsigned AS = 0;
2834 if (auto *MemN = dyn_cast<MemIntrinsicSDNode>(Val: N))
2835 AS = MemN->getAddressSpace();
2836 Type *PtrTy = PointerType::get(C&: *DAG.getContext(), AddressSpace: AS);
2837 Module *M = DAG.getMachineFunction().getFunction().getParent();
2838
2839 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
2840 Fn,
2841 "Intrinsic " +
2842 Intrinsic::getName(Id: N->getConstantOperandVal(Num: 1), OverloadTys: {PtrTy}, M) +
2843 " with value " + Twine(Val) +
2844 " is not supported on the given target.",
2845 DL.getDebugLoc()));
2846 return Op.getOperand(i: 0);
2847}
2848
2849static SDValue lowerTensormapReplaceElemtype(SDValue Op, SelectionDAG &DAG) {
2850 SDNode *N = Op.getNode();
2851 SDLoc DL(N);
2852
2853 // immediate argument representing elemtype
2854 unsigned Val = N->getConstantOperandVal(Num: 3);
2855
2856 if (!DAG.getSubtarget<NVPTXSubtarget>().hasTensormapReplaceElemtypeSupport(
2857 ElemType: Val))
2858 return reportInvalidTensormapReplaceUsage(Op, DAG, Val);
2859
2860 return Op;
2861}
2862
2863static SDValue lowerTensormapReplaceSwizzleMode(SDValue Op, SelectionDAG &DAG) {
2864 SDNode *N = Op.getNode();
2865 SDLoc DL(N);
2866
2867 // immediate argument representing swizzle mode
2868 unsigned Val = N->getConstantOperandVal(Num: 3);
2869
2870 if (!DAG.getSubtarget<NVPTXSubtarget>().hasTensormapReplaceSwizzleModeSupport(
2871 SwizzleMode: Val))
2872 return reportInvalidTensormapReplaceUsage(Op, DAG, Val);
2873
2874 return Op;
2875}
2876
2877static SDValue lowerIntrinsicVoid(SDValue Op, SelectionDAG &DAG) {
2878 SDNode *N = Op.getNode();
2879 SDValue Intrin = N->getOperand(Num: 1);
2880
2881 // Get the intrinsic ID
2882 unsigned IntrinNo = cast<ConstantSDNode>(Val: Intrin.getNode())->getZExtValue();
2883 switch (IntrinNo) {
2884 default:
2885 break;
2886 case Intrinsic::nvvm_st_async:
2887 return lowerStAsyncWithMbarrier(Op, DAG);
2888 case Intrinsic::nvvm_st_async_sys:
2889 case Intrinsic::nvvm_st_async_gpu:
2890 case Intrinsic::nvvm_st_async_mmio_sys:
2891 return lowerStAsyncRelease(Op, DAG);
2892
2893 case Intrinsic::nvvm_tcgen05_st_16x64b_x1:
2894 case Intrinsic::nvvm_tcgen05_st_16x64b_x2:
2895 case Intrinsic::nvvm_tcgen05_st_16x64b_x4:
2896 case Intrinsic::nvvm_tcgen05_st_16x64b_x8:
2897 case Intrinsic::nvvm_tcgen05_st_16x64b_x16:
2898 case Intrinsic::nvvm_tcgen05_st_16x64b_x32:
2899 case Intrinsic::nvvm_tcgen05_st_16x64b_x128:
2900 case Intrinsic::nvvm_tcgen05_st_16x128b_x1:
2901 case Intrinsic::nvvm_tcgen05_st_16x128b_x2:
2902 case Intrinsic::nvvm_tcgen05_st_16x128b_x4:
2903 case Intrinsic::nvvm_tcgen05_st_16x128b_x8:
2904 case Intrinsic::nvvm_tcgen05_st_16x128b_x16:
2905 case Intrinsic::nvvm_tcgen05_st_16x128b_x32:
2906 case Intrinsic::nvvm_tcgen05_st_16x128b_x64:
2907 case Intrinsic::nvvm_tcgen05_st_16x256b_x1:
2908 case Intrinsic::nvvm_tcgen05_st_16x256b_x2:
2909 case Intrinsic::nvvm_tcgen05_st_16x256b_x4:
2910 case Intrinsic::nvvm_tcgen05_st_16x256b_x8:
2911 case Intrinsic::nvvm_tcgen05_st_16x256b_x16:
2912 case Intrinsic::nvvm_tcgen05_st_16x256b_x32:
2913 case Intrinsic::nvvm_tcgen05_st_32x32b_x1:
2914 case Intrinsic::nvvm_tcgen05_st_32x32b_x2:
2915 case Intrinsic::nvvm_tcgen05_st_32x32b_x4:
2916 case Intrinsic::nvvm_tcgen05_st_32x32b_x8:
2917 case Intrinsic::nvvm_tcgen05_st_32x32b_x16:
2918 case Intrinsic::nvvm_tcgen05_st_32x32b_x32:
2919 case Intrinsic::nvvm_tcgen05_st_16x64b_x64:
2920 case Intrinsic::nvvm_tcgen05_st_32x32b_x64:
2921 case Intrinsic::nvvm_tcgen05_st_32x32b_x128:
2922 return lowerTcgen05St(Op, DAG);
2923 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x1:
2924 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x2:
2925 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x4:
2926 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x8:
2927 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x16:
2928 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x32:
2929 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x64:
2930 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x128:
2931 return lowerTcgen05St(Op, DAG, /* hasOffset */ true);
2932 case Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg1:
2933 case Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg2:
2934 case Intrinsic::nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg1:
2935 case Intrinsic::nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg2:
2936 case Intrinsic::nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg1:
2937 case Intrinsic::nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg2:
2938 case Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg1:
2939 case Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg2:
2940 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg1:
2941 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg2:
2942 case Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1:
2943 case Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2:
2944 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1:
2945 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2:
2946 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1:
2947 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2:
2948 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg1_ashift:
2949 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg2_ashift:
2950 case Intrinsic::
2951 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1_ashift:
2952 case Intrinsic::
2953 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2_ashift:
2954 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1_ashift:
2955 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2_ashift:
2956 case Intrinsic::
2957 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1_ashift:
2958 case Intrinsic::
2959 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2_ashift:
2960 case Intrinsic::
2961 nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg1_decompress_b:
2962 case Intrinsic::
2963 nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg2_decompress_b:
2964 case Intrinsic::
2965 nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg1_decompress_b:
2966 case Intrinsic::
2967 nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg2_decompress_b:
2968 return LowerTcgen05MMADisableOutputLane(Op, DAG);
2969 case Intrinsic::nvvm_tensormap_replace_elemtype:
2970 return lowerTensormapReplaceElemtype(Op, DAG);
2971 case Intrinsic::nvvm_tensormap_replace_swizzle_mode:
2972 return lowerTensormapReplaceSwizzleMode(Op, DAG);
2973 }
2974 return Op;
2975}
2976
2977static SDValue LowerClusterLaunchControlQueryCancel(SDValue Op,
2978 SelectionDAG &DAG) {
2979
2980 SDNode *N = Op.getNode();
2981 if (N->getOperand(Num: 1).getValueType() != MVT::i128) {
2982 // return, if the operand is already lowered
2983 return SDValue();
2984 }
2985
2986 unsigned IID =
2987 cast<ConstantSDNode>(Val: N->getOperand(Num: 0).getNode())->getZExtValue();
2988 auto Opcode = [&]() {
2989 switch (IID) {
2990 case Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_is_canceled:
2991 return NVPTXISD::CLUSTERLAUNCHCONTROL_QUERY_CANCEL_IS_CANCELED;
2992 case Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_x:
2993 return NVPTXISD::CLUSTERLAUNCHCONTROL_QUERY_CANCEL_GET_FIRST_CTAID_X;
2994 case Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_y:
2995 return NVPTXISD::CLUSTERLAUNCHCONTROL_QUERY_CANCEL_GET_FIRST_CTAID_Y;
2996 case Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_z:
2997 return NVPTXISD::CLUSTERLAUNCHCONTROL_QUERY_CANCEL_GET_FIRST_CTAID_Z;
2998 default:
2999 llvm_unreachable("unsupported/unhandled intrinsic");
3000 }
3001 }();
3002
3003 SDLoc DL(N);
3004 SDValue TryCancelResponse = N->getOperand(Num: 1);
3005 SDValue Cast = DAG.getNode(Opcode: ISD::BITCAST, DL, VT: MVT::v2i64, Operand: TryCancelResponse);
3006 SDValue TryCancelResponse0 =
3007 DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i64, N1: Cast,
3008 N2: DAG.getIntPtrConstant(Val: 0, DL));
3009 SDValue TryCancelResponse1 =
3010 DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i64, N1: Cast,
3011 N2: DAG.getIntPtrConstant(Val: 1, DL));
3012
3013 return DAG.getNode(Opcode, DL, VTList: N->getVTList(),
3014 Ops: {TryCancelResponse0, TryCancelResponse1});
3015}
3016
3017static SDValue lowerCvtRSIntrinsics(SDValue Op, SelectionDAG &DAG) {
3018 SDNode *N = Op.getNode();
3019 SDLoc DL(N);
3020 SDValue F32Vec = N->getOperand(Num: 1);
3021 SDValue RBits = N->getOperand(Num: 2);
3022
3023 unsigned IntrinsicID = N->getConstantOperandVal(Num: 0);
3024
3025 // Extract the 4 float elements from the vector
3026 SmallVector<SDValue, 6> Ops;
3027 for (unsigned i = 0; i < 4; ++i)
3028 Ops.push_back(Elt: DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::f32, N1: F32Vec,
3029 N2: DAG.getIntPtrConstant(Val: i, DL)));
3030
3031 using NVPTX::PTXCvtMode::CvtMode;
3032
3033 auto [OpCode, RetTy, CvtModeFlag] =
3034 [&]() -> std::tuple<unsigned, MVT::SimpleValueType, uint32_t> {
3035 switch (IntrinsicID) {
3036 case Intrinsic::nvvm_f32x4_to_e4m3x4_rs_relu_satfinite:
3037 return {NVPTXISD::CVT_E4M3X4_F32X4_RS_SF, MVT::v4i8,
3038 CvtMode::RS | CvtMode::RELU_FLAG};
3039 case Intrinsic::nvvm_f32x4_to_e4m3x4_rs_satfinite:
3040 return {NVPTXISD::CVT_E4M3X4_F32X4_RS_SF, MVT::v4i8, CvtMode::RS};
3041 case Intrinsic::nvvm_f32x4_to_e5m2x4_rs_relu_satfinite:
3042 return {NVPTXISD::CVT_E5M2X4_F32X4_RS_SF, MVT::v4i8,
3043 CvtMode::RS | CvtMode::RELU_FLAG};
3044 case Intrinsic::nvvm_f32x4_to_e5m2x4_rs_satfinite:
3045 return {NVPTXISD::CVT_E5M2X4_F32X4_RS_SF, MVT::v4i8, CvtMode::RS};
3046 case Intrinsic::nvvm_f32x4_to_e2m3x4_rs_relu_satfinite:
3047 return {NVPTXISD::CVT_E2M3X4_F32X4_RS_SF, MVT::v4i8,
3048 CvtMode::RS | CvtMode::RELU_FLAG};
3049 case Intrinsic::nvvm_f32x4_to_e2m3x4_rs_satfinite:
3050 return {NVPTXISD::CVT_E2M3X4_F32X4_RS_SF, MVT::v4i8, CvtMode::RS};
3051 case Intrinsic::nvvm_f32x4_to_e3m2x4_rs_relu_satfinite:
3052 return {NVPTXISD::CVT_E3M2X4_F32X4_RS_SF, MVT::v4i8,
3053 CvtMode::RS | CvtMode::RELU_FLAG};
3054 case Intrinsic::nvvm_f32x4_to_e3m2x4_rs_satfinite:
3055 return {NVPTXISD::CVT_E3M2X4_F32X4_RS_SF, MVT::v4i8, CvtMode::RS};
3056 case Intrinsic::nvvm_f32x4_to_e2m1x4_rs_relu_satfinite:
3057 return {NVPTXISD::CVT_E2M1X4_F32X4_RS_SF, MVT::i16,
3058 CvtMode::RS | CvtMode::RELU_FLAG};
3059 case Intrinsic::nvvm_f32x4_to_e2m1x4_rs_satfinite:
3060 return {NVPTXISD::CVT_E2M1X4_F32X4_RS_SF, MVT::i16, CvtMode::RS};
3061 default:
3062 llvm_unreachable("unsupported/unhandled intrinsic");
3063 }
3064 }();
3065
3066 Ops.push_back(Elt: RBits);
3067 Ops.push_back(Elt: DAG.getConstant(Val: CvtModeFlag, DL, VT: MVT::i32));
3068
3069 return DAG.getNode(Opcode: OpCode, DL, VT: RetTy, Ops);
3070}
3071
3072static SDValue lowerPrmtIntrinsic(SDValue Op, SelectionDAG &DAG) {
3073 const unsigned Mode = [&]() {
3074 switch (Op->getConstantOperandVal(Num: 0)) {
3075 case Intrinsic::nvvm_prmt:
3076 return NVPTX::PTXPrmtMode::NONE;
3077 case Intrinsic::nvvm_prmt_b4e:
3078 return NVPTX::PTXPrmtMode::B4E;
3079 case Intrinsic::nvvm_prmt_ecl:
3080 return NVPTX::PTXPrmtMode::ECL;
3081 case Intrinsic::nvvm_prmt_ecr:
3082 return NVPTX::PTXPrmtMode::ECR;
3083 case Intrinsic::nvvm_prmt_f4e:
3084 return NVPTX::PTXPrmtMode::F4E;
3085 case Intrinsic::nvvm_prmt_rc16:
3086 return NVPTX::PTXPrmtMode::RC16;
3087 case Intrinsic::nvvm_prmt_rc8:
3088 return NVPTX::PTXPrmtMode::RC8;
3089 default:
3090 llvm_unreachable("unsupported/unhandled intrinsic");
3091 }
3092 }();
3093 SDLoc DL(Op);
3094 SDValue A = Op->getOperand(Num: 1);
3095 SDValue B = Op.getNumOperands() == 4 ? Op.getOperand(i: 2)
3096 : DAG.getConstant(Val: 0, DL, VT: MVT::i32);
3097 SDValue Selector = (Op->op_end() - 1)->get();
3098 return getPRMT(A, B, Selector, DL, DAG, Mode);
3099}
3100
3101#define TCGEN05_LD_RED_INTR(SHAPE, NUM, TYPE) \
3102 Intrinsic::nvvm_tcgen05_ld_red_##SHAPE##_x##NUM##_##TYPE
3103
3104#define TCGEN05_LD_RED_INST(SHAPE, NUM, TYPE) \
3105 NVPTXISD::TCGEN05_LD_RED_##SHAPE##_X##NUM##_##TYPE
3106
3107static unsigned getTcgen05LdRedID(Intrinsic::ID IID) {
3108 switch (IID) {
3109 case TCGEN05_LD_RED_INTR(32x32b, 2, f32):
3110 return TCGEN05_LD_RED_INST(32x32b, 2, F32);
3111 case TCGEN05_LD_RED_INTR(32x32b, 4, f32):
3112 return TCGEN05_LD_RED_INST(32x32b, 4, F32);
3113 case TCGEN05_LD_RED_INTR(32x32b, 8, f32):
3114 return TCGEN05_LD_RED_INST(32x32b, 8, F32);
3115 case TCGEN05_LD_RED_INTR(32x32b, 16, f32):
3116 return TCGEN05_LD_RED_INST(32x32b, 16, F32);
3117 case TCGEN05_LD_RED_INTR(32x32b, 32, f32):
3118 return TCGEN05_LD_RED_INST(32x32b, 32, F32);
3119 case TCGEN05_LD_RED_INTR(32x32b, 64, f32):
3120 return TCGEN05_LD_RED_INST(32x32b, 64, F32);
3121 case TCGEN05_LD_RED_INTR(32x32b, 128, f32):
3122 return TCGEN05_LD_RED_INST(32x32b, 128, F32);
3123 case TCGEN05_LD_RED_INTR(16x32bx2, 2, f32):
3124 return TCGEN05_LD_RED_INST(16x32bx2, 2, F32);
3125 case TCGEN05_LD_RED_INTR(16x32bx2, 4, f32):
3126 return TCGEN05_LD_RED_INST(16x32bx2, 4, F32);
3127 case TCGEN05_LD_RED_INTR(16x32bx2, 8, f32):
3128 return TCGEN05_LD_RED_INST(16x32bx2, 8, F32);
3129 case TCGEN05_LD_RED_INTR(16x32bx2, 16, f32):
3130 return TCGEN05_LD_RED_INST(16x32bx2, 16, F32);
3131 case TCGEN05_LD_RED_INTR(16x32bx2, 32, f32):
3132 return TCGEN05_LD_RED_INST(16x32bx2, 32, F32);
3133 case TCGEN05_LD_RED_INTR(16x32bx2, 64, f32):
3134 return TCGEN05_LD_RED_INST(16x32bx2, 64, F32);
3135 case TCGEN05_LD_RED_INTR(16x32bx2, 128, f32):
3136 return TCGEN05_LD_RED_INST(16x32bx2, 128, F32);
3137 case TCGEN05_LD_RED_INTR(32x32b, 2, i32):
3138 return TCGEN05_LD_RED_INST(32x32b, 2, I32);
3139 case TCGEN05_LD_RED_INTR(32x32b, 4, i32):
3140 return TCGEN05_LD_RED_INST(32x32b, 4, I32);
3141 case TCGEN05_LD_RED_INTR(32x32b, 8, i32):
3142 return TCGEN05_LD_RED_INST(32x32b, 8, I32);
3143 case TCGEN05_LD_RED_INTR(32x32b, 16, i32):
3144 return TCGEN05_LD_RED_INST(32x32b, 16, I32);
3145 case TCGEN05_LD_RED_INTR(32x32b, 32, i32):
3146 return TCGEN05_LD_RED_INST(32x32b, 32, I32);
3147 case TCGEN05_LD_RED_INTR(32x32b, 64, i32):
3148 return TCGEN05_LD_RED_INST(32x32b, 64, I32);
3149 case TCGEN05_LD_RED_INTR(32x32b, 128, i32):
3150 return TCGEN05_LD_RED_INST(32x32b, 128, I32);
3151 case TCGEN05_LD_RED_INTR(16x32bx2, 2, i32):
3152 return TCGEN05_LD_RED_INST(16x32bx2, 2, I32);
3153 case TCGEN05_LD_RED_INTR(16x32bx2, 4, i32):
3154 return TCGEN05_LD_RED_INST(16x32bx2, 4, I32);
3155 case TCGEN05_LD_RED_INTR(16x32bx2, 8, i32):
3156 return TCGEN05_LD_RED_INST(16x32bx2, 8, I32);
3157 case TCGEN05_LD_RED_INTR(16x32bx2, 16, i32):
3158 return TCGEN05_LD_RED_INST(16x32bx2, 16, I32);
3159 case TCGEN05_LD_RED_INTR(16x32bx2, 32, i32):
3160 return TCGEN05_LD_RED_INST(16x32bx2, 32, I32);
3161 case TCGEN05_LD_RED_INTR(16x32bx2, 64, i32):
3162 return TCGEN05_LD_RED_INST(16x32bx2, 64, I32);
3163 case TCGEN05_LD_RED_INTR(16x32bx2, 128, i32):
3164 return TCGEN05_LD_RED_INST(16x32bx2, 128, I32);
3165 default:
3166 llvm_unreachable("Invalid tcgen05.ld.red intrinsic ID");
3167 }
3168}
3169
3170// Lower vector return type of tcgen05.ld intrinsics
3171static std::optional<std::tuple<SDValue, SDValue, SDValue>>
3172lowerTcgen05LdRed(SDNode *N, SelectionDAG &DAG) {
3173 SDLoc DL(N);
3174 EVT ResVT = N->getValueType(ResNo: 0);
3175 if (!ResVT.isVector())
3176 return {}; // already legalized.
3177
3178 const unsigned NumElts = ResVT.getVectorNumElements();
3179
3180 // Create the return type of the instructions
3181 // +1 represents the reduction value
3182 SmallVector<EVT, 132> ListVTs{
3183 NumElts + 1,
3184 ResVT.getVectorElementType().isFloatingPoint() ? MVT::f32 : MVT::i32};
3185
3186 ListVTs.push_back(Elt: MVT::Other); // Chain
3187
3188 SDVTList ResVTs = DAG.getVTList(VTs: ListVTs);
3189
3190 // Prepare the Operands
3191 SmallVector<SDValue, 8> Ops{N->getOperand(Num: 0)}; // Chain
3192
3193 // skip IID at index 1
3194 for (unsigned i = 2; i < N->getNumOperands(); i++)
3195 Ops.push_back(Elt: N->getOperand(Num: i));
3196
3197 unsigned IID = cast<ConstantSDNode>(Val: N->getOperand(Num: 1))->getZExtValue();
3198 MemIntrinsicSDNode *MemSD = cast<MemIntrinsicSDNode>(Val: N);
3199 SDValue NewNode =
3200 DAG.getMemIntrinsicNode(Opcode: getTcgen05LdRedID(IID), dl: DL, VTList: ResVTs, Ops,
3201 MemVT: MemSD->getMemoryVT(), MMO: MemSD->getMemOperand());
3202
3203 // Split vector result
3204 SmallVector<SDValue, 132> ScalarRes;
3205 for (unsigned i = 0; i < NumElts; ++i) {
3206 SDValue Res = NewNode.getValue(R: i);
3207 ScalarRes.push_back(Elt: Res);
3208 }
3209
3210 SDValue BuildVector = DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL, VT: ResVT, Ops: ScalarRes);
3211 SDValue RedResult = NewNode.getValue(R: NumElts);
3212 SDValue Chain = NewNode.getValue(R: NumElts + 1);
3213 return {{BuildVector, RedResult, Chain}};
3214}
3215
3216static SDValue lowerIntrinsicWChain(SDValue Op, SelectionDAG &DAG) {
3217 switch (Op->getConstantOperandVal(Num: 1)) {
3218 default:
3219 return Op;
3220
3221 // These tcgen05 intrinsics return a v2i32, which is legal, so we have to
3222 // lower them through LowerOperation() instead of ReplaceNodeResults().
3223 case Intrinsic::nvvm_tcgen05_ld_16x64b_x2:
3224 case Intrinsic::nvvm_tcgen05_ld_16x128b_x1:
3225 case Intrinsic::nvvm_tcgen05_ld_32x32b_x2:
3226 if (auto Res = lowerTcgen05Ld(N: Op.getNode(), DAG))
3227 return DAG.getMergeValues(Ops: {Res->first, Res->second}, dl: SDLoc(Op));
3228 return SDValue();
3229
3230 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x2:
3231 if (auto Res = lowerTcgen05Ld(N: Op.getNode(), DAG, /*HasOffset=*/true))
3232 return DAG.getMergeValues(Ops: {Res->first, Res->second}, dl: SDLoc(Op));
3233 return SDValue();
3234
3235 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x2_f32:
3236 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x2_i32:
3237 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x2_f32:
3238 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x2_i32:
3239 if (auto Res = lowerTcgen05LdRed(N: Op.getNode(), DAG))
3240 return DAG.getMergeValues(
3241 Ops: {std::get<0>(t&: *Res), std::get<1>(t&: *Res), std::get<2>(t&: *Res)}, dl: SDLoc(Op));
3242 return SDValue();
3243 }
3244}
3245
3246static unsigned getSPNumRegs(EVT VT) {
3247 return divideCeil(Numerator: VT.getSizeInBits(), Denominator: 32);
3248}
3249
3250static EVT getSPRegType(EVT VT) {
3251 assert(VT.isVector() && "Expected a vector type");
3252 EVT EltVT = VT.getVectorElementType();
3253 if (EltVT == MVT::i8)
3254 return MVT::v4i8;
3255 if (EltVT == MVT::i16)
3256 return MVT::v2i16;
3257 return EltVT;
3258}
3259
3260static void appendSplitSPOperands(ArrayRef<SDUse> Input,
3261 SmallVectorImpl<SDValue> &Output,
3262 SelectionDAG &DAG) {
3263 for (SDValue Operand : Input) {
3264 EVT VT = Operand.getValueType();
3265 SDLoc DL(Operand);
3266 if (!VT.isVector()) {
3267 Output.push_back(Elt: Operand);
3268 continue;
3269 }
3270
3271 // Pack partial .b32 operands at their natural width before extending them
3272 // to avoid redundant cvt/prmt sequences.
3273 if (VT.getSizeInBits() < 32) {
3274 EVT IntVT = EVT::getIntegerVT(Context&: *DAG.getContext(), BitWidth: VT.getSizeInBits());
3275 Output.push_back(
3276 Elt: DAG.getAnyExtOrTrunc(Op: DAG.getBitcast(VT: IntVT, V: Operand), DL, VT: MVT::i32));
3277 continue;
3278 }
3279
3280 EVT RegVT = getSPRegType(VT);
3281 if (!RegVT.isVector()) {
3282 DAG.ExtractVectorElements(Op: Operand, Args&: Output);
3283 continue;
3284 }
3285
3286 unsigned RegElts = RegVT.getVectorNumElements();
3287 for (unsigned I : llvm::seq(Size: getSPNumRegs(VT)))
3288 Output.push_back(Elt: DAG.getBitcast(
3289 VT: MVT::i32, V: DAG.getNode(Opcode: ISD::EXTRACT_SUBVECTOR, DL, VT: RegVT, N1: Operand,
3290 N2: DAG.getVectorIdxConstant(Val: I * RegElts, DL))));
3291 }
3292}
3293
3294static SDValue joinSPRegs(EVT VT, ArrayRef<SDValue> Regs, const SDLoc &DL,
3295 SelectionDAG &DAG) {
3296 assert(Regs.size() == getSPNumRegs(VT) && VT.getSizeInBits() % 32 == 0 &&
3297 "Incorrect number of SP vector registers");
3298 if (!VT.isVector())
3299 return Regs.front();
3300
3301 EVT RegVT = getSPRegType(VT);
3302 SmallVector<SDValue> Parts;
3303 for (SDValue Reg : Regs)
3304 Parts.push_back(Elt: DAG.getBitcast(VT: RegVT, V: Reg));
3305
3306 if (!RegVT.isVector())
3307 return DAG.getBuildVector(VT, DL, Ops: Parts);
3308 if (Parts.size() == 1)
3309 return Parts.front();
3310 return DAG.getNode(Opcode: ISD::CONCAT_VECTORS, DL, VT, Ops: Parts);
3311}
3312
3313static unsigned getSPCompressNodeOpcode(unsigned NumResults) {
3314#define SPCOMPRESS_NODE_CASE(NumResults) \
3315 case NumResults: \
3316 return NVPTXISD::SPCOMPRESS_R##NumResults
3317
3318 switch (NumResults) {
3319 SPCOMPRESS_NODE_CASE(2);
3320 SPCOMPRESS_NODE_CASE(3);
3321 SPCOMPRESS_NODE_CASE(5);
3322 SPCOMPRESS_NODE_CASE(6);
3323 SPCOMPRESS_NODE_CASE(9);
3324 SPCOMPRESS_NODE_CASE(10);
3325 SPCOMPRESS_NODE_CASE(12);
3326 SPCOMPRESS_NODE_CASE(18);
3327 SPCOMPRESS_NODE_CASE(20);
3328 SPCOMPRESS_NODE_CASE(24);
3329 SPCOMPRESS_NODE_CASE(36);
3330 SPCOMPRESS_NODE_CASE(40);
3331 SPCOMPRESS_NODE_CASE(48);
3332 SPCOMPRESS_NODE_CASE(72);
3333 SPCOMPRESS_NODE_CASE(80);
3334 SPCOMPRESS_NODE_CASE(96);
3335 default:
3336 llvm_unreachable("Invalid spcompress result count");
3337 }
3338
3339#undef SPCOMPRESS_NODE_CASE
3340}
3341
3342static unsigned getSPDecompressNodeOpcode(unsigned NumResults) {
3343#define SPDECOMPRESS_NODE_CASE(NumResults) \
3344 case NumResults: \
3345 return NVPTXISD::SPDECOMPRESS_R##NumResults
3346
3347 switch (NumResults) {
3348 SPDECOMPRESS_NODE_CASE(1);
3349 SPDECOMPRESS_NODE_CASE(2);
3350 SPDECOMPRESS_NODE_CASE(4);
3351 SPDECOMPRESS_NODE_CASE(8);
3352 SPDECOMPRESS_NODE_CASE(16);
3353 SPDECOMPRESS_NODE_CASE(32);
3354 SPDECOMPRESS_NODE_CASE(64);
3355 SPDECOMPRESS_NODE_CASE(128);
3356 default:
3357 llvm_unreachable("Invalid spdecompress result count");
3358 }
3359
3360#undef SPDECOMPRESS_NODE_CASE
3361}
3362
3363static SDValue lowerSPIntrinsic(SDValue Op, SelectionDAG &DAG) {
3364 SDNode *N = Op.getNode();
3365 SDLoc DL(N);
3366
3367 auto IID = static_cast<Intrinsic::ID>(N->getConstantOperandVal(Num: 0));
3368 bool IsCompress = IID == Intrinsic::nvvm_spcompress;
3369 assert((IsCompress || IID == Intrinsic::nvvm_spdecompress) &&
3370 "Unexpected SP intrinsic");
3371
3372 EVT DataVT =
3373 IsCompress ? N->getOperand(Num: 1).getValueType() : N->getValueType(ResNo: 0);
3374 EVT CDataVT =
3375 IsCompress ? N->getValueType(ResNo: 1) : N->getOperand(Num: 2).getValueType();
3376
3377 unsigned ElemSize = DataVT.getScalarSizeInBits();
3378 unsigned IdxSize = N->getConstantOperandVal(Num: 3);
3379 unsigned NumTgt = N->getConstantOperandVal(Num: 4);
3380 unsigned NumSrc =
3381 NumTgt * CDataVT.getVectorNumElements() / DataVT.getVectorNumElements();
3382 unsigned RepeatFactor = IsCompress ? getSPNumRegs(VT: DataVT) / 2
3383 : DataVT.getVectorNumElements() / NumTgt;
3384
3385 SmallVector<SDValue, 16> Ops;
3386 for (unsigned Qualifier : {ElemSize, IdxSize, RepeatFactor, NumSrc, NumTgt})
3387 Ops.push_back(Elt: DAG.getTargetConstant(Val: Qualifier, DL, VT: MVT::i32));
3388 appendSplitSPOperands(Input: N->ops().slice(N: 1, M: 2), Output&: Ops, DAG);
3389
3390 SmallVector<EVT, 16> ResultVTs;
3391 for (const EVT VT : N->values())
3392 ResultVTs.append(NumInputs: getSPNumRegs(VT), Elt: MVT::i32);
3393
3394 unsigned Opcode = IsCompress ? getSPCompressNodeOpcode(NumResults: ResultVTs.size())
3395 : getSPDecompressNodeOpcode(NumResults: ResultVTs.size());
3396 SDValue Lowered = DAG.getNode(Opcode, DL, ResultTys: ResultVTs, Ops);
3397
3398 SmallVector<SDValue, 2> Retvals;
3399 unsigned Reg = 0;
3400 for (const EVT VT : N->values()) {
3401 SmallVector<SDValue> Regs;
3402 for (unsigned R : llvm::seq(Size: getSPNumRegs(VT)))
3403 Regs.push_back(Elt: Lowered.getValue(R: Reg + R));
3404 Reg += Regs.size();
3405 Retvals.push_back(Elt: joinSPRegs(VT, Regs, DL, DAG));
3406 }
3407
3408 return DAG.getMergeValues(Ops: Retvals, dl: DL);
3409}
3410
3411static SDValue lowerIntrinsicWOChain(SDValue Op, SelectionDAG &DAG) {
3412 switch (Op->getConstantOperandVal(Num: 0)) {
3413 default:
3414 return Op;
3415 case Intrinsic::nvvm_prmt:
3416 case Intrinsic::nvvm_prmt_b4e:
3417 case Intrinsic::nvvm_prmt_ecl:
3418 case Intrinsic::nvvm_prmt_ecr:
3419 case Intrinsic::nvvm_prmt_f4e:
3420 case Intrinsic::nvvm_prmt_rc16:
3421 case Intrinsic::nvvm_prmt_rc8:
3422 return lowerPrmtIntrinsic(Op, DAG);
3423 case Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_is_canceled:
3424 case Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_x:
3425 case Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_y:
3426 case Intrinsic::nvvm_clusterlaunchcontrol_query_cancel_get_first_ctaid_z:
3427 return LowerClusterLaunchControlQueryCancel(Op, DAG);
3428 case Intrinsic::nvvm_f32x4_to_e4m3x4_rs_satfinite:
3429 case Intrinsic::nvvm_f32x4_to_e4m3x4_rs_relu_satfinite:
3430 case Intrinsic::nvvm_f32x4_to_e5m2x4_rs_satfinite:
3431 case Intrinsic::nvvm_f32x4_to_e5m2x4_rs_relu_satfinite:
3432 case Intrinsic::nvvm_f32x4_to_e2m3x4_rs_satfinite:
3433 case Intrinsic::nvvm_f32x4_to_e2m3x4_rs_relu_satfinite:
3434 case Intrinsic::nvvm_f32x4_to_e3m2x4_rs_satfinite:
3435 case Intrinsic::nvvm_f32x4_to_e3m2x4_rs_relu_satfinite:
3436 case Intrinsic::nvvm_f32x4_to_e2m1x4_rs_satfinite:
3437 case Intrinsic::nvvm_f32x4_to_e2m1x4_rs_relu_satfinite:
3438 return lowerCvtRSIntrinsics(Op, DAG);
3439
3440 case Intrinsic::nvvm_spcompress:
3441 case Intrinsic::nvvm_spdecompress:
3442 return lowerSPIntrinsic(Op, DAG);
3443 }
3444}
3445
3446// In PTX 64-bit CTLZ and CTPOP are supported, but they return a 32-bit value.
3447// Lower these into a node returning the correct type which is zero-extended
3448// back to the correct size.
3449static SDValue lowerCTLZCTPOP(SDValue Op, SelectionDAG &DAG) {
3450 SDValue V = Op->getOperand(Num: 0);
3451 assert(V.getValueType() == MVT::i64 &&
3452 "Unexpected CTLZ/CTPOP type to legalize");
3453
3454 SDLoc DL(Op);
3455 SDValue CT = DAG.getNode(Opcode: Op->getOpcode(), DL, VT: MVT::i32, Operand: V);
3456 return DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL, VT: MVT::i64, Operand: CT, Flags: SDNodeFlags::NonNeg);
3457}
3458
3459static SDValue expandFSH64(SDValue A, SDValue B, SDValue ShiftAmount, SDLoc DL,
3460 unsigned Opcode, SelectionDAG &DAG) {
3461 assert(A.getValueType() == MVT::i64 && B.getValueType() == MVT::i64);
3462
3463 const auto *AmtConst = dyn_cast<ConstantSDNode>(Val&: ShiftAmount);
3464 if (!AmtConst)
3465 return SDValue();
3466 const auto Amt = AmtConst->getZExtValue() & 63;
3467
3468 SDValue UnpackA =
3469 DAG.getNode(Opcode: NVPTXISD::UNPACK_VECTOR, DL, ResultTys: {MVT::i32, MVT::i32}, Ops: A);
3470 SDValue UnpackB =
3471 DAG.getNode(Opcode: NVPTXISD::UNPACK_VECTOR, DL, ResultTys: {MVT::i32, MVT::i32}, Ops: B);
3472
3473 // Arch is Little endiain: 0 = low bits, 1 = high bits
3474 SDValue ALo = UnpackA.getValue(R: 0);
3475 SDValue AHi = UnpackA.getValue(R: 1);
3476 SDValue BLo = UnpackB.getValue(R: 0);
3477 SDValue BHi = UnpackB.getValue(R: 1);
3478
3479 // The bitfeild consists of { AHi : ALo : BHi : BLo }
3480 //
3481 // * FSHL, Amt < 32 - The window will contain { AHi : ALo : BHi }
3482 // * FSHL, Amt >= 32 - The window will contain { ALo : BHi : BLo }
3483 // * FSHR, Amt < 32 - The window will contain { ALo : BHi : BLo }
3484 // * FSHR, Amt >= 32 - The window will contain { AHi : ALo : BHi }
3485 //
3486 // Note that Amt = 0 and Amt = 32 are special cases where 32-bit funnel shifts
3487 // are not needed at all. Amt = 0 is a no-op producing either A or B depending
3488 // on the direction. Amt = 32 can be implemented by a packing and unpacking
3489 // move to select and arrange the 32bit values. For simplicity, these cases
3490 // are not handled here explicitly and instead we rely on DAGCombiner to
3491 // remove the no-op funnel shifts we insert.
3492 auto [High, Mid, Low] = ((Opcode == ISD::FSHL) == (Amt < 32))
3493 ? std::make_tuple(args&: AHi, args&: ALo, args&: BHi)
3494 : std::make_tuple(args&: ALo, args&: BHi, args&: BLo);
3495
3496 SDValue NewAmt = DAG.getConstant(Val: Amt & 31, DL, VT: MVT::i32);
3497 SDValue RHi = DAG.getNode(Opcode, DL, VT: MVT::i32, Ops: {High, Mid, NewAmt});
3498 SDValue RLo = DAG.getNode(Opcode, DL, VT: MVT::i32, Ops: {Mid, Low, NewAmt});
3499
3500 return DAG.getNode(Opcode: NVPTXISD::BUILD_VECTOR, DL, VT: MVT::i64, Ops: {RLo, RHi});
3501}
3502
3503static SDValue lowerFSH(SDValue Op, SelectionDAG &DAG) {
3504 return expandFSH64(A: Op->getOperand(Num: 0), B: Op->getOperand(Num: 1), ShiftAmount: Op->getOperand(Num: 2),
3505 DL: SDLoc(Op), Opcode: Op->getOpcode(), DAG);
3506}
3507
3508static SDValue lowerROT(SDValue Op, SelectionDAG &DAG) {
3509 unsigned Opcode = Op->getOpcode() == ISD::ROTL ? ISD::FSHL : ISD::FSHR;
3510 return expandFSH64(A: Op->getOperand(Num: 0), B: Op->getOperand(Num: 0), ShiftAmount: Op->getOperand(Num: 1),
3511 DL: SDLoc(Op), Opcode, DAG);
3512}
3513
3514static SDValue lowerSELECT(SDValue Op, SelectionDAG &DAG) {
3515 assert(Op.getValueType() == MVT::i1 && "Custom lowering enabled only for i1");
3516
3517 SDValue Cond = Op->getOperand(Num: 0);
3518 SDValue TrueVal = Op->getOperand(Num: 1);
3519 SDValue FalseVal = Op->getOperand(Num: 2);
3520 SDLoc DL(Op);
3521
3522 // If both operands are truncated, we push the select through the truncates.
3523 if (TrueVal.getOpcode() == ISD::TRUNCATE &&
3524 FalseVal.getOpcode() == ISD::TRUNCATE) {
3525 TrueVal = TrueVal.getOperand(i: 0);
3526 FalseVal = FalseVal.getOperand(i: 0);
3527
3528 EVT VT = TrueVal.getSimpleValueType().bitsLE(VT: FalseVal.getSimpleValueType())
3529 ? TrueVal.getValueType()
3530 : FalseVal.getValueType();
3531 TrueVal = DAG.getAnyExtOrTrunc(Op: TrueVal, DL, VT);
3532 FalseVal = DAG.getAnyExtOrTrunc(Op: FalseVal, DL, VT);
3533 SDValue Select = DAG.getSelect(DL, VT, Cond, LHS: TrueVal, RHS: FalseVal);
3534 return DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: MVT::i1, Operand: Select);
3535 }
3536
3537 // Otherwise, expand the select into a series of logical operations. These
3538 // often can be folded into other operations either by us or ptxas.
3539 TrueVal = DAG.getFreeze(V: TrueVal);
3540 FalseVal = DAG.getFreeze(V: FalseVal);
3541 SDValue And1 = DAG.getNode(Opcode: ISD::AND, DL, VT: MVT::i1, N1: Cond, N2: TrueVal);
3542 SDValue NotCond = DAG.getNOT(DL, Val: Cond, VT: MVT::i1);
3543 SDValue And2 = DAG.getNode(Opcode: ISD::AND, DL, VT: MVT::i1, N1: NotCond, N2: FalseVal);
3544 SDValue Or = DAG.getNode(Opcode: ISD::OR, DL, VT: MVT::i1, N1: And1, N2: And2);
3545 return Or;
3546}
3547
3548static SDValue lowerMSTORE(SDValue Op, SelectionDAG &DAG) {
3549 SDNode *N = Op.getNode();
3550
3551 SDValue Chain = N->getOperand(Num: 0);
3552 SDValue Val = N->getOperand(Num: 1);
3553 SDValue BasePtr = N->getOperand(Num: 2);
3554 SDValue Offset = N->getOperand(Num: 3);
3555 SDValue Mask = N->getOperand(Num: 4);
3556
3557 SDLoc DL(N);
3558 EVT ValVT = Val.getValueType();
3559 MemSDNode *MemSD = cast<MemSDNode>(Val: N);
3560 assert(ValVT.isVector() && "Masked vector store must have vector type");
3561 assert(MemSD->getAlign() >= DAG.getEVTAlign(ValVT) &&
3562 "Unexpected alignment for masked store");
3563
3564 unsigned Opcode = 0;
3565 switch (ValVT.getSimpleVT().SimpleTy) {
3566 default:
3567 llvm_unreachable("Unexpected masked vector store type");
3568 case MVT::v4i64:
3569 case MVT::v4f64: {
3570 Opcode = NVPTXISD::StoreV4;
3571 break;
3572 }
3573 case MVT::v8i32:
3574 case MVT::v8f32: {
3575 Opcode = NVPTXISD::StoreV8;
3576 break;
3577 }
3578 }
3579
3580 SmallVector<SDValue, 8> Ops;
3581
3582 // Construct the new SDNode. First operand is the chain.
3583 Ops.push_back(Elt: Chain);
3584
3585 // The next N operands are the values to store. Encode the mask into the
3586 // values using the sentinel register 0 to represent a masked-off element.
3587 assert(Mask.getValueType().isVector() &&
3588 Mask.getValueType().getVectorElementType() == MVT::i1 &&
3589 "Mask must be a vector of i1");
3590 assert(Mask.getOpcode() == ISD::BUILD_VECTOR &&
3591 "Mask expected to be a BUILD_VECTOR");
3592 assert(Mask.getValueType().getVectorNumElements() ==
3593 ValVT.getVectorNumElements() &&
3594 "Mask size must be the same as the vector size");
3595 for (auto [I, Op] : enumerate(First: Mask->ops())) {
3596 // Mask elements must be constants.
3597 if (Op.getNode()->getAsZExtVal() == 0) {
3598 // Append a sentinel register 0 to the Ops vector to represent a masked
3599 // off element, this will be handled in tablegen
3600 Ops.push_back(Elt: DAG.getRegister(Reg: MCRegister::NoRegister,
3601 VT: ValVT.getVectorElementType()));
3602 } else {
3603 // Extract the element from the vector to store
3604 SDValue ExtVal =
3605 DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: ValVT.getVectorElementType(),
3606 N1: Val, N2: DAG.getIntPtrConstant(Val: I, DL));
3607 Ops.push_back(Elt: ExtVal);
3608 }
3609 }
3610
3611 // Next, the pointer operand.
3612 Ops.push_back(Elt: BasePtr);
3613
3614 // Finally, the offset operand. We expect this to always be undef, and it will
3615 // be ignored in lowering, but to mirror the handling of the other vector
3616 // store instructions we include it in the new SDNode.
3617 assert(Offset.isUndef() && "Offset operand expected to be undef or poison");
3618 Ops.push_back(Elt: Offset);
3619
3620 SDValue NewSt =
3621 DAG.getMemIntrinsicNode(Opcode, dl: DL, VTList: DAG.getVTList(VT: MVT::Other), Ops,
3622 MemVT: MemSD->getMemoryVT(), MMO: MemSD->getMemOperand());
3623
3624 return NewSt;
3625}
3626
3627SDValue
3628NVPTXTargetLowering::LowerOperation(SDValue Op, SelectionDAG &DAG) const {
3629 switch (Op.getOpcode()) {
3630 case ISD::RETURNADDR:
3631 return SDValue();
3632 case ISD::FRAMEADDR:
3633 return SDValue();
3634 case ISD::ADDRSPACECAST:
3635 return LowerADDRSPACECAST(Op, DAG);
3636 case ISD::INTRINSIC_W_CHAIN:
3637 return lowerIntrinsicWChain(Op, DAG);
3638 case ISD::INTRINSIC_WO_CHAIN:
3639 return lowerIntrinsicWOChain(Op, DAG);
3640 case ISD::INTRINSIC_VOID:
3641 return lowerIntrinsicVoid(Op, DAG);
3642 case ISD::BUILD_VECTOR:
3643 return LowerBUILD_VECTOR(Op, DAG);
3644 case ISD::BITCAST:
3645 return LowerBITCAST(Op, DAG);
3646 case ISD::EXTRACT_SUBVECTOR:
3647 return Op;
3648 case ISD::EXTRACT_VECTOR_ELT:
3649 return LowerEXTRACT_VECTOR_ELT(Op, DAG);
3650 case ISD::INSERT_VECTOR_ELT:
3651 return LowerINSERT_VECTOR_ELT(Op, DAG);
3652 case ISD::VECTOR_SHUFFLE:
3653 return LowerVECTOR_SHUFFLE(Op, DAG);
3654 case ISD::CONCAT_VECTORS:
3655 return LowerCONCAT_VECTORS(Op, DAG);
3656 case ISD::VECREDUCE_FMAX:
3657 case ISD::VECREDUCE_FMIN:
3658 case ISD::VECREDUCE_FMAXIMUM:
3659 case ISD::VECREDUCE_FMINIMUM:
3660 return LowerVECREDUCE(Op, DAG);
3661 case ISD::STORE:
3662 return LowerSTORE(Op, DAG);
3663 case ISD::MSTORE: {
3664 assert(STI.has256BitVectorLoadStore(
3665 cast<MemSDNode>(Op.getNode())->getAddressSpace()) &&
3666 "Masked store vector not supported on subtarget.");
3667 return lowerMSTORE(Op, DAG);
3668 }
3669 case ISD::LOAD:
3670 return LowerLOAD(Op, DAG);
3671 case ISD::MLOAD:
3672 return LowerMLOAD(Op, DAG);
3673 case ISD::SHL_PARTS:
3674 return LowerShiftLeftParts(Op, DAG);
3675 case ISD::SRA_PARTS:
3676 case ISD::SRL_PARTS:
3677 return LowerShiftRightParts(Op, DAG);
3678 case ISD::SELECT:
3679 return lowerSELECT(Op, DAG);
3680 case ISD::FROUND:
3681 return LowerFROUND(Op, DAG);
3682 case ISD::FCOPYSIGN:
3683 return LowerFCOPYSIGN(Op, DAG);
3684 case ISD::SINT_TO_FP:
3685 case ISD::UINT_TO_FP:
3686 return LowerINT_TO_FP(Op, DAG);
3687 case ISD::FP_TO_SINT:
3688 case ISD::FP_TO_UINT:
3689 // fptosi/fptoui to i1 truncate toward zero, so the only defined results
3690 // are {0,-1} (signed) and {0,1} (unsigned); every other input results in
3691 // poison. Thus we can simply lower to `x <= -1.0` or `x >= 1.0`.
3692 if (Op.getValueType() == MVT::i1) {
3693 SDLoc DL(Op);
3694 SDValue X = Op.getOperand(i: 0);
3695 bool IsSigned = Op.getOpcode() == ISD::FP_TO_SINT;
3696 return DAG.getSetCC(
3697 DL, VT: MVT::i1, LHS: X,
3698 RHS: DAG.getConstantFP(Val: IsSigned ? -1.0 : 1.0, DL, VT: X.getValueType()),
3699 Cond: IsSigned ? ISD::SETOLE : ISD::SETOGE);
3700 }
3701 return LowerFP_TO_INT(Op, DAG);
3702 case ISD::FP_ROUND:
3703 return LowerFP_ROUND(Op, DAG);
3704 case ISD::FPTRUNC_ROUND:
3705 return lowerFPTRUNC_ROUND(Op, DAG, STI);
3706 case ISD::FP_EXTEND:
3707 return LowerFP_EXTEND(Op, DAG);
3708 case ISD::VAARG:
3709 return LowerVAARG(Op, DAG);
3710 case ISD::VASTART:
3711 return LowerVASTART(Op, DAG);
3712 case ISD::FSHL:
3713 case ISD::FSHR:
3714 return lowerFSH(Op, DAG);
3715 case ISD::ROTL:
3716 case ISD::ROTR:
3717 return lowerROT(Op, DAG);
3718 case ISD::ABS:
3719 case ISD::ABS_MIN_POISON:
3720 case ISD::SMIN:
3721 case ISD::SMAX:
3722 case ISD::UMIN:
3723 case ISD::UMAX:
3724 case ISD::ADD:
3725 case ISD::SUB:
3726 case ISD::MUL:
3727 case ISD::SHL:
3728 case ISD::SREM:
3729 case ISD::UREM:
3730 return LowerVectorArith(Op, DAG);
3731 case ISD::DYNAMIC_STACKALLOC:
3732 return LowerDYNAMIC_STACKALLOC(Op, DAG);
3733 case ISD::STACKRESTORE:
3734 return LowerSTACKRESTORE(Op, DAG);
3735 case ISD::STACKSAVE:
3736 return LowerSTACKSAVE(Op, DAG);
3737 case ISD::CopyToReg:
3738 return LowerCopyToReg_128(Op, DAG);
3739 case ISD::FADD:
3740 case ISD::FSUB:
3741 case ISD::FMUL:
3742 // Used only for bf16 on SM80, where we select fma for non-ftz operation
3743 return PromoteBinOpIfF32FTZ(Op, DAG);
3744 case ISD::CTPOP:
3745 case ISD::CTLZ:
3746 return lowerCTLZCTPOP(Op, DAG);
3747 case ISD::BSWAP:
3748 return lowerBSWAP(Op, DAG);
3749 default:
3750 llvm_unreachable("Custom lowering not defined for operation");
3751 }
3752}
3753
3754// This will prevent AsmPrinter from trying to print the jump tables itself.
3755unsigned NVPTXTargetLowering::getJumpTableEncoding() const {
3756 return MachineJumpTableInfo::EK_Inline;
3757}
3758
3759SDValue NVPTXTargetLowering::LowerADDRSPACECAST(SDValue Op,
3760 SelectionDAG &DAG) const {
3761 AddrSpaceCastSDNode *N = cast<AddrSpaceCastSDNode>(Val: Op.getNode());
3762 unsigned SrcAS = N->getSrcAddressSpace();
3763 unsigned DestAS = N->getDestAddressSpace();
3764 if (SrcAS != llvm::ADDRESS_SPACE_GENERIC &&
3765 DestAS != llvm::ADDRESS_SPACE_GENERIC) {
3766 // Shared and SharedCluster can be converted to each other through generic
3767 // space
3768 if ((SrcAS == llvm::ADDRESS_SPACE_SHARED &&
3769 DestAS == llvm::ADDRESS_SPACE_SHARED_CLUSTER) ||
3770 (SrcAS == llvm::ADDRESS_SPACE_SHARED_CLUSTER &&
3771 DestAS == llvm::ADDRESS_SPACE_SHARED)) {
3772 SDLoc DL(Op.getNode());
3773 const MVT GenerictVT =
3774 getPointerTy(DL: DAG.getDataLayout(), AS: ADDRESS_SPACE_GENERIC);
3775 SDValue GenericConversion = DAG.getAddrSpaceCast(
3776 dl: DL, VT: GenerictVT, Ptr: Op.getOperand(i: 0), SrcAS, DestAS: ADDRESS_SPACE_GENERIC);
3777 SDValue SharedClusterConversion =
3778 DAG.getAddrSpaceCast(dl: DL, VT: Op.getValueType(), Ptr: GenericConversion,
3779 SrcAS: ADDRESS_SPACE_GENERIC, DestAS);
3780 return SharedClusterConversion;
3781 }
3782
3783 return DAG.getUNDEF(VT: Op.getValueType());
3784 }
3785
3786 return Op;
3787}
3788
3789// This function is almost a copy of SelectionDAG::expandVAArg().
3790// The only diff is that this one produces loads from local address space.
3791SDValue NVPTXTargetLowering::LowerVAARG(SDValue Op, SelectionDAG &DAG) const {
3792 const TargetLowering *TLI = STI.getTargetLowering();
3793 SDLoc DL(Op);
3794
3795 SDNode *Node = Op.getNode();
3796 const Value *V = cast<SrcValueSDNode>(Val: Node->getOperand(Num: 2))->getValue();
3797 EVT VT = Node->getValueType(ResNo: 0);
3798 auto *Ty = VT.getTypeForEVT(Context&: *DAG.getContext());
3799 SDValue Tmp1 = Node->getOperand(Num: 0);
3800 SDValue Tmp2 = Node->getOperand(Num: 1);
3801 const MaybeAlign MA(Node->getConstantOperandVal(Num: 3));
3802
3803 SDValue VAListLoad = DAG.getLoad(VT: TLI->getPointerTy(DL: DAG.getDataLayout()), dl: DL,
3804 Chain: Tmp1, Ptr: Tmp2, PtrInfo: MachinePointerInfo(V));
3805 SDValue VAList = VAListLoad;
3806
3807 if (MA && *MA > TLI->getMinStackArgumentAlignment()) {
3808 VAList = DAG.getNode(
3809 Opcode: ISD::ADD, DL, VT: VAList.getValueType(), N1: VAList,
3810 N2: DAG.getConstant(Val: MA->value() - 1, DL, VT: VAList.getValueType()));
3811
3812 VAList = DAG.getNode(Opcode: ISD::AND, DL, VT: VAList.getValueType(), N1: VAList,
3813 N2: DAG.getSignedConstant(Val: -(int64_t)MA->value(), DL,
3814 VT: VAList.getValueType()));
3815 }
3816
3817 // Increment the pointer, VAList, to the next vaarg
3818 Tmp1 = DAG.getNode(Opcode: ISD::ADD, DL, VT: VAList.getValueType(), N1: VAList,
3819 N2: DAG.getConstant(Val: DAG.getDataLayout().getTypeAllocSize(Ty),
3820 DL, VT: VAList.getValueType()));
3821
3822 // Store the incremented VAList to the legalized pointer
3823 Tmp1 = DAG.getStore(Chain: VAListLoad.getValue(R: 1), dl: DL, Val: Tmp1, Ptr: Tmp2,
3824 PtrInfo: MachinePointerInfo(V));
3825
3826 const Value *SrcV = Constant::getNullValue(
3827 Ty: PointerType::get(C&: *DAG.getContext(), AddressSpace: ADDRESS_SPACE_LOCAL));
3828
3829 // Load the actual argument out of the pointer VAList
3830 return DAG.getLoad(VT, dl: DL, Chain: Tmp1, Ptr: VAList, PtrInfo: MachinePointerInfo(SrcV));
3831}
3832
3833SDValue NVPTXTargetLowering::LowerVASTART(SDValue Op, SelectionDAG &DAG) const {
3834 const TargetLowering *TLI = STI.getTargetLowering();
3835 SDLoc DL(Op);
3836 EVT PtrVT = TLI->getPointerTy(DL: DAG.getDataLayout());
3837
3838 // Store the address of unsized array <function>_vararg[] in the ap object.
3839 SDValue VAReg = getParamSymbolNode(DAG, /* vararg */ I: -1, T: PtrVT);
3840
3841 const Value *SV = cast<SrcValueSDNode>(Val: Op.getOperand(i: 2))->getValue();
3842 return DAG.getStore(Chain: Op.getOperand(i: 0), dl: DL, Val: VAReg, Ptr: Op.getOperand(i: 1),
3843 PtrInfo: MachinePointerInfo(SV));
3844}
3845
3846static std::pair<MemSDNode *, uint32_t>
3847convertMLOADToLoadWithUsedBytesMask(MemSDNode *N, SelectionDAG &DAG,
3848 const NVPTXSubtarget &STI) {
3849 SDValue Chain = N->getOperand(Num: 0);
3850 SDValue BasePtr = N->getOperand(Num: 1);
3851 SDValue Mask = N->getOperand(Num: 3);
3852 [[maybe_unused]] SDValue Passthru = N->getOperand(Num: 4);
3853
3854 SDLoc DL(N);
3855 EVT ResVT = N->getValueType(ResNo: 0);
3856 assert(ResVT.isVector() && "Masked vector load must have vector type");
3857 // While we only expect poison passthru vectors as an input to the backend,
3858 // when the legalization framework splits a poison vector in half, it creates
3859 // two undef vectors, so we can technically expect those too.
3860 assert((Passthru.getOpcode() == ISD::POISON ||
3861 Passthru.getOpcode() == ISD::UNDEF) &&
3862 "Passthru operand expected to be poison or undef");
3863
3864 // Extract the mask and convert it to a uint32_t representing the used bytes
3865 // of the entire vector load
3866 uint32_t UsedBytesMask = 0;
3867 uint32_t ElementSizeInBits = ResVT.getVectorElementType().getSizeInBits();
3868 assert(ElementSizeInBits % 8 == 0 && "Unexpected element size");
3869 uint32_t ElementSizeInBytes = ElementSizeInBits / 8;
3870 uint32_t ElementMask = (1u << ElementSizeInBytes) - 1u;
3871
3872 for (SDValue Op : reverse(C: Mask->ops())) {
3873 // We technically only want to do this shift for every
3874 // iteration *but* the first, but in the first iteration UsedBytesMask is 0,
3875 // so this shift is a no-op.
3876 UsedBytesMask <<= ElementSizeInBytes;
3877
3878 // Mask elements must be constants.
3879 if (Op->getAsZExtVal() != 0)
3880 UsedBytesMask |= ElementMask;
3881 }
3882
3883 assert(UsedBytesMask != 0 && UsedBytesMask != UINT32_MAX &&
3884 "Unexpected masked load with elements masked all on or all off");
3885
3886 // Create a new load sd node to be handled normally by ReplaceLoadVector.
3887 MemSDNode *NewLD = cast<MemSDNode>(
3888 Val: DAG.getLoad(VT: ResVT, dl: DL, Chain, Ptr: BasePtr, MMO: N->getMemOperand()).getNode());
3889
3890 // If our subtarget does not support the used bytes mask pragma, "drop" the
3891 // mask by setting it to UINT32_MAX
3892 if (!STI.hasUsedBytesMaskPragma())
3893 UsedBytesMask = UINT32_MAX;
3894
3895 return {NewLD, UsedBytesMask};
3896}
3897
3898/// replaceLoadVector - Convert vector loads into multi-output scalar loads.
3899static std::optional<std::pair<SDValue, SDValue>>
3900replaceLoadVector(SDNode *N, SelectionDAG &DAG, const NVPTXSubtarget &STI) {
3901 MemSDNode *LD = cast<MemSDNode>(Val: N);
3902 const EVT ResVT = LD->getValueType(ResNo: 0);
3903 const EVT MemVT = LD->getMemoryVT();
3904
3905 // If we're doing sign/zero extension as part of the load, avoid lowering to
3906 // a LoadV node. TODO: consider relaxing this restriction.
3907 if (ResVT != MemVT)
3908 return std::nullopt;
3909
3910 const auto NumEltsAndEltVT =
3911 getVectorLoweringShape(VectorEVT: ResVT, STI, AddressSpace: LD->getAddressSpace());
3912 if (!NumEltsAndEltVT)
3913 return std::nullopt;
3914 const auto [NumElts, EltVT] = NumEltsAndEltVT.value();
3915
3916 Align Alignment = LD->getAlign();
3917 const auto &TD = DAG.getDataLayout();
3918 Align PrefAlign = TD.getPrefTypeAlign(Ty: MemVT.getTypeForEVT(Context&: *DAG.getContext()));
3919 if (Alignment < PrefAlign) {
3920 // This load is not sufficiently aligned, so bail out and let this vector
3921 // load be scalarized. Note that we may still be able to emit smaller
3922 // vector loads. For example, if we are loading a <4 x float> with an
3923 // alignment of 8, this check will fail but the legalizer will try again
3924 // with 2 x <2 x float>, which will succeed with an alignment of 8.
3925 return std::nullopt;
3926 }
3927
3928 // If we have a masked load, convert it to a normal load now
3929 std::optional<uint32_t> UsedBytesMask = std::nullopt;
3930 if (LD->getOpcode() == ISD::MLOAD)
3931 std::tie(args&: LD, args&: UsedBytesMask) =
3932 convertMLOADToLoadWithUsedBytesMask(N: LD, DAG, STI);
3933
3934 // Since LoadV2 is a target node, we cannot rely on DAG type legalization.
3935 // Therefore, we must ensure the type is legal. For i1 and i8, we set the
3936 // loaded type to i16 and propagate the "real" type as the memory type.
3937 const MVT LoadEltVT = (EltVT.getSizeInBits() < 16) ? MVT::i16 : EltVT;
3938
3939 unsigned Opcode;
3940 switch (NumElts) {
3941 default:
3942 return std::nullopt;
3943 case 2:
3944 Opcode = NVPTXISD::LoadV2;
3945 break;
3946 case 4:
3947 Opcode = NVPTXISD::LoadV4;
3948 break;
3949 case 8:
3950 Opcode = NVPTXISD::LoadV8;
3951 break;
3952 }
3953 auto ListVTs = SmallVector<EVT, 9>(NumElts, LoadEltVT);
3954 ListVTs.push_back(Elt: MVT::Other);
3955 SDVTList LdResVTs = DAG.getVTList(VTs: ListVTs);
3956
3957 SDLoc DL(LD);
3958
3959 // Copy regular operands
3960 SmallVector<SDValue, 8> OtherOps(LD->ops());
3961
3962 OtherOps.push_back(
3963 Elt: DAG.getConstant(Val: UsedBytesMask.value_or(UINT32_MAX), DL, VT: MVT::i32));
3964
3965 // The select routine does not have access to the LoadSDNode instance, so
3966 // pass along the extension information
3967 OtherOps.push_back(
3968 Elt: DAG.getIntPtrConstant(Val: cast<LoadSDNode>(Val: LD)->getExtensionType(), DL));
3969
3970 SDValue NewLD = DAG.getMemIntrinsicNode(Opcode, dl: DL, VTList: LdResVTs, Ops: OtherOps, MemVT,
3971 MMO: LD->getMemOperand());
3972
3973 SmallVector<SDValue> ScalarRes;
3974 if (EltVT.isVector()) {
3975 assert(EVT(EltVT.getVectorElementType()) == ResVT.getVectorElementType());
3976 assert(NumElts * EltVT.getVectorNumElements() ==
3977 ResVT.getVectorNumElements());
3978 // Generate EXTRACT_VECTOR_ELTs to split v2[i,f,bf]16/v4i8 subvectors back
3979 // into individual elements.
3980 for (const unsigned I : llvm::seq(Size: NumElts)) {
3981 SDValue SubVector = NewLD.getValue(R: I);
3982 DAG.ExtractVectorElements(Op: SubVector, Args&: ScalarRes);
3983 }
3984 } else {
3985 for (const unsigned I : llvm::seq(Size: NumElts)) {
3986 SDValue Res = NewLD.getValue(R: I);
3987 if (LoadEltVT != EltVT)
3988 Res = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: EltVT, Operand: Res);
3989 ScalarRes.push_back(Elt: Res);
3990 }
3991 }
3992
3993 SDValue LoadChain = NewLD.getValue(R: NumElts);
3994
3995 const MVT BuildVecVT =
3996 MVT::getVectorVT(VT: EltVT.getScalarType(), NumElements: ScalarRes.size());
3997 SDValue BuildVec = DAG.getBuildVector(VT: BuildVecVT, DL, Ops: ScalarRes);
3998 SDValue LoadValue = DAG.getBitcast(VT: ResVT, V: BuildVec);
3999
4000 return {{LoadValue, LoadChain}};
4001}
4002
4003static void replaceLoadVector(SDNode *N, SelectionDAG &DAG,
4004 SmallVectorImpl<SDValue> &Results,
4005 const NVPTXSubtarget &STI) {
4006 if (auto Res = replaceLoadVector(N, DAG, STI))
4007 Results.append(IL: {Res->first, Res->second});
4008}
4009
4010static SDValue lowerLoadVector(SDNode *N, SelectionDAG &DAG,
4011 const NVPTXSubtarget &STI) {
4012 if (auto Res = replaceLoadVector(N, DAG, STI))
4013 return DAG.getMergeValues(Ops: {Res->first, Res->second}, dl: SDLoc(N));
4014 return SDValue();
4015}
4016
4017// v = ld i1* addr
4018// =>
4019// v1 = ld i8* addr (-> i16)
4020// v = trunc i16 to i1
4021static SDValue lowerLOADi1(LoadSDNode *LD, SelectionDAG &DAG) {
4022 SDLoc dl(LD);
4023 assert(LD->getExtensionType() == ISD::NON_EXTLOAD);
4024 assert(LD->getValueType(0) == MVT::i1 && "Custom lowering for i1 load only");
4025 SDValue newLD = DAG.getExtLoad(ExtType: ISD::ZEXTLOAD, dl, VT: MVT::i16, Chain: LD->getChain(),
4026 Ptr: LD->getBasePtr(), PtrInfo: LD->getPointerInfo(),
4027 MemVT: MVT::i8, Alignment: LD->getAlign(),
4028 MMOFlags: LD->getMemOperand()->getFlags());
4029 SDValue result = DAG.getNode(Opcode: ISD::TRUNCATE, DL: dl, VT: MVT::i1, Operand: newLD);
4030 // The legalizer (the caller) is expecting two values from the legalized
4031 // load, so we build a MergeValues node for it. See ExpandUnalignedLoad()
4032 // in LegalizeDAG.cpp which also uses MergeValues.
4033 return DAG.getMergeValues(Ops: {result, LD->getChain()}, dl);
4034}
4035
4036SDValue NVPTXTargetLowering::LowerLOAD(SDValue Op, SelectionDAG &DAG) const {
4037 LoadSDNode *LD = cast<LoadSDNode>(Val&: Op);
4038
4039 if (Op.getValueType() == MVT::i1)
4040 return lowerLOADi1(LD, DAG);
4041
4042 // To improve CodeGen we'll legalize any-extend loads to zext loads. This is
4043 // how they'll be lowered in ISel anyway, and by doing this a little earlier
4044 // we allow for more DAG combine opportunities.
4045 if (LD->getExtensionType() == ISD::EXTLOAD) {
4046 assert(LD->getValueType(0).isInteger() && LD->getMemoryVT().isInteger() &&
4047 "Unexpected fpext-load");
4048 return DAG.getExtLoad(ExtType: ISD::ZEXTLOAD, dl: SDLoc(Op), VT: Op.getValueType(),
4049 Chain: LD->getChain(), Ptr: LD->getBasePtr(), MemVT: LD->getMemoryVT(),
4050 MMO: LD->getMemOperand());
4051 }
4052
4053 llvm_unreachable("Unexpected custom lowering for load");
4054}
4055
4056SDValue NVPTXTargetLowering::LowerMLOAD(SDValue Op, SelectionDAG &DAG) const {
4057 // v2f16/v2bf16/v2i16/v4i8 are legal, so we can't rely on legalizer to handle
4058 // masked loads of these types and have to handle them here.
4059 // v2f32 also needs to be handled here if the subtarget has f32x2
4060 // instructions, making it legal.
4061 //
4062 // Note: misaligned masked loads should never reach this point
4063 // because the override of isLegalMaskedLoad in NVPTXTargetTransformInfo.cpp
4064 // will validate alignment. Therefore, we do not need to special case handle
4065 // them here.
4066 EVT VT = Op.getValueType();
4067 if (NVPTX::isPackedVectorTy(VT)) {
4068 auto Result = convertMLOADToLoadWithUsedBytesMask(
4069 N: cast<MemSDNode>(Val: Op.getNode()), DAG, STI);
4070 MemSDNode *LD = std::get<0>(in&: Result);
4071 uint32_t UsedBytesMask = std::get<1>(in&: Result);
4072
4073 SDLoc DL(LD);
4074
4075 // Copy regular operands
4076 SmallVector<SDValue, 8> OtherOps(LD->ops());
4077
4078 OtherOps.push_back(Elt: DAG.getConstant(Val: UsedBytesMask, DL, VT: MVT::i32));
4079
4080 // We currently are not lowering extending loads, but pass the extension
4081 // type anyway as later handling expects it.
4082 OtherOps.push_back(
4083 Elt: DAG.getIntPtrConstant(Val: cast<LoadSDNode>(Val: LD)->getExtensionType(), DL));
4084 SDValue NewLD =
4085 DAG.getMemIntrinsicNode(Opcode: NVPTXISD::MLoad, dl: DL, VTList: LD->getVTList(), Ops: OtherOps,
4086 MemVT: LD->getMemoryVT(), MMO: LD->getMemOperand());
4087 return NewLD;
4088 }
4089 return SDValue();
4090}
4091
4092static SDValue lowerSTOREVector(SDValue Op, SelectionDAG &DAG,
4093 const NVPTXSubtarget &STI) {
4094 MemSDNode *N = cast<MemSDNode>(Val: Op.getNode());
4095 SDValue Val = N->getOperand(Num: 1);
4096 SDLoc DL(N);
4097 const EVT ValVT = Val.getValueType();
4098 const EVT MemVT = N->getMemoryVT();
4099
4100 // If we're truncating as part of the store, avoid lowering to a StoreV node.
4101 // TODO: consider relaxing this restriction.
4102 if (ValVT != MemVT)
4103 return SDValue();
4104
4105 const auto NumEltsAndEltVT =
4106 getVectorLoweringShape(VectorEVT: ValVT, STI, AddressSpace: N->getAddressSpace());
4107 if (!NumEltsAndEltVT)
4108 return SDValue();
4109 const auto [NumElts, EltVT] = NumEltsAndEltVT.value();
4110
4111 const DataLayout &TD = DAG.getDataLayout();
4112
4113 Align Alignment = N->getAlign();
4114 Align PrefAlign = TD.getPrefTypeAlign(Ty: ValVT.getTypeForEVT(Context&: *DAG.getContext()));
4115 if (Alignment < PrefAlign) {
4116 // This store is not sufficiently aligned, so bail out and let this vector
4117 // store be scalarized. Note that we may still be able to emit smaller
4118 // vector stores. For example, if we are storing a <4 x float> with an
4119 // alignment of 8, this check will fail but the legalizer will try again
4120 // with 2 x <2 x float>, which will succeed with an alignment of 8.
4121 return SDValue();
4122 }
4123
4124 unsigned Opcode;
4125 switch (NumElts) {
4126 default:
4127 return SDValue();
4128 case 2:
4129 Opcode = NVPTXISD::StoreV2;
4130 break;
4131 case 4:
4132 Opcode = NVPTXISD::StoreV4;
4133 break;
4134 case 8:
4135 Opcode = NVPTXISD::StoreV8;
4136 break;
4137 }
4138
4139 SmallVector<SDValue, 8> Ops;
4140
4141 // First is the chain
4142 Ops.push_back(Elt: N->getOperand(Num: 0));
4143
4144 // Then the split values
4145 if (EltVT.isVector()) {
4146 assert(EVT(EltVT.getVectorElementType()) == ValVT.getVectorElementType());
4147 assert(NumElts * EltVT.getVectorNumElements() ==
4148 ValVT.getVectorNumElements());
4149 // Combine individual elements into v2[i,f,bf]16/v4i8 subvectors to be
4150 // stored as b32s
4151 const unsigned NumEltsPerSubVector = EltVT.getVectorNumElements();
4152 for (const unsigned I : llvm::seq(Size: NumElts)) {
4153 SmallVector<SDValue, 4> SubVectorElts;
4154 DAG.ExtractVectorElements(Op: Val, Args&: SubVectorElts, Start: I * NumEltsPerSubVector,
4155 Count: NumEltsPerSubVector);
4156 Ops.push_back(Elt: DAG.getBuildVector(VT: EltVT, DL, Ops: SubVectorElts));
4157 }
4158 } else {
4159 SDValue V = DAG.getBitcast(VT: MVT::getVectorVT(VT: EltVT, NumElements: NumElts), V: Val);
4160 for (const unsigned I : llvm::seq(Size: NumElts)) {
4161 SDValue ExtVal = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT, N1: V,
4162 N2: DAG.getIntPtrConstant(Val: I, DL));
4163
4164 // Since StoreV2 is a target node, we cannot rely on DAG type
4165 // legalization. Therefore, we must ensure the type is legal. For i1 and
4166 // i8, we set the stored type to i16 and propagate the "real" type as the
4167 // memory type.
4168 if (EltVT.getSizeInBits() < 16)
4169 ExtVal = DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT: MVT::i16, Operand: ExtVal);
4170 Ops.push_back(Elt: ExtVal);
4171 }
4172 }
4173
4174 // Then any remaining arguments
4175 Ops.append(in_start: N->op_begin() + 2, in_end: N->op_end());
4176
4177 SDValue NewSt =
4178 DAG.getMemIntrinsicNode(Opcode, dl: DL, VTList: DAG.getVTList(VT: MVT::Other), Ops,
4179 MemVT: N->getMemoryVT(), MMO: N->getMemOperand());
4180
4181 // return DCI.CombineTo(N, NewSt, true);
4182 return NewSt;
4183}
4184
4185SDValue NVPTXTargetLowering::LowerSTORE(SDValue Op, SelectionDAG &DAG) const {
4186 StoreSDNode *Store = cast<StoreSDNode>(Val&: Op);
4187 EVT VT = Store->getMemoryVT();
4188
4189 if (VT == MVT::i1)
4190 return LowerSTOREi1(Op, DAG);
4191
4192 // Lower store of any other vector type, including v2f32 as we want to break
4193 // it apart since this is not a widely-supported type.
4194 return lowerSTOREVector(Op, DAG, STI);
4195}
4196
4197// st i1 v, addr
4198// =>
4199// v1 = zxt v to i16
4200// st.u8 i16, addr
4201SDValue NVPTXTargetLowering::LowerSTOREi1(SDValue Op, SelectionDAG &DAG) const {
4202 SDNode *Node = Op.getNode();
4203 SDLoc dl(Node);
4204 StoreSDNode *ST = cast<StoreSDNode>(Val: Node);
4205 SDValue Tmp1 = ST->getChain();
4206 SDValue Tmp2 = ST->getBasePtr();
4207 SDValue Tmp3 = ST->getValue();
4208 assert(Tmp3.getValueType() == MVT::i1 && "Custom lowering for i1 store only");
4209 Tmp3 = DAG.getNode(Opcode: ISD::ZERO_EXTEND, DL: dl, VT: MVT::i16, Operand: Tmp3);
4210 SDValue Result =
4211 DAG.getTruncStore(Chain: Tmp1, dl, Val: Tmp3, Ptr: Tmp2, PtrInfo: ST->getPointerInfo(), SVT: MVT::i8,
4212 Alignment: ST->getAlign(), MMOFlags: ST->getMemOperand()->getFlags());
4213 return Result;
4214}
4215
4216SDValue NVPTXTargetLowering::LowerCopyToReg_128(SDValue Op,
4217 SelectionDAG &DAG) const {
4218 // Change the CopyToReg to take in two 64-bit operands instead of a 128-bit
4219 // operand so that it can pass the legalization.
4220
4221 assert(Op.getOperand(1).getValueType() == MVT::i128 &&
4222 "Custom lowering for 128-bit CopyToReg only");
4223
4224 SDNode *Node = Op.getNode();
4225 SDLoc DL(Node);
4226
4227 SDValue Cast = DAG.getBitcast(VT: MVT::v2i64, V: Op->getOperand(Num: 2));
4228 SDValue Lo = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i64, N1: Cast,
4229 N2: DAG.getIntPtrConstant(Val: 0, DL));
4230 SDValue Hi = DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i64, N1: Cast,
4231 N2: DAG.getIntPtrConstant(Val: 1, DL));
4232
4233 SmallVector<SDValue, 5> NewOps(Op->getNumOperands() + 1);
4234 SmallVector<EVT, 3> ResultsType(Node->values());
4235
4236 NewOps[0] = Op->getOperand(Num: 0); // Chain
4237 NewOps[1] = Op->getOperand(Num: 1); // Dst Reg
4238 NewOps[2] = Lo; // Lower 64-bit
4239 NewOps[3] = Hi; // Higher 64-bit
4240 if (Op.getNumOperands() == 4)
4241 NewOps[4] = Op->getOperand(Num: 3); // Glue if exists
4242
4243 return DAG.getNode(Opcode: ISD::CopyToReg, DL, ResultTys: ResultsType, Ops: NewOps);
4244}
4245
4246unsigned NVPTXTargetLowering::getNumRegisters(
4247 LLVMContext &Context, EVT VT,
4248 std::optional<MVT> RegisterVT = std::nullopt) const {
4249 if (VT == MVT::i128 && RegisterVT == MVT::i128)
4250 return 1;
4251 return TargetLoweringBase::getNumRegisters(Context, VT, RegisterVT);
4252}
4253
4254bool NVPTXTargetLowering::splitValueIntoRegisterParts(
4255 SelectionDAG &DAG, const SDLoc &DL, SDValue Val, SDValue *Parts,
4256 unsigned NumParts, MVT PartVT, std::optional<CallingConv::ID> CC) const {
4257 if (Val.getValueType() == MVT::i128 && NumParts == 1) {
4258 Parts[0] = Val;
4259 return true;
4260 }
4261 return false;
4262}
4263
4264SDValue NVPTXTargetLowering::getParamSymbolNode(SelectionDAG &DAG, int I,
4265 EVT T) const {
4266 const MachineFunction &MF = DAG.getMachineFunction();
4267 return getSymbolNode(
4268 DAG, Sym: getParamSymbol(Ctx&: MF.getContext(), F: &MF.getFunction(), Idx: I), T);
4269}
4270
4271SDValue NVPTXTargetLowering::getCallParamSymbolNode(SelectionDAG &DAG, int I,
4272 EVT T) const {
4273 return getSymbolNode(DAG, Name: "param" + Twine(I), T);
4274}
4275
4276SDValue NVPTXTargetLowering::LowerFormalArguments(
4277 SDValue Chain, CallingConv::ID CallConv, bool isVarArg,
4278 const SmallVectorImpl<ISD::InputArg> &Ins, const SDLoc &dl,
4279 SelectionDAG &DAG, SmallVectorImpl<SDValue> &InVals) const {
4280 const DataLayout &DL = DAG.getDataLayout();
4281 LLVMContext &Ctx = *DAG.getContext();
4282
4283 const Function &F = DAG.getMachineFunction().getFunction();
4284 const bool IsKernel = isKernelFunction(F);
4285
4286 const MVT PtrVT = getPointerTy(DL, AS: IsKernel ? ADDRESS_SPACE_ENTRY_PARAM
4287 : ADDRESS_SPACE_LOCAL);
4288
4289 SDValue Root = DAG.getRoot();
4290 SmallVector<SDValue, 16> OutChains;
4291
4292 // argTypes.size() (or theArgs.size()) and Ins.size() need not match.
4293 // Ins.size() will be larger
4294 // * if there is an aggregate argument with multiple fields (each field
4295 // showing up separately in Ins)
4296 // * if there is a vector argument with more than typical vector-length
4297 // elements (generally if more than 4) where each vector element is
4298 // individually present in Ins.
4299 // So a different index should be used for indexing into Ins.
4300 // See similar issue in LowerCall.
4301
4302 auto AllIns = ArrayRef(Ins);
4303 const auto NonEmptyArgs = make_filter_range(
4304 Range: F.args(), Pred: [](const Argument &A) { return !A.getType()->isEmptyTy(); });
4305 for (const auto &[ParamI, Arg] : enumerate(First: NonEmptyArgs)) {
4306 const unsigned ArgNo = Arg.getArgNo();
4307 const auto ArgIns =
4308 AllIns.take_while(Pred: [&](auto I) { return I.OrigArgIndex == ArgNo; });
4309 AllIns = AllIns.drop_front(N: ArgIns.size());
4310
4311 Type *Ty = Arg.getType();
4312 assert(!ArgIns.empty() &&
4313 "Non-empty argument produced no parameter values");
4314
4315 if (Arg.use_empty()) {
4316 // argument is dead
4317 for (const auto &In : ArgIns) {
4318 assert(!In.Used && "Arg.use_empty() is true but Arg is used?");
4319 InVals.push_back(Elt: DAG.getUNDEF(VT: In.VT));
4320 }
4321 continue;
4322 }
4323
4324 SDValue ArgSymbol = getParamSymbolNode(DAG, I: ParamI, T: PtrVT);
4325
4326 // In the following cases, assign a node order of "i+1"
4327 // to newly created nodes. The SDNodes for params have to
4328 // appear in the same order as their order of appearance
4329 // in the original function. "i+1" holds that order.
4330 if (Arg.hasByValAttr()) {
4331 // Param has ByVal attribute
4332 // Return MoveParam(param symbol).
4333 // Ideally, the param symbol can be returned directly,
4334 // but when SDNode builder decides to use it in a CopyToReg(),
4335 // machine instruction fails because TargetExternalSymbol
4336 // (not lowered) is target dependent, and CopyToReg assumes
4337 // the source is lowered.
4338 assert(ArgIns.size() == 1 && "ByVal argument must be a pointer");
4339 const auto &ByvalIn = ArgIns[0];
4340 assert(getValueType(DL, Ty) == ByvalIn.VT &&
4341 "Ins type did not match function type");
4342
4343 SDValue P;
4344 if (IsKernel) {
4345 assert(Ty->getPointerAddressSpace() == ADDRESS_SPACE_ENTRY_PARAM &&
4346 "Kernel ByVal argument must be lowered to the param address "
4347 "space by NVPTXLowerArgs");
4348 P = ArgSymbol;
4349 P.getNode()->setIROrder(Arg.getArgNo() + 1);
4350 } else {
4351 P = DAG.getNode(Opcode: NVPTXISD::MoveParam, DL: dl, VT: ArgSymbol.getValueType(),
4352 Operand: ArgSymbol);
4353 P.getNode()->setIROrder(Arg.getArgNo() + 1);
4354 P = DAG.getAddrSpaceCast(dl, VT: ByvalIn.VT, Ptr: P, SrcAS: ADDRESS_SPACE_LOCAL,
4355 DestAS: ADDRESS_SPACE_GENERIC);
4356 }
4357 InVals.push_back(Elt: P);
4358 } else {
4359 SmallVector<EVT, 16> VTs;
4360 SmallVector<uint64_t, 16> Offsets;
4361 ComputePTXValueVTs(TLI: *this, DL, Ctx, CallConv, Ty, ValueVTs&: VTs, Offsets);
4362 assert(VTs.size() == ArgIns.size() && "Size mismatch");
4363 assert(VTs.size() == Offsets.size() && "Size mismatch");
4364
4365 const Align ArgAlign = getPTXParamAlign(
4366 F: &F, Ty, AttrIdx: Arg.getArgNo() + AttributeList::FirstArgIndex, DL);
4367
4368 unsigned I = 0;
4369 const auto VI = VectorizePTXValueVTs(ValueVTs: VTs, Offsets, ParamAlignment: ArgAlign);
4370 for (const unsigned NumElts : VI) {
4371 // i1 is loaded/stored as i8
4372 const EVT LoadVT = VTs[I] == MVT::i1 ? MVT::i8 : VTs[I];
4373 const EVT VecVT = getVectorizedVT(VT: LoadVT, N: NumElts, C&: Ctx);
4374
4375 SDValue VecAddr = DAG.getObjectPtrOffset(
4376 SL: dl, Ptr: ArgSymbol, Offset: TypeSize::getFixed(ExactSize: Offsets[I]));
4377
4378 const Align PartAlign = commonAlignment(A: ArgAlign, Offset: Offsets[I]);
4379 const unsigned AS = IsKernel ? NVPTX::AddressSpace::EntryParam
4380 : NVPTX::AddressSpace::DeviceParam;
4381 SDValue P = DAG.getLoad(VT: VecVT, dl, Chain: Root, Ptr: VecAddr,
4382 PtrInfo: MachinePointerInfo(AS), Alignment: PartAlign,
4383 MMOFlags: MachineMemOperand::MODereferenceable |
4384 MachineMemOperand::MOInvariant);
4385 P.getNode()->setIROrder(Arg.getArgNo() + 1);
4386 for (const unsigned J : llvm::seq(Size: NumElts)) {
4387 SDValue Elt = getExtractVectorizedValue(V: P, I: J, VT: LoadVT, dl, DAG);
4388
4389 Elt = correctParamType(V: Elt, ExpectedVT: ArgIns[I + J].VT, Flags: ArgIns[I + J].Flags,
4390 DAG, dl);
4391 InVals.push_back(Elt);
4392 }
4393 I += NumElts;
4394 }
4395 }
4396 }
4397
4398 if (!OutChains.empty())
4399 DAG.setRoot(DAG.getTokenFactor(DL: dl, Vals&: OutChains));
4400
4401 return Chain;
4402}
4403
4404SDValue
4405NVPTXTargetLowering::LowerReturn(SDValue Chain, CallingConv::ID CallConv,
4406 bool isVarArg,
4407 const SmallVectorImpl<ISD::OutputArg> &Outs,
4408 const SmallVectorImpl<SDValue> &OutVals,
4409 const SDLoc &dl, SelectionDAG &DAG) const {
4410 const Function &F = DAG.getMachineFunction().getFunction();
4411 Type *RetTy = F.getReturnType();
4412
4413 if (RetTy->isVoidTy()) {
4414 assert(OutVals.empty() && Outs.empty() && "Return value expected for void");
4415 return DAG.getNode(Opcode: NVPTXISD::RET_GLUE, DL: dl, VT: MVT::Other, Operand: Chain);
4416 }
4417
4418 const DataLayout &DL = DAG.getDataLayout();
4419 LLVMContext &Ctx = *DAG.getContext();
4420
4421 const SDValue RetSymbol = getSymbolNode(DAG, Name: "func_retval0", T: MVT::i32);
4422 const auto RetAlign =
4423 getPTXParamAlign(F: &F, Ty: RetTy, AttrIdx: AttributeList::ReturnIndex, DL);
4424
4425 // PTX Interoperability Guide 3.3(A): [Integer] Values shorter than
4426 // 32-bits are sign extended or zero extended, depending on whether
4427 // they are signed or unsigned types.
4428 const bool ExtendIntegerRetVal =
4429 RetTy->isIntegerTy() && DL.getTypeAllocSizeInBits(Ty: RetTy) < 32;
4430
4431 SmallVector<EVT, 16> VTs;
4432 SmallVector<uint64_t, 16> Offsets;
4433 ComputePTXValueVTs(TLI: *this, DL, Ctx, CallConv, Ty: RetTy, ValueVTs&: VTs, Offsets);
4434 assert(VTs.size() == OutVals.size() && "Bad return value decomposition");
4435
4436 const auto GetRetVal = [&](unsigned I) -> SDValue {
4437 SDValue RetVal = OutVals[I];
4438 assert(promoteScalarIntegerPTX(RetVal.getValueType()) ==
4439 RetVal.getValueType() &&
4440 "OutVal type should always be legal");
4441
4442 const EVT VTI = promoteScalarIntegerPTX(VT: VTs[I]);
4443 const EVT StoreVT =
4444 ExtendIntegerRetVal ? MVT::i32 : (VTI == MVT::i1 ? MVT::i8 : VTI);
4445 return correctParamType(V: RetVal, ExpectedVT: StoreVT, Flags: Outs[I].Flags, DAG, dl);
4446 };
4447
4448 unsigned I = 0;
4449 const auto VI = VectorizePTXValueVTs(ValueVTs: VTs, Offsets, ParamAlignment: RetAlign);
4450 for (const unsigned NumElts : VI) {
4451 const MaybeAlign CurrentAlign = ExtendIntegerRetVal
4452 ? MaybeAlign(std::nullopt)
4453 : commonAlignment(A: RetAlign, Offset: Offsets[I]);
4454
4455 SDValue Val = getBuildVectorizedValue(
4456 N: NumElts, dl, DAG, GetElement: [&](unsigned K) { return GetRetVal(I + K); });
4457
4458 SDValue Ptr =
4459 DAG.getObjectPtrOffset(SL: dl, Ptr: RetSymbol, Offset: TypeSize::getFixed(ExactSize: Offsets[I]));
4460
4461 Chain = DAG.getStore(Chain, dl, Val, Ptr,
4462 PtrInfo: MachinePointerInfo(NVPTX::AddressSpace::DeviceParam),
4463 Alignment: CurrentAlign);
4464
4465 I += NumElts;
4466 }
4467
4468 return DAG.getNode(Opcode: NVPTXISD::RET_GLUE, DL: dl, VT: MVT::Other, Operand: Chain);
4469}
4470
4471void NVPTXTargetLowering::LowerAsmOperandForConstraint(
4472 SDValue Op, StringRef Constraint, std::vector<SDValue> &Ops,
4473 SelectionDAG &DAG) const {
4474 if (Constraint.size() > 1)
4475 return;
4476 TargetLowering::LowerAsmOperandForConstraint(Op, Constraint, Ops, DAG);
4477}
4478
4479// llvm.ptx.memcpy.const and llvm.ptx.memmove.const need to be modeled as
4480// TgtMemIntrinsic
4481// because we need the information that is only available in the "Value" type
4482// of destination
4483// pointer. In particular, the address space information.
4484void NVPTXTargetLowering::getTgtMemIntrinsic(
4485 SmallVectorImpl<IntrinsicInfo> &Infos, const CallBase &I,
4486 MachineFunction &MF, unsigned Intrinsic) const {
4487 IntrinsicInfo Info;
4488 switch (Intrinsic) {
4489 default:
4490 return;
4491 case Intrinsic::nvvm_match_all_sync_i32p:
4492 case Intrinsic::nvvm_match_all_sync_i64p:
4493 Info.opc = ISD::INTRINSIC_W_CHAIN;
4494 // memVT is bogus. These intrinsics have IntrInaccessibleMemOnly attribute
4495 // in order to model data exchange with other threads, but perform no real
4496 // memory accesses.
4497 Info.memVT = MVT::i1;
4498
4499 // Our result depends on both our and other thread's arguments.
4500 Info.flags = MachineMemOperand::MOLoad | MachineMemOperand::MOStore;
4501 Infos.push_back(Elt: Info);
4502 return;
4503 case Intrinsic::nvvm_wmma_m16n16k16_load_a_f16_col:
4504 case Intrinsic::nvvm_wmma_m16n16k16_load_a_f16_row:
4505 case Intrinsic::nvvm_wmma_m16n16k16_load_a_f16_col_stride:
4506 case Intrinsic::nvvm_wmma_m16n16k16_load_a_f16_row_stride:
4507 case Intrinsic::nvvm_wmma_m16n16k16_load_b_f16_col:
4508 case Intrinsic::nvvm_wmma_m16n16k16_load_b_f16_row:
4509 case Intrinsic::nvvm_wmma_m16n16k16_load_b_f16_col_stride:
4510 case Intrinsic::nvvm_wmma_m16n16k16_load_b_f16_row_stride:
4511 case Intrinsic::nvvm_wmma_m32n8k16_load_a_f16_col:
4512 case Intrinsic::nvvm_wmma_m32n8k16_load_a_f16_row:
4513 case Intrinsic::nvvm_wmma_m32n8k16_load_a_f16_col_stride:
4514 case Intrinsic::nvvm_wmma_m32n8k16_load_a_f16_row_stride:
4515 case Intrinsic::nvvm_wmma_m32n8k16_load_b_f16_col:
4516 case Intrinsic::nvvm_wmma_m32n8k16_load_b_f16_row:
4517 case Intrinsic::nvvm_wmma_m32n8k16_load_b_f16_col_stride:
4518 case Intrinsic::nvvm_wmma_m32n8k16_load_b_f16_row_stride:
4519 case Intrinsic::nvvm_wmma_m8n32k16_load_a_f16_col:
4520 case Intrinsic::nvvm_wmma_m8n32k16_load_a_f16_row:
4521 case Intrinsic::nvvm_wmma_m8n32k16_load_a_f16_col_stride:
4522 case Intrinsic::nvvm_wmma_m8n32k16_load_a_f16_row_stride:
4523 case Intrinsic::nvvm_wmma_m8n32k16_load_b_f16_col:
4524 case Intrinsic::nvvm_wmma_m8n32k16_load_b_f16_row:
4525 case Intrinsic::nvvm_wmma_m8n32k16_load_b_f16_col_stride:
4526 case Intrinsic::nvvm_wmma_m8n32k16_load_b_f16_row_stride: {
4527 Info.opc = ISD::INTRINSIC_W_CHAIN;
4528 Info.memVT = MVT::v8f16;
4529 Info.ptrVal = I.getArgOperand(i: 0);
4530 Info.offset = 0;
4531 Info.flags = MachineMemOperand::MOLoad;
4532 Info.align = Align(16);
4533 Infos.push_back(Elt: Info);
4534 return;
4535 }
4536 case Intrinsic::nvvm_wmma_m16n16k16_load_a_s8_col:
4537 case Intrinsic::nvvm_wmma_m16n16k16_load_a_s8_col_stride:
4538 case Intrinsic::nvvm_wmma_m16n16k16_load_a_u8_col_stride:
4539 case Intrinsic::nvvm_wmma_m16n16k16_load_a_u8_col:
4540 case Intrinsic::nvvm_wmma_m16n16k16_load_a_s8_row:
4541 case Intrinsic::nvvm_wmma_m16n16k16_load_a_s8_row_stride:
4542 case Intrinsic::nvvm_wmma_m16n16k16_load_a_u8_row_stride:
4543 case Intrinsic::nvvm_wmma_m16n16k16_load_a_u8_row:
4544 case Intrinsic::nvvm_wmma_m8n32k16_load_a_bf16_col:
4545 case Intrinsic::nvvm_wmma_m8n32k16_load_a_bf16_col_stride:
4546 case Intrinsic::nvvm_wmma_m8n32k16_load_a_bf16_row:
4547 case Intrinsic::nvvm_wmma_m8n32k16_load_a_bf16_row_stride:
4548 case Intrinsic::nvvm_wmma_m16n16k16_load_b_s8_col:
4549 case Intrinsic::nvvm_wmma_m16n16k16_load_b_s8_col_stride:
4550 case Intrinsic::nvvm_wmma_m16n16k16_load_b_u8_col_stride:
4551 case Intrinsic::nvvm_wmma_m16n16k16_load_b_u8_col:
4552 case Intrinsic::nvvm_wmma_m16n16k16_load_b_s8_row:
4553 case Intrinsic::nvvm_wmma_m16n16k16_load_b_s8_row_stride:
4554 case Intrinsic::nvvm_wmma_m16n16k16_load_b_u8_row_stride:
4555 case Intrinsic::nvvm_wmma_m16n16k16_load_b_u8_row:
4556 case Intrinsic::nvvm_wmma_m32n8k16_load_b_bf16_col:
4557 case Intrinsic::nvvm_wmma_m32n8k16_load_b_bf16_col_stride:
4558 case Intrinsic::nvvm_wmma_m32n8k16_load_b_bf16_row:
4559 case Intrinsic::nvvm_wmma_m32n8k16_load_b_bf16_row_stride: {
4560 Info.opc = ISD::INTRINSIC_W_CHAIN;
4561 Info.memVT = MVT::v2i32;
4562 Info.ptrVal = I.getArgOperand(i: 0);
4563 Info.offset = 0;
4564 Info.flags = MachineMemOperand::MOLoad;
4565 Info.align = Align(8);
4566 Infos.push_back(Elt: Info);
4567 return;
4568 }
4569
4570 case Intrinsic::nvvm_wmma_m32n8k16_load_a_s8_col:
4571 case Intrinsic::nvvm_wmma_m32n8k16_load_a_s8_col_stride:
4572 case Intrinsic::nvvm_wmma_m32n8k16_load_a_u8_col_stride:
4573 case Intrinsic::nvvm_wmma_m32n8k16_load_a_u8_col:
4574 case Intrinsic::nvvm_wmma_m32n8k16_load_a_s8_row:
4575 case Intrinsic::nvvm_wmma_m32n8k16_load_a_s8_row_stride:
4576 case Intrinsic::nvvm_wmma_m32n8k16_load_a_u8_row_stride:
4577 case Intrinsic::nvvm_wmma_m32n8k16_load_a_u8_row:
4578 case Intrinsic::nvvm_wmma_m16n16k16_load_a_bf16_col:
4579 case Intrinsic::nvvm_wmma_m16n16k16_load_a_bf16_col_stride:
4580 case Intrinsic::nvvm_wmma_m16n16k16_load_a_bf16_row:
4581 case Intrinsic::nvvm_wmma_m16n16k16_load_a_bf16_row_stride:
4582 case Intrinsic::nvvm_wmma_m16n16k8_load_a_tf32_col:
4583 case Intrinsic::nvvm_wmma_m16n16k8_load_a_tf32_col_stride:
4584 case Intrinsic::nvvm_wmma_m16n16k8_load_a_tf32_row:
4585 case Intrinsic::nvvm_wmma_m16n16k8_load_a_tf32_row_stride:
4586
4587 case Intrinsic::nvvm_wmma_m8n32k16_load_b_s8_col:
4588 case Intrinsic::nvvm_wmma_m8n32k16_load_b_s8_col_stride:
4589 case Intrinsic::nvvm_wmma_m8n32k16_load_b_u8_col_stride:
4590 case Intrinsic::nvvm_wmma_m8n32k16_load_b_u8_col:
4591 case Intrinsic::nvvm_wmma_m8n32k16_load_b_s8_row:
4592 case Intrinsic::nvvm_wmma_m8n32k16_load_b_s8_row_stride:
4593 case Intrinsic::nvvm_wmma_m8n32k16_load_b_u8_row_stride:
4594 case Intrinsic::nvvm_wmma_m8n32k16_load_b_u8_row:
4595 case Intrinsic::nvvm_wmma_m16n16k16_load_b_bf16_col:
4596 case Intrinsic::nvvm_wmma_m16n16k16_load_b_bf16_col_stride:
4597 case Intrinsic::nvvm_wmma_m16n16k16_load_b_bf16_row:
4598 case Intrinsic::nvvm_wmma_m16n16k16_load_b_bf16_row_stride:
4599 case Intrinsic::nvvm_wmma_m16n16k8_load_b_tf32_col:
4600 case Intrinsic::nvvm_wmma_m16n16k8_load_b_tf32_col_stride:
4601 case Intrinsic::nvvm_wmma_m16n16k8_load_b_tf32_row:
4602 case Intrinsic::nvvm_wmma_m16n16k8_load_b_tf32_row_stride:
4603 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x4_b16:
4604 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x4_trans_b16:
4605 case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8:
4606 case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8x16_b4x16_p64:
4607 case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x2_trans_b8x16_b6x16_p32:
4608 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b4x16_p64:
4609 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x4_b8x16_b6x16_p32:
4610 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x4_s8_s4: {
4611 Info.opc = ISD::INTRINSIC_W_CHAIN;
4612 Info.memVT = MVT::v4i32;
4613 Info.ptrVal = I.getArgOperand(i: 0);
4614 Info.offset = 0;
4615 Info.flags = MachineMemOperand::MOLoad;
4616 Info.align = Align(16);
4617 Infos.push_back(Elt: Info);
4618 return;
4619 }
4620
4621 case Intrinsic::nvvm_wmma_m32n8k16_load_b_s8_col:
4622 case Intrinsic::nvvm_wmma_m32n8k16_load_b_s8_col_stride:
4623 case Intrinsic::nvvm_wmma_m32n8k16_load_b_u8_col_stride:
4624 case Intrinsic::nvvm_wmma_m32n8k16_load_b_u8_col:
4625 case Intrinsic::nvvm_wmma_m32n8k16_load_b_s8_row:
4626 case Intrinsic::nvvm_wmma_m32n8k16_load_b_s8_row_stride:
4627 case Intrinsic::nvvm_wmma_m32n8k16_load_b_u8_row_stride:
4628 case Intrinsic::nvvm_wmma_m32n8k16_load_b_u8_row:
4629
4630 case Intrinsic::nvvm_wmma_m8n32k16_load_a_s8_col:
4631 case Intrinsic::nvvm_wmma_m8n32k16_load_a_s8_col_stride:
4632 case Intrinsic::nvvm_wmma_m8n32k16_load_a_u8_col_stride:
4633 case Intrinsic::nvvm_wmma_m8n32k16_load_a_u8_col:
4634 case Intrinsic::nvvm_wmma_m8n32k16_load_a_s8_row:
4635 case Intrinsic::nvvm_wmma_m8n32k16_load_a_s8_row_stride:
4636 case Intrinsic::nvvm_wmma_m8n32k16_load_a_u8_row_stride:
4637 case Intrinsic::nvvm_wmma_m8n32k16_load_a_u8_row:
4638 case Intrinsic::nvvm_wmma_m8n8k128_load_a_b1_row:
4639 case Intrinsic::nvvm_wmma_m8n8k128_load_a_b1_row_stride:
4640 case Intrinsic::nvvm_wmma_m8n8k128_load_b_b1_col:
4641 case Intrinsic::nvvm_wmma_m8n8k128_load_b_b1_col_stride:
4642 case Intrinsic::nvvm_wmma_m8n8k32_load_a_s4_row:
4643 case Intrinsic::nvvm_wmma_m8n8k32_load_a_s4_row_stride:
4644 case Intrinsic::nvvm_wmma_m8n8k32_load_a_u4_row_stride:
4645 case Intrinsic::nvvm_wmma_m8n8k32_load_a_u4_row:
4646 case Intrinsic::nvvm_wmma_m8n8k32_load_b_s4_col:
4647 case Intrinsic::nvvm_wmma_m8n8k32_load_b_s4_col_stride:
4648 case Intrinsic::nvvm_wmma_m8n8k32_load_b_u4_col_stride:
4649 case Intrinsic::nvvm_wmma_m8n8k32_load_b_u4_col:
4650 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x1_b16:
4651 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x1_trans_b16:
4652 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b4x16_p64:
4653 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x1_b8x16_b6x16_p32:
4654 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x1_s8_s4: {
4655 Info.opc = ISD::INTRINSIC_W_CHAIN;
4656 Info.memVT = MVT::i32;
4657 Info.ptrVal = I.getArgOperand(i: 0);
4658 Info.offset = 0;
4659 Info.flags = MachineMemOperand::MOLoad;
4660 Info.align = Align(4);
4661 Infos.push_back(Elt: Info);
4662 return;
4663 }
4664
4665 case Intrinsic::nvvm_wmma_m16n16k16_load_c_f16_col:
4666 case Intrinsic::nvvm_wmma_m16n16k16_load_c_f16_row:
4667 case Intrinsic::nvvm_wmma_m16n16k16_load_c_f16_col_stride:
4668 case Intrinsic::nvvm_wmma_m16n16k16_load_c_f16_row_stride:
4669 case Intrinsic::nvvm_wmma_m32n8k16_load_c_f16_col:
4670 case Intrinsic::nvvm_wmma_m32n8k16_load_c_f16_row:
4671 case Intrinsic::nvvm_wmma_m32n8k16_load_c_f16_col_stride:
4672 case Intrinsic::nvvm_wmma_m32n8k16_load_c_f16_row_stride:
4673 case Intrinsic::nvvm_wmma_m8n32k16_load_c_f16_col:
4674 case Intrinsic::nvvm_wmma_m8n32k16_load_c_f16_row:
4675 case Intrinsic::nvvm_wmma_m8n32k16_load_c_f16_col_stride:
4676 case Intrinsic::nvvm_wmma_m8n32k16_load_c_f16_row_stride: {
4677 Info.opc = ISD::INTRINSIC_W_CHAIN;
4678 Info.memVT = MVT::v4f16;
4679 Info.ptrVal = I.getArgOperand(i: 0);
4680 Info.offset = 0;
4681 Info.flags = MachineMemOperand::MOLoad;
4682 Info.align = Align(16);
4683 Infos.push_back(Elt: Info);
4684 return;
4685 }
4686
4687 case Intrinsic::nvvm_wmma_m16n16k16_load_c_f32_col:
4688 case Intrinsic::nvvm_wmma_m16n16k16_load_c_f32_row:
4689 case Intrinsic::nvvm_wmma_m16n16k16_load_c_f32_col_stride:
4690 case Intrinsic::nvvm_wmma_m16n16k16_load_c_f32_row_stride:
4691 case Intrinsic::nvvm_wmma_m32n8k16_load_c_f32_col:
4692 case Intrinsic::nvvm_wmma_m32n8k16_load_c_f32_row:
4693 case Intrinsic::nvvm_wmma_m32n8k16_load_c_f32_col_stride:
4694 case Intrinsic::nvvm_wmma_m32n8k16_load_c_f32_row_stride:
4695 case Intrinsic::nvvm_wmma_m8n32k16_load_c_f32_col:
4696 case Intrinsic::nvvm_wmma_m8n32k16_load_c_f32_row:
4697 case Intrinsic::nvvm_wmma_m8n32k16_load_c_f32_col_stride:
4698 case Intrinsic::nvvm_wmma_m8n32k16_load_c_f32_row_stride:
4699 case Intrinsic::nvvm_wmma_m16n16k8_load_c_f32_col:
4700 case Intrinsic::nvvm_wmma_m16n16k8_load_c_f32_row:
4701 case Intrinsic::nvvm_wmma_m16n16k8_load_c_f32_col_stride:
4702 case Intrinsic::nvvm_wmma_m16n16k8_load_c_f32_row_stride: {
4703 Info.opc = ISD::INTRINSIC_W_CHAIN;
4704 Info.memVT = MVT::v8f32;
4705 Info.ptrVal = I.getArgOperand(i: 0);
4706 Info.offset = 0;
4707 Info.flags = MachineMemOperand::MOLoad;
4708 Info.align = Align(16);
4709 Infos.push_back(Elt: Info);
4710 return;
4711 }
4712
4713 case Intrinsic::nvvm_wmma_m32n8k16_load_a_bf16_col:
4714 case Intrinsic::nvvm_wmma_m32n8k16_load_a_bf16_col_stride:
4715 case Intrinsic::nvvm_wmma_m32n8k16_load_a_bf16_row:
4716 case Intrinsic::nvvm_wmma_m32n8k16_load_a_bf16_row_stride:
4717
4718 case Intrinsic::nvvm_wmma_m8n32k16_load_b_bf16_col:
4719 case Intrinsic::nvvm_wmma_m8n32k16_load_b_bf16_col_stride:
4720 case Intrinsic::nvvm_wmma_m8n32k16_load_b_bf16_row:
4721 case Intrinsic::nvvm_wmma_m8n32k16_load_b_bf16_row_stride:
4722
4723 case Intrinsic::nvvm_wmma_m16n16k16_load_c_s32_col:
4724 case Intrinsic::nvvm_wmma_m16n16k16_load_c_s32_col_stride:
4725 case Intrinsic::nvvm_wmma_m16n16k16_load_c_s32_row:
4726 case Intrinsic::nvvm_wmma_m16n16k16_load_c_s32_row_stride:
4727 case Intrinsic::nvvm_wmma_m32n8k16_load_c_s32_col:
4728 case Intrinsic::nvvm_wmma_m32n8k16_load_c_s32_col_stride:
4729 case Intrinsic::nvvm_wmma_m32n8k16_load_c_s32_row:
4730 case Intrinsic::nvvm_wmma_m32n8k16_load_c_s32_row_stride:
4731 case Intrinsic::nvvm_wmma_m8n32k16_load_c_s32_col:
4732 case Intrinsic::nvvm_wmma_m8n32k16_load_c_s32_col_stride:
4733 case Intrinsic::nvvm_wmma_m8n32k16_load_c_s32_row:
4734 case Intrinsic::nvvm_wmma_m8n32k16_load_c_s32_row_stride: {
4735 Info.opc = ISD::INTRINSIC_W_CHAIN;
4736 Info.memVT = MVT::v8i32;
4737 Info.ptrVal = I.getArgOperand(i: 0);
4738 Info.offset = 0;
4739 Info.flags = MachineMemOperand::MOLoad;
4740 Info.align = Align(16);
4741 Infos.push_back(Elt: Info);
4742 return;
4743 }
4744
4745 case Intrinsic::nvvm_wmma_m8n8k128_load_c_s32_col:
4746 case Intrinsic::nvvm_wmma_m8n8k128_load_c_s32_col_stride:
4747 case Intrinsic::nvvm_wmma_m8n8k128_load_c_s32_row:
4748 case Intrinsic::nvvm_wmma_m8n8k128_load_c_s32_row_stride:
4749 case Intrinsic::nvvm_wmma_m8n8k32_load_c_s32_col:
4750 case Intrinsic::nvvm_wmma_m8n8k32_load_c_s32_col_stride:
4751 case Intrinsic::nvvm_wmma_m8n8k32_load_c_s32_row:
4752 case Intrinsic::nvvm_wmma_m8n8k32_load_c_s32_row_stride:
4753 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x2_b16:
4754 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n8_x2_trans_b16:
4755 case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8:
4756 case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8x16_b4x16_p64:
4757 case Intrinsic::nvvm_ldmatrix_sync_aligned_m16n16_x1_trans_b8x16_b6x16_p32:
4758 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b4x16_p64:
4759 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x2_b8x16_b6x16_p32:
4760 case Intrinsic::nvvm_ldmatrix_sync_aligned_m8n16_x2_s8_s4: {
4761 Info.opc = ISD::INTRINSIC_W_CHAIN;
4762 Info.memVT = MVT::v2i32;
4763 Info.ptrVal = I.getArgOperand(i: 0);
4764 Info.offset = 0;
4765 Info.flags = MachineMemOperand::MOLoad;
4766 Info.align = Align(8);
4767 Infos.push_back(Elt: Info);
4768 return;
4769 }
4770
4771 case Intrinsic::nvvm_wmma_m8n8k4_load_a_f64_col:
4772 case Intrinsic::nvvm_wmma_m8n8k4_load_a_f64_col_stride:
4773 case Intrinsic::nvvm_wmma_m8n8k4_load_a_f64_row:
4774 case Intrinsic::nvvm_wmma_m8n8k4_load_a_f64_row_stride:
4775
4776 case Intrinsic::nvvm_wmma_m8n8k4_load_b_f64_col:
4777 case Intrinsic::nvvm_wmma_m8n8k4_load_b_f64_col_stride:
4778 case Intrinsic::nvvm_wmma_m8n8k4_load_b_f64_row:
4779 case Intrinsic::nvvm_wmma_m8n8k4_load_b_f64_row_stride: {
4780 Info.opc = ISD::INTRINSIC_W_CHAIN;
4781 Info.memVT = MVT::f64;
4782 Info.ptrVal = I.getArgOperand(i: 0);
4783 Info.offset = 0;
4784 Info.flags = MachineMemOperand::MOLoad;
4785 Info.align = Align(8);
4786 Infos.push_back(Elt: Info);
4787 return;
4788 }
4789
4790 case Intrinsic::nvvm_wmma_m8n8k4_load_c_f64_col:
4791 case Intrinsic::nvvm_wmma_m8n8k4_load_c_f64_col_stride:
4792 case Intrinsic::nvvm_wmma_m8n8k4_load_c_f64_row:
4793 case Intrinsic::nvvm_wmma_m8n8k4_load_c_f64_row_stride: {
4794 Info.opc = ISD::INTRINSIC_W_CHAIN;
4795 Info.memVT = MVT::v2f64;
4796 Info.ptrVal = I.getArgOperand(i: 0);
4797 Info.offset = 0;
4798 Info.flags = MachineMemOperand::MOLoad;
4799 Info.align = Align(16);
4800 Infos.push_back(Elt: Info);
4801 return;
4802 }
4803
4804 case Intrinsic::nvvm_wmma_m16n16k16_store_d_f16_col:
4805 case Intrinsic::nvvm_wmma_m16n16k16_store_d_f16_row:
4806 case Intrinsic::nvvm_wmma_m16n16k16_store_d_f16_col_stride:
4807 case Intrinsic::nvvm_wmma_m16n16k16_store_d_f16_row_stride:
4808 case Intrinsic::nvvm_wmma_m32n8k16_store_d_f16_col:
4809 case Intrinsic::nvvm_wmma_m32n8k16_store_d_f16_row:
4810 case Intrinsic::nvvm_wmma_m32n8k16_store_d_f16_col_stride:
4811 case Intrinsic::nvvm_wmma_m32n8k16_store_d_f16_row_stride:
4812 case Intrinsic::nvvm_wmma_m8n32k16_store_d_f16_col:
4813 case Intrinsic::nvvm_wmma_m8n32k16_store_d_f16_row:
4814 case Intrinsic::nvvm_wmma_m8n32k16_store_d_f16_col_stride:
4815 case Intrinsic::nvvm_wmma_m8n32k16_store_d_f16_row_stride: {
4816 Info.opc = ISD::INTRINSIC_VOID;
4817 Info.memVT = MVT::v4f16;
4818 Info.ptrVal = I.getArgOperand(i: 0);
4819 Info.offset = 0;
4820 Info.flags = MachineMemOperand::MOStore;
4821 Info.align = Align(16);
4822 Infos.push_back(Elt: Info);
4823 return;
4824 }
4825
4826 case Intrinsic::nvvm_wmma_m16n16k16_store_d_f32_col:
4827 case Intrinsic::nvvm_wmma_m16n16k16_store_d_f32_row:
4828 case Intrinsic::nvvm_wmma_m16n16k16_store_d_f32_col_stride:
4829 case Intrinsic::nvvm_wmma_m16n16k16_store_d_f32_row_stride:
4830 case Intrinsic::nvvm_wmma_m32n8k16_store_d_f32_col:
4831 case Intrinsic::nvvm_wmma_m32n8k16_store_d_f32_row:
4832 case Intrinsic::nvvm_wmma_m32n8k16_store_d_f32_col_stride:
4833 case Intrinsic::nvvm_wmma_m32n8k16_store_d_f32_row_stride:
4834 case Intrinsic::nvvm_wmma_m8n32k16_store_d_f32_col:
4835 case Intrinsic::nvvm_wmma_m8n32k16_store_d_f32_row:
4836 case Intrinsic::nvvm_wmma_m8n32k16_store_d_f32_col_stride:
4837 case Intrinsic::nvvm_wmma_m8n32k16_store_d_f32_row_stride:
4838 case Intrinsic::nvvm_wmma_m16n16k8_store_d_f32_col:
4839 case Intrinsic::nvvm_wmma_m16n16k8_store_d_f32_row:
4840 case Intrinsic::nvvm_wmma_m16n16k8_store_d_f32_col_stride:
4841 case Intrinsic::nvvm_wmma_m16n16k8_store_d_f32_row_stride: {
4842 Info.opc = ISD::INTRINSIC_VOID;
4843 Info.memVT = MVT::v8f32;
4844 Info.ptrVal = I.getArgOperand(i: 0);
4845 Info.offset = 0;
4846 Info.flags = MachineMemOperand::MOStore;
4847 Info.align = Align(16);
4848 Infos.push_back(Elt: Info);
4849 return;
4850 }
4851
4852 case Intrinsic::nvvm_wmma_m16n16k16_store_d_s32_col:
4853 case Intrinsic::nvvm_wmma_m16n16k16_store_d_s32_col_stride:
4854 case Intrinsic::nvvm_wmma_m16n16k16_store_d_s32_row:
4855 case Intrinsic::nvvm_wmma_m16n16k16_store_d_s32_row_stride:
4856 case Intrinsic::nvvm_wmma_m32n8k16_store_d_s32_col:
4857 case Intrinsic::nvvm_wmma_m32n8k16_store_d_s32_col_stride:
4858 case Intrinsic::nvvm_wmma_m32n8k16_store_d_s32_row:
4859 case Intrinsic::nvvm_wmma_m32n8k16_store_d_s32_row_stride:
4860 case Intrinsic::nvvm_wmma_m8n32k16_store_d_s32_col:
4861 case Intrinsic::nvvm_wmma_m8n32k16_store_d_s32_col_stride:
4862 case Intrinsic::nvvm_wmma_m8n32k16_store_d_s32_row:
4863 case Intrinsic::nvvm_wmma_m8n32k16_store_d_s32_row_stride: {
4864 Info.opc = ISD::INTRINSIC_VOID;
4865 Info.memVT = MVT::v8i32;
4866 Info.ptrVal = I.getArgOperand(i: 0);
4867 Info.offset = 0;
4868 Info.flags = MachineMemOperand::MOStore;
4869 Info.align = Align(16);
4870 Infos.push_back(Elt: Info);
4871 return;
4872 }
4873
4874 case Intrinsic::nvvm_wmma_m8n8k128_store_d_s32_col:
4875 case Intrinsic::nvvm_wmma_m8n8k128_store_d_s32_col_stride:
4876 case Intrinsic::nvvm_wmma_m8n8k128_store_d_s32_row:
4877 case Intrinsic::nvvm_wmma_m8n8k128_store_d_s32_row_stride:
4878 case Intrinsic::nvvm_wmma_m8n8k32_store_d_s32_col:
4879 case Intrinsic::nvvm_wmma_m8n8k32_store_d_s32_col_stride:
4880 case Intrinsic::nvvm_wmma_m8n8k32_store_d_s32_row:
4881 case Intrinsic::nvvm_wmma_m8n8k32_store_d_s32_row_stride:
4882 case Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x2_b16:
4883 case Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x2_trans_b16:
4884 case Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x2_trans_b8: {
4885 Info.opc = ISD::INTRINSIC_VOID;
4886 Info.memVT = MVT::v2i32;
4887 Info.ptrVal = I.getArgOperand(i: 0);
4888 Info.offset = 0;
4889 Info.flags = MachineMemOperand::MOStore;
4890 Info.align = Align(8);
4891 Infos.push_back(Elt: Info);
4892 return;
4893 }
4894
4895 case Intrinsic::nvvm_wmma_m8n8k4_store_d_f64_col:
4896 case Intrinsic::nvvm_wmma_m8n8k4_store_d_f64_col_stride:
4897 case Intrinsic::nvvm_wmma_m8n8k4_store_d_f64_row:
4898 case Intrinsic::nvvm_wmma_m8n8k4_store_d_f64_row_stride: {
4899 Info.opc = ISD::INTRINSIC_VOID;
4900 Info.memVT = MVT::v2f64;
4901 Info.ptrVal = I.getArgOperand(i: 0);
4902 Info.offset = 0;
4903 Info.flags = MachineMemOperand::MOStore;
4904 Info.align = Align(16);
4905 Infos.push_back(Elt: Info);
4906 return;
4907 }
4908
4909 case Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x1_b16:
4910 case Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x1_trans_b16:
4911 case Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x1_trans_b8: {
4912 Info.opc = ISD::INTRINSIC_VOID;
4913 Info.memVT = MVT::i32;
4914 Info.ptrVal = I.getArgOperand(i: 0);
4915 Info.offset = 0;
4916 Info.flags = MachineMemOperand::MOStore;
4917 Info.align = Align(4);
4918 Infos.push_back(Elt: Info);
4919 return;
4920 }
4921
4922 case Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x4_b16:
4923 case Intrinsic::nvvm_stmatrix_sync_aligned_m8n8_x4_trans_b16:
4924 case Intrinsic::nvvm_stmatrix_sync_aligned_m16n8_x4_trans_b8: {
4925 Info.opc = ISD::INTRINSIC_VOID;
4926 Info.memVT = MVT::v4i32;
4927 Info.ptrVal = I.getArgOperand(i: 0);
4928 Info.offset = 0;
4929 Info.flags = MachineMemOperand::MOStore;
4930 Info.align = Align(16);
4931 Infos.push_back(Elt: Info);
4932 return;
4933 }
4934
4935 case Intrinsic::nvvm_st_bulk: {
4936 Value *Dst = I.getArgOperand(i: 0);
4937 Value *Val = I.getArgOperand(i: 2);
4938 Info.opc = ISD::INTRINSIC_VOID;
4939 Info.memVT = MVT::getVT(Ty: Val->getType());
4940 Info.ptrVal = Dst;
4941 Info.offset = 0;
4942 Info.flags = MachineMemOperand::MOStore;
4943 Info.align.reset();
4944 Info.size = MemoryLocation::UnknownSize;
4945 Infos.push_back(Elt: Info);
4946 return;
4947 }
4948
4949 case Intrinsic::nvvm_prefetch_tensormap: {
4950 auto &DL = I.getDataLayout();
4951 Info.opc = ISD::INTRINSIC_VOID;
4952 Info.memVT = getPointerTy(DL);
4953 Info.ptrVal = I.getArgOperand(i: 0);
4954 Info.offset = 0;
4955 Info.flags =
4956 MachineMemOperand::MOLoad | MachineMemOperand::MODereferenceable;
4957 Info.align.reset();
4958 Infos.push_back(Elt: Info);
4959 return;
4960 }
4961
4962 case Intrinsic::nvvm_prefetch_L1_32B_valid_addr: {
4963 Info.opc = ISD::INTRINSIC_VOID;
4964 Info.memVT = MVT::i8;
4965 Info.ptrVal = I.getArgOperand(i: 0);
4966 Info.offset = 0;
4967 Info.flags = MachineMemOperand::MOLoad;
4968 Info.align = Align(1);
4969 Infos.push_back(Elt: Info);
4970 return;
4971 }
4972
4973 case Intrinsic::nvvm_mbarrier_init: {
4974 Info.opc = ISD::INTRINSIC_VOID;
4975 Info.memVT = MVT::i64;
4976 Info.ptrVal = I.getArgOperand(i: 0);
4977 Info.offset = 0;
4978 Info.flags = MachineMemOperand::MOStore;
4979 Info.align = Align(8);
4980 Infos.push_back(Elt: Info);
4981 return;
4982 }
4983
4984 case Intrinsic::nvvm_mbarrier_check_layout: {
4985 Info.opc = ISD::INTRINSIC_W_CHAIN;
4986 Info.memVT = MVT::i64;
4987 Info.ptrVal = I.getArgOperand(i: 0);
4988 Info.offset = 0;
4989 Info.flags = MachineMemOperand::MOLoad;
4990 Info.align = Align(8);
4991 Infos.push_back(Elt: Info);
4992 return;
4993 }
4994
4995 case Intrinsic::nvvm_tensormap_replace_global_address:
4996 case Intrinsic::nvvm_tensormap_replace_global_stride: {
4997 Info.opc = ISD::INTRINSIC_VOID;
4998 Info.memVT = MVT::i64;
4999 Info.ptrVal = I.getArgOperand(i: 0);
5000 Info.offset = 0;
5001 Info.flags = MachineMemOperand::MOStore;
5002 Info.align.reset();
5003 Infos.push_back(Elt: Info);
5004 return;
5005 }
5006
5007 case Intrinsic::nvvm_tensormap_replace_rank:
5008 case Intrinsic::nvvm_tensormap_replace_box_dim:
5009 case Intrinsic::nvvm_tensormap_replace_global_dim:
5010 case Intrinsic::nvvm_tensormap_replace_element_stride:
5011 case Intrinsic::nvvm_tensormap_replace_elemtype:
5012 case Intrinsic::nvvm_tensormap_replace_interleave_layout:
5013 case Intrinsic::nvvm_tensormap_replace_swizzle_mode:
5014 case Intrinsic::nvvm_tensormap_replace_swizzle_atomicity:
5015 case Intrinsic::nvvm_tensormap_replace_fill_mode: {
5016 Info.opc = ISD::INTRINSIC_VOID;
5017 Info.memVT = MVT::i32;
5018 Info.ptrVal = I.getArgOperand(i: 0);
5019 Info.offset = 0;
5020 Info.flags = MachineMemOperand::MOStore;
5021 Info.align.reset();
5022 Infos.push_back(Elt: Info);
5023 return;
5024 }
5025
5026 case Intrinsic::nvvm_ldu_global_i:
5027 case Intrinsic::nvvm_ldu_global_f:
5028 case Intrinsic::nvvm_ldu_global_p: {
5029 Info.opc = ISD::INTRINSIC_W_CHAIN;
5030 Info.memVT = getValueType(DL: I.getDataLayout(), Ty: I.getType());
5031 Info.ptrVal = I.getArgOperand(i: 0);
5032 Info.offset = 0;
5033 Info.flags = MachineMemOperand::MOLoad;
5034 Info.align = cast<ConstantInt>(Val: I.getArgOperand(i: 1))->getMaybeAlignValue();
5035
5036 Infos.push_back(Elt: Info);
5037 return;
5038 }
5039 case Intrinsic::nvvm_tex_1d_v4f32_s32:
5040 case Intrinsic::nvvm_tex_1d_v4f32_f32:
5041 case Intrinsic::nvvm_tex_1d_level_v4f32_f32:
5042 case Intrinsic::nvvm_tex_1d_grad_v4f32_f32:
5043 case Intrinsic::nvvm_tex_1d_array_v4f32_s32:
5044 case Intrinsic::nvvm_tex_1d_array_v4f32_f32:
5045 case Intrinsic::nvvm_tex_1d_array_level_v4f32_f32:
5046 case Intrinsic::nvvm_tex_1d_array_grad_v4f32_f32:
5047 case Intrinsic::nvvm_tex_2d_v4f32_s32:
5048 case Intrinsic::nvvm_tex_2d_v4f32_f32:
5049 case Intrinsic::nvvm_tex_2d_level_v4f32_f32:
5050 case Intrinsic::nvvm_tex_2d_grad_v4f32_f32:
5051 case Intrinsic::nvvm_tex_2d_array_v4f32_s32:
5052 case Intrinsic::nvvm_tex_2d_array_v4f32_f32:
5053 case Intrinsic::nvvm_tex_2d_array_level_v4f32_f32:
5054 case Intrinsic::nvvm_tex_2d_array_grad_v4f32_f32:
5055 case Intrinsic::nvvm_tex_3d_v4f32_s32:
5056 case Intrinsic::nvvm_tex_3d_v4f32_f32:
5057 case Intrinsic::nvvm_tex_3d_level_v4f32_f32:
5058 case Intrinsic::nvvm_tex_3d_grad_v4f32_f32:
5059 case Intrinsic::nvvm_tex_cube_v4f32_f32:
5060 case Intrinsic::nvvm_tex_cube_level_v4f32_f32:
5061 case Intrinsic::nvvm_tex_cube_array_v4f32_f32:
5062 case Intrinsic::nvvm_tex_cube_array_level_v4f32_f32:
5063 case Intrinsic::nvvm_tld4_r_2d_v4f32_f32:
5064 case Intrinsic::nvvm_tld4_g_2d_v4f32_f32:
5065 case Intrinsic::nvvm_tld4_b_2d_v4f32_f32:
5066 case Intrinsic::nvvm_tld4_a_2d_v4f32_f32:
5067 case Intrinsic::nvvm_tex_unified_1d_v4f32_s32:
5068 case Intrinsic::nvvm_tex_unified_1d_v4f32_f32:
5069 case Intrinsic::nvvm_tex_unified_1d_level_v4f32_f32:
5070 case Intrinsic::nvvm_tex_unified_1d_grad_v4f32_f32:
5071 case Intrinsic::nvvm_tex_unified_1d_array_v4f32_s32:
5072 case Intrinsic::nvvm_tex_unified_1d_array_v4f32_f32:
5073 case Intrinsic::nvvm_tex_unified_1d_array_level_v4f32_f32:
5074 case Intrinsic::nvvm_tex_unified_1d_array_grad_v4f32_f32:
5075 case Intrinsic::nvvm_tex_unified_2d_v4f32_s32:
5076 case Intrinsic::nvvm_tex_unified_2d_v4f32_f32:
5077 case Intrinsic::nvvm_tex_unified_2d_level_v4f32_f32:
5078 case Intrinsic::nvvm_tex_unified_2d_grad_v4f32_f32:
5079 case Intrinsic::nvvm_tex_unified_2d_array_v4f32_s32:
5080 case Intrinsic::nvvm_tex_unified_2d_array_v4f32_f32:
5081 case Intrinsic::nvvm_tex_unified_2d_array_level_v4f32_f32:
5082 case Intrinsic::nvvm_tex_unified_2d_array_grad_v4f32_f32:
5083 case Intrinsic::nvvm_tex_unified_3d_v4f32_s32:
5084 case Intrinsic::nvvm_tex_unified_3d_v4f32_f32:
5085 case Intrinsic::nvvm_tex_unified_3d_level_v4f32_f32:
5086 case Intrinsic::nvvm_tex_unified_3d_grad_v4f32_f32:
5087 case Intrinsic::nvvm_tex_unified_cube_v4f32_f32:
5088 case Intrinsic::nvvm_tex_unified_cube_level_v4f32_f32:
5089 case Intrinsic::nvvm_tex_unified_cube_array_v4f32_f32:
5090 case Intrinsic::nvvm_tex_unified_cube_array_level_v4f32_f32:
5091 case Intrinsic::nvvm_tex_unified_cube_grad_v4f32_f32:
5092 case Intrinsic::nvvm_tex_unified_cube_array_grad_v4f32_f32:
5093 case Intrinsic::nvvm_tld4_unified_r_2d_v4f32_f32:
5094 case Intrinsic::nvvm_tld4_unified_g_2d_v4f32_f32:
5095 case Intrinsic::nvvm_tld4_unified_b_2d_v4f32_f32:
5096 case Intrinsic::nvvm_tld4_unified_a_2d_v4f32_f32:
5097 Info.opc = ISD::INTRINSIC_W_CHAIN;
5098 Info.memVT = MVT::v4f32;
5099 Info.ptrVal = nullptr;
5100 Info.offset = 0;
5101 Info.flags = MachineMemOperand::MOLoad;
5102 Info.align = Align(16);
5103 Infos.push_back(Elt: Info);
5104 return;
5105
5106 case Intrinsic::nvvm_tex_1d_v4s32_s32:
5107 case Intrinsic::nvvm_tex_1d_v4s32_f32:
5108 case Intrinsic::nvvm_tex_1d_level_v4s32_f32:
5109 case Intrinsic::nvvm_tex_1d_grad_v4s32_f32:
5110 case Intrinsic::nvvm_tex_1d_array_v4s32_s32:
5111 case Intrinsic::nvvm_tex_1d_array_v4s32_f32:
5112 case Intrinsic::nvvm_tex_1d_array_level_v4s32_f32:
5113 case Intrinsic::nvvm_tex_1d_array_grad_v4s32_f32:
5114 case Intrinsic::nvvm_tex_2d_v4s32_s32:
5115 case Intrinsic::nvvm_tex_2d_v4s32_f32:
5116 case Intrinsic::nvvm_tex_2d_level_v4s32_f32:
5117 case Intrinsic::nvvm_tex_2d_grad_v4s32_f32:
5118 case Intrinsic::nvvm_tex_2d_array_v4s32_s32:
5119 case Intrinsic::nvvm_tex_2d_array_v4s32_f32:
5120 case Intrinsic::nvvm_tex_2d_array_level_v4s32_f32:
5121 case Intrinsic::nvvm_tex_2d_array_grad_v4s32_f32:
5122 case Intrinsic::nvvm_tex_3d_v4s32_s32:
5123 case Intrinsic::nvvm_tex_3d_v4s32_f32:
5124 case Intrinsic::nvvm_tex_3d_level_v4s32_f32:
5125 case Intrinsic::nvvm_tex_3d_grad_v4s32_f32:
5126 case Intrinsic::nvvm_tex_cube_v4s32_f32:
5127 case Intrinsic::nvvm_tex_cube_level_v4s32_f32:
5128 case Intrinsic::nvvm_tex_cube_array_v4s32_f32:
5129 case Intrinsic::nvvm_tex_cube_array_level_v4s32_f32:
5130 case Intrinsic::nvvm_tex_cube_v4u32_f32:
5131 case Intrinsic::nvvm_tex_cube_level_v4u32_f32:
5132 case Intrinsic::nvvm_tex_cube_array_v4u32_f32:
5133 case Intrinsic::nvvm_tex_cube_array_level_v4u32_f32:
5134 case Intrinsic::nvvm_tex_1d_v4u32_s32:
5135 case Intrinsic::nvvm_tex_1d_v4u32_f32:
5136 case Intrinsic::nvvm_tex_1d_level_v4u32_f32:
5137 case Intrinsic::nvvm_tex_1d_grad_v4u32_f32:
5138 case Intrinsic::nvvm_tex_1d_array_v4u32_s32:
5139 case Intrinsic::nvvm_tex_1d_array_v4u32_f32:
5140 case Intrinsic::nvvm_tex_1d_array_level_v4u32_f32:
5141 case Intrinsic::nvvm_tex_1d_array_grad_v4u32_f32:
5142 case Intrinsic::nvvm_tex_2d_v4u32_s32:
5143 case Intrinsic::nvvm_tex_2d_v4u32_f32:
5144 case Intrinsic::nvvm_tex_2d_level_v4u32_f32:
5145 case Intrinsic::nvvm_tex_2d_grad_v4u32_f32:
5146 case Intrinsic::nvvm_tex_2d_array_v4u32_s32:
5147 case Intrinsic::nvvm_tex_2d_array_v4u32_f32:
5148 case Intrinsic::nvvm_tex_2d_array_level_v4u32_f32:
5149 case Intrinsic::nvvm_tex_2d_array_grad_v4u32_f32:
5150 case Intrinsic::nvvm_tex_3d_v4u32_s32:
5151 case Intrinsic::nvvm_tex_3d_v4u32_f32:
5152 case Intrinsic::nvvm_tex_3d_level_v4u32_f32:
5153 case Intrinsic::nvvm_tex_3d_grad_v4u32_f32:
5154 case Intrinsic::nvvm_tld4_r_2d_v4s32_f32:
5155 case Intrinsic::nvvm_tld4_g_2d_v4s32_f32:
5156 case Intrinsic::nvvm_tld4_b_2d_v4s32_f32:
5157 case Intrinsic::nvvm_tld4_a_2d_v4s32_f32:
5158 case Intrinsic::nvvm_tld4_r_2d_v4u32_f32:
5159 case Intrinsic::nvvm_tld4_g_2d_v4u32_f32:
5160 case Intrinsic::nvvm_tld4_b_2d_v4u32_f32:
5161 case Intrinsic::nvvm_tld4_a_2d_v4u32_f32:
5162 case Intrinsic::nvvm_tex_unified_1d_v4s32_s32:
5163 case Intrinsic::nvvm_tex_unified_1d_v4s32_f32:
5164 case Intrinsic::nvvm_tex_unified_1d_level_v4s32_f32:
5165 case Intrinsic::nvvm_tex_unified_1d_grad_v4s32_f32:
5166 case Intrinsic::nvvm_tex_unified_1d_array_v4s32_s32:
5167 case Intrinsic::nvvm_tex_unified_1d_array_v4s32_f32:
5168 case Intrinsic::nvvm_tex_unified_1d_array_level_v4s32_f32:
5169 case Intrinsic::nvvm_tex_unified_1d_array_grad_v4s32_f32:
5170 case Intrinsic::nvvm_tex_unified_2d_v4s32_s32:
5171 case Intrinsic::nvvm_tex_unified_2d_v4s32_f32:
5172 case Intrinsic::nvvm_tex_unified_2d_level_v4s32_f32:
5173 case Intrinsic::nvvm_tex_unified_2d_grad_v4s32_f32:
5174 case Intrinsic::nvvm_tex_unified_2d_array_v4s32_s32:
5175 case Intrinsic::nvvm_tex_unified_2d_array_v4s32_f32:
5176 case Intrinsic::nvvm_tex_unified_2d_array_level_v4s32_f32:
5177 case Intrinsic::nvvm_tex_unified_2d_array_grad_v4s32_f32:
5178 case Intrinsic::nvvm_tex_unified_3d_v4s32_s32:
5179 case Intrinsic::nvvm_tex_unified_3d_v4s32_f32:
5180 case Intrinsic::nvvm_tex_unified_3d_level_v4s32_f32:
5181 case Intrinsic::nvvm_tex_unified_3d_grad_v4s32_f32:
5182 case Intrinsic::nvvm_tex_unified_1d_v4u32_s32:
5183 case Intrinsic::nvvm_tex_unified_1d_v4u32_f32:
5184 case Intrinsic::nvvm_tex_unified_1d_level_v4u32_f32:
5185 case Intrinsic::nvvm_tex_unified_1d_grad_v4u32_f32:
5186 case Intrinsic::nvvm_tex_unified_1d_array_v4u32_s32:
5187 case Intrinsic::nvvm_tex_unified_1d_array_v4u32_f32:
5188 case Intrinsic::nvvm_tex_unified_1d_array_level_v4u32_f32:
5189 case Intrinsic::nvvm_tex_unified_1d_array_grad_v4u32_f32:
5190 case Intrinsic::nvvm_tex_unified_2d_v4u32_s32:
5191 case Intrinsic::nvvm_tex_unified_2d_v4u32_f32:
5192 case Intrinsic::nvvm_tex_unified_2d_level_v4u32_f32:
5193 case Intrinsic::nvvm_tex_unified_2d_grad_v4u32_f32:
5194 case Intrinsic::nvvm_tex_unified_2d_array_v4u32_s32:
5195 case Intrinsic::nvvm_tex_unified_2d_array_v4u32_f32:
5196 case Intrinsic::nvvm_tex_unified_2d_array_level_v4u32_f32:
5197 case Intrinsic::nvvm_tex_unified_2d_array_grad_v4u32_f32:
5198 case Intrinsic::nvvm_tex_unified_3d_v4u32_s32:
5199 case Intrinsic::nvvm_tex_unified_3d_v4u32_f32:
5200 case Intrinsic::nvvm_tex_unified_3d_level_v4u32_f32:
5201 case Intrinsic::nvvm_tex_unified_3d_grad_v4u32_f32:
5202 case Intrinsic::nvvm_tex_unified_cube_v4s32_f32:
5203 case Intrinsic::nvvm_tex_unified_cube_level_v4s32_f32:
5204 case Intrinsic::nvvm_tex_unified_cube_array_v4s32_f32:
5205 case Intrinsic::nvvm_tex_unified_cube_array_level_v4s32_f32:
5206 case Intrinsic::nvvm_tex_unified_cube_v4u32_f32:
5207 case Intrinsic::nvvm_tex_unified_cube_level_v4u32_f32:
5208 case Intrinsic::nvvm_tex_unified_cube_array_v4u32_f32:
5209 case Intrinsic::nvvm_tex_unified_cube_array_level_v4u32_f32:
5210 case Intrinsic::nvvm_tex_unified_cube_grad_v4s32_f32:
5211 case Intrinsic::nvvm_tex_unified_cube_grad_v4u32_f32:
5212 case Intrinsic::nvvm_tex_unified_cube_array_grad_v4s32_f32:
5213 case Intrinsic::nvvm_tex_unified_cube_array_grad_v4u32_f32:
5214 case Intrinsic::nvvm_tld4_unified_r_2d_v4s32_f32:
5215 case Intrinsic::nvvm_tld4_unified_g_2d_v4s32_f32:
5216 case Intrinsic::nvvm_tld4_unified_b_2d_v4s32_f32:
5217 case Intrinsic::nvvm_tld4_unified_a_2d_v4s32_f32:
5218 case Intrinsic::nvvm_tld4_unified_r_2d_v4u32_f32:
5219 case Intrinsic::nvvm_tld4_unified_g_2d_v4u32_f32:
5220 case Intrinsic::nvvm_tld4_unified_b_2d_v4u32_f32:
5221 case Intrinsic::nvvm_tld4_unified_a_2d_v4u32_f32:
5222 Info.opc = ISD::INTRINSIC_W_CHAIN;
5223 Info.memVT = MVT::v4i32;
5224 Info.ptrVal = nullptr;
5225 Info.offset = 0;
5226 Info.flags = MachineMemOperand::MOLoad;
5227 Info.align = Align(16);
5228 Infos.push_back(Elt: Info);
5229 return;
5230
5231 case Intrinsic::nvvm_suld_1d_i8_clamp:
5232 case Intrinsic::nvvm_suld_1d_v2i8_clamp:
5233 case Intrinsic::nvvm_suld_1d_v4i8_clamp:
5234 case Intrinsic::nvvm_suld_1d_array_i8_clamp:
5235 case Intrinsic::nvvm_suld_1d_array_v2i8_clamp:
5236 case Intrinsic::nvvm_suld_1d_array_v4i8_clamp:
5237 case Intrinsic::nvvm_suld_2d_i8_clamp:
5238 case Intrinsic::nvvm_suld_2d_v2i8_clamp:
5239 case Intrinsic::nvvm_suld_2d_v4i8_clamp:
5240 case Intrinsic::nvvm_suld_2d_array_i8_clamp:
5241 case Intrinsic::nvvm_suld_2d_array_v2i8_clamp:
5242 case Intrinsic::nvvm_suld_2d_array_v4i8_clamp:
5243 case Intrinsic::nvvm_suld_3d_i8_clamp:
5244 case Intrinsic::nvvm_suld_3d_v2i8_clamp:
5245 case Intrinsic::nvvm_suld_3d_v4i8_clamp:
5246 case Intrinsic::nvvm_suld_1d_i8_trap:
5247 case Intrinsic::nvvm_suld_1d_v2i8_trap:
5248 case Intrinsic::nvvm_suld_1d_v4i8_trap:
5249 case Intrinsic::nvvm_suld_1d_array_i8_trap:
5250 case Intrinsic::nvvm_suld_1d_array_v2i8_trap:
5251 case Intrinsic::nvvm_suld_1d_array_v4i8_trap:
5252 case Intrinsic::nvvm_suld_2d_i8_trap:
5253 case Intrinsic::nvvm_suld_2d_v2i8_trap:
5254 case Intrinsic::nvvm_suld_2d_v4i8_trap:
5255 case Intrinsic::nvvm_suld_2d_array_i8_trap:
5256 case Intrinsic::nvvm_suld_2d_array_v2i8_trap:
5257 case Intrinsic::nvvm_suld_2d_array_v4i8_trap:
5258 case Intrinsic::nvvm_suld_3d_i8_trap:
5259 case Intrinsic::nvvm_suld_3d_v2i8_trap:
5260 case Intrinsic::nvvm_suld_3d_v4i8_trap:
5261 case Intrinsic::nvvm_suld_1d_i8_zero:
5262 case Intrinsic::nvvm_suld_1d_v2i8_zero:
5263 case Intrinsic::nvvm_suld_1d_v4i8_zero:
5264 case Intrinsic::nvvm_suld_1d_array_i8_zero:
5265 case Intrinsic::nvvm_suld_1d_array_v2i8_zero:
5266 case Intrinsic::nvvm_suld_1d_array_v4i8_zero:
5267 case Intrinsic::nvvm_suld_2d_i8_zero:
5268 case Intrinsic::nvvm_suld_2d_v2i8_zero:
5269 case Intrinsic::nvvm_suld_2d_v4i8_zero:
5270 case Intrinsic::nvvm_suld_2d_array_i8_zero:
5271 case Intrinsic::nvvm_suld_2d_array_v2i8_zero:
5272 case Intrinsic::nvvm_suld_2d_array_v4i8_zero:
5273 case Intrinsic::nvvm_suld_3d_i8_zero:
5274 case Intrinsic::nvvm_suld_3d_v2i8_zero:
5275 case Intrinsic::nvvm_suld_3d_v4i8_zero:
5276 Info.opc = ISD::INTRINSIC_W_CHAIN;
5277 Info.memVT = MVT::i8;
5278 Info.ptrVal = nullptr;
5279 Info.offset = 0;
5280 Info.flags = MachineMemOperand::MOLoad;
5281 Info.align = Align(16);
5282 Infos.push_back(Elt: Info);
5283 return;
5284
5285 case Intrinsic::nvvm_suld_1d_i16_clamp:
5286 case Intrinsic::nvvm_suld_1d_v2i16_clamp:
5287 case Intrinsic::nvvm_suld_1d_v4i16_clamp:
5288 case Intrinsic::nvvm_suld_1d_array_i16_clamp:
5289 case Intrinsic::nvvm_suld_1d_array_v2i16_clamp:
5290 case Intrinsic::nvvm_suld_1d_array_v4i16_clamp:
5291 case Intrinsic::nvvm_suld_2d_i16_clamp:
5292 case Intrinsic::nvvm_suld_2d_v2i16_clamp:
5293 case Intrinsic::nvvm_suld_2d_v4i16_clamp:
5294 case Intrinsic::nvvm_suld_2d_array_i16_clamp:
5295 case Intrinsic::nvvm_suld_2d_array_v2i16_clamp:
5296 case Intrinsic::nvvm_suld_2d_array_v4i16_clamp:
5297 case Intrinsic::nvvm_suld_3d_i16_clamp:
5298 case Intrinsic::nvvm_suld_3d_v2i16_clamp:
5299 case Intrinsic::nvvm_suld_3d_v4i16_clamp:
5300 case Intrinsic::nvvm_suld_1d_i16_trap:
5301 case Intrinsic::nvvm_suld_1d_v2i16_trap:
5302 case Intrinsic::nvvm_suld_1d_v4i16_trap:
5303 case Intrinsic::nvvm_suld_1d_array_i16_trap:
5304 case Intrinsic::nvvm_suld_1d_array_v2i16_trap:
5305 case Intrinsic::nvvm_suld_1d_array_v4i16_trap:
5306 case Intrinsic::nvvm_suld_2d_i16_trap:
5307 case Intrinsic::nvvm_suld_2d_v2i16_trap:
5308 case Intrinsic::nvvm_suld_2d_v4i16_trap:
5309 case Intrinsic::nvvm_suld_2d_array_i16_trap:
5310 case Intrinsic::nvvm_suld_2d_array_v2i16_trap:
5311 case Intrinsic::nvvm_suld_2d_array_v4i16_trap:
5312 case Intrinsic::nvvm_suld_3d_i16_trap:
5313 case Intrinsic::nvvm_suld_3d_v2i16_trap:
5314 case Intrinsic::nvvm_suld_3d_v4i16_trap:
5315 case Intrinsic::nvvm_suld_1d_i16_zero:
5316 case Intrinsic::nvvm_suld_1d_v2i16_zero:
5317 case Intrinsic::nvvm_suld_1d_v4i16_zero:
5318 case Intrinsic::nvvm_suld_1d_array_i16_zero:
5319 case Intrinsic::nvvm_suld_1d_array_v2i16_zero:
5320 case Intrinsic::nvvm_suld_1d_array_v4i16_zero:
5321 case Intrinsic::nvvm_suld_2d_i16_zero:
5322 case Intrinsic::nvvm_suld_2d_v2i16_zero:
5323 case Intrinsic::nvvm_suld_2d_v4i16_zero:
5324 case Intrinsic::nvvm_suld_2d_array_i16_zero:
5325 case Intrinsic::nvvm_suld_2d_array_v2i16_zero:
5326 case Intrinsic::nvvm_suld_2d_array_v4i16_zero:
5327 case Intrinsic::nvvm_suld_3d_i16_zero:
5328 case Intrinsic::nvvm_suld_3d_v2i16_zero:
5329 case Intrinsic::nvvm_suld_3d_v4i16_zero:
5330 Info.opc = ISD::INTRINSIC_W_CHAIN;
5331 Info.memVT = MVT::i16;
5332 Info.ptrVal = nullptr;
5333 Info.offset = 0;
5334 Info.flags = MachineMemOperand::MOLoad;
5335 Info.align = Align(16);
5336 Infos.push_back(Elt: Info);
5337 return;
5338
5339 case Intrinsic::nvvm_suld_1d_i32_clamp:
5340 case Intrinsic::nvvm_suld_1d_v2i32_clamp:
5341 case Intrinsic::nvvm_suld_1d_v4i32_clamp:
5342 case Intrinsic::nvvm_suld_1d_array_i32_clamp:
5343 case Intrinsic::nvvm_suld_1d_array_v2i32_clamp:
5344 case Intrinsic::nvvm_suld_1d_array_v4i32_clamp:
5345 case Intrinsic::nvvm_suld_2d_i32_clamp:
5346 case Intrinsic::nvvm_suld_2d_v2i32_clamp:
5347 case Intrinsic::nvvm_suld_2d_v4i32_clamp:
5348 case Intrinsic::nvvm_suld_2d_array_i32_clamp:
5349 case Intrinsic::nvvm_suld_2d_array_v2i32_clamp:
5350 case Intrinsic::nvvm_suld_2d_array_v4i32_clamp:
5351 case Intrinsic::nvvm_suld_3d_i32_clamp:
5352 case Intrinsic::nvvm_suld_3d_v2i32_clamp:
5353 case Intrinsic::nvvm_suld_3d_v4i32_clamp:
5354 case Intrinsic::nvvm_suld_1d_i32_trap:
5355 case Intrinsic::nvvm_suld_1d_v2i32_trap:
5356 case Intrinsic::nvvm_suld_1d_v4i32_trap:
5357 case Intrinsic::nvvm_suld_1d_array_i32_trap:
5358 case Intrinsic::nvvm_suld_1d_array_v2i32_trap:
5359 case Intrinsic::nvvm_suld_1d_array_v4i32_trap:
5360 case Intrinsic::nvvm_suld_2d_i32_trap:
5361 case Intrinsic::nvvm_suld_2d_v2i32_trap:
5362 case Intrinsic::nvvm_suld_2d_v4i32_trap:
5363 case Intrinsic::nvvm_suld_2d_array_i32_trap:
5364 case Intrinsic::nvvm_suld_2d_array_v2i32_trap:
5365 case Intrinsic::nvvm_suld_2d_array_v4i32_trap:
5366 case Intrinsic::nvvm_suld_3d_i32_trap:
5367 case Intrinsic::nvvm_suld_3d_v2i32_trap:
5368 case Intrinsic::nvvm_suld_3d_v4i32_trap:
5369 case Intrinsic::nvvm_suld_1d_i32_zero:
5370 case Intrinsic::nvvm_suld_1d_v2i32_zero:
5371 case Intrinsic::nvvm_suld_1d_v4i32_zero:
5372 case Intrinsic::nvvm_suld_1d_array_i32_zero:
5373 case Intrinsic::nvvm_suld_1d_array_v2i32_zero:
5374 case Intrinsic::nvvm_suld_1d_array_v4i32_zero:
5375 case Intrinsic::nvvm_suld_2d_i32_zero:
5376 case Intrinsic::nvvm_suld_2d_v2i32_zero:
5377 case Intrinsic::nvvm_suld_2d_v4i32_zero:
5378 case Intrinsic::nvvm_suld_2d_array_i32_zero:
5379 case Intrinsic::nvvm_suld_2d_array_v2i32_zero:
5380 case Intrinsic::nvvm_suld_2d_array_v4i32_zero:
5381 case Intrinsic::nvvm_suld_3d_i32_zero:
5382 case Intrinsic::nvvm_suld_3d_v2i32_zero:
5383 case Intrinsic::nvvm_suld_3d_v4i32_zero:
5384 Info.opc = ISD::INTRINSIC_W_CHAIN;
5385 Info.memVT = MVT::i32;
5386 Info.ptrVal = nullptr;
5387 Info.offset = 0;
5388 Info.flags = MachineMemOperand::MOLoad;
5389 Info.align = Align(16);
5390 Infos.push_back(Elt: Info);
5391 return;
5392
5393 case Intrinsic::nvvm_suld_1d_i64_clamp:
5394 case Intrinsic::nvvm_suld_1d_v2i64_clamp:
5395 case Intrinsic::nvvm_suld_1d_array_i64_clamp:
5396 case Intrinsic::nvvm_suld_1d_array_v2i64_clamp:
5397 case Intrinsic::nvvm_suld_2d_i64_clamp:
5398 case Intrinsic::nvvm_suld_2d_v2i64_clamp:
5399 case Intrinsic::nvvm_suld_2d_array_i64_clamp:
5400 case Intrinsic::nvvm_suld_2d_array_v2i64_clamp:
5401 case Intrinsic::nvvm_suld_3d_i64_clamp:
5402 case Intrinsic::nvvm_suld_3d_v2i64_clamp:
5403 case Intrinsic::nvvm_suld_1d_i64_trap:
5404 case Intrinsic::nvvm_suld_1d_v2i64_trap:
5405 case Intrinsic::nvvm_suld_1d_array_i64_trap:
5406 case Intrinsic::nvvm_suld_1d_array_v2i64_trap:
5407 case Intrinsic::nvvm_suld_2d_i64_trap:
5408 case Intrinsic::nvvm_suld_2d_v2i64_trap:
5409 case Intrinsic::nvvm_suld_2d_array_i64_trap:
5410 case Intrinsic::nvvm_suld_2d_array_v2i64_trap:
5411 case Intrinsic::nvvm_suld_3d_i64_trap:
5412 case Intrinsic::nvvm_suld_3d_v2i64_trap:
5413 case Intrinsic::nvvm_suld_1d_i64_zero:
5414 case Intrinsic::nvvm_suld_1d_v2i64_zero:
5415 case Intrinsic::nvvm_suld_1d_array_i64_zero:
5416 case Intrinsic::nvvm_suld_1d_array_v2i64_zero:
5417 case Intrinsic::nvvm_suld_2d_i64_zero:
5418 case Intrinsic::nvvm_suld_2d_v2i64_zero:
5419 case Intrinsic::nvvm_suld_2d_array_i64_zero:
5420 case Intrinsic::nvvm_suld_2d_array_v2i64_zero:
5421 case Intrinsic::nvvm_suld_3d_i64_zero:
5422 case Intrinsic::nvvm_suld_3d_v2i64_zero:
5423 Info.opc = ISD::INTRINSIC_W_CHAIN;
5424 Info.memVT = MVT::i64;
5425 Info.ptrVal = nullptr;
5426 Info.offset = 0;
5427 Info.flags = MachineMemOperand::MOLoad;
5428 Info.align = Align(16);
5429 Infos.push_back(Elt: Info);
5430 return;
5431
5432 case Intrinsic::nvvm_tcgen05_ld_16x64b_x1:
5433 case Intrinsic::nvvm_tcgen05_ld_32x32b_x1:
5434 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x1: {
5435 Info.opc = ISD::INTRINSIC_W_CHAIN;
5436 Info.memVT = MVT::v1i32;
5437 Info.ptrVal = I.getArgOperand(i: 0);
5438 Info.offset = 0;
5439 Info.flags = MachineMemOperand::MOLoad;
5440 Info.align.reset();
5441 Infos.push_back(Elt: Info);
5442 return;
5443 }
5444
5445 case Intrinsic::nvvm_tcgen05_ld_16x64b_x2:
5446 case Intrinsic::nvvm_tcgen05_ld_16x128b_x1:
5447 case Intrinsic::nvvm_tcgen05_ld_32x32b_x2:
5448 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x2:
5449 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x2_i32:
5450 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x2_i32: {
5451 Info.opc = ISD::INTRINSIC_W_CHAIN;
5452 Info.memVT = MVT::v2i32;
5453 Info.ptrVal = I.getArgOperand(i: 0);
5454 Info.offset = 0;
5455 Info.flags = MachineMemOperand::MOLoad;
5456 Info.align.reset();
5457 Infos.push_back(Elt: Info);
5458 return;
5459 }
5460
5461 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x2_f32:
5462 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x2_f32: {
5463 Info.opc = ISD::INTRINSIC_W_CHAIN;
5464 Info.memVT = MVT::v2f32;
5465 Info.ptrVal = I.getArgOperand(i: 0);
5466 Info.offset = 0;
5467 Info.flags = MachineMemOperand::MOLoad;
5468 Info.align.reset();
5469 Infos.push_back(Elt: Info);
5470 return;
5471 }
5472
5473 case Intrinsic::nvvm_tcgen05_ld_16x64b_x4:
5474 case Intrinsic::nvvm_tcgen05_ld_16x128b_x2:
5475 case Intrinsic::nvvm_tcgen05_ld_32x32b_x4:
5476 case Intrinsic::nvvm_tcgen05_ld_16x256b_x1:
5477 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x4:
5478 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x4_i32:
5479 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x4_i32: {
5480 Info.opc = ISD::INTRINSIC_W_CHAIN;
5481 Info.memVT = MVT::v4i32;
5482 Info.ptrVal = I.getArgOperand(i: 0);
5483 Info.offset = 0;
5484 Info.flags = MachineMemOperand::MOLoad;
5485 Info.align.reset();
5486 Infos.push_back(Elt: Info);
5487 return;
5488 }
5489
5490 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x4_f32:
5491 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x4_f32: {
5492 Info.opc = ISD::INTRINSIC_W_CHAIN;
5493 Info.memVT = MVT::v4f32;
5494 Info.ptrVal = I.getArgOperand(i: 0);
5495 Info.offset = 0;
5496 Info.flags = MachineMemOperand::MOLoad;
5497 Info.align.reset();
5498 Infos.push_back(Elt: Info);
5499 return;
5500 }
5501
5502 case Intrinsic::nvvm_tcgen05_ld_16x64b_x8:
5503 case Intrinsic::nvvm_tcgen05_ld_16x128b_x4:
5504 case Intrinsic::nvvm_tcgen05_ld_16x256b_x2:
5505 case Intrinsic::nvvm_tcgen05_ld_32x32b_x8:
5506 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x8:
5507 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x8_i32:
5508 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x8_i32: {
5509 Info.opc = ISD::INTRINSIC_W_CHAIN;
5510 Info.memVT = MVT::v8i32;
5511 Info.ptrVal = I.getArgOperand(i: 0);
5512 Info.offset = 0;
5513 Info.flags = MachineMemOperand::MOLoad;
5514 Info.align.reset();
5515 Infos.push_back(Elt: Info);
5516 return;
5517 }
5518
5519 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x8_f32:
5520 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x8_f32: {
5521 Info.opc = ISD::INTRINSIC_W_CHAIN;
5522 Info.memVT = MVT::v8f32;
5523 Info.ptrVal = I.getArgOperand(i: 0);
5524 Info.offset = 0;
5525 Info.flags = MachineMemOperand::MOLoad;
5526 Info.align.reset();
5527 Infos.push_back(Elt: Info);
5528 return;
5529 }
5530
5531 case Intrinsic::nvvm_tcgen05_ld_16x64b_x16:
5532 case Intrinsic::nvvm_tcgen05_ld_16x128b_x8:
5533 case Intrinsic::nvvm_tcgen05_ld_16x256b_x4:
5534 case Intrinsic::nvvm_tcgen05_ld_32x32b_x16:
5535 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x16:
5536 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x16_i32:
5537 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x16_i32: {
5538 Info.opc = ISD::INTRINSIC_W_CHAIN;
5539 Info.memVT = MVT::v16i32;
5540 Info.ptrVal = I.getArgOperand(i: 0);
5541 Info.offset = 0;
5542 Info.flags = MachineMemOperand::MOLoad;
5543 Info.align.reset();
5544 Infos.push_back(Elt: Info);
5545 return;
5546 }
5547
5548 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x16_f32:
5549 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x16_f32: {
5550 Info.opc = ISD::INTRINSIC_W_CHAIN;
5551 Info.memVT = MVT::v16f32;
5552 Info.ptrVal = I.getArgOperand(i: 0);
5553 Info.offset = 0;
5554 Info.flags = MachineMemOperand::MOLoad;
5555 Info.align.reset();
5556 Infos.push_back(Elt: Info);
5557 return;
5558 }
5559
5560 case Intrinsic::nvvm_tcgen05_ld_16x64b_x32:
5561 case Intrinsic::nvvm_tcgen05_ld_16x128b_x16:
5562 case Intrinsic::nvvm_tcgen05_ld_16x256b_x8:
5563 case Intrinsic::nvvm_tcgen05_ld_32x32b_x32:
5564 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x32:
5565 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x32_i32:
5566 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x32_i32: {
5567 Info.opc = ISD::INTRINSIC_W_CHAIN;
5568 Info.memVT = MVT::v32i32;
5569 Info.ptrVal = I.getArgOperand(i: 0);
5570 Info.offset = 0;
5571 Info.flags = MachineMemOperand::MOLoad;
5572 Info.align.reset();
5573 Infos.push_back(Elt: Info);
5574 return;
5575 }
5576
5577 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x32_f32:
5578 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x32_f32: {
5579 Info.opc = ISD::INTRINSIC_W_CHAIN;
5580 Info.memVT = MVT::v32f32;
5581 Info.ptrVal = I.getArgOperand(i: 0);
5582 Info.offset = 0;
5583 Info.flags = MachineMemOperand::MOLoad;
5584 Info.align.reset();
5585 Infos.push_back(Elt: Info);
5586 return;
5587 }
5588
5589 case Intrinsic::nvvm_tcgen05_ld_16x64b_x64:
5590 case Intrinsic::nvvm_tcgen05_ld_16x128b_x32:
5591 case Intrinsic::nvvm_tcgen05_ld_16x256b_x16:
5592 case Intrinsic::nvvm_tcgen05_ld_32x32b_x64:
5593 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x64:
5594 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x64_i32:
5595 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x64_i32: {
5596 Info.opc = ISD::INTRINSIC_W_CHAIN;
5597 Info.memVT = MVT::v64i32;
5598 Info.ptrVal = I.getArgOperand(i: 0);
5599 Info.offset = 0;
5600 Info.flags = MachineMemOperand::MOLoad;
5601 Info.align.reset();
5602 Infos.push_back(Elt: Info);
5603 return;
5604 }
5605
5606 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x64_f32:
5607 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x64_f32: {
5608 Info.opc = ISD::INTRINSIC_W_CHAIN;
5609 Info.memVT = MVT::v64f32;
5610 Info.ptrVal = I.getArgOperand(i: 0);
5611 Info.offset = 0;
5612 Info.flags = MachineMemOperand::MOLoad;
5613 Info.align.reset();
5614 Infos.push_back(Elt: Info);
5615 return;
5616 }
5617
5618 case Intrinsic::nvvm_tcgen05_ld_16x64b_x128:
5619 case Intrinsic::nvvm_tcgen05_ld_16x128b_x64:
5620 case Intrinsic::nvvm_tcgen05_ld_16x256b_x32:
5621 case Intrinsic::nvvm_tcgen05_ld_32x32b_x128:
5622 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x128:
5623 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x128_i32:
5624 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x128_i32: {
5625 Info.opc = ISD::INTRINSIC_W_CHAIN;
5626 Info.memVT = MVT::v128i32;
5627 Info.ptrVal = I.getArgOperand(i: 0);
5628 Info.offset = 0;
5629 Info.flags = MachineMemOperand::MOLoad;
5630 Info.align.reset();
5631 Infos.push_back(Elt: Info);
5632 return;
5633 }
5634
5635 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x128_f32:
5636 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x128_f32: {
5637 Info.opc = ISD::INTRINSIC_W_CHAIN;
5638 Info.memVT = MVT::v128f32;
5639 Info.ptrVal = I.getArgOperand(i: 0);
5640 Info.offset = 0;
5641 Info.flags = MachineMemOperand::MOLoad;
5642 Info.align.reset();
5643 Infos.push_back(Elt: Info);
5644 return;
5645 }
5646
5647 case Intrinsic::nvvm_tcgen05_st_16x64b_x1:
5648 case Intrinsic::nvvm_tcgen05_st_32x32b_x1:
5649 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x1: {
5650 Info.opc = ISD::INTRINSIC_VOID;
5651 Info.memVT = MVT::v1i32;
5652 Info.ptrVal = I.getArgOperand(i: 0);
5653 Info.offset = 0;
5654 Info.flags = MachineMemOperand::MOStore;
5655 Info.align.reset();
5656 Infos.push_back(Elt: Info);
5657 return;
5658 }
5659
5660 case Intrinsic::nvvm_tcgen05_st_16x64b_x2:
5661 case Intrinsic::nvvm_tcgen05_st_16x128b_x1:
5662 case Intrinsic::nvvm_tcgen05_st_32x32b_x2:
5663 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x2: {
5664 Info.opc = ISD::INTRINSIC_VOID;
5665 Info.memVT = MVT::v2i32;
5666 Info.ptrVal = I.getArgOperand(i: 0);
5667 Info.offset = 0;
5668 Info.flags = MachineMemOperand::MOStore;
5669 Info.align.reset();
5670 Infos.push_back(Elt: Info);
5671 return;
5672 }
5673
5674 case Intrinsic::nvvm_tcgen05_st_16x64b_x4:
5675 case Intrinsic::nvvm_tcgen05_st_16x128b_x2:
5676 case Intrinsic::nvvm_tcgen05_st_16x256b_x1:
5677 case Intrinsic::nvvm_tcgen05_st_32x32b_x4:
5678 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x4: {
5679 Info.opc = ISD::INTRINSIC_VOID;
5680 Info.memVT = MVT::v4i32;
5681 Info.ptrVal = I.getArgOperand(i: 0);
5682 Info.offset = 0;
5683 Info.flags = MachineMemOperand::MOStore;
5684 Info.align.reset();
5685 Infos.push_back(Elt: Info);
5686 return;
5687 }
5688
5689 case Intrinsic::nvvm_tcgen05_st_16x64b_x8:
5690 case Intrinsic::nvvm_tcgen05_st_16x128b_x4:
5691 case Intrinsic::nvvm_tcgen05_st_16x256b_x2:
5692 case Intrinsic::nvvm_tcgen05_st_32x32b_x8:
5693 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x8: {
5694 Info.opc = ISD::INTRINSIC_VOID;
5695 Info.memVT = MVT::v8i32;
5696 Info.ptrVal = I.getArgOperand(i: 0);
5697 Info.offset = 0;
5698 Info.flags = MachineMemOperand::MOStore;
5699 Info.align.reset();
5700 Infos.push_back(Elt: Info);
5701 return;
5702 }
5703
5704 case Intrinsic::nvvm_tcgen05_st_16x64b_x16:
5705 case Intrinsic::nvvm_tcgen05_st_16x128b_x8:
5706 case Intrinsic::nvvm_tcgen05_st_16x256b_x4:
5707 case Intrinsic::nvvm_tcgen05_st_32x32b_x16:
5708 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x16: {
5709 Info.opc = ISD::INTRINSIC_VOID;
5710 Info.memVT = MVT::v16i32;
5711 Info.ptrVal = I.getArgOperand(i: 0);
5712 Info.offset = 0;
5713 Info.flags = MachineMemOperand::MOStore;
5714 Info.align.reset();
5715 Infos.push_back(Elt: Info);
5716 return;
5717 }
5718
5719 case Intrinsic::nvvm_tcgen05_st_16x64b_x32:
5720 case Intrinsic::nvvm_tcgen05_st_16x128b_x16:
5721 case Intrinsic::nvvm_tcgen05_st_16x256b_x8:
5722 case Intrinsic::nvvm_tcgen05_st_32x32b_x32:
5723 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x32: {
5724 Info.opc = ISD::INTRINSIC_VOID;
5725 Info.memVT = MVT::v32i32;
5726 Info.ptrVal = I.getArgOperand(i: 0);
5727 Info.offset = 0;
5728 Info.flags = MachineMemOperand::MOStore;
5729 Info.align.reset();
5730 Infos.push_back(Elt: Info);
5731 return;
5732 }
5733
5734 case Intrinsic::nvvm_tcgen05_st_16x64b_x64:
5735 case Intrinsic::nvvm_tcgen05_st_16x128b_x32:
5736 case Intrinsic::nvvm_tcgen05_st_16x256b_x16:
5737 case Intrinsic::nvvm_tcgen05_st_32x32b_x64:
5738 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x64: {
5739 Info.opc = ISD::INTRINSIC_VOID;
5740 Info.memVT = MVT::v64i32;
5741 Info.ptrVal = I.getArgOperand(i: 0);
5742 Info.offset = 0;
5743 Info.flags = MachineMemOperand::MOStore;
5744 Info.align.reset();
5745 Infos.push_back(Elt: Info);
5746 return;
5747 }
5748
5749 case Intrinsic::nvvm_tcgen05_st_16x64b_x128:
5750 case Intrinsic::nvvm_tcgen05_st_16x128b_x64:
5751 case Intrinsic::nvvm_tcgen05_st_16x256b_x32:
5752 case Intrinsic::nvvm_tcgen05_st_32x32b_x128:
5753 case Intrinsic::nvvm_tcgen05_st_16x32bx2_x128: {
5754 Info.opc = ISD::INTRINSIC_VOID;
5755 Info.memVT = MVT::v128i32;
5756 Info.ptrVal = I.getArgOperand(i: 0);
5757 Info.offset = 0;
5758 Info.flags = MachineMemOperand::MOStore;
5759 Info.align.reset();
5760 Infos.push_back(Elt: Info);
5761 return;
5762 }
5763 case Intrinsic::
5764 nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg1_decompress_b:
5765 case Intrinsic::
5766 nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg1_decompress_b:
5767 case Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg1:
5768 case Intrinsic::nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg1:
5769 case Intrinsic::nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg1:
5770 case Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg1:
5771 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg1:
5772 case Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1:
5773 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg1_ashift:
5774 case Intrinsic::
5775 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg1_ashift:
5776 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1:
5777 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1:
5778 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg1_ashift:
5779 case Intrinsic::
5780 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg1_ashift: {
5781 // We are reading and writing back to TMem
5782 Info.opc = ISD::INTRINSIC_VOID;
5783 Info.memVT = MVT::v4i32;
5784 Info.ptrVal = I.getArgOperand(i: 0);
5785 Info.offset = 0;
5786 Info.flags = MachineMemOperand::MOLoad | MachineMemOperand::MOStore;
5787 Info.align = Align(16);
5788 Infos.push_back(Elt: Info);
5789 return;
5790 }
5791
5792 case Intrinsic::
5793 nvvm_tcgen05_mma_shared_f8f6f4_disable_output_lane_cg2_decompress_b:
5794 case Intrinsic::
5795 nvvm_tcgen05_mma_tensor_f8f6f4_disable_output_lane_cg2_decompress_b:
5796 case Intrinsic::nvvm_tcgen05_mma_shared_disable_output_lane_cg2:
5797 case Intrinsic::nvvm_tcgen05_mma_shared_scale_d_disable_output_lane_cg2:
5798 case Intrinsic::nvvm_tcgen05_mma_sp_shared_disable_output_lane_cg2:
5799 case Intrinsic::nvvm_tcgen05_mma_sp_shared_scale_d_disable_output_lane_cg2:
5800 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg2:
5801 case Intrinsic::nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2:
5802 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2:
5803 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2:
5804 case Intrinsic::nvvm_tcgen05_mma_tensor_disable_output_lane_cg2_ashift:
5805 case Intrinsic::
5806 nvvm_tcgen05_mma_tensor_scale_d_disable_output_lane_cg2_ashift:
5807 case Intrinsic::nvvm_tcgen05_mma_sp_tensor_disable_output_lane_cg2_ashift:
5808 case Intrinsic::
5809 nvvm_tcgen05_mma_sp_tensor_scale_d_disable_output_lane_cg2_ashift: {
5810 // We are reading and writing back to TMem
5811 Info.opc = ISD::INTRINSIC_VOID;
5812 Info.memVT = MVT::v8i32;
5813 Info.ptrVal = I.getArgOperand(i: 0);
5814 Info.offset = 0;
5815 Info.flags = MachineMemOperand::MOLoad | MachineMemOperand::MOStore;
5816 Info.align = Align(16);
5817 Infos.push_back(Elt: Info);
5818 return;
5819 }
5820 case Intrinsic::nvvm_tcgen05_alloc_cg1:
5821 case Intrinsic::nvvm_tcgen05_alloc_cg2:
5822 Info.opc = ISD::INTRINSIC_VOID;
5823 Info.memVT = MVT::i32;
5824 Info.ptrVal = I.getArgOperand(i: 0);
5825 Info.offset = 0;
5826 Info.flags = MachineMemOperand::MOStore;
5827 Info.align = Align(4);
5828 Infos.push_back(Elt: Info);
5829 return;
5830 }
5831}
5832
5833// Helper for getting a function parameter symbol. Its name is composed from
5834// the function name and the parameter index. Negative index corresponds to the
5835// special parameter (unsized array) used for passing variable arguments.
5836MCSymbol *NVPTXTargetLowering::getParamSymbol(MCContext &Ctx, const Function *F,
5837 int Idx) const {
5838 const StringRef FuncName = getTargetMachine().getSymbol(GV: F)->getName();
5839 if (Idx < 0)
5840 return Ctx.getOrCreateSymbol(Name: FuncName + "_vararg");
5841 return Ctx.getOrCreateSymbol(Name: FuncName + "_param_" + Twine(Idx));
5842}
5843
5844/// isLegalAddressingMode - Return true if the addressing mode represented
5845/// by AM is legal for this target, for a load/store of the specified type.
5846/// Used to guide target specific optimizations, like loop strength reduction
5847/// (LoopStrengthReduce.cpp) and memory optimization for address mode
5848/// (CodeGenPrepare.cpp)
5849bool NVPTXTargetLowering::isLegalAddressingMode(const DataLayout &DL,
5850 const AddrMode &AM, Type *Ty,
5851 unsigned AS, Instruction *I) const {
5852 // AddrMode - This represents an addressing mode of:
5853 // BaseGV + BaseOffs + BaseReg + Scale*ScaleReg
5854 //
5855 // The legal address modes are
5856 // - [avar]
5857 // - [areg]
5858 // - [areg+immoff]
5859 // - [immAddr]
5860
5861 // immoff must fit in a signed 32-bit int
5862 if (!APInt(64, AM.BaseOffs).isSignedIntN(N: 32))
5863 return false;
5864
5865 if (AM.BaseGV)
5866 return !AM.BaseOffs && !AM.HasBaseReg && !AM.Scale;
5867
5868 switch (AM.Scale) {
5869 case 0: // "r", "r+i" or "i" is allowed
5870 break;
5871 case 1:
5872 if (AM.HasBaseReg) // "r+r+i" or "r+r" is not allowed.
5873 return false;
5874 // Otherwise we have r+i.
5875 break;
5876 default:
5877 // No scale > 1 is allowed
5878 return false;
5879 }
5880 return true;
5881}
5882
5883//===----------------------------------------------------------------------===//
5884// NVPTX Inline Assembly Support
5885//===----------------------------------------------------------------------===//
5886
5887/// getConstraintType - Given a constraint letter, return the type of
5888/// constraint it is for this target.
5889NVPTXTargetLowering::ConstraintType
5890NVPTXTargetLowering::getConstraintType(StringRef Constraint) const {
5891 if (Constraint.size() == 1) {
5892 switch (Constraint[0]) {
5893 default:
5894 break;
5895 case 'b':
5896 case 'r':
5897 case 'h':
5898 case 'c':
5899 case 'l':
5900 case 'f':
5901 case 'd':
5902 case 'q':
5903 case '0':
5904 case 'N':
5905 return C_RegisterClass;
5906 }
5907 }
5908 return TargetLowering::getConstraintType(Constraint);
5909}
5910
5911std::pair<unsigned, const TargetRegisterClass *>
5912NVPTXTargetLowering::getRegForInlineAsmConstraint(const TargetRegisterInfo *TRI,
5913 StringRef Constraint,
5914 MVT VT) const {
5915 if (Constraint.size() == 1) {
5916 switch (Constraint[0]) {
5917 case 'b':
5918 return std::make_pair(x: 0U, y: &NVPTX::B1RegClass);
5919 case 'c':
5920 case 'h':
5921 return std::make_pair(x: 0U, y: &NVPTX::B16RegClass);
5922 case 'r':
5923 case 'f':
5924 return std::make_pair(x: 0U, y: &NVPTX::B32RegClass);
5925 case 'l':
5926 case 'N':
5927 case 'd':
5928 return std::make_pair(x: 0U, y: &NVPTX::B64RegClass);
5929 case 'q': {
5930 if (!STI.hasFeature(Feature: NVPTX::SM70))
5931 report_fatal_error(reason: "Inline asm with 128 bit operands is only "
5932 "supported for sm_70 and higher!");
5933 return std::make_pair(x: 0U, y: &NVPTX::B128RegClass);
5934 }
5935 }
5936 }
5937 return TargetLowering::getRegForInlineAsmConstraint(TRI, Constraint, VT);
5938}
5939
5940//===----------------------------------------------------------------------===//
5941// NVPTX DAG Combining
5942//===----------------------------------------------------------------------===//
5943
5944bool NVPTXTargetLowering::allowFMA(MachineFunction &MF,
5945 CodeGenOptLevel OptLevel) const {
5946 // Always honor command-line argument
5947 if (FMAContractLevelOpt.getNumOccurrences() > 0)
5948 return FMAContractLevelOpt > 0;
5949
5950 // Do not contract if we're not optimizing the code.
5951 if (OptLevel == CodeGenOptLevel::None)
5952 return false;
5953
5954 return false;
5955}
5956
5957static bool isConstZero(const SDValue &Operand) {
5958 const auto *Const = dyn_cast<ConstantSDNode>(Val: Operand);
5959 return Const && Const->getZExtValue() == 0;
5960}
5961
5962/// PerformADDCombineWithOperands - Try DAG combinations for an ADD with
5963/// operands N0 and N1. This is a helper for PerformADDCombine that is
5964/// called with the default operands, and if that fails, with commuted
5965/// operands.
5966static SDValue
5967PerformADDCombineWithOperands(SDNode *N, SDValue N0, SDValue N1,
5968 TargetLowering::DAGCombinerInfo &DCI) {
5969 EVT VT = N0.getValueType();
5970
5971 // Since integer multiply-add costs the same as integer multiply
5972 // but is more costly than integer add, do the fusion only when
5973 // the mul is only used in the add.
5974 // TODO: this may not be true for later architectures, consider relaxing this
5975 if (!N0.getNode()->hasOneUse())
5976 return SDValue();
5977
5978 // fold (add (select cond, 0, (mul a, b)), c)
5979 // -> (select cond, c, (add (mul a, b), c))
5980 //
5981 if (N0.getOpcode() == ISD::SELECT) {
5982 unsigned ZeroOpNum;
5983 if (isConstZero(Operand: N0->getOperand(Num: 1)))
5984 ZeroOpNum = 1;
5985 else if (isConstZero(Operand: N0->getOperand(Num: 2)))
5986 ZeroOpNum = 2;
5987 else
5988 return SDValue();
5989
5990 SDValue M = N0->getOperand(Num: (ZeroOpNum == 1) ? 2 : 1);
5991 if (M->getOpcode() != ISD::MUL || !M.getNode()->hasOneUse())
5992 return SDValue();
5993
5994 SDLoc DL(N);
5995 SDValue Mul =
5996 DCI.DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: M->getOperand(Num: 0), N2: M->getOperand(Num: 1));
5997 SDValue MAD = DCI.DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Mul, N2: N1);
5998 return DCI.DAG.getSelect(DL: SDLoc(N), VT, Cond: N0->getOperand(Num: 0),
5999 LHS: ((ZeroOpNum == 1) ? N1 : MAD),
6000 RHS: ((ZeroOpNum == 1) ? MAD : N1));
6001 }
6002
6003 return SDValue();
6004}
6005
6006SDValue NVPTXTargetLowering::performFADDCombineWithOperands(
6007 SDNode *N, SDValue N0, SDValue N1, TargetLowering::DAGCombinerInfo &DCI,
6008 CodeGenOptLevel OptLevel) const {
6009 EVT VT = N0.getValueType();
6010 if (N0.getOpcode() == ISD::FMUL) {
6011 if (!(allowFMA(MF&: DCI.DAG.getMachineFunction(), OptLevel) ||
6012 (N->getFlags().hasAllowContract() &&
6013 N0->getFlags().hasAllowContract())))
6014 return SDValue();
6015
6016 // For floating point:
6017 // Do the fusion only when the mul has less than 5 uses and all
6018 // are add.
6019 // The heuristic is that if a use is not an add, then that use
6020 // cannot be fused into fma, therefore mul is still needed anyway.
6021 // If there are more than 4 uses, even if they are all add, fusing
6022 // them will increase register pressue.
6023 //
6024 int numUses = 0;
6025 int nonAddCount = 0;
6026 for (const SDNode *User : N0.getNode()->users()) {
6027 numUses++;
6028 if (User->getOpcode() != ISD::FADD)
6029 ++nonAddCount;
6030 if (numUses >= 5)
6031 return SDValue();
6032 }
6033 if (nonAddCount) {
6034 int orderNo = N->getIROrder();
6035 int orderNo2 = N0.getNode()->getIROrder();
6036 // simple heuristics here for considering potential register
6037 // pressure, the logics here is that the differnce are used
6038 // to measure the distance between def and use, the longer distance
6039 // more likely cause register pressure.
6040 if (orderNo - orderNo2 < 500)
6041 return SDValue();
6042
6043 // Now, check if at least one of the FMUL's operands is live beyond the
6044 // node N, which guarantees that the FMA will not increase register
6045 // pressure at node N.
6046 bool opIsLive = false;
6047 const SDNode *left = N0.getOperand(i: 0).getNode();
6048 const SDNode *right = N0.getOperand(i: 1).getNode();
6049
6050 if (isa<ConstantSDNode>(Val: left) || isa<ConstantSDNode>(Val: right))
6051 opIsLive = true;
6052
6053 if (!opIsLive)
6054 for (const SDNode *User : left->users()) {
6055 int orderNo3 = User->getIROrder();
6056 if (orderNo3 > orderNo) {
6057 opIsLive = true;
6058 break;
6059 }
6060 }
6061
6062 if (!opIsLive)
6063 for (const SDNode *User : right->users()) {
6064 int orderNo3 = User->getIROrder();
6065 if (orderNo3 > orderNo) {
6066 opIsLive = true;
6067 break;
6068 }
6069 }
6070
6071 if (!opIsLive)
6072 return SDValue();
6073 }
6074
6075 return DCI.DAG.getNode(Opcode: ISD::FMA, DL: SDLoc(N), VT, N1: N0.getOperand(i: 0),
6076 N2: N0.getOperand(i: 1), N3: N1);
6077 }
6078
6079 return SDValue();
6080}
6081
6082/// Fold unpacking movs into a load by increasing the number of return values.
6083///
6084/// ex:
6085/// L: v2f16,ch = load <p>
6086/// a: f16 = extractelt L:0, 0
6087/// b: f16 = extractelt L:0, 1
6088/// use(a, b)
6089///
6090/// ...is turned into...
6091///
6092/// L: f16,f16,ch = LoadV2 <p>
6093/// use(L:0, L:1)
6094static SDValue
6095combineUnpackingMovIntoLoad(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
6096 // Don't run this optimization before the legalizer
6097 if (!DCI.isAfterLegalizeDAG())
6098 return SDValue();
6099
6100 EVT ElementVT = N->getValueType(ResNo: 0);
6101 // Avoid non-packed types and v4i8
6102 if (!NVPTX::isPackedVectorTy(VT: ElementVT) || ElementVT == MVT::v4i8)
6103 return SDValue();
6104
6105 // Check whether all outputs are either used by an extractelt or are
6106 // glue/chain nodes
6107 if (!all_of(Range: N->uses(), P: [&](SDUse &U) {
6108 // Skip glue, chain nodes
6109 if (U.getValueType() == MVT::Glue || U.getValueType() == MVT::Other)
6110 return true;
6111 if (U.getUser()->getOpcode() == ISD::EXTRACT_VECTOR_ELT) {
6112 if (N->getOpcode() != ISD::LOAD)
6113 return true;
6114 // Since this is an ISD::LOAD, check all extractelts are used. If
6115 // any are not used, we don't want to defeat another optimization that
6116 // will narrow the load.
6117 //
6118 // For example:
6119 //
6120 // L: v2f16,ch = load <p>
6121 // e0: f16 = extractelt L:0, 0
6122 // e1: f16 = extractelt L:0, 1 <-- unused
6123 // store e0
6124 //
6125 // Can be optimized by DAGCombiner to:
6126 //
6127 // L: f16,ch = load <p>
6128 // store L:0
6129 return !U.getUser()->use_empty();
6130 }
6131
6132 // Otherwise, this use prevents us from splitting a value.
6133 return false;
6134 }))
6135 return SDValue();
6136
6137 auto *LD = cast<MemSDNode>(Val: N);
6138 SDLoc DL(LD);
6139
6140 // the new opcode after we double the number of operands
6141 unsigned Opcode;
6142 SmallVector<SDValue> Operands(LD->ops());
6143 unsigned OldNumOutputs; // non-glue, non-chain outputs
6144 switch (LD->getOpcode()) {
6145 case ISD::LOAD:
6146 OldNumOutputs = 1;
6147 // Any packed type is legal, so the legalizer will not have lowered
6148 // ISD::LOAD -> NVPTXISD::Load (unless it's under-aligned). We have to do it
6149 // here.
6150 Opcode = NVPTXISD::LoadV2;
6151 // append a "full" used bytes mask operand right before the extension type
6152 // operand, signifying that all bytes are used.
6153 Operands.push_back(Elt: DCI.DAG.getConstant(UINT32_MAX, DL, VT: MVT::i32));
6154 Operands.push_back(Elt: DCI.DAG.getIntPtrConstant(
6155 Val: cast<LoadSDNode>(Val: LD)->getExtensionType(), DL));
6156 break;
6157 case NVPTXISD::LoadV2:
6158 OldNumOutputs = 2;
6159 Opcode = NVPTXISD::LoadV4;
6160 break;
6161 case NVPTXISD::LoadV4:
6162 // V8 is only supported for f32/i32. Don't forget, we're not changing the
6163 // load size here. This is already a 256-bit load.
6164 if (ElementVT != MVT::v2f32 && ElementVT != MVT::v2i32)
6165 return SDValue();
6166 OldNumOutputs = 4;
6167 Opcode = NVPTXISD::LoadV8;
6168 break;
6169 case NVPTXISD::LoadV8:
6170 // PTX doesn't support the next doubling of outputs
6171 return SDValue();
6172 }
6173
6174 // the non-glue, non-chain outputs in the new load
6175 const unsigned NewNumOutputs = OldNumOutputs * 2;
6176 SmallVector<EVT> NewVTs(NewNumOutputs, ElementVT.getVectorElementType());
6177 // add remaining chain and glue values
6178 NewVTs.append(in_start: LD->value_begin() + OldNumOutputs, in_end: LD->value_end());
6179
6180 // Create the new load
6181 SDValue NewLoad = DCI.DAG.getMemIntrinsicNode(
6182 Opcode, dl: DL, VTList: DCI.DAG.getVTList(VTs: NewVTs), Ops: Operands, MemVT: LD->getMemoryVT(),
6183 MMO: LD->getMemOperand());
6184
6185 // Now we use a combination of BUILD_VECTORs and a MERGE_VALUES node to keep
6186 // the outputs the same. These nodes will be optimized away in later
6187 // DAGCombiner iterations.
6188 SmallVector<SDValue> Results;
6189 for (unsigned I : seq(Size: OldNumOutputs))
6190 Results.push_back(Elt: DCI.DAG.getBuildVector(
6191 VT: ElementVT, DL, Ops: {NewLoad.getValue(R: I * 2), NewLoad.getValue(R: I * 2 + 1)}));
6192 // Add remaining chain and glue nodes
6193 for (unsigned I : seq(Size: NewLoad->getNumValues() - NewNumOutputs))
6194 Results.push_back(Elt: NewLoad.getValue(R: NewNumOutputs + I));
6195
6196 return DCI.DAG.getMergeValues(Ops: Results, dl: DL);
6197}
6198
6199/// Fold packing movs into a store.
6200///
6201/// ex:
6202/// v1: v2f16 = BUILD_VECTOR a:f16, b:f16
6203/// v2: v2f16 = BUILD_VECTOR c:f16, d:f16
6204/// StoreV2 v1, v2
6205///
6206/// ...is turned into...
6207///
6208/// StoreV4 a, b, c, d
6209static SDValue combinePackingMovIntoStore(SDNode *N,
6210 TargetLowering::DAGCombinerInfo &DCI,
6211 unsigned Front, unsigned Back) {
6212 // We want to run this as late as possible since other optimizations may
6213 // eliminate the BUILD_VECTORs.
6214 if (!DCI.isAfterLegalizeDAG())
6215 return SDValue();
6216
6217 // Get the type of the operands being stored.
6218 EVT ElementVT = N->getOperand(Num: Front).getValueType();
6219
6220 // Avoid non-packed types and v4i8
6221 if (!NVPTX::isPackedVectorTy(VT: ElementVT) || ElementVT == MVT::v4i8)
6222 return SDValue();
6223
6224 auto *ST = cast<MemSDNode>(Val: N);
6225
6226 // The new opcode after we double the number of operands.
6227 unsigned Opcode;
6228 switch (N->getOpcode()) {
6229 case ISD::STORE:
6230 // Any packed type is legal, so the legalizer will not have lowered
6231 // ISD::STORE -> NVPTXISD::Store (unless it's under-aligned). We have to do
6232 // it here.
6233 Opcode = NVPTXISD::StoreV2;
6234 break;
6235 case NVPTXISD::StoreV2:
6236 Opcode = NVPTXISD::StoreV4;
6237 break;
6238 case NVPTXISD::StoreV4:
6239 // V8 is only supported for f32/i32. Don't forget, we're not changing the
6240 // store size here. This is already a 256-bit store.
6241 if (ElementVT != MVT::v2f32 && ElementVT != MVT::v2i32)
6242 return SDValue();
6243 Opcode = NVPTXISD::StoreV8;
6244 break;
6245 case NVPTXISD::StoreV8:
6246 // PTX doesn't support the next doubling of operands
6247 return SDValue();
6248 default:
6249 llvm_unreachable("Unhandled store opcode");
6250 }
6251
6252 // Scan the operands and if they're all BUILD_VECTORs, we'll have gathered
6253 // their elements.
6254 SmallVector<SDValue, 4> Operands(N->ops().take_front(N: Front));
6255 for (SDValue BV : N->ops().drop_front(N: Front).drop_back(N: Back)) {
6256 if (BV.getOpcode() != ISD::BUILD_VECTOR)
6257 return SDValue();
6258
6259 // If the operand has multiple uses, this optimization can increase register
6260 // pressure.
6261 if (!BV.hasOneUse())
6262 return SDValue();
6263
6264 // DAGCombiner visits nodes bottom-up. Check the BUILD_VECTOR operands for
6265 // any signs they may be folded by some other pattern or rule.
6266 for (SDValue Op : BV->ops()) {
6267 // Peek through bitcasts
6268 if (Op.getOpcode() == ISD::BITCAST)
6269 Op = Op.getOperand(i: 0);
6270
6271 // This may be folded into a PRMT.
6272 if (Op.getValueType() == MVT::i16 && Op.getOpcode() == ISD::TRUNCATE &&
6273 Op->getOperand(Num: 0).getValueType() == MVT::i32)
6274 return SDValue();
6275
6276 // This may be folded into cvt.bf16x2
6277 if (Op.getOpcode() == ISD::FP_ROUND)
6278 return SDValue();
6279 }
6280 Operands.append(IL: {BV.getOperand(i: 0), BV.getOperand(i: 1)});
6281 }
6282 Operands.append(in_start: N->op_end() - Back, in_end: N->op_end());
6283
6284 // Now we replace the store
6285 return DCI.DAG.getMemIntrinsicNode(Opcode, dl: SDLoc(N), VTList: N->getVTList(), Ops: Operands,
6286 MemVT: ST->getMemoryVT(), MMO: ST->getMemOperand());
6287}
6288
6289static SDValue combineSTORE(SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
6290 const NVPTXSubtarget &STI) {
6291
6292 if (DCI.isBeforeLegalize() && N->getOpcode() == ISD::STORE) {
6293 // Here is our chance to custom lower a store with a non-simple type.
6294 // Unfortunately, we can't do this in the legalizer because there is no
6295 // way to setOperationAction for an non-simple type.
6296 StoreSDNode *ST = cast<StoreSDNode>(Val: N);
6297 if (!ST->getValue().getValueType().isSimple())
6298 return lowerSTOREVector(Op: SDValue(ST, 0), DAG&: DCI.DAG, STI);
6299 }
6300
6301 return combinePackingMovIntoStore(N, DCI, Front: 1, Back: 2);
6302}
6303
6304static SDValue combineLOAD(SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
6305 const NVPTXSubtarget &STI) {
6306 if (DCI.isBeforeLegalize() && N->getOpcode() == ISD::LOAD) {
6307 // Here is our chance to custom lower a load with a non-simple type.
6308 // Unfortunately, we can't do this in the legalizer because there is no
6309 // way to setOperationAction for an non-simple type.
6310 if (!N->getValueType(ResNo: 0).isSimple())
6311 return lowerLoadVector(N, DAG&: DCI.DAG, STI);
6312 }
6313
6314 return combineUnpackingMovIntoLoad(N, DCI);
6315}
6316
6317/// PerformADDCombine - Target-specific dag combine xforms for ISD::ADD.
6318///
6319static SDValue PerformADDCombine(SDNode *N,
6320 TargetLowering::DAGCombinerInfo &DCI,
6321 CodeGenOptLevel OptLevel) {
6322 if (OptLevel == CodeGenOptLevel::None)
6323 return SDValue();
6324
6325 SDValue N0 = N->getOperand(Num: 0);
6326 SDValue N1 = N->getOperand(Num: 1);
6327
6328 // Skip non-integer, non-scalar case
6329 EVT VT = N0.getValueType();
6330 if (VT.isVector() || VT != MVT::i32)
6331 return SDValue();
6332
6333 // First try with the default operand order.
6334 if (SDValue Result = PerformADDCombineWithOperands(N, N0, N1, DCI))
6335 return Result;
6336
6337 // If that didn't work, try again with the operands commuted.
6338 return PerformADDCombineWithOperands(N, N0: N1, N1: N0, DCI);
6339}
6340
6341/// Check if a v2f32 BUILD_VECTOR provably packs values from non-adjacent
6342/// register pairs (non-coalescable).
6343static bool isNonCoalescableBuildVector(const SDValue &BV) {
6344 if (BV.getOpcode() != ISD::BUILD_VECTOR || BV.getValueType() != MVT::v2f32)
6345 return false;
6346
6347 SDValue Elt0 = BV.getOperand(i: 0);
6348 SDValue Elt1 = BV.getOperand(i: 1);
6349
6350 bool IsExt0 = Elt0.getOpcode() == ISD::EXTRACT_VECTOR_ELT;
6351 bool IsExt1 = Elt1.getOpcode() == ISD::EXTRACT_VECTOR_ELT;
6352
6353 // If neither element is an EXTRACT_VECTOR_ELT they are free-standing
6354 // scalars and the register allocator can still place them side-by-side.
6355 if (!IsExt0 && !IsExt1)
6356 return false;
6357
6358 // If exactly one element is an EXTRACT_VECTOR_ELT, the other is a scalar
6359 // that cannot generally occupy the adjacent register slot.
6360 if (IsExt0 != IsExt1)
6361 return true;
6362
6363 // At this point both sources are extracting from vectors. If they are from
6364 // different vectors, then the BUILD_VECTOR is non-coalescable.
6365 SDValue Src0 = Elt0.getOperand(i: 0);
6366 SDValue Src1 = Elt1.getOperand(i: 0);
6367 if (Src0 != Src1)
6368 return true;
6369
6370 auto *Idx0 = dyn_cast<ConstantSDNode>(Val: Elt0.getOperand(i: 1));
6371 auto *Idx1 = dyn_cast<ConstantSDNode>(Val: Elt1.getOperand(i: 1));
6372 // If both indices are dynamic they will be lowered to
6373 // loads and the vector will be spilled to local memory. The register
6374 // allocator can easily place the results in adjacent registers.
6375 if (!Idx0 && !Idx1)
6376 return false;
6377
6378 // If one index is dynamic and the other is constant, the value from the
6379 // constant load will result in an additional register to pair with the result
6380 // from the dynamic load. We consider this non-coalescable.
6381 if ((Idx0 && !Idx1) || (!Idx0 && Idx1))
6382 return true;
6383
6384 // Both are constant, adjacent pairs are coalescable
6385 return std::abs(i: Idx0->getSExtValue() - Idx1->getSExtValue()) != 1;
6386}
6387
6388/// Return true if FMUL v2f32 node \p N may be scalarized to fold each lane's
6389/// product into a scalar FMA.
6390bool NVPTXTargetLowering::mayFoldFMULIntoFMA(SDNode *N, MachineFunction &MF,
6391 CodeGenOptLevel OptLevel) const {
6392 if (N->getOpcode() != ISD::FMUL || N->getValueType(ResNo: 0) != MVT::v2f32)
6393 return false;
6394 const bool GlobalFMA = allowFMA(MF, OptLevel);
6395 if (!N->getFlags().hasAllowContract() && !GlobalFMA)
6396 return false;
6397
6398 const SDNode *FirstFAdd = nullptr;
6399 unsigned NumScalarFAdd = 0;
6400
6401 // Both lanes must feed unique FADDs
6402 for (SDNode *EE : N->users()) {
6403 if (NumScalarFAdd == 2)
6404 return false;
6405
6406 if (EE->getOpcode() != ISD::EXTRACT_VECTOR_ELT || !EE->hasOneUse() ||
6407 !isa<ConstantSDNode>(Val: EE->getOperand(Num: 1)))
6408 return false;
6409
6410 const SDNode *const FAdd = *EE->users().begin();
6411 if (FAdd->getOpcode() != ISD::FADD ||
6412 (!GlobalFMA && !FAdd->getFlags().hasAllowContract()))
6413 return false;
6414
6415 if (!FirstFAdd)
6416 FirstFAdd = FAdd;
6417 else if (FAdd == FirstFAdd)
6418 return false;
6419
6420 NumScalarFAdd++;
6421 }
6422
6423 return NumScalarFAdd == 2;
6424}
6425
6426/// Scalarize a v2f32 arithmetic node (FADD, FMUL, FSUB, FMA) when at least
6427/// one operand is a BUILD_VECTOR that repacks values from non-adjacent register
6428/// pairs. Without this combine the BUILD_VECTOR forces allocation of a
6429/// temporary 64-bit register, increasing register pressure.
6430///
6431/// Example - before:
6432/// t0: v2f32,v2f32,ch = LoadV2 ...
6433/// t1: f32 = extract_vector_elt t0, 0
6434/// t2: f32 = extract_vector_elt t0:1, 0
6435/// t3: v2f32 = BUILD_VECTOR t1, t2 ;; non-coalescable repack
6436/// t4: v2f32 = fma t_a, t3, t_c
6437///
6438/// After:
6439/// t0: v2f32,v2f32,ch = LoadV2 ...
6440/// t1: f32 = extract_vector_elt t0, 0
6441/// t2: f32 = extract_vector_elt t0:1, 0
6442/// a0: f32 = extract_vector_elt t_a, 0
6443/// a1: f32 = extract_vector_elt t_a, 1
6444/// c0: f32 = extract_vector_elt t_c, 0
6445/// c1: f32 = extract_vector_elt t_c, 1
6446/// r0: f32 = fma a0, t1, c0
6447/// r1: f32 = fma a1, t2, c1
6448/// t4: v2f32 = BUILD_VECTOR r0, r1
6449///
6450/// Also scalarizes an FMUL when all output lanes feed into scalar FADDs
6451/// to enable scalar FMA combining.
6452SDValue NVPTXTargetLowering::performScalarizeV2F32Op(
6453 SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
6454 CodeGenOptLevel OptLevel) const {
6455 EVT VT = N->getValueType(ResNo: 0);
6456 if (VT != MVT::v2f32)
6457 return SDValue();
6458
6459 if (none_of(Range: N->ops(), P: isNonCoalescableBuildVector) &&
6460 !mayFoldFMULIntoFMA(N, MF&: DCI.DAG.getMachineFunction(), OptLevel))
6461 return SDValue();
6462
6463 SelectionDAG &DAG = DCI.DAG;
6464 SDLoc DL(N);
6465 EVT EltVT = VT.getVectorElementType();
6466 unsigned Opc = N->getOpcode();
6467
6468 // For each operand, get the scalar element at the given index: if the operand
6469 // is a BUILD_VECTOR, grab the element directly; otherwise, emit an
6470 // EXTRACT_VECTOR_ELT.
6471 auto GetElement = [&](SDValue Op, unsigned Index) -> SDValue {
6472 if (Op.getOpcode() == ISD::BUILD_VECTOR)
6473 return Op.getOperand(i: Index);
6474 return DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: EltVT, N1: Op,
6475 N2: DAG.getVectorIdxConstant(Val: Index, DL));
6476 };
6477
6478 // Build scalar operand lists for element 0 and element 1.
6479 SmallVector<SDValue, 3> Ops0, Ops1;
6480 for (const SDValue &Op : N->ops()) {
6481 Ops0.push_back(Elt: GetElement(Op, 0));
6482 Ops1.push_back(Elt: GetElement(Op, 1));
6483 }
6484
6485 SDValue Res0 = DAG.getNode(Opcode: Opc, DL, VT: EltVT, Ops: Ops0, Flags: N->getFlags());
6486 SDValue Res1 = DAG.getNode(Opcode: Opc, DL, VT: EltVT, Ops: Ops1, Flags: N->getFlags());
6487
6488 return DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL, VT, N1: Res0, N2: Res1);
6489}
6490
6491/// Target-specific dag combine xforms for ISD::FADD.
6492SDValue
6493NVPTXTargetLowering::performFADDCombine(SDNode *N,
6494 TargetLowering::DAGCombinerInfo &DCI,
6495 CodeGenOptLevel OptLevel) const {
6496 if (SDValue Result = performScalarizeV2F32Op(N, DCI, OptLevel))
6497 return Result;
6498
6499 SDValue N0 = N->getOperand(Num: 0);
6500 SDValue N1 = N->getOperand(Num: 1);
6501
6502 EVT VT = N0.getValueType();
6503 if (VT.isVector() || !(VT == MVT::f32 || VT == MVT::f64))
6504 return SDValue();
6505
6506 // First try with the default operand order.
6507 if (SDValue Result = performFADDCombineWithOperands(N, N0, N1, DCI, OptLevel))
6508 return Result;
6509
6510 // If that didn't work, try again with the operands commuted.
6511 return performFADDCombineWithOperands(N, N0: N1, N1: N0, DCI, OptLevel);
6512}
6513
6514/// Get 3-input version of a 2-input min/max opcode
6515static unsigned getMinMax3Opcode(unsigned MinMax2Opcode) {
6516 switch (MinMax2Opcode) {
6517 case ISD::FMAXNUM:
6518 case ISD::FMAXIMUMNUM:
6519 return NVPTXISD::FMAXNUM3;
6520 case ISD::FMINNUM:
6521 case ISD::FMINIMUMNUM:
6522 return NVPTXISD::FMINNUM3;
6523 case ISD::FMAXIMUM:
6524 return NVPTXISD::FMAXIMUM3;
6525 case ISD::FMINIMUM:
6526 return NVPTXISD::FMINIMUM3;
6527 default:
6528 llvm_unreachable("Invalid 2-input min/max opcode");
6529 }
6530}
6531
6532/// PerformFMinMaxCombine - Combine (fmaxnum (fmaxnum a, b), c) into
6533/// (fmaxnum3 a, b, c). Also covers other llvm min/max intrinsics.
6534static SDValue PerformFMinMaxCombine(SDNode *N,
6535 TargetLowering::DAGCombinerInfo &DCI,
6536 const NVPTXSubtarget &STI) {
6537
6538 // 3-input min/max requires PTX 8.8+ and SM_100+, and only supports f32s
6539 EVT VT = N->getValueType(ResNo: 0);
6540 if (VT != MVT::f32 || !STI.hasFeature(Feature: NVPTX::PTX88) ||
6541 !STI.hasFeature(Feature: NVPTX::SM100))
6542 return SDValue();
6543
6544 SDValue Op0 = N->getOperand(Num: 0);
6545 SDValue Op1 = N->getOperand(Num: 1);
6546 unsigned MinMaxOp2 = N->getOpcode();
6547 unsigned MinMaxOp3 = getMinMax3Opcode(MinMax2Opcode: MinMaxOp2);
6548
6549 if (Op0.getOpcode() == MinMaxOp2 && Op0.hasOneUse()) {
6550 // (maxnum (maxnum a, b), c) -> (maxnum3 a, b, c)
6551 SDValue A = Op0.getOperand(i: 0);
6552 SDValue B = Op0.getOperand(i: 1);
6553 SDValue C = Op1;
6554 return DCI.DAG.getNode(Opcode: MinMaxOp3, DL: SDLoc(N), VT, N1: A, N2: B, N3: C, Flags: N->getFlags());
6555 } else if (Op1.getOpcode() == MinMaxOp2 && Op1.hasOneUse()) {
6556 // (maxnum a, (maxnum b, c)) -> (maxnum3 a, b, c)
6557 SDValue A = Op0;
6558 SDValue B = Op1.getOperand(i: 0);
6559 SDValue C = Op1.getOperand(i: 1);
6560 return DCI.DAG.getNode(Opcode: MinMaxOp3, DL: SDLoc(N), VT, N1: A, N2: B, N3: C, Flags: N->getFlags());
6561 }
6562 return SDValue();
6563}
6564
6565// sext (mul.iN nsw x, y) => mul.wide.sN x, y
6566// zext (mul.iN nuw x, y) => mul.wide.uN x, y
6567// sext (shl.iN nsw x, const) => mul.wide.sN x, (1 << const)
6568// zext (shl.iN nuw x, const) => mul.wide.uN x, (1 << const)
6569static SDValue combineSZExtToMulWide(SDNode *N,
6570 TargetLowering::DAGCombinerInfo &DCI,
6571 CodeGenOptLevel OptLevel) {
6572 assert(N->getOpcode() == ISD::SIGN_EXTEND ||
6573 N->getOpcode() == ISD::ZERO_EXTEND);
6574
6575 if (OptLevel == CodeGenOptLevel::None)
6576 return SDValue();
6577
6578 SDValue Op = N->getOperand(Num: 0);
6579 if (!Op.hasOneUse())
6580 return SDValue();
6581
6582 EVT ToVT = N->getValueType(ResNo: 0);
6583 EVT FromVT = Op.getValueType();
6584 if (!((ToVT == MVT::i32 && FromVT == MVT::i16) ||
6585 (ToVT == MVT::i64 && FromVT == MVT::i32)))
6586 return SDValue();
6587
6588 bool IsSigned = N->getOpcode() == ISD::SIGN_EXTEND;
6589 if ((IsSigned && !Op->getFlags().hasNoSignedWrap()) ||
6590 (!IsSigned && !Op->getFlags().hasNoUnsignedWrap()))
6591 return SDValue();
6592
6593 SDLoc DL(N);
6594 SDValue LHS = Op.getOperand(i: 0);
6595 SDValue RHS = Op.getOperand(i: 1);
6596 unsigned MulWideOpcode =
6597 IsSigned ? NVPTXISD::MUL_WIDE_SIGNED : NVPTXISD::MUL_WIDE_UNSIGNED;
6598 if (Op.getOpcode() == ISD::MUL) {
6599 return DCI.DAG.getNode(Opcode: MulWideOpcode, DL, VT: ToVT, N1: LHS, N2: RHS);
6600 } else if (Op.getOpcode() == ISD::SHL && isa<ConstantSDNode>(Val: RHS)) {
6601 const auto ShiftAmt = Op.getConstantOperandVal(i: 1);
6602 const auto MulVal = APInt(FromVT.getSizeInBits(), 1) << ShiftAmt;
6603
6604 // Note that the sext (shl nsw ...) case doesn't work if 1 << const
6605 // overflows to a negative value! The only valid input values in this
6606 // case are 0 and -1 (all other values yield poison because of the nsw),
6607 // and mul.wide.sN would give us the wrong sign for -1. We could use
6608 // mul.wide.uN, but since this is a weird case anyway, we might as well not
6609 // apply this transformation at all.
6610 if (IsSigned && MulVal.isNegative())
6611 return SDValue();
6612
6613 RHS = DCI.DAG.getConstant(Val: MulVal, DL, VT: FromVT);
6614 return DCI.DAG.getNode(Opcode: MulWideOpcode, DL, VT: ToVT, N1: LHS, N2: RHS);
6615 }
6616
6617 return SDValue();
6618}
6619
6620enum OperandSignedness {
6621 Signed = 0,
6622 Unsigned,
6623 Unknown
6624};
6625
6626/// IsMulWideOperandDemotable - Checks if the provided DAG node is an operand
6627/// that can be demoted to \p OptSize bits without loss of information. The
6628/// signedness of the operand, if determinable, is placed in \p S.
6629static bool IsMulWideOperandDemotable(SDValue Op,
6630 unsigned OptSize,
6631 OperandSignedness &S) {
6632 S = Unknown;
6633
6634 if (Op.getOpcode() == ISD::SIGN_EXTEND ||
6635 Op.getOpcode() == ISD::SIGN_EXTEND_INREG) {
6636 EVT OrigVT = Op.getOperand(i: 0).getValueType();
6637 if (OrigVT.getFixedSizeInBits() <= OptSize) {
6638 S = Signed;
6639 return true;
6640 }
6641 } else if (Op.getOpcode() == ISD::ZERO_EXTEND) {
6642 EVT OrigVT = Op.getOperand(i: 0).getValueType();
6643 if (OrigVT.getFixedSizeInBits() <= OptSize) {
6644 S = Unsigned;
6645 return true;
6646 }
6647 }
6648
6649 return false;
6650}
6651
6652/// AreMulWideOperandsDemotable - Checks if the given LHS and RHS operands can
6653/// be demoted to \p OptSize bits without loss of information. If the operands
6654/// contain a constant, it should appear as the RHS operand. The signedness of
6655/// the operands is placed in \p IsSigned.
6656static bool AreMulWideOperandsDemotable(SDValue LHS, SDValue RHS,
6657 unsigned OptSize,
6658 bool &IsSigned) {
6659 OperandSignedness LHSSign;
6660
6661 // The LHS operand must be a demotable op
6662 if (!IsMulWideOperandDemotable(Op: LHS, OptSize, S&: LHSSign))
6663 return false;
6664
6665 // We should have been able to determine the signedness from the LHS
6666 if (LHSSign == Unknown)
6667 return false;
6668
6669 IsSigned = (LHSSign == Signed);
6670
6671 // The RHS can be a demotable op or a constant
6672 if (ConstantSDNode *CI = dyn_cast<ConstantSDNode>(Val&: RHS)) {
6673 const APInt &Val = CI->getAPIntValue();
6674 if (LHSSign == Unsigned) {
6675 return Val.isIntN(N: OptSize);
6676 } else {
6677 return Val.isSignedIntN(N: OptSize);
6678 }
6679 } else {
6680 OperandSignedness RHSSign;
6681 if (!IsMulWideOperandDemotable(Op: RHS, OptSize, S&: RHSSign))
6682 return false;
6683
6684 return LHSSign == RHSSign;
6685 }
6686}
6687
6688/// TryMULWIDECombine - Attempt to replace a multiply of M bits with a multiply
6689/// of M/2 bits that produces an M-bit result (i.e. mul.wide). This transform
6690/// works on both multiply DAG nodes and SHL DAG nodes with a constant shift
6691/// amount.
6692static SDValue TryMULWIDECombine(SDNode *N,
6693 TargetLowering::DAGCombinerInfo &DCI) {
6694 EVT MulType = N->getValueType(ResNo: 0);
6695 if (MulType != MVT::i32 && MulType != MVT::i64) {
6696 return SDValue();
6697 }
6698
6699 SDLoc DL(N);
6700 unsigned OptSize = MulType.getSizeInBits() >> 1;
6701 SDValue LHS = N->getOperand(Num: 0);
6702 SDValue RHS = N->getOperand(Num: 1);
6703
6704 // Canonicalize the multiply so the constant (if any) is on the right
6705 if (N->getOpcode() == ISD::MUL) {
6706 if (isa<ConstantSDNode>(Val: LHS)) {
6707 std::swap(a&: LHS, b&: RHS);
6708 }
6709 }
6710
6711 // If we have a SHL, determine the actual multiply amount
6712 if (N->getOpcode() == ISD::SHL) {
6713 ConstantSDNode *ShlRHS = dyn_cast<ConstantSDNode>(Val&: RHS);
6714 if (!ShlRHS) {
6715 return SDValue();
6716 }
6717
6718 APInt ShiftAmt = ShlRHS->getAPIntValue();
6719 unsigned BitWidth = MulType.getSizeInBits();
6720 if (ShiftAmt.sge(RHS: 0) && ShiftAmt.slt(RHS: BitWidth)) {
6721 APInt MulVal = APInt(BitWidth, 1) << ShiftAmt;
6722 RHS = DCI.DAG.getConstant(Val: MulVal, DL, VT: MulType);
6723 } else {
6724 return SDValue();
6725 }
6726 }
6727
6728 bool Signed;
6729 // Verify that our operands are demotable
6730 if (!AreMulWideOperandsDemotable(LHS, RHS, OptSize, IsSigned&: Signed)) {
6731 return SDValue();
6732 }
6733
6734 EVT DemotedVT;
6735 if (MulType == MVT::i32) {
6736 DemotedVT = MVT::i16;
6737 } else {
6738 DemotedVT = MVT::i32;
6739 }
6740
6741 // Truncate the operands to the correct size. Note that these are just for
6742 // type consistency and will (likely) be eliminated in later phases.
6743 SDValue TruncLHS =
6744 DCI.DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: DemotedVT, Operand: LHS);
6745 SDValue TruncRHS =
6746 DCI.DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: DemotedVT, Operand: RHS);
6747
6748 unsigned Opc;
6749 if (Signed) {
6750 Opc = NVPTXISD::MUL_WIDE_SIGNED;
6751 } else {
6752 Opc = NVPTXISD::MUL_WIDE_UNSIGNED;
6753 }
6754
6755 return DCI.DAG.getNode(Opcode: Opc, DL, VT: MulType, N1: TruncLHS, N2: TruncRHS);
6756}
6757
6758static bool isConstOne(const SDValue &Operand) {
6759 const auto *Const = dyn_cast<ConstantSDNode>(Val: Operand);
6760 return Const && Const->getZExtValue() == 1;
6761}
6762
6763static SDValue matchMADConstOnePattern(SDValue Add) {
6764 if (Add->getOpcode() != ISD::ADD)
6765 return SDValue();
6766
6767 if (isConstOne(Operand: Add->getOperand(Num: 0)))
6768 return Add->getOperand(Num: 1);
6769
6770 if (isConstOne(Operand: Add->getOperand(Num: 1)))
6771 return Add->getOperand(Num: 0);
6772
6773 return SDValue();
6774}
6775
6776static SDValue combineMADConstOne(SDValue X, SDValue Add, EVT VT, SDLoc DL,
6777 TargetLowering::DAGCombinerInfo &DCI) {
6778
6779 if (SDValue Y = matchMADConstOnePattern(Add)) {
6780 SDValue Mul = DCI.DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: X, N2: Y);
6781 return DCI.DAG.getNode(Opcode: ISD::ADD, DL, VT, N1: Mul, N2: X);
6782 }
6783
6784 return SDValue();
6785}
6786
6787static SDValue combineMulSelectConstOne(SDValue X, SDValue Select, EVT VT,
6788 SDLoc DL,
6789 TargetLowering::DAGCombinerInfo &DCI) {
6790 if (Select->getOpcode() != ISD::SELECT)
6791 return SDValue();
6792
6793 SDValue Cond = Select->getOperand(Num: 0);
6794
6795 unsigned ConstOpNo;
6796 if (isConstOne(Operand: Select->getOperand(Num: 1)))
6797 ConstOpNo = 1;
6798 else if (isConstOne(Operand: Select->getOperand(Num: 2)))
6799 ConstOpNo = 2;
6800 else
6801 return SDValue();
6802
6803 SDValue Y = Select->getOperand(Num: (ConstOpNo == 1) ? 2 : 1);
6804
6805 // Do not combine if the resulting sequence is not obviously profitable.
6806 if (!matchMADConstOnePattern(Add: Y))
6807 return SDValue();
6808
6809 SDValue NewMul = DCI.DAG.getNode(Opcode: ISD::MUL, DL, VT, N1: X, N2: Y);
6810
6811 return DCI.DAG.getNode(Opcode: ISD::SELECT, DL, VT, N1: Cond,
6812 N2: (ConstOpNo == 1) ? X : NewMul,
6813 N3: (ConstOpNo == 1) ? NewMul : X);
6814}
6815
6816static SDValue
6817PerformMULCombineWithOperands(SDNode *N, SDValue N0, SDValue N1,
6818 TargetLowering::DAGCombinerInfo &DCI) {
6819
6820 EVT VT = N0.getValueType();
6821 if (VT.isVector())
6822 return SDValue();
6823
6824 if (VT != MVT::i16 && VT != MVT::i32 && VT != MVT::i64)
6825 return SDValue();
6826
6827 SDLoc DL(N);
6828
6829 // (mul x, (add y, 1)) -> (add (mul x, y), x)
6830 if (SDValue Res = combineMADConstOne(X: N0, Add: N1, VT, DL, DCI))
6831 return Res;
6832 if (SDValue Res = combineMADConstOne(X: N1, Add: N0, VT, DL, DCI))
6833 return Res;
6834
6835 // (mul x, (select y, 1)) -> (select (mul x, y), x)
6836 if (SDValue Res = combineMulSelectConstOne(X: N0, Select: N1, VT, DL, DCI))
6837 return Res;
6838 if (SDValue Res = combineMulSelectConstOne(X: N1, Select: N0, VT, DL, DCI))
6839 return Res;
6840
6841 return SDValue();
6842}
6843
6844/// PerformMULCombine - Runs PTX-specific DAG combine patterns on MUL nodes.
6845static SDValue PerformMULCombine(SDNode *N,
6846 TargetLowering::DAGCombinerInfo &DCI,
6847 CodeGenOptLevel OptLevel) {
6848 if (OptLevel == CodeGenOptLevel::None)
6849 return SDValue();
6850
6851 if (SDValue Ret = TryMULWIDECombine(N, DCI))
6852 return Ret;
6853
6854 SDValue N0 = N->getOperand(Num: 0);
6855 SDValue N1 = N->getOperand(Num: 1);
6856 return PerformMULCombineWithOperands(N, N0, N1, DCI);
6857}
6858
6859/// Commute SHL with a bitwise logic operation when doing so exposes a common
6860/// shifted operand. For example:
6861///
6862/// Before:
6863/// N = shl (zext (LogicOp X, C)), ShiftAmount
6864/// OtherShift = shl (zext (OtherLogicOp X, OtherC)), ShiftAmount
6865///
6866/// After:
6867/// ShiftedX = shl (zext X), ShiftAmount
6868/// N = LogicOp ShiftedX, ShiftedC
6869/// OtherShift = OtherLogicOp ShiftedX, ShiftedOtherC
6870///
6871/// ShiftedC = (zext C) << ShiftAmount and ShiftedOtherC =
6872/// (zext OtherC) << ShiftAmount are folded constants. This replaces two
6873/// variable shifts with the single shared ShiftedX. Requiring another matching
6874/// shift avoids disrupting isolated address calculations where a shift may be
6875/// folded into the addressing mode.
6876static SDValue combineShiftOfLogicOp(SDNode *N,
6877 TargetLowering::DAGCombinerInfo &DCI) {
6878 using namespace SDPatternMatch;
6879
6880 struct ShiftOfLogicOp {
6881 SDNode *Shift;
6882 SDValue LogicOp;
6883 SDValue X;
6884 SDValue Constant;
6885 unsigned ExtendOpcode;
6886 };
6887
6888 // Match a logic operation, with an optional extension, inside a SHL.
6889 auto matchShiftOfLogicOp =
6890 [&](SDNode *Shift) -> std::optional<ShiftOfLogicOp> {
6891 if (Shift->getOpcode() != ISD::SHL || !Shift->getOperand(Num: 0).hasOneUse())
6892 return std::nullopt;
6893 ShiftOfLogicOp Match;
6894 Match.Shift = Shift;
6895 Match.LogicOp = Shift->getOperand(Num: 0);
6896 Match.ExtendOpcode = 0;
6897 if (ISD::isExtOpcode(Opcode: Match.LogicOp.getOpcode())) {
6898 Match.ExtendOpcode = Match.LogicOp.getOpcode();
6899 Match.LogicOp = Match.LogicOp.getOperand(i: 0);
6900 }
6901
6902 if (!sd_match(N: Match.LogicOp, P: m_OneUse(P: m_BitwiseLogic(
6903 L: m_Value(N&: Match.X),
6904 R: m_Value(N&: Match.Constant, P: m_ConstInt())))))
6905 return std::nullopt;
6906
6907 return Match;
6908 };
6909
6910 // Match N as the root shift-of-logic; bail if it does not fit the pattern.
6911 const std::optional<ShiftOfLogicOp> Root = matchShiftOfLogicOp(N);
6912 if (!Root)
6913 return SDValue();
6914
6915 // Only profitable for a constant shift amount: the per-op constant shift then
6916 // folds away instead of becoming an extra variable shift.
6917 if (!isConstOrConstSplat(N: N->getOperand(Num: 1)))
6918 return SDValue();
6919
6920 // Collect candidate shifts that share X. Reached through another user of X,
6921 // the logic result feeds the shift directly or through an optional extend.
6922 SmallVector<SDNode *, 4> CandidateShifts;
6923 for (const SDNode *CandidateLogicOp : Root->X->users()) {
6924 if (CandidateLogicOp == Root->LogicOp.getNode())
6925 continue;
6926 for (SDNode *LogicUser : CandidateLogicOp->users()) {
6927 if (ISD::isExtOpcode(Opcode: LogicUser->getOpcode())) {
6928 // shl (ext (logic X, C)): step through the extend to find the shift.
6929 for (SDNode *ExtendUser : LogicUser->users())
6930 if (ExtendUser->getOpcode() == ISD::SHL)
6931 CandidateShifts.push_back(Elt: ExtendUser);
6932 } else if (LogicUser->getOpcode() == ISD::SHL) {
6933 // shl (logic X, C): the user is already the shift.
6934 CandidateShifts.push_back(Elt: LogicUser);
6935 }
6936 }
6937 }
6938
6939 // Verify each candidate against the root's pattern: the same X, extension,
6940 // type, and shift amount.
6941 const EVT VT = N->getValueType(ResNo: 0);
6942 const SDValue ShiftAmount = N->getOperand(Num: 1);
6943 SmallVector<ShiftOfLogicOp, 4> Matches;
6944 for (SDNode *CandidateShift : CandidateShifts) {
6945 const std::optional<ShiftOfLogicOp> Candidate =
6946 matchShiftOfLogicOp(CandidateShift);
6947 if (Candidate && Candidate->X == Root->X &&
6948 Candidate->ExtendOpcode == Root->ExtendOpcode &&
6949 CandidateShift->getValueType(ResNo: 0) == VT &&
6950 CandidateShift->getOperand(Num: 1) == ShiftAmount)
6951 Matches.push_back(Elt: *Candidate);
6952 }
6953 if (Matches.empty())
6954 return SDValue();
6955
6956 // Build the shared shifted X once, then rewrite the root and every match
6957 // into a logic op over it so the shift is CSE'd.
6958 SelectionDAG &DAG = DCI.DAG;
6959 const SDValue ShiftedX =
6960 DAG.getNode(Opcode: ISD::SHL, DL: SDLoc(N), VT,
6961 N1: Root->ExtendOpcode
6962 ? DAG.getNode(Opcode: Root->ExtendOpcode, DL: SDLoc(N), VT, Operand: Root->X)
6963 : Root->X,
6964 N2: ShiftAmount);
6965
6966 // Rebuild the logic op from shared ShiftedX and a folded constant shift.
6967 auto buildCommutedLogicOp = [&](const SDValue LogicOp, SDValue C,
6968 const SDLoc &DL) {
6969 if (Root->ExtendOpcode)
6970 C = DAG.getNode(Opcode: Root->ExtendOpcode, DL, VT, Operand: C);
6971 const SDValue ShiftedC = DAG.getNode(Opcode: ISD::SHL, DL, VT, N1: C, N2: ShiftAmount);
6972 return DAG.getNode(Opcode: LogicOp.getOpcode(), DL, VT, N1: ShiftedX, N2: ShiftedC,
6973 Flags: LogicOp->getFlags());
6974 };
6975
6976 for (const ShiftOfLogicOp &Match : Matches)
6977 DCI.CombineTo(N: Match.Shift,
6978 Res: buildCommutedLogicOp(Match.LogicOp, Match.Constant,
6979 SDLoc(Match.Shift)));
6980 return buildCommutedLogicOp(Root->LogicOp, Root->Constant, SDLoc(N));
6981}
6982
6983/// PerformSHLCombine - Runs PTX-specific DAG combine patterns on SHL nodes.
6984static SDValue PerformSHLCombine(SDNode *N,
6985 TargetLowering::DAGCombinerInfo &DCI,
6986 CodeGenOptLevel OptLevel) {
6987 if (OptLevel > CodeGenOptLevel::None) {
6988 // Expose a shared shifted operand for CSE before mul.wide folding, which
6989 // would otherwise consume the shift.
6990 if (SDValue Ret = combineShiftOfLogicOp(N, DCI))
6991 return Ret;
6992
6993 // Try mul.wide combining at OptLevel > 0
6994 if (SDValue Ret = TryMULWIDECombine(N, DCI))
6995 return Ret;
6996 }
6997
6998 return SDValue();
6999}
7000
7001static SDValue PerformSETCCCombine(SDNode *N,
7002 TargetLowering::DAGCombinerInfo &DCI,
7003 const NVPTXSubtarget &STI) {
7004 EVT CCType = N->getValueType(ResNo: 0);
7005 SDValue A = N->getOperand(Num: 0);
7006 SDValue B = N->getOperand(Num: 1);
7007
7008 EVT AType = A.getValueType();
7009 if (!(CCType == MVT::v2i1 && (AType == MVT::v2f16 || AType == MVT::v2bf16)))
7010 return SDValue();
7011
7012 if (A.getValueType() == MVT::v2bf16 && !STI.hasFeature(Feature: NVPTX::SM90))
7013 return SDValue();
7014
7015 SDLoc DL(N);
7016 // setp.f16x2 returns two scalar predicates, which we need to
7017 // convert back to v2i1. The returned result will be scalarized by
7018 // the legalizer, but the comparison will remain a single vector
7019 // instruction.
7020 SDValue CCNode = DCI.DAG.getNode(
7021 Opcode: A.getValueType() == MVT::v2f16 ? NVPTXISD::SETP_F16X2
7022 : NVPTXISD::SETP_BF16X2,
7023 DL, VTList: DCI.DAG.getVTList(VT1: MVT::i1, VT2: MVT::i1), Ops: {A, B, N->getOperand(Num: 2)});
7024 return DCI.DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL, VT: CCType, N1: CCNode.getValue(R: 0),
7025 N2: CCNode.getValue(R: 1));
7026}
7027
7028static SDValue PerformEXTRACTCombine(SDNode *N,
7029 TargetLowering::DAGCombinerInfo &DCI) {
7030 SDValue Vector = peekThroughFreeze(V: N->getOperand(Num: 0));
7031 SDLoc DL(N);
7032 EVT VectorVT = Vector.getValueType();
7033 if (Vector->getOpcode() == ISD::LOAD && VectorVT.isSimple() &&
7034 IsPTXVectorType(VT: VectorVT.getSimpleVT()))
7035 return SDValue(); // Native vector loads already combine nicely w/
7036 // extract_vector_elt.
7037 // Don't mess with singletons or packed types (v2*32, v2*16, v4i8 and v8i8),
7038 // we already handle them OK.
7039 if (VectorVT.getVectorNumElements() == 1 ||
7040 NVPTX::isPackedVectorTy(VT: VectorVT) || VectorVT == MVT::v8i8)
7041 return SDValue();
7042
7043 // Don't mess with undef values as sra may be simplified to 0, not undef.
7044 if (Vector->isUndef() || ISD::allOperandsUndef(N: Vector.getNode()))
7045 return SDValue();
7046
7047 uint64_t VectorBits = VectorVT.getSizeInBits();
7048 // We only handle the types we can extract in-register.
7049 if (!(VectorBits == 16 || VectorBits == 32 || VectorBits == 64))
7050 return SDValue();
7051
7052 ConstantSDNode *Index = dyn_cast<ConstantSDNode>(Val: N->getOperand(Num: 1));
7053 // Index == 0 is handled by generic DAG combiner.
7054 if (!Index || Index->getZExtValue() == 0)
7055 return SDValue();
7056
7057 MVT IVT = MVT::getIntegerVT(BitWidth: VectorBits);
7058 EVT EltVT = VectorVT.getVectorElementType();
7059 EVT EltIVT = EltVT.changeTypeToInteger();
7060 uint64_t EltBits = EltVT.getScalarSizeInBits();
7061
7062 SDValue Result = DCI.DAG.getNode(
7063 Opcode: ISD::TRUNCATE, DL, VT: EltIVT,
7064 Operand: DCI.DAG.getNode(
7065 Opcode: ISD::SRA, DL, VT: IVT, N1: DCI.DAG.getNode(Opcode: ISD::BITCAST, DL, VT: IVT, Operand: Vector),
7066 N2: DCI.DAG.getConstant(Val: Index->getZExtValue() * EltBits, DL, VT: IVT)));
7067
7068 // If element has non-integer type, bitcast it back to the expected type.
7069 if (EltVT != EltIVT)
7070 Result = DCI.DAG.getNode(Opcode: ISD::BITCAST, DL, VT: EltVT, Operand: Result);
7071 // Past legalizer, we may need to extent i8 -> i16 to match the register type.
7072 if (EltVT != N->getValueType(ResNo: 0))
7073 Result = DCI.DAG.getNode(Opcode: ISD::ANY_EXTEND, DL, VT: N->getValueType(ResNo: 0), Operand: Result);
7074
7075 return Result;
7076}
7077
7078/// Transform patterns like:
7079/// (select (ugt shift_amt, BitWidth-1), 0, (srl/shl x, shift_amt))
7080/// (select (ult shift_amt, BitWidth), (srl/shl x, shift_amt), 0)
7081/// Into:
7082/// (NVPTXISD::SRL_CLAMP x, shift_amt) or (NVPTXISD::SHL_CLAMP x, shift_amt)
7083///
7084/// These patterns arise from code like `s >= 32 ? 0 : x >> s`. In LLVM,
7085/// over-shifting a value results in poison, but PTX shr/shl instructions clamp
7086/// the shift amount to BitWidth, making the guard redundant.
7087///
7088/// Note: We only handle SRL and SHL, not SRA, because arithmetic right shifts
7089/// can produce 0 or -1 when shift >= BitWidth.
7090/// Note: We don't handle uge or ule. These don't appear because of
7091/// canonicalization.
7092static SDValue PerformSELECTShiftCombine(SDNode *N,
7093 TargetLowering::DAGCombinerInfo &DCI) {
7094 if (!DCI.isAfterLegalizeDAG())
7095 return SDValue();
7096
7097 using namespace SDPatternMatch;
7098 unsigned BitWidth = N->getValueType(ResNo: 0).getSizeInBits();
7099 SDValue ShiftAmt, ShiftOp;
7100
7101 // Match logical shifts where the shift amount in the guard matches the shift
7102 // amount in the operation.
7103 auto LogicalShift = m_Value(
7104 N&: ShiftOp, P: m_AnyOf(preds: m_Srl(L: m_Value(), R: m_TruncOrSelf(Op: m_Deferred(V&: ShiftAmt))),
7105 preds: m_Shl(L: m_Value(), R: m_TruncOrSelf(Op: m_Deferred(V&: ShiftAmt)))));
7106
7107 // shift_amt > BitWidth-1 ? 0 : shift_op
7108 bool MatchedUGT = sd_match(
7109 N, P: m_Select(Cond: m_SpecificSetCC(CC: ISD::SETUGT, LHS: m_Value(N&: ShiftAmt),
7110 RHS: m_SpecificInt(V: APInt(BitWidth, BitWidth - 1))),
7111 T: m_Zero(), F: LogicalShift));
7112 // shift_amt < BitWidth ? shift_op : 0
7113 bool MatchedULT =
7114 !MatchedUGT &&
7115 sd_match(
7116 N, P: m_Select(Cond: m_SpecificSetCC(CC: ISD::SETULT, LHS: m_Value(N&: ShiftAmt),
7117 RHS: m_SpecificInt(V: APInt(BitWidth, BitWidth))),
7118 T: LogicalShift, F: m_Zero()));
7119
7120 if (!MatchedUGT && !MatchedULT)
7121 return SDValue();
7122
7123 // In LLVM IR, the shift amount and the value-to-be-shifted are the same
7124 // type, whereas in PTX the shift amount is always i32. Therefore when
7125 // shifting types larger than i32, we can only do this transformation if we
7126 // know that the upper bits of the shift amount are known zero.
7127 SDValue ClampAmt = ShiftOp.getOperand(i: 1);
7128 unsigned ClampAmtBits = ClampAmt.getValueSizeInBits();
7129 if (ShiftAmt.getValueSizeInBits() > ClampAmtBits &&
7130 DCI.DAG.computeKnownBits(Op: ShiftAmt).countMaxActiveBits() > ClampAmtBits)
7131 return SDValue();
7132
7133 // Return a clamp shift operation, which has the same semantics as PTX shift.
7134 unsigned ClampOpc = ShiftOp.getOpcode() == ISD::SRL ? NVPTXISD::SRL_CLAMP
7135 : NVPTXISD::SHL_CLAMP;
7136 return DCI.DAG.getNode(Opcode: ClampOpc, DL: SDLoc(N), VT: ShiftOp.getValueType(),
7137 N1: ShiftOp.getOperand(i: 0), N2: ClampAmt);
7138}
7139
7140static SDValue PerformVSELECTCombine(SDNode *N,
7141 TargetLowering::DAGCombinerInfo &DCI) {
7142 SDValue VA = N->getOperand(Num: 1);
7143 EVT VectorVT = VA.getValueType();
7144 if (VectorVT != MVT::v4i8)
7145 return SDValue();
7146
7147 // We need to split vselect into individual per-element operations Because we
7148 // use BFE/BFI instruction for byte extraction/insertion, we do end up with
7149 // 32-bit values, so we may as well do comparison as i32 to avoid conversions
7150 // to/from i16 normally used for i8 values.
7151 SmallVector<SDValue, 4> E;
7152 SDLoc DL(N);
7153 SDValue VCond = N->getOperand(Num: 0);
7154 SDValue VB = N->getOperand(Num: 2);
7155 for (int I = 0; I < 4; ++I) {
7156 SDValue C = DCI.DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i1, N1: VCond,
7157 N2: DCI.DAG.getConstant(Val: I, DL, VT: MVT::i32));
7158 SDValue EA = DCI.DAG.getAnyExtOrTrunc(
7159 Op: DCI.DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i8, N1: VA,
7160 N2: DCI.DAG.getConstant(Val: I, DL, VT: MVT::i32)),
7161 DL, VT: MVT::i32);
7162 SDValue EB = DCI.DAG.getAnyExtOrTrunc(
7163 Op: DCI.DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL, VT: MVT::i8, N1: VB,
7164 N2: DCI.DAG.getConstant(Val: I, DL, VT: MVT::i32)),
7165 DL, VT: MVT::i32);
7166 E.push_back(Elt: DCI.DAG.getAnyExtOrTrunc(
7167 Op: DCI.DAG.getNode(Opcode: ISD::SELECT, DL, VT: MVT::i32, N1: C, N2: EA, N3: EB), DL, VT: MVT::i8));
7168 }
7169 return DCI.DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL, VT: MVT::v4i8, Ops: E);
7170}
7171
7172static SDValue
7173PerformBUILD_VECTORCombine(SDNode *N, TargetLowering::DAGCombinerInfo &DCI) {
7174 auto VT = N->getValueType(ResNo: 0);
7175 if (!DCI.isAfterLegalizeDAG() ||
7176 // only process v2*16 types
7177 !(NVPTX::isPackedVectorTy(VT) && VT.is32BitVector() &&
7178 VT.getVectorNumElements() == 2))
7179 return SDValue();
7180
7181 auto Op0 = N->getOperand(Num: 0);
7182 auto Op1 = N->getOperand(Num: 1);
7183
7184 // Start out by assuming we want to take the lower 2 bytes of each i32
7185 // operand.
7186 uint64_t Op0Bytes = 0x10;
7187 uint64_t Op1Bytes = 0x54;
7188
7189 std::pair<SDValue *, uint64_t *> OpData[2] = {{&Op0, &Op0Bytes},
7190 {&Op1, &Op1Bytes}};
7191
7192 // Check that each operand is an i16, truncated from an i32 operand. We'll
7193 // select individual bytes from those original operands. Optionally, fold in a
7194 // shift right of that original operand.
7195 for (auto &[Op, OpBytes] : OpData) {
7196 // Eat up any bitcast
7197 if (Op->getOpcode() == ISD::BITCAST)
7198 *Op = Op->getOperand(i: 0);
7199
7200 if (!(Op->getValueType() == MVT::i16 && Op->getOpcode() == ISD::TRUNCATE &&
7201 Op->getOperand(i: 0).getValueType() == MVT::i32))
7202 return SDValue();
7203
7204 // If the truncate has multiple uses, this optimization can increase
7205 // register pressure
7206 if (!Op->hasOneUse())
7207 return SDValue();
7208
7209 *Op = Op->getOperand(i: 0);
7210
7211 // Optionally, fold in a shift-right of the original operand and let permute
7212 // pick the two higher bytes of the original value directly.
7213 if (Op->getOpcode() == ISD::SRL && isa<ConstantSDNode>(Val: Op->getOperand(i: 1))) {
7214 if (cast<ConstantSDNode>(Val: Op->getOperand(i: 1))->getZExtValue() == 16) {
7215 // Shift the PRMT byte selector to pick upper bytes from each respective
7216 // value, instead of the lower ones: 0x10 -> 0x32, 0x54 -> 0x76
7217 assert((*OpBytes == 0x10 || *OpBytes == 0x54) &&
7218 "PRMT selector values out of range");
7219 *OpBytes += 0x22;
7220 *Op = Op->getOperand(i: 0);
7221 }
7222 }
7223 }
7224
7225 SDLoc DL(N);
7226 auto &DAG = DCI.DAG;
7227
7228 auto PRMT =
7229 getPRMT(A: DAG.getBitcast(VT: MVT::i32, V: Op0), B: DAG.getBitcast(VT: MVT::i32, V: Op1),
7230 Selector: (Op1Bytes << 8) | Op0Bytes, DL, DAG);
7231 return DAG.getBitcast(VT, V: PRMT);
7232}
7233
7234static SDValue combineADDRSPACECAST(SDNode *N,
7235 TargetLowering::DAGCombinerInfo &DCI) {
7236 auto *ASCN1 = cast<AddrSpaceCastSDNode>(Val: N);
7237
7238 if (auto *ASCN2 = dyn_cast<AddrSpaceCastSDNode>(Val: ASCN1->getOperand(Num: 0))) {
7239 assert(ASCN2->getDestAddressSpace() == ASCN1->getSrcAddressSpace());
7240
7241 // Fold asc[B -> A](asc[A -> B](x)) -> x
7242 if (ASCN1->getDestAddressSpace() == ASCN2->getSrcAddressSpace())
7243 return ASCN2->getOperand(Num: 0);
7244 }
7245
7246 return SDValue();
7247}
7248
7249// Given a constant selector value and a prmt mode, return the selector value
7250// normalized to the generic prmt mode. See the PTX ISA documentation for more
7251// details:
7252// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#data-movement-and-conversion-instructions-prmt
7253static APInt getPRMTSelector(const APInt &Selector, unsigned Mode) {
7254 assert(Selector.getBitWidth() == 32 && "PRMT must have i32 operands");
7255
7256 if (Mode == NVPTX::PTXPrmtMode::NONE)
7257 return Selector;
7258
7259 const unsigned V = Selector.trunc(width: 2).getZExtValue();
7260
7261 const auto GetSelector = [](unsigned S0, unsigned S1, unsigned S2,
7262 unsigned S3) {
7263 return APInt(32, S0 | (S1 << 4) | (S2 << 8) | (S3 << 12));
7264 };
7265
7266 switch (Mode) {
7267 case NVPTX::PTXPrmtMode::F4E:
7268 return GetSelector(V, V + 1, V + 2, V + 3);
7269 case NVPTX::PTXPrmtMode::B4E:
7270 return GetSelector(V, (V - 1) & 7, (V - 2) & 7, (V - 3) & 7);
7271 case NVPTX::PTXPrmtMode::RC8:
7272 return GetSelector(V, V, V, V);
7273 case NVPTX::PTXPrmtMode::ECL:
7274 return GetSelector(V, std::max(a: V, b: 1U), std::max(a: V, b: 2U), 3U);
7275 case NVPTX::PTXPrmtMode::ECR:
7276 return GetSelector(0, std::min(a: V, b: 1U), std::min(a: V, b: 2U), V);
7277 case NVPTX::PTXPrmtMode::RC16: {
7278 unsigned V1 = (V & 1) << 1;
7279 return GetSelector(V1, V1 + 1, V1, V1 + 1);
7280 }
7281 default:
7282 llvm_unreachable("Invalid PRMT mode");
7283 }
7284}
7285
7286static APInt computePRMT(APInt A, APInt B, APInt Selector, unsigned Mode) {
7287 assert(A.getBitWidth() == 32 && B.getBitWidth() == 32 &&
7288 Selector.getBitWidth() == 32 && "PRMT must have i32 operands");
7289 // {b, a} = {{b7, b6, b5, b4}, {b3, b2, b1, b0}}
7290 APInt BitField = B.concat(NewLSB: A);
7291 APInt SelectorVal = getPRMTSelector(Selector, Mode);
7292 APInt Result(32, 0);
7293 for (unsigned I : llvm::seq(Size: 4U)) {
7294 APInt Sel = SelectorVal.extractBits(numBits: 4, bitPosition: I * 4);
7295 unsigned Idx = Sel.getLoBits(numBits: 3).getZExtValue();
7296 unsigned Sign = Sel.getHiBits(numBits: 1).getZExtValue();
7297 APInt Byte = BitField.extractBits(numBits: 8, bitPosition: Idx * 8);
7298 if (Sign)
7299 Byte = Byte.ashr(ShiftAmt: 8);
7300 Result.insertBits(SubBits: Byte, bitPosition: I * 8);
7301 }
7302 return Result;
7303}
7304
7305static SDValue combinePRMT(SDNode *N, TargetLowering::DAGCombinerInfo &DCI,
7306 CodeGenOptLevel OptLevel) {
7307 if (OptLevel == CodeGenOptLevel::None)
7308 return SDValue();
7309
7310 // Constant fold PRMT
7311 if (isa<ConstantSDNode>(Val: N->getOperand(Num: 0)) &&
7312 isa<ConstantSDNode>(Val: N->getOperand(Num: 1)) &&
7313 isa<ConstantSDNode>(Val: N->getOperand(Num: 2)))
7314 return DCI.DAG.getConstant(Val: computePRMT(A: N->getConstantOperandAPInt(Num: 0),
7315 B: N->getConstantOperandAPInt(Num: 1),
7316 Selector: N->getConstantOperandAPInt(Num: 2),
7317 Mode: N->getConstantOperandVal(Num: 3)),
7318 DL: SDLoc(N), VT: N->getValueType(ResNo: 0));
7319 return SDValue();
7320}
7321
7322// During call lowering we wrap the return values in a ProxyReg node which
7323// depend on the chain value produced by the completed call. This ensures that
7324// the full call is emitted in cases where libcalls are used to legalize
7325// operations. To improve the functioning of other DAG combines we pull all
7326// operations we can through one of these nodes, ensuring that the ProxyReg
7327// directly wraps a load. That is:
7328//
7329// (ProxyReg (zext (load retval0))) => (zext (ProxyReg (load retval0)))
7330//
7331static SDValue sinkProxyReg(SDValue R, SDValue Chain,
7332 TargetLowering::DAGCombinerInfo &DCI) {
7333 switch (R.getOpcode()) {
7334 case ISD::TRUNCATE:
7335 case ISD::ANY_EXTEND:
7336 case ISD::SIGN_EXTEND:
7337 case ISD::ZERO_EXTEND:
7338 case ISD::BITCAST: {
7339 if (SDValue V = sinkProxyReg(R: R.getOperand(i: 0), Chain, DCI))
7340 return DCI.DAG.getNode(Opcode: R.getOpcode(), DL: SDLoc(R), VT: R.getValueType(), Operand: V);
7341 return SDValue();
7342 }
7343 case ISD::SHL:
7344 case ISD::SRL:
7345 case ISD::SRA:
7346 case ISD::OR: {
7347 if (SDValue A = sinkProxyReg(R: R.getOperand(i: 0), Chain, DCI))
7348 if (SDValue B = sinkProxyReg(R: R.getOperand(i: 1), Chain, DCI))
7349 return DCI.DAG.getNode(Opcode: R.getOpcode(), DL: SDLoc(R), VT: R.getValueType(), N1: A, N2: B);
7350 return SDValue();
7351 }
7352 case ISD::Constant:
7353 return R;
7354 case ISD::LOAD:
7355 case NVPTXISD::LoadV2:
7356 case NVPTXISD::LoadV4: {
7357 return DCI.DAG.getNode(Opcode: NVPTXISD::ProxyReg, DL: SDLoc(R), VT: R.getValueType(),
7358 Ops: {Chain, R});
7359 }
7360 case ISD::BUILD_VECTOR: {
7361 if (DCI.isBeforeLegalize())
7362 return SDValue();
7363
7364 SmallVector<SDValue, 16> Ops;
7365 for (auto &Op : R->ops()) {
7366 SDValue V = sinkProxyReg(R: Op, Chain, DCI);
7367 if (!V)
7368 return SDValue();
7369 Ops.push_back(Elt: V);
7370 }
7371 return DCI.DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL: SDLoc(R), VT: R.getValueType(), Ops);
7372 }
7373 case ISD::EXTRACT_VECTOR_ELT: {
7374 if (DCI.isBeforeLegalize())
7375 return SDValue();
7376
7377 if (SDValue V = sinkProxyReg(R: R.getOperand(i: 0), Chain, DCI))
7378 return DCI.DAG.getNode(Opcode: ISD::EXTRACT_VECTOR_ELT, DL: SDLoc(R),
7379 VT: R.getValueType(), N1: V, N2: R.getOperand(i: 1));
7380 return SDValue();
7381 }
7382 default:
7383 return SDValue();
7384 }
7385}
7386
7387static unsigned getFAddWithNegOpcode(EVT VT, Intrinsic::ID IID,
7388 APFloat::roundingMode RoundingMode) {
7389 const bool IsFTZ = nvvm::FPArithShouldFTZ(IntrinsicID: IID);
7390 const bool IsSat = nvvm::FPArithIsSaturating(IntrinsicID: IID);
7391 switch (VT.getScalarType().getSimpleVT().SimpleTy) {
7392 case MVT::f16: {
7393 static constexpr unsigned SubRNOpcodes[2][2] = {
7394 {NVPTXISD::SUB_RN, NVPTXISD::SUB_RN_SAT},
7395 {NVPTXISD::SUB_RN_FTZ, NVPTXISD::SUB_RN_FTZ_SAT}};
7396 return SubRNOpcodes[IsFTZ][IsSat];
7397 }
7398 case MVT::bf16:
7399 return NVPTXISD::SUB_RN;
7400 case MVT::f32: {
7401 // for f32x2 inputs
7402 if (!VT.isVector() || IsSat)
7403 return 0;
7404 static constexpr unsigned SubF32x2Opcodes[4][2] = {
7405 {NVPTXISD::SUB_RZ, NVPTXISD::SUB_RZ_FTZ}, // RZ
7406 {NVPTXISD::SUB_RN, NVPTXISD::SUB_RN_FTZ}, // RN
7407 {NVPTXISD::SUB_RP, NVPTXISD::SUB_RP_FTZ}, // RP
7408 {NVPTXISD::SUB_RM, NVPTXISD::SUB_RM_FTZ}}; // RM
7409 return SubF32x2Opcodes[static_cast<unsigned>(RoundingMode)][IsFTZ];
7410 }
7411 default:
7412 return 0;
7413 }
7414}
7415
7416static SDValue combineFAddWithNeg(SDNode *N, SelectionDAG &DAG,
7417 Intrinsic::ID AddIntrinsicID,
7418 APFloat::roundingMode RoundingMode) {
7419 const EVT VT = N->getValueType(ResNo: 0);
7420 const unsigned Opc = getFAddWithNegOpcode(VT, IID: AddIntrinsicID, RoundingMode);
7421 if (!Opc)
7422 return SDValue();
7423
7424 SDValue Op1 = N->getOperand(Num: 1);
7425 SDValue Op2 = N->getOperand(Num: 2);
7426
7427 SDValue SubOp1, SubOp2;
7428
7429 if (Op1.getOpcode() == ISD::FNEG) {
7430 SubOp1 = Op2;
7431 SubOp2 = Op1.getOperand(i: 0);
7432 } else if (Op2.getOpcode() == ISD::FNEG) {
7433 SubOp1 = Op1;
7434 SubOp2 = Op2.getOperand(i: 0);
7435 } else {
7436 return SDValue();
7437 }
7438
7439 return DAG.getNode(Opcode: Opc, DL: SDLoc(N), VT, N1: SubOp1, N2: SubOp2);
7440}
7441
7442// TODO: Remove the type-legality checks here once
7443// https://github.com/llvm/llvm-project/pull/172442 lands, adding support for
7444// explicit type constraints for overloaded intrinsics in tablegen.
7445static bool isSupportedFPArith(SDNode *N, const NVPTXSubtarget &STI,
7446 unsigned ISDOpcode, Intrinsic::ID IID,
7447 APFloat::roundingMode RoundingMode) {
7448 const EVT VT = N->getValueType(ResNo: 0);
7449 if (VT.isVector() && VT.getVectorElementCount() != ElementCount::getFixed(MinVal: 2))
7450 return false;
7451
7452 const bool IsRN = RoundingMode == APFloat::rmNearestTiesToEven;
7453 const bool IsFTZ = nvvm::FPArithShouldFTZ(IntrinsicID: IID);
7454 const bool IsSat = nvvm::FPArithIsSaturating(IntrinsicID: IID);
7455 switch (VT.getScalarType().getSimpleVT().SimpleTy) {
7456 case MVT::f16:
7457 return IsRN;
7458 case MVT::bf16:
7459 return IsRN && !IsSat && !IsFTZ && STI.hasNativeBF16Support(Opcode: ISDOpcode);
7460 case MVT::f32:
7461 return !VT.isVector() || (!IsSat && STI.hasF32x2Instructions());
7462 case MVT::f64:
7463 return !VT.isVector() && !IsSat && !IsFTZ;
7464 default:
7465 return false;
7466 }
7467}
7468
7469static SDValue diagnoseUnsupportedFPArith(SDNode *N, SelectionDAG &DAG,
7470 Intrinsic::ID IID,
7471 APFloat::roundingMode RoundingMode) {
7472 const EVT VT = N->getValueType(ResNo: 0);
7473 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
7474 DAG.getMachineFunction().getFunction(),
7475 Twine(Intrinsic::getBaseName(id: IID)) + " with rounding mode " +
7476 nvvm::GetRoundingModeName(RM: RoundingMode) + " and operand type " +
7477 VT.getEVTString() + " is not supported on this target",
7478 SDLoc(N).getDebugLoc()));
7479 return DAG.getPOISON(VT);
7480}
7481
7482static SDValue combineIntrinsicWOChain(SDNode *N,
7483 TargetLowering::DAGCombinerInfo &DCI,
7484 const NVPTXSubtarget &STI) {
7485 const Intrinsic::ID IID =
7486 static_cast<Intrinsic::ID>(N->getConstantOperandVal(Num: 0));
7487
7488 switch (IID) {
7489 default:
7490 break;
7491 case Intrinsic::nvvm_fadd:
7492 case Intrinsic::nvvm_fadd_ftz:
7493 case Intrinsic::nvvm_fadd_sat:
7494 case Intrinsic::nvvm_fadd_ftz_sat: {
7495 const auto RoundingMode = static_cast<APFloat::roundingMode>(
7496 N->getConstantOperandAPInt(Num: 3).getSExtValue());
7497 if (!isSupportedFPArith(N, STI, ISDOpcode: ISD::FADD, IID, RoundingMode))
7498 return diagnoseUnsupportedFPArith(N, DAG&: DCI.DAG, IID, RoundingMode);
7499 return combineFAddWithNeg(N, DAG&: DCI.DAG, AddIntrinsicID: IID, RoundingMode);
7500 }
7501 case Intrinsic::nvvm_spcompress:
7502 case Intrinsic::nvvm_spdecompress:
7503 return lowerSPIntrinsic(Op: SDValue(N, 0), DAG&: DCI.DAG);
7504 case Intrinsic::nvvm_fmul:
7505 case Intrinsic::nvvm_fmul_ftz:
7506 case Intrinsic::nvvm_fmul_sat:
7507 case Intrinsic::nvvm_fmul_ftz_sat: {
7508 const auto RoundingMode = static_cast<APFloat::roundingMode>(
7509 N->getConstantOperandAPInt(Num: 3).getSExtValue());
7510 if (!isSupportedFPArith(N, STI, ISDOpcode: ISD::FMUL, IID, RoundingMode))
7511 return diagnoseUnsupportedFPArith(N, DAG&: DCI.DAG, IID, RoundingMode);
7512 break;
7513 }
7514 }
7515 return SDValue();
7516}
7517
7518static SDValue combineProxyReg(SDNode *N,
7519 TargetLowering::DAGCombinerInfo &DCI) {
7520
7521 SDValue Chain = N->getOperand(Num: 0);
7522 SDValue Reg = N->getOperand(Num: 1);
7523
7524 // If the ProxyReg is not wrapping a load, try to pull the operations through
7525 // the ProxyReg.
7526 if (Reg.getOpcode() != ISD::LOAD) {
7527 if (SDValue V = sinkProxyReg(R: Reg, Chain, DCI))
7528 return V;
7529 }
7530
7531 return SDValue();
7532}
7533
7534SDValue NVPTXTargetLowering::PerformDAGCombine(SDNode *N,
7535 DAGCombinerInfo &DCI) const {
7536 CodeGenOptLevel OptLevel = getTargetMachine().getOptLevel();
7537 switch (N->getOpcode()) {
7538 default:
7539 break;
7540 case ISD::ADD:
7541 return PerformADDCombine(N, DCI, OptLevel);
7542 case ISD::ADDRSPACECAST:
7543 return combineADDRSPACECAST(N, DCI);
7544 case ISD::SIGN_EXTEND:
7545 case ISD::ZERO_EXTEND:
7546 return combineSZExtToMulWide(N, DCI, OptLevel);
7547 case ISD::BUILD_VECTOR:
7548 return PerformBUILD_VECTORCombine(N, DCI);
7549 case ISD::EXTRACT_VECTOR_ELT:
7550 return PerformEXTRACTCombine(N, DCI);
7551 case ISD::FADD:
7552 return performFADDCombine(N, DCI, OptLevel);
7553 case ISD::FMA:
7554 case ISD::FMUL:
7555 case ISD::FSUB:
7556 return performScalarizeV2F32Op(N, DCI, OptLevel);
7557 case ISD::FMAXNUM:
7558 case ISD::FMINNUM:
7559 case ISD::FMAXIMUM:
7560 case ISD::FMINIMUM:
7561 case ISD::FMAXIMUMNUM:
7562 case ISD::FMINIMUMNUM:
7563 return PerformFMinMaxCombine(N, DCI, STI);
7564 case ISD::LOAD:
7565 case NVPTXISD::LoadV2:
7566 case NVPTXISD::LoadV4:
7567 return combineLOAD(N, DCI, STI);
7568 case ISD::MUL:
7569 return PerformMULCombine(N, DCI, OptLevel);
7570 case NVPTXISD::PRMT:
7571 return combinePRMT(N, DCI, OptLevel);
7572 case NVPTXISD::ProxyReg:
7573 return combineProxyReg(N, DCI);
7574 case ISD::SETCC:
7575 return PerformSETCCCombine(N, DCI, STI);
7576 case ISD::SHL:
7577 return PerformSHLCombine(N, DCI, OptLevel);
7578 case ISD::STORE:
7579 case NVPTXISD::StoreV2:
7580 case NVPTXISD::StoreV4:
7581 return combineSTORE(N, DCI, STI);
7582 case ISD::SELECT:
7583 return PerformSELECTShiftCombine(N, DCI);
7584 case ISD::VSELECT:
7585 return PerformVSELECTCombine(N, DCI);
7586 case ISD::INTRINSIC_WO_CHAIN:
7587 return combineIntrinsicWOChain(N, DCI, STI);
7588 }
7589 return SDValue();
7590}
7591
7592static void ReplaceBITCAST(SDNode *Node, SelectionDAG &DAG,
7593 SmallVectorImpl<SDValue> &Results) {
7594 // Handle bitcasting to v2i8 without hitting the default promotion
7595 // strategy which goes through stack memory.
7596 SDValue Op(Node, 0);
7597 EVT ToVT = Op->getValueType(ResNo: 0);
7598 if (ToVT != MVT::v2i8) {
7599 return;
7600 }
7601
7602 // Bitcast to i16 and unpack elements into a vector
7603 SDLoc DL(Node);
7604 SDValue AsInt = DAG.getBitcast(VT: MVT::i16, V: Op->getOperand(Num: 0));
7605 SDValue Vec0 = DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: MVT::i8, Operand: AsInt);
7606 SDValue Const8 = DAG.getConstant(Val: 8, DL, VT: MVT::i16);
7607 SDValue Vec1 =
7608 DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: MVT::i8,
7609 Operand: DAG.getNode(Opcode: ISD::SRL, DL, VT: MVT::i16, Ops: {AsInt, Const8}));
7610 Results.push_back(
7611 Elt: DAG.getNode(Opcode: ISD::BUILD_VECTOR, DL, VT: MVT::v2i8, Ops: {Vec0, Vec1}));
7612}
7613
7614static void ReplaceINTRINSIC_W_CHAIN(SDNode *N, SelectionDAG &DAG,
7615 SmallVectorImpl<SDValue> &Results) {
7616 SDValue Chain = N->getOperand(Num: 0);
7617 SDValue Intrin = N->getOperand(Num: 1);
7618 SDLoc DL(N);
7619
7620 // Get the intrinsic ID
7621 unsigned IntrinNo = Intrin.getNode()->getAsZExtVal();
7622 switch (IntrinNo) {
7623 default:
7624 return;
7625 case Intrinsic::nvvm_ldu_global_i:
7626 case Intrinsic::nvvm_ldu_global_f:
7627 case Intrinsic::nvvm_ldu_global_p: {
7628 EVT ResVT = N->getValueType(ResNo: 0);
7629
7630 if (ResVT.isVector()) {
7631 // Vector LDG/LDU
7632
7633 unsigned NumElts = ResVT.getVectorNumElements();
7634 EVT EltVT = ResVT.getVectorElementType();
7635
7636 // Since LDU/LDG are target nodes, we cannot rely on DAG type
7637 // legalization.
7638 // Therefore, we must ensure the type is legal. For i1 and i8, we set the
7639 // loaded type to i16 and propagate the "real" type as the memory type.
7640 bool NeedTrunc = false;
7641 if (EltVT.getSizeInBits() < 16) {
7642 EltVT = MVT::i16;
7643 NeedTrunc = true;
7644 }
7645
7646 unsigned Opcode = 0;
7647 SDVTList LdResVTs;
7648
7649 switch (NumElts) {
7650 default:
7651 return;
7652 case 2:
7653 Opcode = NVPTXISD::LDUV2;
7654 LdResVTs = DAG.getVTList(VT1: EltVT, VT2: EltVT, VT3: MVT::Other);
7655 break;
7656 case 4: {
7657 Opcode = NVPTXISD::LDUV4;
7658 EVT ListVTs[] = { EltVT, EltVT, EltVT, EltVT, MVT::Other };
7659 LdResVTs = DAG.getVTList(VTs: ListVTs);
7660 break;
7661 }
7662 }
7663
7664 SmallVector<SDValue, 8> OtherOps;
7665
7666 // Copy regular operands
7667
7668 OtherOps.push_back(Elt: Chain); // Chain
7669 // Skip operand 1 (intrinsic ID)
7670 // Others
7671 OtherOps.append(in_start: N->op_begin() + 2, in_end: N->op_end());
7672
7673 MemIntrinsicSDNode *MemSD = cast<MemIntrinsicSDNode>(Val: N);
7674
7675 SDValue NewLD = DAG.getMemIntrinsicNode(Opcode, dl: DL, VTList: LdResVTs, Ops: OtherOps,
7676 MemVT: MemSD->getMemoryVT(),
7677 MMO: MemSD->getMemOperand());
7678
7679 SmallVector<SDValue, 4> ScalarRes;
7680
7681 for (unsigned i = 0; i < NumElts; ++i) {
7682 SDValue Res = NewLD.getValue(R: i);
7683 if (NeedTrunc)
7684 Res =
7685 DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: ResVT.getVectorElementType(), Operand: Res);
7686 ScalarRes.push_back(Elt: Res);
7687 }
7688
7689 SDValue LoadChain = NewLD.getValue(R: NumElts);
7690
7691 SDValue BuildVec =
7692 DAG.getBuildVector(VT: ResVT, DL, Ops: ScalarRes);
7693
7694 Results.push_back(Elt: BuildVec);
7695 Results.push_back(Elt: LoadChain);
7696 } else {
7697 // i8 LDG/LDU
7698 assert(ResVT.isSimple() && ResVT.getSimpleVT().SimpleTy == MVT::i8 &&
7699 "Custom handling of non-i8 ldu/ldg?");
7700
7701 // Just copy all operands as-is
7702 SmallVector<SDValue, 4> Ops(N->ops());
7703
7704 // Force output to i16
7705 SDVTList LdResVTs = DAG.getVTList(VT1: MVT::i16, VT2: MVT::Other);
7706
7707 MemIntrinsicSDNode *MemSD = cast<MemIntrinsicSDNode>(Val: N);
7708
7709 // We make sure the memory type is i8, which will be used during isel
7710 // to select the proper instruction.
7711 SDValue NewLD =
7712 DAG.getMemIntrinsicNode(Opcode: ISD::INTRINSIC_W_CHAIN, dl: DL, VTList: LdResVTs, Ops,
7713 MemVT: MVT::i8, MMO: MemSD->getMemOperand());
7714
7715 Results.push_back(Elt: DAG.getNode(Opcode: ISD::TRUNCATE, DL, VT: MVT::i8,
7716 Operand: NewLD.getValue(R: 0)));
7717 Results.push_back(Elt: NewLD.getValue(R: 1));
7718 }
7719 return;
7720 }
7721
7722 case Intrinsic::nvvm_tcgen05_ld_16x64b_x1:
7723 case Intrinsic::nvvm_tcgen05_ld_16x64b_x4:
7724 case Intrinsic::nvvm_tcgen05_ld_16x64b_x8:
7725 case Intrinsic::nvvm_tcgen05_ld_16x64b_x16:
7726 case Intrinsic::nvvm_tcgen05_ld_16x64b_x32:
7727 case Intrinsic::nvvm_tcgen05_ld_16x64b_x64:
7728 case Intrinsic::nvvm_tcgen05_ld_16x64b_x128:
7729 case Intrinsic::nvvm_tcgen05_ld_32x32b_x1:
7730 case Intrinsic::nvvm_tcgen05_ld_32x32b_x4:
7731 case Intrinsic::nvvm_tcgen05_ld_32x32b_x8:
7732 case Intrinsic::nvvm_tcgen05_ld_32x32b_x16:
7733 case Intrinsic::nvvm_tcgen05_ld_32x32b_x32:
7734 case Intrinsic::nvvm_tcgen05_ld_32x32b_x64:
7735 case Intrinsic::nvvm_tcgen05_ld_32x32b_x128:
7736 case Intrinsic::nvvm_tcgen05_ld_16x128b_x2:
7737 case Intrinsic::nvvm_tcgen05_ld_16x128b_x4:
7738 case Intrinsic::nvvm_tcgen05_ld_16x128b_x8:
7739 case Intrinsic::nvvm_tcgen05_ld_16x128b_x16:
7740 case Intrinsic::nvvm_tcgen05_ld_16x128b_x32:
7741 case Intrinsic::nvvm_tcgen05_ld_16x128b_x64:
7742 case Intrinsic::nvvm_tcgen05_ld_16x256b_x1:
7743 case Intrinsic::nvvm_tcgen05_ld_16x256b_x2:
7744 case Intrinsic::nvvm_tcgen05_ld_16x256b_x4:
7745 case Intrinsic::nvvm_tcgen05_ld_16x256b_x8:
7746 case Intrinsic::nvvm_tcgen05_ld_16x256b_x16:
7747 case Intrinsic::nvvm_tcgen05_ld_16x256b_x32:
7748 if (auto Res = lowerTcgen05Ld(N, DAG)) {
7749 Results.push_back(Elt: Res->first);
7750 Results.push_back(Elt: Res->second);
7751 }
7752 return;
7753
7754 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x1:
7755 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x4:
7756 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x8:
7757 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x16:
7758 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x32:
7759 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x64:
7760 case Intrinsic::nvvm_tcgen05_ld_16x32bx2_x128:
7761 if (auto Res = lowerTcgen05Ld(N, DAG, /*HasOffset=*/true)) {
7762 Results.push_back(Elt: Res->first);
7763 Results.push_back(Elt: Res->second);
7764 }
7765 return;
7766
7767 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x8_i32:
7768 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x8_f32:
7769 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x64_i32:
7770 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x64_f32:
7771 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x4_i32:
7772 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x4_f32:
7773 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x32_i32:
7774 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x32_f32:
7775 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x16_i32:
7776 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x16_f32:
7777 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x128_i32:
7778 case Intrinsic::nvvm_tcgen05_ld_red_32x32b_x128_f32:
7779 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x8_i32:
7780 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x8_f32:
7781 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x64_i32:
7782 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x64_f32:
7783 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x4_i32:
7784 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x4_f32:
7785 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x32_i32:
7786 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x32_f32:
7787 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x16_i32:
7788 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x16_f32:
7789 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x128_i32:
7790 case Intrinsic::nvvm_tcgen05_ld_red_16x32bx2_x128_f32:
7791 if (auto Res = lowerTcgen05LdRed(N, DAG)) {
7792 Results.push_back(Elt: std::get<0>(t&: *Res));
7793 Results.push_back(Elt: std::get<1>(t&: *Res));
7794 Results.push_back(Elt: std::get<2>(t&: *Res));
7795 }
7796 return;
7797 }
7798}
7799
7800static void ReplaceCopyFromReg_128(SDNode *N, SelectionDAG &DAG,
7801 SmallVectorImpl<SDValue> &Results) {
7802 // Change the CopyFromReg to output 2 64-bit results instead of a 128-bit
7803 // result so that it can pass the legalization
7804 SDLoc DL(N);
7805 SDValue Chain = N->getOperand(Num: 0);
7806 SDValue Reg = N->getOperand(Num: 1);
7807 SDValue Glue = N->getOperand(Num: 2);
7808
7809 assert(Reg.getValueType() == MVT::i128 &&
7810 "Custom lowering for CopyFromReg with 128-bit reg only");
7811 SmallVector<EVT, 4> ResultsType = {MVT::i64, MVT::i64, N->getValueType(ResNo: 1),
7812 N->getValueType(ResNo: 2)};
7813 SmallVector<SDValue, 3> NewOps = {Chain, Reg, Glue};
7814
7815 SDValue NewValue = DAG.getNode(Opcode: ISD::CopyFromReg, DL, ResultTys: ResultsType, Ops: NewOps);
7816 SDValue Pair = DAG.getNode(Opcode: ISD::BUILD_PAIR, DL, VT: MVT::i128,
7817 Ops: {NewValue.getValue(R: 0), NewValue.getValue(R: 1)});
7818
7819 Results.push_back(Elt: Pair);
7820 Results.push_back(Elt: NewValue.getValue(R: 2));
7821 Results.push_back(Elt: NewValue.getValue(R: 3));
7822}
7823
7824static void replaceProxyReg(SDNode *N, SelectionDAG &DAG,
7825 const TargetLowering &TLI,
7826 SmallVectorImpl<SDValue> &Results) {
7827 SDValue Chain = N->getOperand(Num: 0);
7828 SDValue Reg = N->getOperand(Num: 1);
7829
7830 MVT VT = TLI.getRegisterType(Context&: *DAG.getContext(), VT: Reg.getValueType());
7831
7832 SDValue NewReg = DAG.getAnyExtOrTrunc(Op: Reg, DL: SDLoc(N), VT);
7833 SDValue NewProxy =
7834 DAG.getNode(Opcode: NVPTXISD::ProxyReg, DL: SDLoc(N), VT, Ops: {Chain, NewReg});
7835 SDValue Res = DAG.getAnyExtOrTrunc(Op: NewProxy, DL: SDLoc(N), VT: N->getValueType(ResNo: 0));
7836
7837 Results.push_back(Elt: Res);
7838}
7839
7840static void replaceAtomicSwap128(SDNode *N, SelectionDAG &DAG,
7841 const NVPTXSubtarget &STI,
7842 SmallVectorImpl<SDValue> &Results) {
7843 assert(N->getValueType(0) == MVT::i128 &&
7844 "Custom lowering for atomic128 only supports i128");
7845
7846 AtomicSDNode *AN = cast<AtomicSDNode>(Val: N);
7847 SDLoc dl(N);
7848
7849 if (!STI.hasAtomSwap128()) {
7850 DAG.getContext()->diagnose(DI: DiagnosticInfoUnsupported(
7851 DAG.getMachineFunction().getFunction(),
7852 "Support for b128 atomics introduced in PTX ISA version 8.3 and "
7853 "requires target sm_90.",
7854 dl.getDebugLoc()));
7855
7856 Results.push_back(Elt: DAG.getUNDEF(VT: MVT::i128));
7857 Results.push_back(Elt: AN->getOperand(Num: 0)); // Chain
7858 return;
7859 }
7860
7861 SmallVector<SDValue, 6> Ops;
7862 Ops.push_back(Elt: AN->getOperand(Num: 0)); // Chain
7863 Ops.push_back(Elt: AN->getOperand(Num: 1)); // Ptr
7864 for (const auto &Op : AN->ops().drop_front(N: 2)) {
7865 // Low part
7866 Ops.push_back(Elt: DAG.getNode(Opcode: ISD::EXTRACT_ELEMENT, DL: dl, VT: MVT::i64, N1: Op,
7867 N2: DAG.getIntPtrConstant(Val: 0, DL: dl)));
7868 // High part
7869 Ops.push_back(Elt: DAG.getNode(Opcode: ISD::EXTRACT_ELEMENT, DL: dl, VT: MVT::i64, N1: Op,
7870 N2: DAG.getIntPtrConstant(Val: 1, DL: dl)));
7871 }
7872 unsigned Opcode = N->getOpcode() == ISD::ATOMIC_SWAP
7873 ? NVPTXISD::ATOMIC_SWAP_B128
7874 : NVPTXISD::ATOMIC_CMP_SWAP_B128;
7875 SDVTList Tys = DAG.getVTList(VT1: MVT::i64, VT2: MVT::i64, VT3: MVT::Other);
7876 SDValue Result = DAG.getMemIntrinsicNode(Opcode, dl, VTList: Tys, Ops, MemVT: MVT::i128,
7877 MMO: AN->getMemOperand());
7878 Results.push_back(Elt: DAG.getNode(Opcode: ISD::BUILD_PAIR, DL: dl, VT: MVT::i128,
7879 Ops: {Result.getValue(R: 0), Result.getValue(R: 1)}));
7880 Results.push_back(Elt: Result.getValue(R: 2));
7881}
7882
7883void NVPTXTargetLowering::ReplaceNodeResults(
7884 SDNode *N, SmallVectorImpl<SDValue> &Results, SelectionDAG &DAG) const {
7885 switch (N->getOpcode()) {
7886 default:
7887 report_fatal_error(reason: "Unhandled custom legalization");
7888 case ISD::BITCAST:
7889 ReplaceBITCAST(Node: N, DAG, Results);
7890 return;
7891 case ISD::LOAD:
7892 case ISD::MLOAD:
7893 replaceLoadVector(N, DAG, Results, STI);
7894 return;
7895 case ISD::INTRINSIC_W_CHAIN:
7896 ReplaceINTRINSIC_W_CHAIN(N, DAG, Results);
7897 return;
7898 case ISD::CopyFromReg:
7899 ReplaceCopyFromReg_128(N, DAG, Results);
7900 return;
7901 case NVPTXISD::ProxyReg:
7902 replaceProxyReg(N, DAG, TLI: *this, Results);
7903 return;
7904 case ISD::ATOMIC_CMP_SWAP:
7905 case ISD::ATOMIC_SWAP:
7906 replaceAtomicSwap128(N, DAG, STI, Results);
7907 return;
7908 }
7909}
7910
7911NVPTXTargetLowering::AtomicExpansionKind
7912NVPTXTargetLowering::shouldExpandAtomicRMWInIR(const AtomicRMWInst *AI) const {
7913 Type *Ty = AI->getValOperand()->getType();
7914
7915 // Try to lower LLVM atomicrmw fadd/fsub to PTX atomic.add. Fsub is first
7916 // expanded to an fadd with a negated operand. This is complicated by the
7917 // weird FTZ behavior PTX atom.add has:
7918 // - atom.add.f32 on global memory flushes denormals
7919 // - atom.add.f32 on shared memory does not flush denormals
7920 // - atom.add.f16 and atomic.add.bf16 never flush denormals
7921 //
7922 // We lower to atom.add only if the function's FTZ behavior matches that of
7923 // atom.add; otherwise, we lower to a CAS loop. But we always allow
7924 // atomic.add.bf16; even though it never flushes denormals, we never flush
7925 // bf16 denormals when doing regular arithmetic, even when FTZ is enabled.
7926 if (AI->isFloatingPointOperation() &&
7927 (AI->getOperation() == AtomicRMWInst::BinOp::FAdd ||
7928 AI->getOperation() == AtomicRMWInst::BinOp::FSub)) {
7929 const Function *F = AI->getFunction();
7930 AtomicExpansionKind ExpansionKind =
7931 AI->getOperation() == AtomicRMWInst::BinOp::FSub
7932 ? AtomicExpansionKind::Expand
7933 : AtomicExpansionKind::None;
7934
7935 // Both the -nvptx-allow-ftz-atomics option and per-instruction
7936 // !atomic.ignore.denormal.mode say that denormal handling is insignificant
7937 // here, so atom.add may be used even when its FTZ behavior disagrees with
7938 // the function's.
7939 const bool IgnoreFTZMismatch =
7940 AllowFTZAtomics ||
7941 AI->hasMetadata(KindID: LLVMContext::MD_atomic_ignore_denormal_mode);
7942
7943 if (Ty->isFloatTy()) {
7944 const bool FTZ = F->getDenormalMode(FPType: APFloat::IEEEsingle()).Output ==
7945 DenormalMode::PreserveSign;
7946 bool UseNative = IgnoreFTZMismatch;
7947 switch (AI->getPointerAddressSpace()) {
7948 case llvm::ADDRESS_SPACE_GLOBAL:
7949 UseNative |= FTZ;
7950 break;
7951 case llvm::ADDRESS_SPACE_SHARED:
7952 case llvm::ADDRESS_SPACE_SHARED_CLUSTER:
7953 UseNative |= !FTZ;
7954 break;
7955 }
7956 if (UseNative)
7957 return ExpansionKind;
7958 }
7959
7960 if (Ty->isHalfTy()) {
7961 // atom.add.f16 never flushes denormals, so it only agrees with a
7962 // function that is not in FTZ mode for f16.
7963 const bool FTZ = F->getDenormalMode(FPType: APFloat::IEEEhalf()).Output ==
7964 DenormalMode::PreserveSign;
7965 if ((!FTZ || IgnoreFTZMismatch) && STI.hasFeature(Feature: NVPTX::SM70) &&
7966 STI.hasFeature(Feature: NVPTX::PTX63))
7967 return ExpansionKind;
7968 }
7969
7970 if (Ty->isBFloatTy() && STI.hasFeature(Feature: NVPTX::SM90))
7971 return ExpansionKind;
7972
7973 if (Ty->isDoubleTy() && STI.hasAtomAddF64())
7974 return ExpansionKind;
7975 }
7976
7977 // PTX's only atomic fp op is `add`; all other ops expand to a CAS loop.
7978 if (AI->isFloatingPointOperation())
7979 return AtomicExpansionKind::CmpXChg;
7980
7981 if (Ty->isVectorTy())
7982 return AtomicExpansionKind::CmpXChg;
7983
7984 assert(Ty->isIntegerTy() && "Ty should be integer at this point");
7985 const unsigned BitWidth = cast<IntegerType>(Val: Ty)->getBitWidth();
7986
7987 switch (AI->getOperation()) {
7988 default:
7989 return AtomicExpansionKind::CmpXChg;
7990 case AtomicRMWInst::BinOp::Xchg:
7991 if (BitWidth == 128)
7992 return AtomicExpansionKind::None;
7993 [[fallthrough]];
7994 case AtomicRMWInst::BinOp::Add:
7995 case AtomicRMWInst::BinOp::Sub: {
7996 AtomicExpansionKind ExpansionKind =
7997 AI->getOperation() == AtomicRMWInst::BinOp::Sub
7998 ? AtomicExpansionKind::Expand
7999 : AtomicExpansionKind::None;
8000 switch (BitWidth) {
8001 case 8:
8002 case 16:
8003 return AtomicExpansionKind::CmpXChg;
8004 case 32:
8005 case 64:
8006 return ExpansionKind;
8007 case 128:
8008 return AtomicExpansionKind::CmpXChg;
8009 default:
8010 llvm_unreachable("unsupported width encountered");
8011 }
8012 }
8013 case AtomicRMWInst::BinOp::And:
8014 case AtomicRMWInst::BinOp::Or:
8015 case AtomicRMWInst::BinOp::Xor:
8016 case AtomicRMWInst::BinOp::Max:
8017 case AtomicRMWInst::BinOp::Min:
8018 case AtomicRMWInst::BinOp::UMax:
8019 case AtomicRMWInst::BinOp::UMin:
8020 switch (BitWidth) {
8021 case 8:
8022 case 16:
8023 return AtomicExpansionKind::CmpXChg;
8024 case 32:
8025 return AtomicExpansionKind::None;
8026 case 64:
8027 if (STI.hasAtomMinMaxAndOrXor())
8028 return AtomicExpansionKind::None;
8029 return AtomicExpansionKind::CmpXChg;
8030 case 128:
8031 return AtomicExpansionKind::CmpXChg;
8032 default:
8033 llvm_unreachable("unsupported width encountered");
8034 }
8035 case AtomicRMWInst::BinOp::UIncWrap:
8036 case AtomicRMWInst::BinOp::UDecWrap:
8037 switch (BitWidth) {
8038 case 32:
8039 return AtomicExpansionKind::None;
8040 case 8:
8041 case 16:
8042 case 64:
8043 case 128:
8044 return AtomicExpansionKind::CmpXChg;
8045 default:
8046 llvm_unreachable("unsupported width encountered");
8047 }
8048 }
8049
8050 return AtomicExpansionKind::CmpXChg;
8051}
8052
8053bool NVPTXTargetLowering::shouldInsertFencesForAtomic(
8054 const Instruction *I) const {
8055 // This function returns true iff the operation is emulated using a CAS-loop,
8056 // the target does not support memory-order qualifiers, or the operation has
8057 // the memory order seq_cst (which is not natively supported in the PTX
8058 // `atom` instruction).
8059 //
8060 // atomicrmw and cmpxchg instructions not efficiently supported by PTX
8061 // are lowered to CAS emulation loops that preserve their memory order,
8062 // syncscope, and volatile semantics. For PTX, it is more efficient to use
8063 // atom.cas.relaxed.sco instructions within the loop, and fences before and
8064 // after the loop to restore order.
8065 //
8066 // On targets with memory-order qualifiers, atomic instructions efficiently
8067 // supported by PTX are lowered to `atom.<op>.<sem>.<scope>` instructions with
8068 // their corresponding memory order and scope. Since PTX does not support
8069 // seq_cst, we emulate it by lowering to a fence.sc followed by an atom
8070 // according to the PTX atomics ABI.
8071 // https://docs.nvidia.com/cuda/ptx-writers-guide-to-interoperability/atomic-abi.html
8072 if (auto *CI = dyn_cast<AtomicCmpXchgInst>(Val: I))
8073 return (cast<IntegerType>(Val: CI->getCompareOperand()->getType())
8074 ->getBitWidth() < STI.getMinCmpXchgSizeInBits()) ||
8075 !STI.hasMemoryOrdering() ||
8076 CI->getMergedOrdering() == AtomicOrdering::SequentiallyConsistent;
8077 if (auto *RI = dyn_cast<AtomicRMWInst>(Val: I))
8078 return shouldExpandAtomicRMWInIR(AI: RI) == AtomicExpansionKind::CmpXChg ||
8079 !STI.hasMemoryOrdering() ||
8080 RI->getOrdering() == AtomicOrdering::SequentiallyConsistent;
8081 return false;
8082}
8083
8084AtomicOrdering NVPTXTargetLowering::atomicOperationOrderAfterFenceSplit(
8085 const Instruction *I) const {
8086 // If the operation is emulated by a CAS-loop, or the target does not support
8087 // memory-order qualifiers, we set its IR ordering to monotonic.
8088 // AtomicExpandPass inserts fences to enforce the original memory ordering.
8089 //
8090 // When the operation is not emulated, but the memory order is seq_cst,
8091 // we must lower to "fence.sc.<scope>; atom.<op>.acquire.<scope>;" to conform
8092 // to the PTX atomics ABI.
8093 // https://docs.nvidia.com/cuda/ptx-writers-guide-to-interoperability/atomic-abi.html
8094 // For such cases, emitLeadingFence() will separately insert the leading
8095 // "fence.sc.<scope>;". Here, we only set the memory order to acquire.
8096 //
8097 // Otherwise, the operation is not emulated, and the memory order is not
8098 // seq_cst. In this case, the LLVM memory order is natively supported by the
8099 // PTX `atom` instruction, and we just lower to the corresponding
8100 // `atom.<op>.relaxed|acquire|release|acq_rel". For such cases, this function
8101 // will NOT be called.
8102 // prerequisite: shouldInsertFencesForAtomic() should have returned `true` for
8103 // I before its memory order was modified.
8104 if (!STI.hasMemoryOrdering())
8105 return AtomicOrdering::Monotonic;
8106
8107 if (auto *CI = dyn_cast<AtomicCmpXchgInst>(Val: I);
8108 CI && CI->getMergedOrdering() == AtomicOrdering::SequentiallyConsistent &&
8109 cast<IntegerType>(Val: CI->getCompareOperand()->getType())->getBitWidth() >=
8110 STI.getMinCmpXchgSizeInBits())
8111 return AtomicOrdering::Acquire;
8112 else if (auto *RI = dyn_cast<AtomicRMWInst>(Val: I);
8113 RI && RI->getOrdering() == AtomicOrdering::SequentiallyConsistent) {
8114 AtomicExpansionKind ExpansionKind = shouldExpandAtomicRMWInIR(AI: RI);
8115 if (ExpansionKind == AtomicExpansionKind::None ||
8116 ExpansionKind == AtomicExpansionKind::Expand)
8117 return AtomicOrdering::Acquire;
8118 }
8119
8120 return AtomicOrdering::Monotonic;
8121}
8122
8123Instruction *NVPTXTargetLowering::emitLeadingFence(IRBuilderBase &Builder,
8124 Instruction *Inst,
8125 AtomicOrdering Ord) const {
8126 // prerequisite: shouldInsertFencesForAtomic() should have returned `true` for
8127 // `Inst` before its memory order was modified. We cannot enforce this with an
8128 // assert, because AtomicExpandPass will have modified the memory order
8129 // between the initial call to shouldInsertFencesForAtomic() and the call to
8130 // this function.
8131 if (!isa<AtomicCmpXchgInst>(Val: Inst) && !isa<AtomicRMWInst>(Val: Inst))
8132 return TargetLoweringBase::emitLeadingFence(Builder, Inst, Ord);
8133
8134 // Specialize for cmpxchg and atomicrmw
8135 auto SSID = getAtomicSyncScopeID(I: Inst);
8136 assert(SSID.has_value() && "Expected an atomic operation");
8137
8138 if (isReleaseOrStronger(AO: Ord))
8139 return Builder.CreateFence(Ordering: Ord == AtomicOrdering::SequentiallyConsistent
8140 ? AtomicOrdering::SequentiallyConsistent
8141 : AtomicOrdering::Release,
8142 SSID: SSID.value());
8143
8144 return nullptr;
8145}
8146
8147Instruction *NVPTXTargetLowering::emitTrailingFence(IRBuilderBase &Builder,
8148 Instruction *Inst,
8149 AtomicOrdering Ord) const {
8150 // prerequisite: shouldInsertFencesForAtomic() should have returned `true` for
8151 // `Inst` before its memory order was modified. See `emitLeadingFence` for why
8152 // this cannot be enforced with an assert. Specialize for cmpxchg and
8153 // atomicrmw
8154 auto *CI = dyn_cast<AtomicCmpXchgInst>(Val: Inst);
8155 auto *RI = dyn_cast<AtomicRMWInst>(Val: Inst);
8156 if (!CI && !RI)
8157 return TargetLoweringBase::emitTrailingFence(Builder, Inst, Ord);
8158
8159 auto SSID = getAtomicSyncScopeID(I: Inst);
8160 assert(SSID.has_value() && "Expected an atomic operation");
8161
8162 bool IsEmulated =
8163 !STI.hasMemoryOrdering() ||
8164 (CI ? cast<IntegerType>(Val: CI->getCompareOperand()->getType())
8165 ->getBitWidth() < STI.getMinCmpXchgSizeInBits()
8166 : shouldExpandAtomicRMWInIR(AI: RI) == AtomicExpansionKind::CmpXChg);
8167
8168 if (isAcquireOrStronger(AO: Ord) && IsEmulated)
8169 return Builder.CreateFence(Ordering: AtomicOrdering::Acquire, SSID: SSID.value());
8170
8171 return nullptr;
8172}
8173
8174// Rather than default to SINT when both UINT and SINT are custom, we only
8175// change the opcode when UINT is not legal and SINT is. UINT is preferred when
8176// both are custom since unsigned CVT instructions can lead to slightly better
8177// SASS code with fewer instructions.
8178unsigned NVPTXTargetLowering::getPreferredFPToIntOpcode(unsigned Op, EVT FromVT,
8179 EVT ToVT) const {
8180 if (isOperationLegal(Op, VT: ToVT))
8181 return Op;
8182 switch (Op) {
8183 case ISD::FP_TO_UINT:
8184 if (isOperationLegal(Op: ISD::FP_TO_SINT, VT: ToVT))
8185 return ISD::FP_TO_SINT;
8186 break;
8187 case ISD::STRICT_FP_TO_UINT:
8188 if (isOperationLegal(Op: ISD::STRICT_FP_TO_SINT, VT: ToVT))
8189 return ISD::STRICT_FP_TO_SINT;
8190 break;
8191 default:
8192 break;
8193 }
8194 return Op;
8195}
8196
8197// Pin NVPTXTargetObjectFile's vtables to this file.
8198NVPTXTargetObjectFile::~NVPTXTargetObjectFile() = default;
8199
8200MCSection *NVPTXTargetObjectFile::SelectSectionForGlobal(
8201 const GlobalObject *GO, SectionKind Kind, const TargetMachine &TM) const {
8202 return getDataSection();
8203}
8204
8205static void computeKnownBitsForPRMT(const SDValue Op, KnownBits &Known,
8206 const SelectionDAG &DAG, unsigned Depth) {
8207 SDValue A = Op.getOperand(i: 0);
8208 SDValue B = Op.getOperand(i: 1);
8209 ConstantSDNode *Selector = dyn_cast<ConstantSDNode>(Val: Op.getOperand(i: 2));
8210 unsigned Mode = Op.getConstantOperandVal(i: 3);
8211
8212 if (!Selector)
8213 return;
8214
8215 KnownBits AKnown = DAG.computeKnownBits(Op: A, Depth);
8216 KnownBits BKnown = DAG.computeKnownBits(Op: B, Depth);
8217
8218 // {b, a} = {{b7, b6, b5, b4}, {b3, b2, b1, b0}}
8219 assert(AKnown.getBitWidth() == 32 && BKnown.getBitWidth() == 32 &&
8220 "PRMT must have i32 operands");
8221 assert(Known.getBitWidth() == 32 && "PRMT must have i32 result");
8222 KnownBits BitField = BKnown.concat(Lo: AKnown);
8223
8224 APInt SelectorVal = getPRMTSelector(Selector: Selector->getAPIntValue(), Mode);
8225 for (unsigned I : llvm::seq(Size: 4)) {
8226 APInt Sel = SelectorVal.extractBits(numBits: 4, bitPosition: I * 4);
8227 unsigned Idx = Sel.getLoBits(numBits: 3).getZExtValue();
8228 unsigned Sign = Sel.getHiBits(numBits: 1).getZExtValue();
8229 KnownBits Byte = BitField.extractBits(NumBits: 8, BitPosition: Idx * 8);
8230 if (Sign)
8231 Byte = KnownBits::ashr(LHS: Byte, RHS: KnownBits::makeConstant(C: APInt(8, 7)));
8232 Known.insertBits(SubBits: Byte, BitPosition: I * 8);
8233 }
8234}
8235
8236static void computeKnownBitsForLoadV(const SDValue Op, KnownBits &Known) {
8237 MemSDNode *LD = cast<MemSDNode>(Val: Op);
8238
8239 // We can't do anything without knowing the sign bit.
8240 auto ExtType = LD->getConstantOperandVal(Num: LD->getNumOperands() - 1);
8241 if (ExtType == ISD::SEXTLOAD)
8242 return;
8243
8244 // ExtLoading to vector types is weird and may not work well with known bits.
8245 auto DestVT = LD->getValueType(ResNo: 0);
8246 if (DestVT.isVector())
8247 return;
8248
8249 assert(Known.getBitWidth() == DestVT.getSizeInBits());
8250 auto ElementBitWidth = getFromTypeWidthForLoad(Mem: LD);
8251 Known.Zero.setHighBits(Known.getBitWidth() - ElementBitWidth);
8252}
8253
8254void NVPTXTargetLowering::computeKnownBitsForTargetNode(
8255 const SDValue Op, KnownBits &Known, const APInt &DemandedElts,
8256 const SelectionDAG &DAG, unsigned Depth) const {
8257 Known.resetAll();
8258
8259 switch (Op.getOpcode()) {
8260 case NVPTXISD::PRMT:
8261 computeKnownBitsForPRMT(Op, Known, DAG, Depth);
8262 break;
8263 case NVPTXISD::LoadV2:
8264 case NVPTXISD::LoadV4:
8265 case NVPTXISD::LoadV8:
8266 computeKnownBitsForLoadV(Op, Known);
8267 break;
8268 default:
8269 break;
8270 }
8271}
8272
8273static std::pair<APInt, APInt> getPRMTDemandedBits(const APInt &SelectorVal,
8274 const APInt &DemandedBits) {
8275 APInt DemandedLHS = APInt(32, 0);
8276 APInt DemandedRHS = APInt(32, 0);
8277
8278 for (unsigned I : llvm::seq(Size: 4)) {
8279 if (DemandedBits.extractBits(numBits: 8, bitPosition: I * 8).isZero())
8280 continue;
8281
8282 APInt Sel = SelectorVal.extractBits(numBits: 4, bitPosition: I * 4);
8283 unsigned Idx = Sel.getLoBits(numBits: 3).getZExtValue();
8284 unsigned Sign = Sel.getHiBits(numBits: 1).getZExtValue();
8285
8286 APInt &Src = Idx < 4 ? DemandedLHS : DemandedRHS;
8287 unsigned ByteStart = (Idx % 4) * 8;
8288 if (Sign)
8289 Src.setBit(ByteStart + 7);
8290 else
8291 Src.setBits(loBit: ByteStart, hiBit: ByteStart + 8);
8292 }
8293
8294 return {DemandedLHS, DemandedRHS};
8295}
8296
8297// Replace undef with 0 as this is easier for other optimizations such as
8298// known bits.
8299static SDValue canonicalizePRMTInput(SDValue Op, SelectionDAG &DAG) {
8300 if (!Op)
8301 return SDValue();
8302 if (Op.isUndef())
8303 return DAG.getConstant(Val: 0, DL: SDLoc(), VT: MVT::i32);
8304 return Op;
8305}
8306
8307static SDValue simplifyDemandedBitsForPRMT(SDValue PRMT,
8308 const APInt &DemandedBits,
8309 SelectionDAG &DAG,
8310 const TargetLowering &TLI,
8311 unsigned Depth) {
8312 assert(PRMT.getOpcode() == NVPTXISD::PRMT);
8313 SDValue Op0 = PRMT.getOperand(i: 0);
8314 SDValue Op1 = PRMT.getOperand(i: 1);
8315 auto *SelectorConst = dyn_cast<ConstantSDNode>(Val: PRMT.getOperand(i: 2));
8316 if (!SelectorConst)
8317 return SDValue();
8318
8319 unsigned Mode = PRMT.getConstantOperandVal(i: 3);
8320 const APInt Selector = getPRMTSelector(Selector: SelectorConst->getAPIntValue(), Mode);
8321
8322 // Try to simplify the PRMT to one of the inputs if the used bytes are all
8323 // from the same input in the correct order.
8324 const unsigned LeadingBytes = DemandedBits.countLeadingZeros() / 8;
8325 const unsigned SelBits = (4 - LeadingBytes) * 4;
8326 if (Selector.getLoBits(numBits: SelBits) == APInt(32, 0x3210).getLoBits(numBits: SelBits))
8327 return Op0;
8328 if (Selector.getLoBits(numBits: SelBits) == APInt(32, 0x7654).getLoBits(numBits: SelBits))
8329 return Op1;
8330
8331 auto [DemandedLHS, DemandedRHS] = getPRMTDemandedBits(SelectorVal: Selector, DemandedBits);
8332
8333 // Attempt to avoid multi-use ops if we don't need anything from them.
8334 SDValue DemandedOp0 =
8335 TLI.SimplifyMultipleUseDemandedBits(Op: Op0, DemandedBits: DemandedLHS, DAG, Depth: Depth + 1);
8336 SDValue DemandedOp1 =
8337 TLI.SimplifyMultipleUseDemandedBits(Op: Op1, DemandedBits: DemandedRHS, DAG, Depth: Depth + 1);
8338
8339 DemandedOp0 = canonicalizePRMTInput(Op: DemandedOp0, DAG);
8340 DemandedOp1 = canonicalizePRMTInput(Op: DemandedOp1, DAG);
8341 if ((DemandedOp0 && DemandedOp0 != Op0) ||
8342 (DemandedOp1 && DemandedOp1 != Op1)) {
8343 Op0 = DemandedOp0 ? DemandedOp0 : Op0;
8344 Op1 = DemandedOp1 ? DemandedOp1 : Op1;
8345 return getPRMT(A: Op0, B: Op1, Selector: Selector.getZExtValue(), DL: SDLoc(PRMT), DAG);
8346 }
8347
8348 return SDValue();
8349}
8350
8351bool NVPTXTargetLowering::SimplifyDemandedBitsForTargetNode(
8352 SDValue Op, const APInt &DemandedBits, const APInt &DemandedElts,
8353 KnownBits &Known, TargetLoweringOpt &TLO, unsigned Depth) const {
8354 Known.resetAll();
8355
8356 switch (Op.getOpcode()) {
8357 case NVPTXISD::PRMT:
8358 if (SDValue Result = simplifyDemandedBitsForPRMT(PRMT: Op, DemandedBits, DAG&: TLO.DAG,
8359 TLI: *this, Depth)) {
8360 TLO.CombineTo(O: Op, N: Result);
8361 return true;
8362 }
8363 break;
8364 default:
8365 break;
8366 }
8367
8368 computeKnownBitsForTargetNode(Op, Known, DemandedElts, DAG: TLO.DAG, Depth);
8369 return false;
8370}
8371