[llvm] 9610d9c - [offload] Split strictness for threads and blocks (#211400)
via llvm-commits
llvm-commits at lists.llvm.org
Wed Aug 19 23:31:17 PDT 2026
Author: Kevin Sala Penades
Date: 2026-08-19T23:31:10-07:00
New Revision: 9610d9c895df5b8a451fb66955fdc2b14f41eff7
URL: https://github.com/llvm/llvm-project/commit/9610d9c895df5b8a451fb66955fdc2b14f41eff7
DIFF: https://github.com/llvm/llvm-project/commit/9610d9c895df5b8a451fb66955fdc2b14f41eff7.diff
LOG: [offload] Split strictness for threads and blocks (#211400)
This commit splits the strictness for the number of threads and blocks. This will be needed to support `dims` modifier in OpenMP 6.1.
Added:
Modified:
clang/lib/CodeGen/CGOpenMPRuntime.cpp
clang/test/OpenMP/target_teams_codegen.cpp
llvm/include/llvm/Frontend/OpenMP/OMPIRBuilder.h
llvm/lib/Frontend/OpenMP/OMPIRBuilder.cpp
offload/include/Shared/APITypes.h
offload/liboffload/src/OffloadImpl.cpp
offload/libomptarget/KernelLanguage/API.cpp
offload/libomptarget/device.cpp
offload/libomptarget/omptarget.cpp
offload/plugins-nextgen/common/include/PluginInterface.h
offload/plugins-nextgen/common/src/PluginInterface.cpp
Removed:
################################################################################
diff --git a/clang/lib/CodeGen/CGOpenMPRuntime.cpp b/clang/lib/CodeGen/CGOpenMPRuntime.cpp
index 409c222da87fc..62d763b3b8ff9 100644
--- a/clang/lib/CodeGen/CGOpenMPRuntime.cpp
+++ b/clang/lib/CodeGen/CGOpenMPRuntime.cpp
@@ -11074,8 +11074,8 @@ static void emitTargetCallKernelLaunch(
llvm::OpenMPIRBuilder::TargetKernelArgs Args(
NumTargetItems, RTArgs, NumIterations, NumTeams, NumThreads,
- DynCGroupMem, HasNoWait, /*StrictBlocksAndThreads=*/IsBare,
- DynCGroupMemFallback);
+ DynCGroupMem, HasNoWait, /*StrictBlocks=*/IsBare,
+ /*StrictThreads=*/IsBare, DynCGroupMemFallback);
llvm::OpenMPIRBuilder::InsertPointTy AfterIP =
cantFail(OMPRuntime->getOMPBuilder().emitKernelLaunch(
diff --git a/clang/test/OpenMP/target_teams_codegen.cpp b/clang/test/OpenMP/target_teams_codegen.cpp
index 243f533ed3b7e..82cdb8bd58bcc 100644
--- a/clang/test/OpenMP/target_teams_codegen.cpp
+++ b/clang/test/OpenMP/target_teams_codegen.cpp
@@ -634,7 +634,7 @@ int bar(int n){
// CHECK1-NEXT: [[TMP127:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS21]], i32 0, i32 8
// CHECK1-NEXT: store i64 0, ptr [[TMP127]], align 8
// CHECK1-NEXT: [[TMP128:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS21]], i32 0, i32 9
-// CHECK1-NEXT: store i64 64, ptr [[TMP128]], align 8
+// CHECK1-NEXT: store i64 192, ptr [[TMP128]], align 8
// CHECK1-NEXT: [[TMP129:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS21]], i32 0, i32 10
// CHECK1-NEXT: store [3 x i32] [i32 1, i32 0, i32 0], ptr [[TMP129]], align 4
// CHECK1-NEXT: [[TMP130:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS21]], i32 0, i32 11
@@ -693,7 +693,7 @@ int bar(int n){
// CHECK1-NEXT: [[TMP157:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS29]], i32 0, i32 8
// CHECK1-NEXT: store i64 0, ptr [[TMP157]], align 8
// CHECK1-NEXT: [[TMP158:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS29]], i32 0, i32 9
-// CHECK1-NEXT: store i64 64, ptr [[TMP158]], align 8
+// CHECK1-NEXT: store i64 192, ptr [[TMP158]], align 8
// CHECK1-NEXT: [[TMP159:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS29]], i32 0, i32 10
// CHECK1-NEXT: store [3 x i32] [i32 1, i32 2, i32 0], ptr [[TMP159]], align 4
// CHECK1-NEXT: [[TMP160:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS29]], i32 0, i32 11
@@ -752,7 +752,7 @@ int bar(int n){
// CHECK1-NEXT: [[TMP187:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS37]], i32 0, i32 8
// CHECK1-NEXT: store i64 0, ptr [[TMP187]], align 8
// CHECK1-NEXT: [[TMP188:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS37]], i32 0, i32 9
-// CHECK1-NEXT: store i64 64, ptr [[TMP188]], align 8
+// CHECK1-NEXT: store i64 192, ptr [[TMP188]], align 8
// CHECK1-NEXT: [[TMP189:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS37]], i32 0, i32 10
// CHECK1-NEXT: store [3 x i32] [i32 1, i32 2, i32 3], ptr [[TMP189]], align 4
// CHECK1-NEXT: [[TMP190:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS37]], i32 0, i32 11
@@ -2531,7 +2531,7 @@ int bar(int n){
// CHECK3-NEXT: [[TMP125:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS21]], i32 0, i32 8
// CHECK3-NEXT: store i64 0, ptr [[TMP125]], align 8
// CHECK3-NEXT: [[TMP126:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS21]], i32 0, i32 9
-// CHECK3-NEXT: store i64 64, ptr [[TMP126]], align 8
+// CHECK3-NEXT: store i64 192, ptr [[TMP126]], align 8
// CHECK3-NEXT: [[TMP127:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS21]], i32 0, i32 10
// CHECK3-NEXT: store [3 x i32] [i32 1, i32 0, i32 0], ptr [[TMP127]], align 4
// CHECK3-NEXT: [[TMP128:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS21]], i32 0, i32 11
@@ -2590,7 +2590,7 @@ int bar(int n){
// CHECK3-NEXT: [[TMP155:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS29]], i32 0, i32 8
// CHECK3-NEXT: store i64 0, ptr [[TMP155]], align 8
// CHECK3-NEXT: [[TMP156:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS29]], i32 0, i32 9
-// CHECK3-NEXT: store i64 64, ptr [[TMP156]], align 8
+// CHECK3-NEXT: store i64 192, ptr [[TMP156]], align 8
// CHECK3-NEXT: [[TMP157:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS29]], i32 0, i32 10
// CHECK3-NEXT: store [3 x i32] [i32 1, i32 2, i32 0], ptr [[TMP157]], align 4
// CHECK3-NEXT: [[TMP158:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS29]], i32 0, i32 11
@@ -2649,7 +2649,7 @@ int bar(int n){
// CHECK3-NEXT: [[TMP185:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS37]], i32 0, i32 8
// CHECK3-NEXT: store i64 0, ptr [[TMP185]], align 8
// CHECK3-NEXT: [[TMP186:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS37]], i32 0, i32 9
-// CHECK3-NEXT: store i64 64, ptr [[TMP186]], align 8
+// CHECK3-NEXT: store i64 192, ptr [[TMP186]], align 8
// CHECK3-NEXT: [[TMP187:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS37]], i32 0, i32 10
// CHECK3-NEXT: store [3 x i32] [i32 1, i32 2, i32 3], ptr [[TMP187]], align 4
// CHECK3-NEXT: [[TMP188:%.*]] = getelementptr inbounds nuw [[STRUCT___TGT_KERNEL_ARGUMENTS]], ptr [[KERNEL_ARGS37]], i32 0, i32 11
diff --git a/llvm/include/llvm/Frontend/OpenMP/OMPIRBuilder.h b/llvm/include/llvm/Frontend/OpenMP/OMPIRBuilder.h
index 3560cfef096fe..f932224eed661 100644
--- a/llvm/include/llvm/Frontend/OpenMP/OMPIRBuilder.h
+++ b/llvm/include/llvm/Frontend/OpenMP/OMPIRBuilder.h
@@ -2890,7 +2890,8 @@ class OpenMPIRBuilder {
bool HasNoWait = false;
/// True if the kernel strictly requires the number of blocks and threads
/// above to run.
- bool StrictBlocksAndThreads = false;
+ bool StrictBlocks = false;
+ bool StrictThreads = false;
/// The fallback mechanism for the shared memory.
omp::OMPDynGroupprivateFallbackType DynCGroupMemFallback =
omp::OMPDynGroupprivateFallbackType::Abort;
@@ -2900,12 +2901,13 @@ class OpenMPIRBuilder {
TargetKernelArgs(unsigned NumTargetItems, TargetDataRTArgs RTArgs,
Value *NumIterations, ArrayRef<Value *> NumTeams,
ArrayRef<Value *> NumThreads, Value *DynCGroupMem,
- bool HasNoWait, bool StrictBlocksAndThreads,
+ bool HasNoWait, bool StrictBlocks, bool StrictThreads,
omp::OMPDynGroupprivateFallbackType DynCGroupMemFallback)
: NumTargetItems(NumTargetItems), RTArgs(RTArgs),
NumIterations(NumIterations), NumTeams(NumTeams),
NumThreads(NumThreads), DynCGroupMem(DynCGroupMem),
- HasNoWait(HasNoWait), StrictBlocksAndThreads(StrictBlocksAndThreads),
+ HasNoWait(HasNoWait), StrictBlocks(StrictBlocks),
+ StrictThreads(StrictThreads),
DynCGroupMemFallback(DynCGroupMemFallback) {}
};
diff --git a/llvm/lib/Frontend/OpenMP/OMPIRBuilder.cpp b/llvm/lib/Frontend/OpenMP/OMPIRBuilder.cpp
index 1c5f84c65ddd5..e6fac2744a8ae 100644
--- a/llvm/lib/Frontend/OpenMP/OMPIRBuilder.cpp
+++ b/llvm/lib/Frontend/OpenMP/OMPIRBuilder.cpp
@@ -663,11 +663,15 @@ void OpenMPIRBuilder::getKernelArgsVector(TargetKernelArgs &KernelArgs,
Builder.getInt64(static_cast<uint64_t>(KernelArgs.DynCGroupMemFallback));
DynCGroupMemFallbackFlag = Builder.CreateShl(DynCGroupMemFallbackFlag, 2);
- Value *StrictFlag = Builder.getInt64(KernelArgs.StrictBlocksAndThreads);
- StrictFlag = Builder.CreateShl(StrictFlag, 6);
+ Value *StrictBlocksFlag = Builder.getInt64(KernelArgs.StrictBlocks);
+ Value *StrictThreadsFlag = Builder.getInt64(KernelArgs.StrictThreads);
+
+ StrictBlocksFlag = Builder.CreateShl(StrictBlocksFlag, 6);
+ StrictThreadsFlag = Builder.CreateShl(StrictThreadsFlag, 7);
Value *Flags = Builder.CreateOr(HasNoWaitFlag, DynCGroupMemFallbackFlag);
- Flags = Builder.CreateOr(Flags, StrictFlag);
+ Flags = Builder.CreateOr(Flags, StrictBlocksFlag);
+ Flags = Builder.CreateOr(Flags, StrictThreadsFlag);
assert(!KernelArgs.NumTeams.empty() && !KernelArgs.NumThreads.empty());
@@ -10166,7 +10170,8 @@ static void emitTargetCall(
KArgs = OpenMPIRBuilder::TargetKernelArgs(
NumTargetItems, RTArgs, TripCount, NumTeamsC, NumThreadsC, DynCGroupMem,
- HasNoWait, /*StrictBlocksAndThreads=*/false, DynCGroupMemFallback);
+ HasNoWait, /*StrictBlocks=*/false, /*StrictThreads=*/false,
+ DynCGroupMemFallback);
// Assume no error was returned because TaskBodyCB and
// EmitTargetCallFallbackCB don't produce any.
diff --git a/offload/include/Shared/APITypes.h b/offload/include/Shared/APITypes.h
index 71cf6773437d1..6392eb1472e57 100644
--- a/offload/include/Shared/APITypes.h
+++ b/offload/include/Shared/APITypes.h
@@ -106,10 +106,11 @@ struct KernelArgsTy {
uint64_t DynCGroupMemFallback : 2; // The fallback for dynamic cgroup mem.
uint64_t Cooperative : 1; // Was this kernel spawned as cooperative.
uint64_t IsPtrArgs : 1; // Arguments are laid out as an array of pointers.
- uint64_t StrictBlocksAndThreads
- : 1; // The user-requested number of blocks and threads are strict.
- uint64_t Unused : 57;
- } Flags = {0, 0, 0, 0, 0, 0, 0};
+ uint64_t StrictBlocks : 1; // The user-requested number of blocks is strict.
+ uint64_t StrictThreads
+ : 1; // The user-requested number of threads is strict.
+ uint64_t Unused : 56;
+ } Flags = {0, 0, 0, 0, 0, 0, 0, 0};
// User-requested number of blocks (for x,y,z dimension).
uint32_t UserNumBlocks[3] = {0, 0, 0};
// User-requested number of threads (for x,y,z dimension).
diff --git a/offload/liboffload/src/OffloadImpl.cpp b/offload/liboffload/src/OffloadImpl.cpp
index 48feac0b6c780..e59fed4b30c34 100644
--- a/offload/liboffload/src/OffloadImpl.cpp
+++ b/offload/liboffload/src/OffloadImpl.cpp
@@ -1276,7 +1276,8 @@ Error olLaunchKernel_impl(ol_queue_handle_t Queue, ol_device_handle_t Device,
LaunchArgs.UserThreadLimit[1] = LaunchSizeArgs->GroupSize.y;
LaunchArgs.UserThreadLimit[2] = LaunchSizeArgs->GroupSize.z;
LaunchArgs.DynCGroupMem = LaunchSizeArgs->DynSharedMemory;
- LaunchArgs.Flags.StrictBlocksAndThreads = true;
+ LaunchArgs.Flags.StrictBlocks = true;
+ LaunchArgs.Flags.StrictThreads = true;
while (Properties && Properties->type != OL_KERNEL_LAUNCH_PROP_TYPE_NONE) {
switch (Properties->type) {
diff --git a/offload/libomptarget/KernelLanguage/API.cpp b/offload/libomptarget/KernelLanguage/API.cpp
index 50f9b695bed6a..9b6533fd30c87 100644
--- a/offload/libomptarget/KernelLanguage/API.cpp
+++ b/offload/libomptarget/KernelLanguage/API.cpp
@@ -68,7 +68,8 @@ unsigned llvmLaunchKernel(const void *func, dim3 gridDim, dim3 blockDim,
Args.UserThreadLimit[2] = blockDim.z;
Args.ArgPtrs = reinterpret_cast<void **>(args);
Args.Flags.IsCUDA = true;
- Args.Flags.StrictBlocksAndThreads = true;
+ Args.Flags.StrictBlocks = true;
+ Args.Flags.StrictThreads = true;
return __tgt_target_kernel(nullptr, 0, gridDim.x, blockDim.x, func, &Args);
}
}
diff --git a/offload/libomptarget/device.cpp b/offload/libomptarget/device.cpp
index 59877f2ac3642..01e20d7e8e8c5 100644
--- a/offload/libomptarget/device.cpp
+++ b/offload/libomptarget/device.cpp
@@ -405,8 +405,8 @@ int32_t DeviceTy::launchKernel(void *TgtEntryPtr, void **TgtVarsPtr,
llvm::copy(KernelArgs.UserNumBlocks, LaunchArgs.UserNumBlocks);
llvm::copy(KernelArgs.UserThreadLimit, LaunchArgs.UserThreadLimit);
LaunchArgs.Flags.Cooperative = KernelArgs.Flags.Cooperative;
- LaunchArgs.Flags.StrictBlocksAndThreads =
- KernelArgs.Flags.StrictBlocksAndThreads;
+ LaunchArgs.Flags.StrictBlocks = KernelArgs.Flags.StrictBlocks;
+ LaunchArgs.Flags.StrictThreads = KernelArgs.Flags.StrictThreads;
LaunchArgs.Flags.DynCGroupMemFallback = KernelArgs.Flags.DynCGroupMemFallback;
if (KernelArgs.Flags.IsCUDA) {
diff --git a/offload/libomptarget/omptarget.cpp b/offload/libomptarget/omptarget.cpp
index 973e949cc2e6d..9e93bca77b292 100644
--- a/offload/libomptarget/omptarget.cpp
+++ b/offload/libomptarget/omptarget.cpp
@@ -2504,7 +2504,8 @@ int target_replay(ident_t *Loc, DeviceTy &Device, void *HostPtr,
KernelArgs.UserThreadLimit[1] = 1;
KernelArgs.UserThreadLimit[2] = 1;
KernelArgs.DynCGroupMem = SharedMemorySize;
- KernelArgs.Flags.StrictBlocksAndThreads = true;
+ KernelArgs.Flags.StrictBlocks = true;
+ KernelArgs.Flags.StrictThreads = true;
int Ret = Device.launchKernel(Symbols[0].DevPtr, TgtArgs, TgtOffsets,
KernelArgs, ReplayOutcome, AsyncInfo);
diff --git a/offload/plugins-nextgen/common/include/PluginInterface.h b/offload/plugins-nextgen/common/include/PluginInterface.h
index 30c79e28f2ea4..80adec46b0972 100644
--- a/offload/plugins-nextgen/common/include/PluginInterface.h
+++ b/offload/plugins-nextgen/common/include/PluginInterface.h
@@ -456,11 +456,12 @@ struct KernelLaunchArgsTy {
uint32_t UserThreadLimit[3] = {0, 0, 0};
struct {
uint64_t Cooperative : 1; // Was this kernel spawned as cooperative.
- uint64_t StrictBlocksAndThreads
- : 1; // The user-requested number of blocks and threads are strict.
+ uint64_t StrictBlocks : 1; // The user-requested number of blocks is strict.
+ uint64_t StrictThreads
+ : 1; // The user-requested number of threads is strict.
uint64_t DynCGroupMemFallback : 2; // The fallback for dynamic cgroup mem.
uint64_t Unused : 60;
- } Flags = {0, 0, 0, 0};
+ } Flags = {0, 0, 0, 0, 0};
/// Set by the caller when replaying a previously recorded kernel launch, so
/// the plugin can report the outcome back; null for a normal launch.
KernelReplayOutcomeTy *ReplayOutcome = nullptr;
@@ -600,6 +601,7 @@ struct GenericKernelTy {
uint32_t getEffectiveNumBlocks(GenericDeviceTy &GenericDevice,
uint32_t UserNumBlocks, uint64_t LoopTripCount,
uint32_t &EffectiveNumThreads,
+ bool IsNumThreadsStrict,
bool IsNumThreadsFromUser) const;
/// Indicate if the kernel works in Generic SPMD, Generic, No-Loop
diff --git a/offload/plugins-nextgen/common/src/PluginInterface.cpp b/offload/plugins-nextgen/common/src/PluginInterface.cpp
index 9e32bcce02ba3..1cf83aa651d7c 100644
--- a/offload/plugins-nextgen/common/src/PluginInterface.cpp
+++ b/offload/plugins-nextgen/common/src/PluginInterface.cpp
@@ -259,21 +259,25 @@ Error GenericKernelTy::launch(GenericDeviceTy &GenericDevice,
"Non-bare mode should only use the first thread and block "
"dimensions");
- assert(!LaunchArgs.Flags.StrictBlocksAndThreads ||
+ assert(!LaunchArgs.Flags.StrictBlocks ||
+ EffectiveNumBlocks[0] > 0 && EffectiveNumBlocks[1] > 0 &&
+ EffectiveNumBlocks[2] > 0 &&
+ "Strict requires number of blocks greater than zero");
+ assert(!LaunchArgs.Flags.StrictThreads ||
EffectiveNumThreads[0] > 0 && EffectiveNumThreads[1] > 0 &&
- EffectiveNumThreads[2] > 0 && EffectiveNumBlocks[0] > 0 &&
- EffectiveNumBlocks[1] > 0 && EffectiveNumBlocks[2] > 0 &&
- "Strict requires number of blocks and threads greater than zero");
+ EffectiveNumThreads[2] > 0 &&
+ "Strict requires number of threads greater than zero");
// Calculate or adjust the effective number of threads and blocks if needed.
- if (!LaunchArgs.Flags.StrictBlocksAndThreads) {
+ if (!LaunchArgs.Flags.StrictThreads)
EffectiveNumThreads[0] =
getEffectiveNumThreads(GenericDevice, EffectiveNumThreads[0]);
+ if (!LaunchArgs.Flags.StrictBlocks)
EffectiveNumBlocks[0] = getEffectiveNumBlocks(
GenericDevice, EffectiveNumBlocks[0], LaunchArgs.Tripcount,
- EffectiveNumThreads[0], LaunchArgs.UserThreadLimit[0] > 0);
- }
+ EffectiveNumThreads[0], LaunchArgs.Flags.StrictThreads,
+ LaunchArgs.UserThreadLimit[0] > 0);
auto DynBlockMemConfOrErr = prepareBlockMemory(
GenericDevice, LaunchArgs,
@@ -349,7 +353,7 @@ GenericKernelTy::getEffectiveNumThreads(GenericDeviceTy &GenericDevice,
uint32_t GenericKernelTy::getEffectiveNumBlocks(
GenericDeviceTy &GenericDevice, uint32_t UserNumBlocks,
uint64_t LoopTripCount, uint32_t &EffectiveNumThreads,
- bool IsNumThreadsFromUser) const {
+ bool IsNumThreadsStrict, bool IsNumThreadsFromUser) const {
assert(!isBareMode() && "bare kernel should not call this function");
// NOTE: This clamps the user-requested number of blocks to the device limit
@@ -380,7 +384,7 @@ uint32_t GenericKernelTy::getEffectiveNumBlocks(
// Honor the thread_limit clause; only lower the number of threads.
[[maybe_unused]] auto OldNumThreads = EffectiveNumThreads;
if (LoopTripCount >= DefaultNumBlocks * EffectiveNumThreads ||
- IsNumThreadsFromUser) {
+ IsNumThreadsFromUser || IsNumThreadsStrict) {
// Enough parallelism for blocks and threads.
TripCountNumBlocks = ((LoopTripCount - 1) / EffectiveNumThreads) + 1;
assert(IsNumThreadsFromUser ||
More information about the llvm-commits
mailing list