rocr: Generalize AMD::MemoryRegion Allocate and Free

Remove KFD-specific Allocate/Free calls from the AMD::MemoryRegion.
The KFD-driver-specific Allocate/Free calls are now implemented in
the KfdDriver. Future changes will migrate the remaining KFD-specific
calls out of AMD::MemoryRegion.

This allows the MemoryRegion to be used across AMD drivers like the
XDNA driver.

Change-Id: Ib6a2a9e5e1a15e61644d2592beb3a8e6578c3010
This commit is contained in:
Tony Gutierrez
2024-08-19 15:46:36 +00:00
parent c42ff44a6a
commit 68669f4e1a
11 changed files with 420 additions and 298 deletions
+12 -10
View File
@@ -49,11 +49,12 @@
#include <vector>
#include "core/inc/checked.h"
#include "core/inc/driver.h"
#include "core/inc/isa.h"
#include "core/inc/queue.h"
#include "core/inc/memory_region.h"
#include "core/util/utils.h"
#include "core/inc/queue.h"
#include "core/util/locks.h"
#include "core/util/utils.h"
namespace rocr {
@@ -117,19 +118,18 @@ class Agent : public Checked<0xF6BC25EB17E6F917> {
// @brief Agent class contructor.
//
// @param [in] type CPU or GPU or other.
explicit Agent(uint32_t node_id, DeviceType type)
: node_id_(node_id),
device_type_(uint32_t(type)),
profiling_enabled_(false),
enabled_(false) {
explicit Agent(DriverType drv_type, uint32_t node_id, DeviceType type)
: driver_type(drv_type), node_id_(node_id), device_type_(uint32_t(type)),
profiling_enabled_(false), enabled_(false) {
public_handle_ = Convert(this);
}
// @brief Agent class contructor.
//
// @param [in] type CPU or GPU or other.
explicit Agent(uint32_t node_id, uint32_t type)
: node_id_(node_id), device_type_(type), profiling_enabled_(false) {
explicit Agent(DriverType drv_type, uint32_t node_id, uint32_t type)
: driver_type(drv_type), node_id_(node_id), device_type_(type),
profiling_enabled_(false) {
public_handle_ = Convert(this);
}
@@ -315,7 +315,9 @@ class Agent : public Checked<0xF6BC25EB17E6F917> {
for (auto region : regions()) region->Trim();
}
protected:
const DriverType driver_type;
protected:
// Intention here is to have a polymorphic update procedure for public_handle_
// which is callable on any Agent* but only from some class dervied from
// Agent*. do_set_public_handle should remain protected or private in all
+125 -112
View File
@@ -51,15 +51,16 @@
#include "hsakmt/hsakmt.h"
#include "core/inc/runtime.h"
#include "core/inc/agent.h"
#include "core/inc/blit.h"
#include "core/inc/signal.h"
#include "core/inc/cache.h"
#include "core/inc/driver.h"
#include "core/inc/runtime.h"
#include "core/inc/scratch_cache.h"
#include "core/util/small_heap.h"
#include "core/util/locks.h"
#include "core/inc/signal.h"
#include "core/util/lazy_ptr.h"
#include "core/util/locks.h"
#include "core/util/small_heap.h"
#include "pcs/pcs_runtime.h"
namespace rocr {
@@ -72,142 +73,154 @@ typedef ScratchCache::ScratchInfo ScratchInfo;
class GpuAgentInt : public core::Agent {
public:
// @brief Constructor
GpuAgentInt(uint32_t node_id)
: core::Agent(node_id,core::Agent::DeviceType::kAmdGpuDevice) {}
GpuAgentInt(uint32_t node_id)
: core::Agent(core::DriverType::KFD, node_id,
core::Agent::DeviceType::kAmdGpuDevice) {}
// @brief Ensure blits are ready (performance hint).
virtual void PreloadBlits() {}
// @brief Ensure blits are ready (performance hint).
virtual void PreloadBlits() {}
// @brief Initialization hook invoked after tools library has loaded,
// to allow tools interception of interface functions.
//
// @retval HSA_STATUS_SUCCESS if initialization is successful.
virtual hsa_status_t PostToolsInit() = 0;
// @brief Initialization hook invoked after tools library has loaded,
// to allow tools interception of interface functions.
//
// @retval HSA_STATUS_SUCCESS if initialization is successful.
virtual hsa_status_t PostToolsInit() = 0;
// @brief Invoke the user provided callback for each region accessible by
// this agent.
//
// @param [in] include_peer If true, the callback will be also invoked on each
// peer memory region accessible by this agent. If false, only invoke the
// callback on memory region owned by this agent.
// @param [in] callback User provided callback function.
// @param [in] data User provided pointer as input for @p callback.
//
// @retval ::HSA_STATUS_SUCCESS if the callback function for each traversed
// region returns ::HSA_STATUS_SUCCESS.
virtual hsa_status_t VisitRegion(bool include_peer,
hsa_status_t (*callback)(hsa_region_t region,
void* data),
void* data) const = 0;
// @brief Invoke the user provided callback for each region accessible by
// this agent.
//
// @param [in] include_peer If true, the callback will be also invoked on
// each peer memory region accessible by this agent. If false, only invoke
// the callback on memory region owned by this agent.
// @param [in] callback User provided callback function.
// @param [in] data User provided pointer as input for @p callback.
//
// @retval ::HSA_STATUS_SUCCESS if the callback function for each traversed
// region returns ::HSA_STATUS_SUCCESS.
virtual hsa_status_t
VisitRegion(bool include_peer,
hsa_status_t (*callback)(hsa_region_t region, void *data),
void *data) const = 0;
// @brief Carve scratch memory for main from scratch pool.
//
// @param [in/out] scratch Structure to be populated with the carved memory
// information.
virtual void AcquireQueueMainScratch(ScratchInfo& scratch) = 0;
// @brief Carve scratch memory for main from scratch pool.
//
// @param [in/out] scratch Structure to be populated with the carved memory
// information.
virtual void AcquireQueueMainScratch(ScratchInfo &scratch) = 0;
// @brief Carve scratch memory for alt from scratch pool.
//
// @param [in/out] scratch Structure to be populated with the carved memory
// information.
virtual void AcquireQueueAltScratch(ScratchInfo& scratch) = 0;
// @brief Carve scratch memory for alt from scratch pool.
//
// @param [in/out] scratch Structure to be populated with the carved memory
// information.
virtual void AcquireQueueAltScratch(ScratchInfo &scratch) = 0;
// @brief Release scratch memory from main back to scratch pool.
//
// @param [in/out] scratch Scratch memory previously acquired with call to
// ::AcquireQueueMainScratch.
virtual void ReleaseQueueMainScratch(ScratchInfo& base) = 0;
// @brief Release scratch memory from main back to scratch pool.
//
// @param [in/out] scratch Scratch memory previously acquired with call to
// ::AcquireQueueMainScratch.
virtual void ReleaseQueueMainScratch(ScratchInfo &base) = 0;
// @brief Release scratch memory back from alternate to scratch pool.
//
// @param [in/out] scratch Scratch memory previously acquired with call to
// ::AcquireQueueAltcratch.
virtual void ReleaseQueueAltScratch(ScratchInfo& base) = 0;
// @brief Release scratch memory back from alternate to scratch pool.
//
// @param [in/out] scratch Scratch memory previously acquired with call to
// ::AcquireQueueAltcratch.
virtual void ReleaseQueueAltScratch(ScratchInfo &base) = 0;
// @brief Translate the kernel start and end dispatch timestamp from agent
// domain to host domain.
//
// @param [in] signal Pointer to signal that provides the dispatch timing.
// @param [out] time Structure to be populated with the host domain value.
virtual void TranslateTime(core::Signal* signal,
hsa_amd_profiling_dispatch_time_t& time) = 0;
// @brief Translate the kernel start and end dispatch timestamp from agent
// domain to host domain.
//
// @param [in] signal Pointer to signal that provides the dispatch timing.
// @param [out] time Structure to be populated with the host domain value.
virtual void TranslateTime(core::Signal *signal,
hsa_amd_profiling_dispatch_time_t &time) = 0;
// @brief Translate the async copy start and end timestamp from agent
// domain to host domain.
//
// @param [in] signal Pointer to signal that provides the async copy timing.
// @param [out] time Structure to be populated with the host domain value.
virtual void TranslateTime(core::Signal* signal, hsa_amd_profiling_async_copy_time_t& time) = 0;
// @brief Translate the async copy start and end timestamp from agent
// domain to host domain.
//
// @param [in] signal Pointer to signal that provides the async copy timing.
// @param [out] time Structure to be populated with the host domain value.
virtual void TranslateTime(core::Signal *signal,
hsa_amd_profiling_async_copy_time_t &time) = 0;
// @brief Translate timestamp agent domain to host domain.
//
// @param [out] time Timestamp in agent domain.
virtual uint64_t TranslateTime(uint64_t tick) = 0;
// @brief Translate timestamp agent domain to host domain.
//
// @param [out] time Timestamp in agent domain.
virtual uint64_t TranslateTime(uint64_t tick) = 0;
// @brief Invalidate caches on the agent which may hold code object data.
virtual void InvalidateCodeCaches() = 0;
// @brief Invalidate caches on the agent which may hold code object data.
virtual void InvalidateCodeCaches() = 0;
// @brief Sets the coherency type of this agent.
//
// @param [in] type New coherency type.
//
// @retval true The new coherency type is set successfuly.
virtual bool current_coherency_type(hsa_amd_coherency_type_t type) = 0;
// @brief Sets the coherency type of this agent.
//
// @param [in] type New coherency type.
//
// @retval true The new coherency type is set successfuly.
virtual bool current_coherency_type(hsa_amd_coherency_type_t type) = 0;
// @brief Returns the current coherency type of this agent.
//
// @retval Coherency type.
virtual hsa_amd_coherency_type_t current_coherency_type() const = 0;
// @brief Returns the current coherency type of this agent.
//
// @retval Coherency type.
virtual hsa_amd_coherency_type_t current_coherency_type() const = 0;
virtual void RegisterGangPeer(core::Agent& gang_peer, unsigned int bandwidth_factor) = 0;
virtual void RegisterGangPeer(core::Agent &gang_peer,
unsigned int bandwidth_factor) = 0;
virtual void RegisterRecSdmaEngIdMaskPeer(core::Agent& gang_peer, uint32_t rec_sdma_eng_id_mask) = 0;
virtual void RegisterRecSdmaEngIdMaskPeer(core::Agent &gang_peer,
uint32_t rec_sdma_eng_id_mask) = 0;
// @brief Query if agent represent Kaveri GPU.
//
// @retval true if agent is Kaveri GPU.
virtual bool is_kv_device() const = 0;
// @brief Query if agent represent Kaveri GPU.
//
// @retval true if agent is Kaveri GPU.
virtual bool is_kv_device() const = 0;
// @brief Query the agent HSA profile.
//
// @retval HSA profile.
virtual hsa_profile_t profile() const = 0;
// @brief Query the agent HSA profile.
//
// @retval HSA profile.
virtual hsa_profile_t profile() const = 0;
// @brief Query the agent memory bus width in bit.
//
// @retval Bus width in bit.
virtual uint32_t memory_bus_width() const = 0;
// @brief Query the agent memory bus width in bit.
//
// @retval Bus width in bit.
virtual uint32_t memory_bus_width() const = 0;
// @brief Query the agent memory maximum frequency in MHz.
//
// @retval Bus width in MHz.
virtual uint32_t memory_max_frequency() const = 0;
// @brief Query the agent memory maximum frequency in MHz.
//
// @retval Bus width in MHz.
virtual uint32_t memory_max_frequency() const = 0;
// @brief Whether agent supports asynchronous scratch reclaim. Depends on CP FW
virtual bool AsyncScratchReclaimEnabled() const = 0;
// @brief Whether agent supports asynchronous scratch reclaim. Depends on CP
// FW
virtual bool AsyncScratchReclaimEnabled() const = 0;
// @brief Update the agent's scratch use-once threshold.
// Only valid when async scratch reclaim is supported
// @retval HSA_STATUS_SUCCESS if successful
virtual hsa_status_t SetAsyncScratchThresholds(size_t use_once_limit) = 0;
// @brief Update the agent's scratch use-once threshold.
// Only valid when async scratch reclaim is supported
// @retval HSA_STATUS_SUCCESS if successful
virtual hsa_status_t SetAsyncScratchThresholds(size_t use_once_limit) = 0;
// @brief Iterate through supported PC Sampling configurations
// @retval HSA_STATUS_SUCCESS if successful
virtual hsa_status_t PcSamplingIterateConfig(hsa_ven_amd_pcs_iterate_configuration_callback_t cb,
void* cb_data) = 0;
// @brief Iterate through supported PC Sampling configurations
// @retval HSA_STATUS_SUCCESS if successful
virtual hsa_status_t
PcSamplingIterateConfig(hsa_ven_amd_pcs_iterate_configuration_callback_t cb,
void *cb_data) = 0;
virtual hsa_status_t PcSamplingCreate(pcs::PcsRuntime::PcSamplingSession& session) = 0;
virtual hsa_status_t
PcSamplingCreate(pcs::PcsRuntime::PcSamplingSession &session) = 0;
virtual hsa_status_t PcSamplingCreateFromId(HsaPcSamplingTraceId pcsId,
pcs::PcsRuntime::PcSamplingSession& session) = 0;
virtual hsa_status_t
PcSamplingCreateFromId(HsaPcSamplingTraceId pcsId,
pcs::PcsRuntime::PcSamplingSession &session) = 0;
virtual hsa_status_t PcSamplingDestroy(pcs::PcsRuntime::PcSamplingSession& session) = 0;
virtual hsa_status_t
PcSamplingDestroy(pcs::PcsRuntime::PcSamplingSession &session) = 0;
virtual hsa_status_t PcSamplingStart(pcs::PcsRuntime::PcSamplingSession& session) = 0;
virtual hsa_status_t
PcSamplingStart(pcs::PcsRuntime::PcSamplingSession &session) = 0;
virtual hsa_status_t PcSamplingStop(pcs::PcsRuntime::PcSamplingSession& session) = 0;
virtual hsa_status_t
PcSamplingStop(pcs::PcsRuntime::PcSamplingSession &session) = 0;
virtual hsa_status_t PcSamplingFlush(pcs::PcsRuntime::PcSamplingSession& session) = 0;
virtual hsa_status_t
PcSamplingFlush(pcs::PcsRuntime::PcSamplingSession &session) = 0;
};
class GpuAgent : public GpuAgentInt {
+37 -7
View File
@@ -43,11 +43,21 @@
#ifndef HSA_RUNTIME_CORE_INC_AMD_KFD_DRIVER_H_
#define HSA_RUNTIME_CORE_INC_AMD_KFD_DRIVER_H_
#include "core/inc/driver.h"
#include <string>
#include "hsakmt/hsakmt.h"
#include "core/inc/driver.h"
#include "core/inc/memory_region.h"
namespace rocr {
namespace core {
class Queue;
}
namespace AMD {
class KfdDriver : public core::Driver {
@@ -57,13 +67,33 @@ public:
static hsa_status_t DiscoverDriver();
hsa_status_t QueryKernelModeDriver(core::DriverQuery query) override;
hsa_status_t GetMemoryProperties(uint32_t node_id,
core::MemProperties &mprops) const override;
hsa_status_t AllocateMemory(void **mem, size_t size, uint32_t node_id,
core::MemFlags flags) override;
hsa_status_t FreeMemory(void *mem, uint32_t node_id) override;
hsa_status_t
GetMemoryProperties(uint32_t node_id,
core::MemoryRegion &mem_region) const override;
hsa_status_t AllocateMemory(const core::MemoryRegion &mem_region,
core::MemoryRegion::AllocateFlags alloc_flags,
void **mem, size_t size,
uint32_t node_id) override;
hsa_status_t FreeMemory(void *mem, size_t size) override;
hsa_status_t CreateQueue(core::Queue &queue) override;
hsa_status_t DestroyQueue(core::Queue &queue) const override;
private:
/// @brief Allocate agent accessible memory (system / local memory).
static void *AllocateKfdMemory(const HsaMemFlags &flags, uint32_t node_id,
size_t size);
/// @brief Free agent accessible memory (system / local memory).
static bool FreeKfdMemory(void *mem, size_t size);
/// @brief Pin memory.
static bool MakeKfdMemoryResident(size_t num_node, const uint32_t *nodes,
const void *mem, size_t size,
uint64_t *alternate_va,
HsaMemMapFlags map_flag);
/// @brief Unpin memory.
static void MakeKfdMemoryUnresident(const void *mem);
};
} // namespace AMD
@@ -77,13 +77,6 @@ class MemoryRegion : public core::MemoryRegion {
return reinterpret_cast<MemoryRegion*>(region.handle);
}
/// @brief Allocate agent accessible memory (system / local memory).
static void* AllocateKfdMemory(const HsaMemFlags& flag, HSAuint32 node_id,
size_t size);
/// @brief Free agent accessible memory (system / local memory).
static bool FreeKfdMemory(void* ptr, size_t size);
static bool RegisterMemory(void* ptr, size_t size, const HsaMemFlags& MemFlags);
static void DeregisterMemory(void* ptr);
@@ -175,7 +168,15 @@ class MemoryRegion : public core::MemoryRegion {
__forceinline size_t GetPageSize() const { return kPageSize_; }
private:
__forceinline const HsaMemFlags &mem_flags() const { return mem_flag_; }
__forceinline const HsaMemMapFlags &map_flags() const { return map_flag_; }
void *fragment_alloc(size_t size) const {
return fragment_allocator_.alloc(size);
}
bool fragment_free(void *mem) const { return fragment_allocator_.free(mem); }
private:
const HsaMemoryProperties mem_props_;
HsaMemFlags mem_flag_;
+13 -5
View File
@@ -45,8 +45,13 @@
#include <memory>
#include "core/inc/driver.h"
#include "core/inc/memory_region.h"
namespace rocr {
namespace core {
class Queue;
}
namespace AMD {
class XdnaDriver : public core::Driver {
@@ -57,11 +62,14 @@ public:
static hsa_status_t DiscoverDriver();
hsa_status_t QueryKernelModeDriver(core::DriverQuery query) override;
hsa_status_t GetMemoryProperties(uint32_t node_id,
core::MemProperties &mprops) const override;
hsa_status_t AllocateMemory(void **mem, size_t size, uint32_t node_id,
core::MemFlags flags) override;
hsa_status_t FreeMemory(void *mem, uint32_t node_id) override;
hsa_status_t
GetMemoryProperties(uint32_t node_id,
core::MemoryRegion &mem_region) const override;
hsa_status_t AllocateMemory(const core::MemoryRegion &mem_region,
core::MemoryRegion::AllocateFlags alloc_flags,
void **mem, size_t size,
uint32_t node_id) override;
hsa_status_t FreeMemory(void *mem, size_t size) override;
hsa_status_t CreateQueue(core::Queue &queue) override;
hsa_status_t DestroyQueue(core::Queue &queue) const override;
+18 -13
View File
@@ -46,20 +46,13 @@
#include <limits>
#include <string>
#include "core/inc/agent.h"
#include "core/inc/memory_region.h"
#include "inc/hsa.h"
namespace rocr {
namespace core {
using MemFlags = uint32_t;
struct MemProperties {
MemFlags flags_;
size_t size_bytes_;
uint64_t virtual_base_addr_;
};
class Queue;
struct DriverVersionInfo {
uint32_t major;
@@ -85,17 +78,27 @@ class Driver {
/// @retval HSA_STATUS_SUCCESS if the kernel-model driver query was
/// successful.
virtual hsa_status_t QueryKernelModeDriver(DriverQuery query) = 0;
/// @brief Open a connection to the driver using name_.
/// @retval HSA_STATUS_SUCCESS if the driver was opened successfully.
hsa_status_t Open();
/// @brief Close a connection to the open driver using fd_.
/// @retval HSA_STATUS_SUCCESS if the driver was opened successfully.
hsa_status_t Close();
/// @brief Get driver version information.
/// @retval DriverVersionInfo containing the driver's version information.
DriverVersionInfo Version() const { return version_; }
const DriverVersionInfo &Version() const { return version_; }
virtual hsa_status_t GetMemoryProperties(uint32_t node_id, MemProperties &mprops) const = 0;
/// @brief Get the memory properties of a specific node.
/// @param node_id Node ID of the agent
/// @param[in, out] mem_region MemoryRegion object whose properties will be
/// retrieved.
/// @retval HSA_STATUS_SUCCESS if the driver sucessfully returns the node's
/// memory properties.
virtual hsa_status_t GetMemoryProperties(uint32_t node_id,
MemoryRegion &mem_region) const = 0;
/// @brief Allocate agent-accessible memory (system or agent-local memory).
///
@@ -103,10 +106,12 @@ class Driver {
///
/// @retval HSA_STATUS_SUCCESS if memory was successfully allocated or
/// hsa_status_t error code if the memory allocation failed.
virtual hsa_status_t AllocateMemory(void** mem, size_t size, uint32_t node_id,
MemFlags flags) = 0;
virtual hsa_status_t AllocateMemory(const MemoryRegion &mem_region,
MemoryRegion::AllocateFlags alloc_flags,
void **mem, size_t size,
uint32_t node_id) = 0;
virtual hsa_status_t FreeMemory(void* mem, uint32_t node_id) = 0;
virtual hsa_status_t FreeMemory(void *mem, size_t size) = 0;
virtual hsa_status_t CreateQueue(Queue &queue) = 0;