From 563462713028233db459d424f79dd8814a5fb1d4 Mon Sep 17 00:00:00 2001 From: Andreas Karatzas Date: Sun, 5 Apr 2026 02:59:20 +0000 Subject: [PATCH] UCX plugin: defer packRkey to prevent GPU memory pin on ROCm --- src/plugins/ucx/ucx_backend.cpp | 22 ++++++++++++++++------ src/plugins/ucx/ucx_backend.h | 4 ++-- 2 files changed, 18 insertions(+), 8 deletions(-) diff --git a/src/plugins/ucx/ucx_backend.cpp b/src/plugins/ucx/ucx_backend.cpp index 65a3bf73..a086efb1 100644 --- a/src/plugins/ucx/ucx_backend.cpp +++ b/src/plugins/ucx/ucx_backend.cpp @@ -1239,11 +1239,9 @@ nixl_status_t nixlUcxEngine::registerMem (const nixlBlobDesc &mem, if (ret) { return NIXL_ERR_BACKEND; } - priv->rkeyStr = uc->packRkey(priv->mem); - - if (priv->rkeyStr.empty()) { - return NIXL_ERR_BACKEND; - } + // Defer packRkey to getPublicData/loadLocalMD. Packing eagerly here + // triggers hsa_amd_ipc_memory_create (via ucp_rkey_pack) which + // permanently pins GPU memory on ROCm with no release API. out = priv.release(); return NIXL_SUCCESS; } @@ -1258,7 +1256,13 @@ nixl_status_t nixlUcxEngine::deregisterMem (nixlBackendMD* meta) nixl_status_t nixlUcxEngine::getPublicData (const nixlBackendMD* meta, std::string &str) const { - const nixlUcxPrivateMetadata *priv = (nixlUcxPrivateMetadata*) meta; + const nixlUcxPrivateMetadata *priv = (const nixlUcxPrivateMetadata *)meta; + if (priv->rkeyStr.empty()) { + priv->rkeyStr = uc->packRkey(priv->mem); + if (priv->rkeyStr.empty()) { + return NIXL_ERR_BACKEND; + } + } str = priv->get(); return NIXL_SUCCESS; } @@ -1303,6 +1307,12 @@ nixlUcxEngine::loadLocalMD (nixlBackendMD* input, nixlBackendMD* &output) { nixlUcxPrivateMetadata* input_md = (nixlUcxPrivateMetadata*) input; + if (input_md->rkeyStr.empty()) { + input_md->rkeyStr = uc->packRkey(input_md->mem); + if (input_md->rkeyStr.empty()) { + return NIXL_ERR_BACKEND; + } + } return internalMDHelper(input_md->rkeyStr, localAgent, output); } diff --git a/src/plugins/ucx/ucx_backend.h b/src/plugins/ucx/ucx_backend.h index a62d5d63..cd86e393 100644 --- a/src/plugins/ucx/ucx_backend.h +++ b/src/plugins/ucx/ucx_backend.h @@ -58,8 +58,8 @@ using ucx_connection_ptr_t = std::shared_ptr; // A private metadata has to implement get, and has all the metadata class nixlUcxPrivateMetadata : public nixlBackendMD { private: - nixlUcxMem mem; - nixl_blob_t rkeyStr; + mutable nixlUcxMem mem; + mutable nixl_blob_t rkeyStr; public: nixlUcxPrivateMetadata() : nixlBackendMD(true) {