brintos

brintos / llvm-project-archived public Read only

0
0
Text · 3.9 KiB · 860448a Raw
145 lines · plain
1//===----------------------------------------------------------------------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8///9/// \file10/// This file contains the definition of wrappers that manage device resources11/// like buffers, binaries, and kernels.12///13//===----------------------------------------------------------------------===//14 15#ifndef MATHTEST_DEVICERESOURCES_HPP16#define MATHTEST_DEVICERESOURCES_HPP17 18#include "mathtest/OffloadForward.hpp"19 20#include "llvm/ADT/ArrayRef.h"21 22#include <cstddef>23#include <memory>24#include <utility>25 26namespace mathtest {27 28class DeviceContext;29 30namespace detail {31 32void freeDeviceMemory(void *Address) noexcept;33} // namespace detail34 35//===----------------------------------------------------------------------===//36// ManagedBuffer37//===----------------------------------------------------------------------===//38 39template <typename T> class [[nodiscard]] ManagedBuffer {40public:41  ~ManagedBuffer() noexcept {42    if (Address)43      detail::freeDeviceMemory(Address);44  }45 46  ManagedBuffer(const ManagedBuffer &) = delete;47  ManagedBuffer &operator=(const ManagedBuffer &) = delete;48 49  ManagedBuffer(ManagedBuffer &&Other) noexcept50      : Address(Other.Address), Size(Other.Size) {51    Other.Address = nullptr;52    Other.Size = 0;53  }54 55  ManagedBuffer &operator=(ManagedBuffer &&Other) noexcept {56    if (this == &Other)57      return *this;58 59    if (Address)60      detail::freeDeviceMemory(Address);61 62    Address = Other.Address;63    Size = Other.Size;64 65    Other.Address = nullptr;66    Other.Size = 0;67 68    return *this;69  }70 71  [[nodiscard]] T *data() noexcept { return Address; }72 73  [[nodiscard]] const T *data() const noexcept { return Address; }74 75  [[nodiscard]] std::size_t getSize() const noexcept { return Size; }76 77  [[nodiscard]] operator llvm::MutableArrayRef<T>() noexcept {78    return llvm::MutableArrayRef<T>(data(), getSize());79  }80 81  [[nodiscard]] operator llvm::ArrayRef<T>() const noexcept {82    return llvm::ArrayRef<T>(data(), getSize());83  }84 85private:86  friend class DeviceContext;87 88  explicit ManagedBuffer(T *Address, std::size_t Size) noexcept89      : Address(Address), Size(Size) {}90 91  T *Address = nullptr;92  std::size_t Size = 0;93};94 95//===----------------------------------------------------------------------===//96// DeviceImage97//===----------------------------------------------------------------------===//98 99class [[nodiscard]] DeviceImage {100public:101  ~DeviceImage() noexcept;102  DeviceImage &operator=(DeviceImage &&Other) noexcept;103 104  DeviceImage(const DeviceImage &) = delete;105  DeviceImage &operator=(const DeviceImage &) = delete;106 107  DeviceImage(DeviceImage &&Other) noexcept;108 109private:110  friend class DeviceContext;111 112  explicit DeviceImage(ol_device_handle_t DeviceHandle,113                       ol_program_handle_t Handle) noexcept;114 115  ol_device_handle_t DeviceHandle = nullptr;116  ol_program_handle_t Handle = nullptr;117};118 119//===----------------------------------------------------------------------===//120// DeviceKernel121//===----------------------------------------------------------------------===//122 123template <typename KernelSignature> class [[nodiscard]] DeviceKernel {124public:125  DeviceKernel() = delete;126 127  DeviceKernel(const DeviceKernel &) = default;128  DeviceKernel &operator=(const DeviceKernel &) = default;129  DeviceKernel(DeviceKernel &&) noexcept = default;130  DeviceKernel &operator=(DeviceKernel &&) noexcept = default;131 132private:133  friend class DeviceContext;134 135  explicit DeviceKernel(std::shared_ptr<DeviceImage> Image,136                        ol_symbol_handle_t Kernel)137      : Image(std::move(Image)), Handle(Kernel) {}138 139  std::shared_ptr<DeviceImage> Image;140  ol_symbol_handle_t Handle = nullptr;141};142} // namespace mathtest143 144#endif // MATHTEST_DEVICERESOURCES_HPP145