diff --git a/flang/include/flang/Runtime/CUDA/allocator.h b/flang/include/flang/Runtime/CUDA/allocator.h index 849785cf991f..8f5204769d7a 100644 --- a/flang/include/flang/Runtime/CUDA/allocator.h +++ b/flang/include/flang/Runtime/CUDA/allocator.h @@ -36,5 +36,8 @@ void CUFFreeDevice(void *); void *CUFAllocManaged(std::size_t); void CUFFreeManaged(void *); +void *CUFAllocUnified(std::size_t); +void CUFFreeUnified(void *); + } // namespace Fortran::runtime::cuda #endif // FORTRAN_RUNTIME_CUDA_ALLOCATOR_H_ diff --git a/flang/include/flang/Runtime/allocator-registry.h b/flang/include/flang/Runtime/allocator-registry.h index 209b4d2e44e9..acfada506faf 100644 --- a/flang/include/flang/Runtime/allocator-registry.h +++ b/flang/include/flang/Runtime/allocator-registry.h @@ -19,8 +19,9 @@ static constexpr unsigned kDefaultAllocator = 0; static constexpr unsigned kPinnedAllocatorPos = 1; static constexpr unsigned kDeviceAllocatorPos = 2; static constexpr unsigned kManagedAllocatorPos = 3; +static constexpr unsigned kUnifiedAllocatorPos = 4; -#define MAX_ALLOCATOR 5 +#define MAX_ALLOCATOR 7 // 3 bits are reserved in the descriptor. namespace Fortran::runtime { diff --git a/flang/lib/Lower/ConvertVariable.cpp b/flang/lib/Lower/ConvertVariable.cpp index 45389091b816..ffbbea238647 100644 --- a/flang/lib/Lower/ConvertVariable.cpp +++ b/flang/lib/Lower/ConvertVariable.cpp @@ -1860,9 +1860,10 @@ static unsigned getAllocatorIdx(const Fortran::semantics::Symbol &sym) { return kPinnedAllocatorPos; if (*cudaAttr == Fortran::common::CUDADataAttr::Device) return kDeviceAllocatorPos; - if (*cudaAttr == Fortran::common::CUDADataAttr::Managed || - *cudaAttr == Fortran::common::CUDADataAttr::Unified) + if (*cudaAttr == Fortran::common::CUDADataAttr::Managed) return kManagedAllocatorPos; + if (*cudaAttr == Fortran::common::CUDADataAttr::Unified) + return kUnifiedAllocatorPos; } return kDefaultAllocator; } diff --git a/flang/runtime/CUDA/allocator.cpp b/flang/runtime/CUDA/allocator.cpp index 08fae9efb3e9..cd00d40361d2 100644 --- a/flang/runtime/CUDA/allocator.cpp +++ b/flang/runtime/CUDA/allocator.cpp @@ -26,6 +26,8 @@ void CUFRegisterAllocator() { kDeviceAllocatorPos, {&CUFAllocDevice, CUFFreeDevice}); allocatorRegistry.Register( kManagedAllocatorPos, {&CUFAllocManaged, CUFFreeManaged}); + allocatorRegistry.Register( + kUnifiedAllocatorPos, {&CUFAllocUnified, CUFFreeUnified}); } void *CUFAllocPinned(std::size_t sizeInBytes) { @@ -57,4 +59,14 @@ void CUFFreeManaged(void *p) { CUDA_REPORT_IF_ERROR(cuMemFree(reinterpret_cast(p))); } +void *CUFAllocUnified(std::size_t sizeInBytes) { + // Call alloc managed for the time being. + return CUFAllocManaged(sizeInBytes); +} + +void CUFFreeUnified(void *p) { + // Call free managed for the time being. + CUFFreeManaged(p); +} + } // namespace Fortran::runtime::cuda