From f6818543102c66c8d1a453fc56010413866751b9 Mon Sep 17 00:00:00 2001 From: Anusha GodavarthySurya Date: Tue, 10 Jan 2023 13:32:27 +0000 Subject: [PATCH] SWDEV-376407 - Track manual nodes added during capture Change-Id: Ic44f842c303eed0eff08b6ee5922976d96a52c71 [ROCm/clr commit: d422ad0ec1e250a316f65463d223a87242fde705] --- projects/clr/hipamd/src/hip_graph.cpp | 75 ++++++++++++------- .../clr/hipamd/src/hip_graph_internal.hpp | 9 +++ projects/clr/hipamd/src/hip_internal.hpp | 2 +- 3 files changed, 57 insertions(+), 29 deletions(-) diff --git a/projects/clr/hipamd/src/hip_graph.cpp b/projects/clr/hipamd/src/hip_graph.cpp index d7eba08b69..a7dec1489f 100644 --- a/projects/clr/hipamd/src/hip_graph.cpp +++ b/projects/clr/hipamd/src/hip_graph.cpp @@ -33,7 +33,8 @@ static amd::Monitor g_streamSetLock{"StreamCaptureset"}; std::unordered_set g_allCapturingStreams; inline hipError_t ihipGraphAddNode(hipGraphNode_t graphNode, hipGraph_t graph, - const hipGraphNode_t* pDependencies, size_t numDependencies) { + const hipGraphNode_t* pDependencies, size_t numDependencies, + bool capture = true) { graph->AddNode(graphNode); for (size_t i = 0; i < numDependencies; i++) { if ((!hipGraphNode::isNodeValid(pDependencies[i])) || @@ -42,12 +43,23 @@ inline hipError_t ihipGraphAddNode(hipGraphNode_t graphNode, hipGraph_t graph, } pDependencies[i]->AddEdge(graphNode); } + if (capture == false) { + { + amd::ScopedLock lock(g_streamSetLock); + for (auto stream : g_allCapturingStreams) { + if (stream->GetCaptureGraph() == graph) { + graph->AddManualNodeDuringCapture(graphNode); + break; + } + } + } + } return hipSuccess; } hipError_t ihipGraphAddKernelNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, const hipGraphNode_t* pDependencies, size_t numDependencies, - const hipKernelNodeParams* pNodeParams) { + const hipKernelNodeParams* pNodeParams, bool capture = true) { if (pGraphNode == nullptr || graph == nullptr || (numDependencies > 0 && pDependencies == nullptr) || pNodeParams == nullptr || pNodeParams->func == nullptr) { @@ -79,13 +91,13 @@ hipError_t ihipGraphAddKernelNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, } *pGraphNode = new hipGraphKernelNode(pNodeParams); - status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, capture); return status; } hipError_t ihipGraphAddMemcpyNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, const hipGraphNode_t* pDependencies, size_t numDependencies, - const hipMemcpy3DParms* pCopyParams) { + const hipMemcpy3DParms* pCopyParams, bool capture = true) { if (pGraphNode == nullptr || graph == nullptr || (numDependencies > 0 && pDependencies == nullptr) || pCopyParams == nullptr) { return hipErrorInvalidValue; @@ -95,13 +107,14 @@ hipError_t ihipGraphAddMemcpyNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, return status; } *pGraphNode = new hipGraphMemcpyNode(pCopyParams); - status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, capture); return status; } hipError_t ihipGraphAddMemcpyNode1D(hipGraphNode_t* pGraphNode, hipGraph_t graph, const hipGraphNode_t* pDependencies, size_t numDependencies, - void* dst, const void* src, size_t count, hipMemcpyKind kind) { + void* dst, const void* src, size_t count, hipMemcpyKind kind, + bool capture = true) { if (pGraphNode == nullptr || graph == nullptr || (numDependencies > 0 && pDependencies == nullptr) || count ==0) { return hipErrorInvalidValue; @@ -111,13 +124,13 @@ hipError_t ihipGraphAddMemcpyNode1D(hipGraphNode_t* pGraphNode, hipGraph_t graph return status; } *pGraphNode = new hipGraphMemcpyNode1D(dst, src, count, kind); - status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, capture); return status; } hipError_t ihipGraphAddMemsetNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, const hipGraphNode_t* pDependencies, size_t numDependencies, - const hipMemsetParams* pMemsetParams) { + const hipMemsetParams* pMemsetParams, bool capture = true) { if (pGraphNode == nullptr || graph == nullptr || pMemsetParams == nullptr || (numDependencies > 0 && pDependencies == nullptr) || pMemsetParams->height == 0) { return hipErrorInvalidValue; @@ -151,7 +164,7 @@ hipError_t ihipGraphAddMemsetNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, return status; } *pGraphNode = new hipGraphMemsetNode(pMemsetParams); - status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, capture); return status; } @@ -188,7 +201,7 @@ hipError_t ihipExtLaunchKernel(hipStream_t stream, hipFunction_t f, uint32_t glo uint32_t localWorkSizeX, uint32_t localWorkSizeY, uint32_t localWorkSizeZ, size_t sharedMemBytes, void** kernelParams, void** extra, hipEvent_t startEvent, hipEvent_t stopEvent, - uint32_t flags) { + uint32_t flags, bool capture = true) { if (!hip::isValid(stream)) { return hipErrorContextIsDestroyed; } @@ -199,7 +212,7 @@ hipError_t ihipExtLaunchKernel(hipStream_t stream, hipFunction_t f, uint32_t glo if (startEvent != nullptr) { pGraphNode = new hipGraphEventRecordNode(startEvent); status = ihipGraphAddNode(pGraphNode, s->GetCaptureGraph(), s->GetLastCapturedNodes().data(), - s->GetLastCapturedNodes().size()); + s->GetLastCapturedNodes().size(), capture); if (status != hipSuccess) { return status; } @@ -813,7 +826,7 @@ hipError_t capturehipStreamWaitEvent(hipEvent_t& event, hipStream_t& stream, uns } hipError_t capturehipLaunchHostFunc(hipStream_t& stream, hipHostFn_t& fn, void*& userData) { - ClPrint(amd::LOG_INFO, amd::LOG_API, "[hipGraph] current capture node Memset2D on stream : %p", + ClPrint(amd::LOG_INFO, amd::LOG_API, "[hipGraph] current capture node host on stream : %p", stream); if (fn == nullptr) { return hipErrorInvalidValue; @@ -1032,6 +1045,10 @@ hipError_t hipStreamEndCapture_common(hipStream_t stream, hipGraph_t* pGraph) { if (s->GetCaptureGraph()->GetLeafNodeCount() > 1) { std::vector leafNodes = s->GetCaptureGraph()->GetLeafNodes(); + std::unordered_set nodes = s->GetCaptureGraph()->GetManualNodesDuringCapture(); + for (auto node : nodes) { + leafNodes.erase(std::find(leafNodes.begin(), leafNodes.end(), node)); + } const std::vector& removedDepNodes = s->GetRemovedDependencies(); bool foundInRemovedDep = false; for (auto leafNode : leafNodes) { @@ -1043,7 +1060,8 @@ hipError_t hipStreamEndCapture_common(hipStream_t stream, hipGraph_t* pGraph) { } // remove temporary node s->GetCaptureGraph()->RemoveNode(pGraphNode); - if (foundInRemovedDep == false) { + s->GetCaptureGraph()->RemoveManualNodesDuringCapture(); + if (leafNodes.size() > 1 && foundInRemovedDep == false) { return hipErrorStreamCaptureUnjoined; } } else { @@ -1098,8 +1116,8 @@ hipError_t hipGraphAddKernelNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, HIP_RETURN(hipErrorInvalidValue); } - HIP_RETURN_DURATION( - ihipGraphAddKernelNode(pGraphNode, graph, pDependencies, numDependencies, pNodeParams)); + HIP_RETURN_DURATION(ihipGraphAddKernelNode(pGraphNode, graph, pDependencies, numDependencies, + pNodeParams, false)); } hipError_t hipGraphAddMemcpyNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, @@ -1112,8 +1130,8 @@ hipError_t hipGraphAddMemcpyNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, HIP_RETURN(hipErrorInvalidValue); } - HIP_RETURN_DURATION( - ihipGraphAddMemcpyNode(pGraphNode, graph, pDependencies, numDependencies, pCopyParams)); + HIP_RETURN_DURATION(ihipGraphAddMemcpyNode(pGraphNode, graph, pDependencies, numDependencies, + pCopyParams, false)); } hipError_t hipGraphAddMemcpyNode1D(hipGraphNode_t* pGraphNode, hipGraph_t graph, @@ -1127,7 +1145,7 @@ hipError_t hipGraphAddMemcpyNode1D(hipGraphNode_t* pGraphNode, hipGraph_t graph, } HIP_RETURN_DURATION(ihipGraphAddMemcpyNode1D(pGraphNode, graph, pDependencies, numDependencies, - dst, src, count, kind)); + dst, src, count, kind, false)); } hipError_t hipGraphMemcpyNodeSetParams1D(hipGraphNode_t node, void* dst, const void* src, @@ -1167,8 +1185,8 @@ hipError_t hipGraphAddMemsetNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, HIP_RETURN(hipErrorInvalidValue); } - HIP_RETURN_DURATION( - ihipGraphAddMemsetNode(pGraphNode, graph, pDependencies, numDependencies, pMemsetParams)); + HIP_RETURN_DURATION(ihipGraphAddMemsetNode(pGraphNode, graph, pDependencies, numDependencies, + pMemsetParams, false)); } hipError_t hipGraphAddEmptyNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, @@ -1179,7 +1197,7 @@ hipError_t hipGraphAddEmptyNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, HIP_RETURN(hipErrorInvalidValue); } *pGraphNode = new hipGraphEmptyNode(); - hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, false); HIP_RETURN(status); } @@ -1192,7 +1210,7 @@ hipError_t hipGraphAddChildGraphNode(hipGraphNode_t* pGraphNode, hipGraph_t grap HIP_RETURN(hipErrorInvalidValue); } *pGraphNode = new hipChildGraphNode(childGraph); - hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, false); HIP_RETURN(status); } @@ -1432,7 +1450,8 @@ hipError_t hipGraphMemsetNodeSetParams(hipGraphNode_t node, const hipMemsetParam if (!hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) { HIP_RETURN(hipErrorInvalidValue); } - if (pNodeParams->height > 1 && pNodeParams->pitch < (pNodeParams->width * pNodeParams->elementSize)) { + if (pNodeParams->height > 1 && + pNodeParams->pitch < (pNodeParams->width * pNodeParams->elementSize)) { return hipErrorInvalidValue; } HIP_RETURN(reinterpret_cast(node)->SetParams(pNodeParams)); @@ -1852,7 +1871,7 @@ hipError_t hipGraphAddMemcpyNodeFromSymbol(hipGraphNode_t* pGraphNode, hipGraph_ HIP_RETURN(status); } *pGraphNode = new hipGraphMemcpyNodeFromSymbol(dst, symbol, count, offset, kind); - status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, false); HIP_RETURN(status); } @@ -1911,7 +1930,7 @@ hipError_t hipGraphAddMemcpyNodeToSymbol(hipGraphNode_t* pGraphNode, hipGraph_t if (*pGraphNode == nullptr) { HIP_RETURN(hipErrorInvalidValue); } - status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, false); HIP_RETURN(status); } @@ -1962,7 +1981,7 @@ hipError_t hipGraphAddEventRecordNode(hipGraphNode_t* pGraphNode, hipGraph_t gra HIP_RETURN(hipErrorInvalidValue); } *pGraphNode = new hipGraphEventRecordNode(event); - hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, false); HIP_RETURN(status); } @@ -2006,7 +2025,7 @@ hipError_t hipGraphAddEventWaitNode(hipGraphNode_t* pGraphNode, hipGraph_t graph HIP_RETURN(hipErrorInvalidValue); } *pGraphNode = new hipGraphEventWaitNode(event); - hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, false); HIP_RETURN(status); } @@ -2051,7 +2070,7 @@ hipError_t hipGraphAddHostNode(hipGraphNode_t* pGraphNode, hipGraph_t graph, } *pGraphNode = new hipGraphHostNode(pNodeParams); - hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies); + hipError_t status = ihipGraphAddNode(*pGraphNode, graph, pDependencies, numDependencies, false); HIP_RETURN(status); } diff --git a/projects/clr/hipamd/src/hip_graph_internal.hpp b/projects/clr/hipamd/src/hip_graph_internal.hpp index 838312e890..6462535d41 100644 --- a/projects/clr/hipamd/src/hip_graph_internal.hpp +++ b/projects/clr/hipamd/src/hip_graph_internal.hpp @@ -380,6 +380,7 @@ struct ihipGraph { static int nextID; hip::Device* device_; //!< HIP device object hip::MemoryPool* mem_pool_; //!< Memory pool, associated with this graph + std::unordered_set capturedNodes_; public: ihipGraph(hip::Device* device, const ihipGraph* original = nullptr) @@ -415,6 +416,14 @@ struct ihipGraph { }; + void AddManualNodeDuringCapture(hipGraphNode* node) { capturedNodes_.insert(node); } + + std::unordered_set GetManualNodesDuringCapture() { return capturedNodes_; } + + void RemoveManualNodesDuringCapture() { + capturedNodes_.erase(capturedNodes_.begin(), capturedNodes_.end()); + } + /// Return graph unique ID int GetID() const { return id_; } diff --git a/projects/clr/hipamd/src/hip_internal.hpp b/projects/clr/hipamd/src/hip_internal.hpp index be90fe5b1c..99bdf00d49 100644 --- a/projects/clr/hipamd/src/hip_internal.hpp +++ b/projects/clr/hipamd/src/hip_internal.hpp @@ -168,7 +168,7 @@ static amd::Monitor g_hipInitlock{"hipInit lock"}; } \ amd::ScopedLock lock(g_captureStreamsLock); \ if (g_captureStreams.size() != 0) { \ - HIP_RETURN(hipErrorStreamCaptureUnsupported); \ + HIP_RETURN(hipErrorStreamCaptureUnsupported); \ } \ }