[llvm] 8035ae5 - [NVPTX] Cleanup and refactor param align computation, addressing a few minor bugs and discrepancies (#188588)
via llvm-commits
llvm-commits at lists.llvm.org
Wed Jun 10 08:25:07 PDT 2026
Author: Alex MacLean
Date: 2026-06-10T08:25:02-07:00
New Revision: 8035ae5f42959a0325d5438e63011587950f43a7
URL: https://github.com/llvm/llvm-project/commit/8035ae5f42959a0325d5438e63011587950f43a7
DIFF: https://github.com/llvm/llvm-project/commit/8035ae5f42959a0325d5438e63011587950f43a7.diff
LOG: [NVPTX] Cleanup and refactor param align computation, addressing a few minor bugs and discrepancies (#188588)
Added:
llvm/test/CodeGen/NVPTX/ret-align-mismatch.ll
Modified:
llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp
llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
llvm/lib/Target/NVPTX/NVPTXSetByValParamAlign.cpp
llvm/lib/Target/NVPTX/NVPTXUtilities.cpp
llvm/lib/Target/NVPTX/NVPTXUtilities.h
llvm/lib/Target/NVPTX/NVVMProperties.cpp
llvm/lib/Target/NVPTX/NVVMProperties.h
llvm/test/CodeGen/NVPTX/nvvm-annotations-D120129.ll
llvm/test/CodeGen/NVPTX/param-overalign.ll
Removed:
################################################################################
diff --git a/llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp b/llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp
index d7b34abb0127c..b2efcb0f0d2b6 100644
--- a/llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXAsmPrinter.cpp
@@ -300,7 +300,7 @@ void NVPTXAsmPrinter::printReturnValStr(const Function *F, raw_ostream &O) {
if (shouldPassAsArray(Ty)) {
const unsigned TotalSize = DL.getTypeAllocSize(Ty);
const Align RetAlignment =
- getFunctionArgumentAlignment(F, Ty, AttributeList::ReturnIndex, DL);
+ getPTXParamAlign(F, Ty, AttributeList::ReturnIndex, DL);
O << ".param .align " << RetAlignment.value() << " .b8 func_retval0["
<< TotalSize << "]";
} else if (Ty->isFloatingPointTy()) {
@@ -1406,17 +1406,6 @@ void NVPTXAsmPrinter::emitFunctionParamList(const Function *F, raw_ostream &O) {
}
}
- auto GetOptimalAlignForParam = [&DL, F, &Arg](Type *Ty) -> Align {
- if (MaybeAlign StackAlign =
- getAlign(*F, Arg.getArgNo() + AttributeList::FirstArgIndex))
- return StackAlign.value();
-
- Align TypeAlign = getFunctionParamOptimizedAlign(F, Ty, DL);
- MaybeAlign ParamAlign =
- Arg.hasByValAttr() ? Arg.getParamAlign() : MaybeAlign();
- return std::max(TypeAlign, ParamAlign.valueOrOne());
- };
-
if (Arg.hasByValAttr()) {
// param has byVal attribute.
Type *ETy = Arg.getParamByValType();
@@ -1427,9 +1416,11 @@ void NVPTXAsmPrinter::emitFunctionParamList(const Function *F, raw_ostream &O) {
// PAL.getParamAlignment
// size = typeallocsize of element type
const Align OptimalAlign =
- IsKernelFunc ? GetOptimalAlignForParam(ETy)
- : getFunctionByValParamAlign(
- F, ETy, Arg.getParamAlign().valueOrOne(), DL);
+ IsKernelFunc
+ ? getPTXParamAlign(
+ F, ETy, Arg.getArgNo() + AttributeList::FirstArgIndex, DL)
+ : getDeviceByValParamAlign(F, ETy,
+ Arg.getParamAlign().valueOrOne(), DL);
O << "\t.param .align " << OptimalAlign.value() << " .b8 " << ParamSym
<< "[" << DL.getTypeAllocSize(ETy) << "]";
@@ -1441,7 +1432,8 @@ void NVPTXAsmPrinter::emitFunctionParamList(const Function *F, raw_ostream &O) {
// <a> = optimal alignment for the element type; always multiple of
// PAL.getParamAlignment
// size = typeallocsize of element type
- Align OptimalAlign = GetOptimalAlignForParam(Ty);
+ Align OptimalAlign = getPTXParamAlign(
+ F, Ty, Arg.getArgNo() + AttributeList::FirstArgIndex, DL);
O << "\t.param .align " << OptimalAlign.value() << " .b8 " << ParamSym
<< "[" << DL.getTypeAllocSize(Ty) << "]";
diff --git a/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp b/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
index e13911f87eed5..aada012ff2bd9 100644
--- a/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
@@ -1203,9 +1203,6 @@ SDValue NVPTXTargetLowering::getSqrtEstimate(SDValue Operand, SelectionDAG &DAG,
}
}
-static Align getArgumentAlignment(const CallBase *CB, Type *Ty, unsigned Idx,
- const DataLayout &DL);
-
std::string NVPTXTargetLowering::getPrototype(
const DataLayout &DL, Type *RetTy, const ArgListTy &Args,
const SmallVectorImpl<ISD::OutputArg> &Outs,
@@ -1222,7 +1219,8 @@ std::string NVPTXTargetLowering::getPrototype(
} else {
O << "(";
if (shouldPassAsArray(RetTy)) {
- const Align RetAlign = getArgumentAlignment(&CB, RetTy, 0, DL);
+ const Align RetAlign =
+ getPTXParamAlign(&CB, RetTy, AttributeList::ReturnIndex, DL);
O << ".param .align " << RetAlign.value() << " .b8 _["
<< DL.getTypeAllocSize(RetTy) << "]";
} else if (RetTy->isFloatingPointTy() || RetTy->isIntegerTy()) {
@@ -1270,14 +1268,14 @@ std::string NVPTXTargetLowering::getPrototype(
Type *ETy = Args[I].IndirectType;
Align InitialAlign = ArgOuts[0].Flags.getNonZeroByValAlign();
Align ParamByValAlign =
- getFunctionByValParamAlign(/*F=*/nullptr, ETy, InitialAlign, DL);
+ getDeviceByValParamAlign(/*F=*/nullptr, ETy, InitialAlign, DL);
O << ".param .align " << ParamByValAlign.value() << " .b8 _["
<< ArgOuts[0].Flags.getByValSize() << "]";
} else {
if (shouldPassAsArray(Ty)) {
Align ParamAlign =
- getArgumentAlignment(&CB, Ty, I + AttributeList::FirstArgIndex, DL);
+ getPTXParamAlign(&CB, Ty, I + AttributeList::FirstArgIndex, DL);
O << ".param .align " << ParamAlign.value() << " .b8 _["
<< DL.getTypeAllocSize(Ty) << "]";
continue;
@@ -1310,37 +1308,6 @@ std::string NVPTXTargetLowering::getPrototype(
return Prototype;
}
-static Align getArgumentAlignment(const CallBase *CB, Type *Ty, unsigned Idx,
- const DataLayout &DL) {
- if (!CB) {
- // CallSite is zero, fallback to ABI type alignment
- return DL.getABITypeAlign(Ty);
- }
-
- const Function *DirectCallee = CB->getCalledFunction();
-
- if (!DirectCallee) {
- // We don't have a direct function symbol, but that may be because of
- // constant cast instructions in the call.
-
- // With bitcast'd call targets, the instruction will be the call
- if (const auto *CI = dyn_cast<CallInst>(CB)) {
- // Check if we have call alignment metadata
- if (MaybeAlign StackAlign = getAlign(*CI, Idx))
- return StackAlign.value();
- }
- DirectCallee = getMaybeBitcastedCallee(CB);
- }
-
- // Check for function alignment information if we found that the
- // ultimate target is a Function
- if (DirectCallee)
- return getFunctionArgumentAlignment(DirectCallee, Ty, Idx, DL);
-
- // Call is indirect, fall back to the ABI type alignment
- return DL.getABITypeAlign(Ty);
-}
-
static MachinePointerInfo refinePtrAS(SDValue &Ptr, SelectionDAG &DAG,
const DataLayout &DL,
const TargetLowering &TL) {
@@ -1507,10 +1474,11 @@ SDValue NVPTXTargetLowering::LowerCall(TargetLowering::CallLoweringInfo &CLI,
// so we don't need to worry whether it's naturally aligned or not.
// See TargetLowering::LowerCallTo().
const Align InitialAlign = ArgOuts[0].Flags.getNonZeroByValAlign();
- return getFunctionByValParamAlign(CB->getCalledFunction(), ETy,
- InitialAlign, DL);
+ return getDeviceByValParamAlign(CB->getCalledFunction(), ETy,
+ InitialAlign, DL);
}
- return getArgumentAlignment(CB, Arg.Ty, ArgI + 1, DL);
+ return getPTXParamAlign(CB, Arg.Ty, ArgI + AttributeList::FirstArgIndex,
+ DL);
}();
const unsigned TySize = DL.getTypeAllocSize(ETy);
@@ -1644,7 +1612,8 @@ SDValue NVPTXTargetLowering::LowerCall(TargetLowering::CallLoweringInfo &CLI,
const SDValue RetSymbol = DAG.getExternalSymbol("retval0", MVT::i32);
const unsigned ResultSize = DL.getTypeAllocSize(RetTy);
if (shouldPassAsArray(RetTy)) {
- const Align RetAlign = getArgumentAlignment(CB, RetTy, 0, DL);
+ const Align RetAlign =
+ getPTXParamAlign(CB, RetTy, AttributeList::ReturnIndex, DL);
MakeDeclareArrayParam(RetSymbol, RetAlign, ResultSize);
} else {
MakeDeclareScalarParam(RetSymbol, ResultSize);
@@ -1737,7 +1706,8 @@ SDValue NVPTXTargetLowering::LowerCall(TargetLowering::CallLoweringInfo &CLI,
ComputePTXValueVTs(*this, DL, Ctx, CLI.CallConv, RetTy, VTs, Offsets);
assert(VTs.size() == Ins.size() && "Bad value decomposition");
- const Align RetAlign = getArgumentAlignment(CB, RetTy, 0, DL);
+ const Align RetAlign =
+ getPTXParamAlign(CB, RetTy, AttributeList::ReturnIndex, DL);
const SDValue RetSymbol = DAG.getExternalSymbol("retval0", MVT::i32);
// PTX Interoperability Guide 3.3(A): [Integer] Values shorter than
@@ -4168,7 +4138,7 @@ SDValue NVPTXTargetLowering::LowerFormalArguments(
assert(VTs.size() == ArgIns.size() && "Size mismatch");
assert(VTs.size() == Offsets.size() && "Size mismatch");
- const Align ArgAlign = getFunctionArgumentAlignment(
+ const Align ArgAlign = getPTXParamAlign(
&F, Ty, Arg.getArgNo() + AttributeList::FirstArgIndex, DL);
unsigned I = 0;
@@ -4225,7 +4195,8 @@ NVPTXTargetLowering::LowerReturn(SDValue Chain, CallingConv::ID CallConv,
LLVMContext &Ctx = *DAG.getContext();
const SDValue RetSymbol = DAG.getExternalSymbol("func_retval0", MVT::i32);
- const auto RetAlign = getFunctionParamOptimizedAlign(&F, RetTy, DL);
+ const auto RetAlign =
+ getPTXParamAlign(&F, RetTy, AttributeList::ReturnIndex, DL);
// PTX Interoperability Guide 3.3(A): [Integer] Values shorter than
// 32-bits are sign extended or zero extended, depending on whether
diff --git a/llvm/lib/Target/NVPTX/NVPTXSetByValParamAlign.cpp b/llvm/lib/Target/NVPTX/NVPTXSetByValParamAlign.cpp
index 214078c1967ab..bd5cbfb4ac6b6 100644
--- a/llvm/lib/Target/NVPTX/NVPTXSetByValParamAlign.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXSetByValParamAlign.cpp
@@ -64,7 +64,7 @@ static Align setByValParamAlign(Argument *Arg) {
Type *ByValType = Arg->getParamByValType();
const DataLayout &DL = F->getDataLayout();
- const Align OptimizedAlign = getFunctionParamOptimizedAlign(F, ByValType, DL);
+ const Align OptimizedAlign = getPTXPromotedParamTypeAlign(F, ByValType, DL);
const Align CurrentAlign = Arg->getParamAlign().valueOrOne();
if (CurrentAlign >= OptimizedAlign)
diff --git a/llvm/lib/Target/NVPTX/NVPTXUtilities.cpp b/llvm/lib/Target/NVPTX/NVPTXUtilities.cpp
index 6cf808bc9c858..2cae67ea0d3c2 100644
--- a/llvm/lib/Target/NVPTX/NVPTXUtilities.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXUtilities.cpp
@@ -14,6 +14,7 @@
#include "NVPTX.h"
#include "NVPTXTargetMachine.h"
#include "NVVMProperties.h"
+#include "llvm/IR/Attributes.h"
#include "llvm/IR/DataLayout.h"
#include "llvm/IR/Function.h"
#include "llvm/Support/Alignment.h"
@@ -32,8 +33,8 @@ Function *getMaybeBitcastedCallee(const CallBase *CB) {
return dyn_cast<Function>(CB->getCalledOperand()->stripPointerCasts());
}
-Align getFunctionParamOptimizedAlign(const Function *F, Type *ArgTy,
- const DataLayout &DL) {
+Align getPTXPromotedParamTypeAlign(const Function *F, Type *ArgTy,
+ const DataLayout &DL) {
// Capping the alignment to 128 bytes as that is the maximum alignment
// supported by PTX.
const Align ABITypeAlign = std::min(Align(128), DL.getABITypeAlign(ArgTy));
@@ -41,27 +42,21 @@ Align getFunctionParamOptimizedAlign(const Function *F, Type *ArgTy,
// If a function has linkage
diff erent from internal or private, we
// must use default ABI alignment as external users rely on it. Same
// for a function that may be called from a function pointer.
- if (!F || !F->hasLocalLinkage() ||
- F->hasAddressTaken(/*Users=*/nullptr,
- /*IgnoreCallbackUses=*/false,
- /*IgnoreAssumeLikeCalls=*/true,
- /*IgnoreLLVMUsed=*/true))
- return ABITypeAlign;
-
- assert(!isKernelFunction(*F) && "Expect kernels to have non-local linkage");
- return std::max(Align(16), ABITypeAlign);
+ const bool MayOptimizeAlign =
+ F && F->hasLocalLinkage() &&
+ !F->hasAddressTaken(/*Users=*/nullptr,
+ /*IgnoreCallbackUses=*/false,
+ /*IgnoreAssumeLikeCalls=*/true,
+ /*IgnoreLLVMUsed=*/true);
+ assert(!(MayOptimizeAlign && isKernelFunction(*F)) &&
+ "Expect kernels to have non-local linkage");
+ const Align OptimizedAlign = MayOptimizeAlign ? Align(16) : Align(1);
+ return std::max(OptimizedAlign, ABITypeAlign);
}
-Align getFunctionArgumentAlignment(const Function *F, Type *Ty, unsigned Idx,
- const DataLayout &DL) {
- return getAlign(*F, Idx).value_or(getFunctionParamOptimizedAlign(F, Ty, DL));
-}
-
-Align getFunctionByValParamAlign(const Function *F, Type *ArgTy,
- Align InitialAlign, const DataLayout &DL) {
- Align ArgAlign = InitialAlign;
- if (F)
- ArgAlign = std::max(ArgAlign, getFunctionParamOptimizedAlign(F, ArgTy, DL));
+Align getDeviceByValParamAlign(const Function *F, Type *ArgTy,
+ Align InitialAlign, const DataLayout &DL) {
+ const Align OptimizedAlign = getPTXPromotedParamTypeAlign(F, ArgTy, DL);
// Old ptx versions have a bug. When PTX code takes address of
// byval parameter with alignment < 4, ptxas generates code to
@@ -72,10 +67,40 @@ Align getFunctionByValParamAlign(const Function *F, Type *ArgTy,
// ptxas > 9.0.
// TODO: remove this after verifying the bug is not reproduced
// on non-deprecated ptxas versions.
- if (ForceMinByValParamAlign)
- ArgAlign = std::max(ArgAlign, Align(4));
+ const bool ShouldForceMinAlign =
+ ForceMinByValParamAlign && (!F || !isKernelFunction(*F));
+ const Align AlignFloor = ShouldForceMinAlign ? Align(4) : Align(1);
+
+ return std::max({InitialAlign, OptimizedAlign, AlignFloor});
+}
+
+Align getPTXParamAlign(const Function *F, Type *Ty, unsigned AttrIdx,
+ const DataLayout &DL) {
+ if (F)
+ if (MaybeAlign StackAlign = getStackAlign(*F, AttrIdx))
+ return StackAlign.value();
+
+ Align TypeAlign = getPTXPromotedParamTypeAlign(F, Ty, DL);
+ if (F && AttrIdx >= AttributeList::FirstArgIndex) {
+ unsigned ArgNo = AttrIdx - AttributeList::FirstArgIndex;
+ if (F->getAttributes().hasParamAttr(ArgNo, Attribute::ByVal))
+ return std::max(TypeAlign, F->getParamAlign(ArgNo).valueOrOne());
+ }
+ return TypeAlign;
+}
+
+Align getPTXParamAlign(const CallBase *CB, Type *Ty, unsigned Idx,
+ const DataLayout &DL) {
+ const Function *DirectCallee = CB ? CB->getCalledFunction() : nullptr;
+
+ if (!DirectCallee && CB) {
+ if (MaybeAlign StackAlign = getStackAlign(*CB, Idx))
+ return StackAlign.value();
+
+ DirectCallee = getMaybeBitcastedCallee(CB);
+ }
- return ArgAlign;
+ return getPTXParamAlign(DirectCallee, Ty, Idx, DL);
}
bool shouldEmitPTXNoReturn(const Value *V, const TargetMachine &TM) {
diff --git a/llvm/lib/Target/NVPTX/NVPTXUtilities.h b/llvm/lib/Target/NVPTX/NVPTXUtilities.h
index b6e18bf998897..6785883e5af1e 100644
--- a/llvm/lib/Target/NVPTX/NVPTXUtilities.h
+++ b/llvm/lib/Target/NVPTX/NVPTXUtilities.h
@@ -38,14 +38,26 @@ Function *getMaybeBitcastedCallee(const CallBase *CB);
/// function has internal or private linkage as for other linkage types callers
/// may already rely on default alignment. To allow using 128-bit vectorized
/// loads/stores, this function ensures that alignment is 16 or greater.
-Align getFunctionParamOptimizedAlign(const Function *F, Type *ArgTy,
- const DataLayout &DL);
-
-Align getFunctionArgumentAlignment(const Function *F, Type *Ty, unsigned Idx,
+Align getPTXPromotedParamTypeAlign(const Function *F, Type *ArgTy,
const DataLayout &DL);
-Align getFunctionByValParamAlign(const Function *F, Type *ArgTy,
- Align InitialAlign, const DataLayout &DL);
+Align getDeviceByValParamAlign(const Function *F, Type *ArgTy,
+ Align InitialAlign, const DataLayout &DL);
+
+/// Get the alignment for a function parameter or return value.
+/// \p AttrIdx is the AttributeList index (e.g. FirstArgIndex + argNo, or
+/// ReturnIndex for return values). Checks for an explicit alignment attribute,
+/// then falls back to getPromotedParamTypeAlign, incorporating byval param
+/// alignment when applicable.
+Align getPTXParamAlign(const Function *F, Type *Ty, unsigned AttrIdx,
+ const DataLayout &DL);
+
+/// Get the alignment for a call-site argument or return value. Resolves the
+/// callee and delegates to the Function overload of getParamAlign. For
+/// indirect calls with no resolvable callee, falls back to
+/// getPromotedParamTypeAlign.
+Align getPTXParamAlign(const CallBase *CB, Type *Ty, unsigned AttrIdx,
+ const DataLayout &DL);
// PTX ABI requires all scalar argument/return values to have
// bit-size as a power of two of at least 32 bits.
diff --git a/llvm/lib/Target/NVPTX/NVVMProperties.cpp b/llvm/lib/Target/NVPTX/NVVMProperties.cpp
index d68c5aaf4fe5f..012d0863c183d 100644
--- a/llvm/lib/Target/NVPTX/NVVMProperties.cpp
+++ b/llvm/lib/Target/NVPTX/NVVMProperties.cpp
@@ -317,7 +317,7 @@ bool isParamGridConstant(const Argument &Arg) {
return Arg.hasAttribute(NVVMAttr::GridConstant);
}
-MaybeAlign getAlign(const CallInst &I, unsigned Index) {
+MaybeAlign getStackAlign(const CallBase &I, unsigned Index) {
// First check the alignstack metadata.
if (MaybeAlign StackAlign =
I.getAttributes().getAttributes(Index).getStackAlignment())
diff --git a/llvm/lib/Target/NVPTX/NVVMProperties.h b/llvm/lib/Target/NVPTX/NVVMProperties.h
index 6ccd6f8a20075..7187b18d3cbf7 100644
--- a/llvm/lib/Target/NVPTX/NVVMProperties.h
+++ b/llvm/lib/Target/NVPTX/NVVMProperties.h
@@ -24,7 +24,7 @@
namespace llvm {
class Argument;
-class CallInst;
+class CallBase;
class GlobalVariable;
class Module;
class Value;
@@ -59,10 +59,10 @@ bool hasBlocksAreClusters(const Function &);
bool isParamGridConstant(const Argument &);
-inline MaybeAlign getAlign(const Function &F, unsigned Index) {
+inline MaybeAlign getStackAlign(const Function &F, unsigned Index) {
return F.getAttributes().getAttributes(Index).getStackAlignment();
}
-MaybeAlign getAlign(const CallInst &, unsigned);
+MaybeAlign getStackAlign(const CallBase &, unsigned);
} // namespace llvm
diff --git a/llvm/test/CodeGen/NVPTX/nvvm-annotations-D120129.ll b/llvm/test/CodeGen/NVPTX/nvvm-annotations-D120129.ll
index 0b8d247b0bca6..bbc83cfc3b7c8 100644
--- a/llvm/test/CodeGen/NVPTX/nvvm-annotations-D120129.ll
+++ b/llvm/test/CodeGen/NVPTX/nvvm-annotations-D120129.ll
@@ -1,7 +1,7 @@
; RUN: llc < %s -mtriple=nvptx64-unknown-unknown | FileCheck %s
; RUN: %if ptxas %{ llc < %s -mtriple=nvptx64-unknown-unknown | %ptxas-verify %}
;
-; NVPTXTargetLowering::getFunctionParamOptimizedAlign, which was introduces in
+; NVPTXTargetLowering::getPromotedParamTypeAlign, which was introduces in
; D120129, contained a poorly designed assertion checking that a function with
; internal or private linkage is not a kernel. It relied on invariants that
; were not actually guaranteed, and that resulted in compiler crash with some
diff --git a/llvm/test/CodeGen/NVPTX/param-overalign.ll b/llvm/test/CodeGen/NVPTX/param-overalign.ll
index 2ee749fb3b0cb..83add8a89b07c 100644
--- a/llvm/test/CodeGen/NVPTX/param-overalign.ll
+++ b/llvm/test/CodeGen/NVPTX/param-overalign.ll
@@ -106,10 +106,9 @@ define alignstack(8) %struct.float2 @aligned_return(%struct.float2 %a ) {
; CHECK-NEXT: .reg .b32 %r<3>;
; CHECK-EMPTY:
; CHECK-NEXT: // %bb.0:
-; CHECK-NEXT: ld.param.b32 %r1, [aligned_return_param_0];
-; CHECK-NEXT: ld.param.b32 %r2, [aligned_return_param_0+4];
-; CHECK-NEXT: st.param.b32 [func_retval0+4], %r2;
-; CHECK-NEXT: st.param.b32 [func_retval0], %r1;
+; CHECK-NEXT: ld.param.b32 %r1, [aligned_return_param_0+4];
+; CHECK-NEXT: ld.param.b32 %r2, [aligned_return_param_0];
+; CHECK-NEXT: st.param.v2.b32 [func_retval0], {%r2, %r1};
; CHECK-NEXT: ret;
ret %struct.float2 %a
}
diff --git a/llvm/test/CodeGen/NVPTX/ret-align-mismatch.ll b/llvm/test/CodeGen/NVPTX/ret-align-mismatch.ll
new file mode 100644
index 0000000000000..23545d6cd9c31
--- /dev/null
+++ b/llvm/test/CodeGen/NVPTX/ret-align-mismatch.ll
@@ -0,0 +1,82 @@
+; RUN: llc < %s -mtriple=nvptx64 | FileCheck %s
+
+; Verify that return value alignment is consistent between the callee
+; (LowerReturn), the declaration (printReturnValStr), and the caller
+; (LowerCall). All three should honor alignstack on the return index.
+
+target triple = "nvptx64-nvidia-cuda"
+
+%struct.big = type { i32, i32, i32, i32, i32 }
+
+; alignstack(4) on the return forces align 4 everywhere: the declaration,
+; the callee stores, and the caller loads all use scalar b32 ops.
+; CHECK-LABEL: .func (.param .align 4 .b8 func_retval0[20]) internal_ret_align4()
+; CHECK-NOT: st.param.v4
+; CHECK: st.param.b32 [func_retval0+16], 5
+; CHECK: st.param.b32 [func_retval0+12], 4
+; CHECK: st.param.b32 [func_retval0+8], 3
+; CHECK: st.param.b32 [func_retval0+4], 2
+; CHECK: st.param.b32 [func_retval0], 1
+
+define internal alignstack(4) %struct.big @internal_ret_align4() {
+ ret %struct.big { i32 1, i32 2, i32 3, i32 4, i32 5 }
+}
+
+; The caller also reads the return value with align 4 (scalar loads).
+; CHECK-LABEL: .visible .func (.param .align 4 .b8 func_retval0[20]) caller_align4()
+; CHECK: .param .align 4 .b8 retval0[20];
+; CHECK: call.uni (retval0), internal_ret_align4
+; CHECK: ld.param.b32 {{%r[0-9]+}}, [retval0+16];
+; CHECK: ld.param.b32 {{%r[0-9]+}}, [retval0+12];
+; CHECK: ld.param.b32 {{%r[0-9]+}}, [retval0+8];
+; CHECK: ld.param.b32 {{%r[0-9]+}}, [retval0+4];
+; CHECK: ld.param.b32 {{%r[0-9]+}}, [retval0];
+
+define %struct.big @caller_align4() {
+ %r = call %struct.big @internal_ret_align4()
+ ret %struct.big %r
+}
+
+; alignstack(16) permits 128-bit vectorization. The declaration, callee
+; stores, and caller loads all agree on align 16 and use a v4.b32 op for
+; the first four elements.
+; CHECK-LABEL: .func (.param .align 16 .b8 func_retval0[20]) internal_ret_align16()
+; CHECK: st.param.b32 [func_retval0+16], 5
+; CHECK: st.param.v4.b32 [func_retval0], {1, 2, 3, 4}
+
+define internal alignstack(16) %struct.big @internal_ret_align16() {
+ ret %struct.big { i32 1, i32 2, i32 3, i32 4, i32 5 }
+}
+
+; CHECK-LABEL: .visible .func (.param .align 4 .b8 func_retval0[20]) caller_align16()
+; CHECK: .param .align 16 .b8 retval0[20];
+; CHECK: call.uni (retval0), internal_ret_align16
+; CHECK: ld.param.b32 {{%r[0-9]+}}, [retval0+16];
+; CHECK: ld.param.v4.b32 {{{%r[0-9]+, %r[0-9]+, %r[0-9]+, %r[0-9]+}}}, [retval0];
+
+define %struct.big @caller_align16() {
+ %r = call %struct.big @internal_ret_align16()
+ ret %struct.big %r
+}
+
+; With no explicit alignstack, an internal-linkage callee gets its return
+; alignment bumped to 16 by the param-align optimization, so vectorization
+; still kicks in on both sides of the call.
+; CHECK-LABEL: .func (.param .align 16 .b8 func_retval0[20]) internal_ret_default()
+; CHECK: st.param.b32 [func_retval0+16], 5
+; CHECK: st.param.v4.b32 [func_retval0], {1, 2, 3, 4}
+
+define internal %struct.big @internal_ret_default() {
+ ret %struct.big { i32 1, i32 2, i32 3, i32 4, i32 5 }
+}
+
+; CHECK-LABEL: .visible .func (.param .align 4 .b8 func_retval0[20]) caller_default()
+; CHECK: .param .align 16 .b8 retval0[20];
+; CHECK: call.uni (retval0), internal_ret_default
+; CHECK: ld.param.b32 {{%r[0-9]+}}, [retval0+16];
+; CHECK: ld.param.v4.b32 {{{%r[0-9]+, %r[0-9]+, %r[0-9]+, %r[0-9]+}}}, [retval0];
+
+define %struct.big @caller_default() {
+ %r = call %struct.big @internal_ret_default()
+ ret %struct.big %r
+}
More information about the llvm-commits
mailing list