[llvm] [offload] add context to olMemAlloc* (PR #222677)
Alex Duran via llvm-commits
llvm-commits at lists.llvm.org
Fri Sep 11 05:35:16 PDT 2026
================
@@ -701,176 +711,133 @@ TargetAllocTy convertOlToPluginAllocTy(ol_alloc_type_t Type) {
}
}
-constexpr size_t MAX_ALLOC_TRIES = 50;
-Error olMemAllocImplHelper(ol_device_handle_t Device, ol_alloc_type_t Type,
- size_t Size, size_t Alignment,
- void **AllocationOut) {
- SmallVector<void *> Rejects;
-
- // Repeat the allocation up to a certain amount of times. If it happens to
- // already be allocated (e.g. by a device from another vendor) throw it away
- // and try again.
- for (size_t Count = 0; Count < MAX_ALLOC_TRIES; Count++) {
- auto NewAlloc = Device->Device->dataAlloc(
- Size, nullptr, convertOlToPluginAllocTy(Type), Alignment);
- if (!NewAlloc)
- return NewAlloc.takeError();
-
- void *NewEnd = &static_cast<char *>(*NewAlloc)[Size];
- auto &AllocBases = OffloadContext::get().AllocBases;
- auto &AllocInfoMap = OffloadContext::get().AllocInfoMap;
- {
- std::lock_guard<std::mutex> Lock(OffloadContext::get().AllocInfoMapMutex);
-
- // Check that this memory region doesn't overlap another one
- // That is, the start of this allocation needs to be after another
- // allocation's end point, and the end of this allocation needs to be
- // before the next one's start.
- // `Gap` is the first alloc who ends after the new alloc's start point.
- auto Gap =
- std::lower_bound(AllocBases.begin(), AllocBases.end(), *NewAlloc,
- [&](const void *Iter, const void *Val) {
- return AllocInfoMap.at(Iter).End <= Val;
- });
- if (Gap == AllocBases.end() || NewEnd <= AllocInfoMap.at(*Gap).Start) {
- // Success, no conflict
- AllocInfoMap.insert_or_assign(
- *NewAlloc, AllocInfo{Device, Type, *NewAlloc, NewEnd});
- AllocBases.insert(
- std::lower_bound(AllocBases.begin(), AllocBases.end(), *NewAlloc),
- *NewAlloc);
- *AllocationOut = *NewAlloc;
-
- for (void *R : Rejects)
- if (auto Err =
- Device->Device->dataDelete(R, convertOlToPluginAllocTy(Type)))
- return Err;
- return Error::success();
- }
-
- // To avoid the next attempt allocating the same memory we just freed, we
- // hold onto it until we complete the allocation
- Rejects.push_back(*NewAlloc);
- }
+ol_alloc_type_t convertPluginToOlAllocTy(TargetAllocTy Kind) {
+ switch (Kind) {
+ case TARGET_ALLOC_HOST:
+ return OL_ALLOC_TYPE_HOST;
+ case TARGET_ALLOC_SHARED:
+ return OL_ALLOC_TYPE_MANAGED;
+ case TARGET_ALLOC_DEVICE:
+ case TARGET_ALLOC_DEFAULT:
+ return OL_ALLOC_TYPE_DEVICE;
}
-
- // We've tried multiple times, and can't allocate a non-overlapping region.
- return createOffloadError(ErrorCode::BACKEND_FAILURE,
- "failed to allocate non-overlapping memory");
+ llvm_unreachable("unhandled TargetAllocTy");
}
-Error olMemAlloc_impl(ol_device_handle_t Device, ol_alloc_type_t Type,
- size_t Size, void **AllocationOut) {
+Error olMemAlloc_impl(ol_context_handle_t Context, ol_device_handle_t Device,
+ ol_alloc_type_t Type, size_t Size, void **AllocationOut) {
if (Type == OL_ALLOC_TYPE_HOST)
return createOffloadError(ErrorCode::INVALID_ENUMERATION,
"use olMemAllocHost for host allocations");
- return olMemAllocImplHelper(Device, Type, Size, /*Alignment=*/0,
- AllocationOut);
+ auto AllocOrErr =
+ Context->allocate(Device, static_cast<int64_t>(Size),
+ convertOlToPluginAllocTy(Type), /*Alignment=*/0);
+ if (!AllocOrErr)
+ return AllocOrErr.takeError();
+ *AllocationOut = *AllocOrErr;
+ return Error::success();
}
-Error olMemAllocHost_impl(ol_device_handle_t Device, size_t Size,
+Error olMemAllocHost_impl(ol_context_handle_t Context,
+ ol_device_handle_t Device, size_t Size,
void **AllocationOut) {
- return olMemAllocImplHelper(Device, OL_ALLOC_TYPE_HOST, Size,
- /*Alignment=*/0, AllocationOut);
+ auto AllocOrErr = Context->allocate(Device, static_cast<int64_t>(Size),
+ TARGET_ALLOC_HOST, /*Alignment=*/0);
+ if (!AllocOrErr)
+ return AllocOrErr.takeError();
+ *AllocationOut = *AllocOrErr;
+ return Error::success();
}
-Error olMemAllocAligned_impl(ol_device_handle_t Device, ol_alloc_type_t Type,
+Error olMemAllocAligned_impl(ol_context_handle_t Context,
+ ol_device_handle_t Device, ol_alloc_type_t Type,
size_t Size, size_t Alignment,
void **AllocationOut) {
if (Type == OL_ALLOC_TYPE_HOST)
return createOffloadError(ErrorCode::INVALID_ENUMERATION,
"use olMemAllocAlignedHost for host allocations");
- return olMemAllocImplHelper(Device, Type, Size, Alignment, AllocationOut);
+ auto AllocOrErr =
+ Context->allocate(Device, static_cast<int64_t>(Size),
+ convertOlToPluginAllocTy(Type), Alignment);
+ if (!AllocOrErr)
+ return AllocOrErr.takeError();
+ *AllocationOut = *AllocOrErr;
+ return Error::success();
}
-Error olMemAllocAlignedHost_impl(ol_device_handle_t Device, size_t Size,
+Error olMemAllocAlignedHost_impl(ol_context_handle_t Context,
+ ol_device_handle_t Device, size_t Size,
size_t Alignment, void **AllocationOut) {
- return olMemAllocImplHelper(Device, OL_ALLOC_TYPE_HOST, Size, Alignment,
- AllocationOut);
+ auto AllocOrErr = Context->allocate(Device, static_cast<int64_t>(Size),
+ TARGET_ALLOC_HOST, Alignment);
+ if (!AllocOrErr)
+ return AllocOrErr.takeError();
+ *AllocationOut = *AllocOrErr;
+ return Error::success();
}
-Error olMemFree_impl(void *Address) {
- ol_device_handle_t Device;
- ol_alloc_type_t Type;
- {
- std::lock_guard<std::mutex> Lock(OffloadContext::get().AllocInfoMapMutex);
- if (!OffloadContext::get().AllocInfoMap.contains(Address))
- return createOffloadError(ErrorCode::INVALID_ARGUMENT,
- "address is not a known allocation");
-
- auto AllocInfo = OffloadContext::get().AllocInfoMap.at(Address);
- Device = AllocInfo.Device;
- Type = AllocInfo.Type;
- OffloadContext::get().AllocInfoMap.erase(Address);
-
- auto &Bases = OffloadContext::get().AllocBases;
- Bases.erase(std::lower_bound(Bases.begin(), Bases.end(), Address));
- }
-
- if (auto Res =
- Device->Device->dataDelete(Address, convertOlToPluginAllocTy(Type)))
- return Res;
-
- return Error::success();
+Error olMemFree_impl(ol_context_handle_t Context, void *Address) {
+ return Context->deallocate(Address);
----------------
adurang wrote:
How is this supposed to work without knowing the allocation type as before?
https://github.com/llvm/llvm-project/pull/222677
More information about the llvm-commits
mailing list