diff --git a/offload/include/omptarget.h b/offload/include/omptarget.h index db9590844b2fd3..a925d63e21cca9 100644 --- a/offload/include/omptarget.h +++ b/offload/include/omptarget.h @@ -412,8 +412,22 @@ void __tgt_target_data_update_nowait_mapper( // same action as data_end above. The following types are used; this // function returns 0 if it was able to transfer the execution to a // target and an int different from zero otherwise. +struct __tgt_async_info_handle; + +__tgt_async_info_handle *__tgt_async_info_create(int64_t DeviceId); + +int __tgt_async_info_synchronize(__tgt_async_info_handle *Handle); + +void __tgt_async_info_destroy(__tgt_async_info_handle *Handle); + int __tgt_target_kernel(ident_t *Loc, int64_t DeviceId, int32_t NumTeams, - int32_t ThreadLimit, void *HostPtr, KernelArgsTy *Args); + int32_t ThreadLimit, void *HostPtr, + KernelArgsTy *Args); + +int __tgt_target_kernel_async(ident_t *Loc, int64_t DeviceId, + int32_t NumTeams, int32_t ThreadLimit, + void *HostPtr, KernelArgsTy *Args, + __tgt_async_info_handle *Handle); // Non-blocking synchronization for target nowait regions. This function // acquires the asynchronous context from task data of the current task being diff --git a/offload/libomptarget/exports b/offload/libomptarget/exports index 1831c43cc5f29c..496aecc7a05d14 100644 --- a/offload/libomptarget/exports +++ b/offload/libomptarget/exports @@ -10,6 +10,10 @@ VERS1.0 { __tgt_target_data_end; __tgt_target_data_update; __tgt_target; + __tgt_target_kernel_async; + __tgt_async_info_create; + __tgt_async_info_synchronize; + __tgt_async_info_destroy; __tgt_target_teams; __tgt_target_data_begin_nowait; __tgt_target_data_end_nowait; diff --git a/offload/libomptarget/interface.cpp b/offload/libomptarget/interface.cpp index 5d7d948711b999..b6463907500c19 100644 --- a/offload/libomptarget/interface.cpp +++ b/offload/libomptarget/interface.cpp @@ -75,6 +75,125 @@ bool checkDevice(int64_t &DeviceID, ident_t *Loc) { return false; } +struct __tgt_async_info_handle { + int64_t DeviceId; + AsyncInfoTy AsyncInfo; + + __tgt_async_info_handle(int64_t DeviceId, DeviceTy &Device) + : DeviceId(DeviceId), AsyncInfo(Device) {} +}; + +EXTERN __tgt_async_info_handle * +__tgt_async_info_create(int64_t DeviceId) { + assert(PM && "Runtime not initialized"); + + if (checkDevice(DeviceId, nullptr)) + return nullptr; + + auto DeviceOrErr = PM->getDevice(DeviceId); + if (!DeviceOrErr) + return nullptr; + + return new __tgt_async_info_handle(DeviceId, *DeviceOrErr); +} + +EXTERN int __tgt_async_info_synchronize( + __tgt_async_info_handle *Handle) { + if (!Handle) + return OFFLOAD_SUCCESS; + + return Handle->AsyncInfo.synchronize(); +} + +EXTERN void __tgt_async_info_destroy( + __tgt_async_info_handle *Handle) { + if (!Handle) + return; + + Handle->AsyncInfo.synchronize(); + delete Handle; +} + +template +static inline int +targetKernel(ident_t *Loc, int64_t DeviceId, int32_t NumTeams, + int32_t ThreadLimit, void *HostPtr, KernelArgsTy *KernelArgs, + AsyncInfoTy *ExternalAsyncInfo = nullptr) { + assert(PM && "Runtime not initialized"); + static_assert(std::is_convertible_v, + "Target AsyncInfoTy must be convertible to AsyncInfoTy."); + + if (checkDevice(DeviceId, Loc)) + return OMP_TGT_FAIL; + + if (KernelArgs->Version > OMP_KERNEL_ARG_VERSION) + FATAL_MESSAGE(DeviceId, "Unsupported kernel argument version %u", + KernelArgs->Version); + + KernelArgs = upgradeKernelArgs(KernelArgs, NumTeams, ThreadLimit); + + auto DeviceOrErr = PM->getDevice(DeviceId); + if (!DeviceOrErr) + FATAL_MESSAGE(DeviceId, "%s", + toString(DeviceOrErr.takeError()).c_str()); + + std::optional OwnedAsyncInfo; + + if (!ExternalAsyncInfo) { + OwnedAsyncInfo.emplace(*DeviceOrErr); + ExternalAsyncInfo = &*OwnedAsyncInfo; + } + + AsyncInfoTy &AsyncInfo = *ExternalAsyncInfo; + + OMPT_IF_BUILT(InterfaceRAII TargetRAII( + RegionInterface.getCallbacks(), DeviceId, + /*CodePtr=*/HostPtr)); + + int Rc = target(Loc, *DeviceOrErr, HostPtr, *KernelArgs, AsyncInfo); + + if (!OwnedAsyncInfo) { + handleTargetOutcome(Rc == OFFLOAD_SUCCESS, Loc); + return Rc; + } + + { + TIMESCOPE_WITH_DETAILS_AND_IDENT("Runtime: synchronize", "", Loc); + if (Rc == OFFLOAD_SUCCESS) + Rc = AsyncInfo.synchronize(); + } + + handleTargetOutcome(Rc == OFFLOAD_SUCCESS, Loc); + return Rc; +} + +EXTERN int __tgt_target_kernel_async( + ident_t *Loc, int64_t DeviceId, int32_t NumTeams, int32_t ThreadLimit, + void *HostPtr, KernelArgsTy *KernelArgs, + __tgt_async_info_handle *Handle) { + OMPT_IF_BUILT(ReturnAddressSetterRAII RA(__builtin_return_address(0))); + + if (!Handle || !KernelArgs) { + handleTargetOutcome(false, Loc); + return OMP_TGT_FAIL; + } + + int64_t ResolvedDeviceId = DeviceId; + + if (checkDevice(ResolvedDeviceId, Loc)) + return OMP_TGT_FAIL; + + if (ResolvedDeviceId != Handle->DeviceId) { + handleTargetOutcome(false, Loc); + return OMP_TGT_FAIL; + } + + return targetKernel( + Loc, ResolvedDeviceId, NumTeams, ThreadLimit, HostPtr, KernelArgs, + &Handle->AsyncInfo); +} + + //////////////////////////////////////////////////////////////////////////////// /// adds requires flags EXTERN void __tgt_register_requires(int64_t Flags) {