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