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 || _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
