From 98d6d268a0496c017b3d7d0e97c928e6f502b668 Mon Sep 17 00:00:00 2001 From: Jaydeep <106300970+jaydeeppatel1111@users.noreply.github.com> Date: Mon, 6 Oct 2025 13:26:12 +0530 Subject: [PATCH] SWDEV-556588 - Handle graph node set params and disabled nodes for AQL packet batching. (#1099) --- projects/clr/hipamd/src/hip_graph.cpp | 5 + .../clr/hipamd/src/hip_graph_internal.cpp | 175 ++++++++++++++---- .../clr/hipamd/src/hip_graph_internal.hpp | 26 ++- projects/clr/rocclr/device/blit.cpp | 9 + 4 files changed, 175 insertions(+), 40 deletions(-) diff --git a/projects/clr/hipamd/src/hip_graph.cpp b/projects/clr/hipamd/src/hip_graph.cpp index b730605e29..b23952e984 100644 --- a/projects/clr/hipamd/src/hip_graph.cpp +++ b/projects/clr/hipamd/src/hip_graph.cpp @@ -3084,6 +3084,11 @@ hipError_t hipGraphNodeSetEnabled(hipGraphExec_t hGraphExec, hipGraphNode_t hNod HIP_RETURN(hipErrorInvalidValue); } clonedNode->SetEnabled(isEnabled); + // Update packet batches when node is enabled/disabled + hipError_t status = graphExec->UpdatePacketBatchesForNodeEnableDisable(clonedNode, isEnabled != 0); + if (status != hipSuccess) { + HIP_RETURN(status); + } HIP_RETURN(hipSuccess); } diff --git a/projects/clr/hipamd/src/hip_graph_internal.cpp b/projects/clr/hipamd/src/hip_graph_internal.cpp index 5a63202e53..7363f73f58 100644 --- a/projects/clr/hipamd/src/hip_graph_internal.cpp +++ b/projects/clr/hipamd/src/hip_graph_internal.cpp @@ -398,6 +398,30 @@ void GraphExec::GetKernelArgSizeForGraph(std::unordered_map& kernAr } } } +// ================================================================================================ +// Enable or disable a graph node's packets in the batch +// Simply updates the enabled state and count of disabled nodes +// ================================================================================================ +void GraphExec::PacketBatch::setEnabled(GraphNode* node, bool enabled) { + auto it = nodeToRangeIndex.find(node); + if (it == nodeToRangeIndex.end()) { + return; + } + NodeRange& range = nodeRanges[it->second]; + // Early return if state hasn't changed + if (range.enabled == enabled) { + return; + } + // Update counter based on state change + if (enabled) { + // Node being enabled: decrement counter + disabledNodeCount--; + } else { + // Node being disabled: increment counter + disabledNodeCount++; + } + range.enabled = enabled; +} // ================================================================================================ hipError_t GraphExec::CaptureAndFormPacketsForGraph() { @@ -424,32 +448,49 @@ hipError_t GraphExec::CaptureAndFormPacketsForGraph() { } if (node->GraphCaptureEnabled()) { - // Start of a potential batch - try to capture packets for this node - std::vector currentBatch; - std::vector currentKernelNames; + // Start of a new batch + PacketBatch newBatch; + size_t j = i; // Collect packets from consecutive captured nodes - size_t j = i; - size_t capturedNodeCount = 0; while (j < topoOrder_.size() && topoOrder_[j]->GraphCaptureEnabled()) { auto& currentNode = topoOrder_[j]; - status = currentNode->CaptureAndFormPacket(GetKernelArgManager(), ¤tBatch, - ¤tKernelNames); - if (status != hipSuccess || currentBatch.empty()) { + // Capture packets for this node + std::vector nodePackets; + std::vector nodeKernelNames; + status = currentNode->CaptureAndFormPacket(GetKernelArgManager(), &nodePackets, + &nodeKernelNames); + + if (status != hipSuccess || nodePackets.empty()) { LogError("Packet capture failed"); return status; } + + // Create NodeRange for this node + PacketBatch::NodeRange range; + range.startIndex = newBatch.dispatchPackets.size(); + range.packetCount = nodePackets.size(); + range.enabled = true; + + // Add to dispatch lists (initially all enabled) + newBatch.dispatchPackets.insert(newBatch.dispatchPackets.end(), + nodePackets.begin(), nodePackets.end()); + newBatch.dispatchKernelNames.insert(newBatch.dispatchKernelNames.end(), + nodeKernelNames.begin(), nodeKernelNames.end()); + + // Store node mapping + newBatch.nodeRanges.push_back(range); + newBatch.nodeToRangeIndex[currentNode] = newBatch.nodeRanges.size() - 1; + // Mark this node as successfully captured nodeCaptureStatus_[j] = true; ++j; - ++capturedNodeCount; } // Add the batch if it has packets - if (!currentBatch.empty()) { - packetBatches_.emplace_back(std::move(currentBatch), std::move(currentKernelNames), - capturedNodeCount); + if (!newBatch.dispatchPackets.empty()) { + packetBatches_.emplace_back(std::move(newBatch)); } // Skip the nodes we just processed, the index will be incremented by the loop @@ -466,7 +507,6 @@ hipError_t GraphExec::CaptureAndFormPacketsForGraph() { } } } - return status; } @@ -512,11 +552,53 @@ hipError_t GraphExec::CaptureAQLPackets() { // ================================================================================================ hipError_t GraphExec::UpdateAQLPacket(hip::GraphNode* node) { - hipError_t status = hipSuccess; - if (max_streams_ == 1 && node->GraphCaptureEnabled()) { - status = node->CaptureAndFormPacket(kernArgManager_); + if (max_streams_ != 1 || !node->GraphCaptureEnabled()) { + return hipSuccess; } - return status; + + // Find which batch contains this node and update it + for (auto& batch : packetBatches_) { + auto it = batch.nodeToRangeIndex.find(node); + if (it != batch.nodeToRangeIndex.end()) { + // Found the batch containing this node - update packets + PacketBatch::NodeRange& range = batch.nodeRanges[it->second]; + + // Capture new packets for this node + std::vector newPackets; + std::vector newKernelNames; + hipError_t status = node->CaptureAndFormPacket(kernArgManager_, &newPackets, &newKernelNames); + if (status != hipSuccess) { + return status; + } + // Update dispatch packets (always update regardless of enabled state) + // The enabled/disabled check happens during dispatch, not here + for (size_t i = 0; i < range.packetCount && i < newPackets.size(); ++i) { + size_t packetIndex = range.startIndex + i; + batch.dispatchPackets[packetIndex] = newPackets[i]; + batch.dispatchKernelNames[packetIndex] = newKernelNames[i]; + } + return hipSuccess; + } + } + return hipSuccess; // Node not in any batch +} + +// ================================================================================================ +hipError_t GraphExec::UpdatePacketBatchesForNodeEnableDisable(hip::GraphNode* node, bool isEnabled) { + if (max_streams_ != 1 || !node->GraphCaptureEnabled()) { + // Only handle single stream case with captured nodes + return hipSuccess; + } + // Find which batch contains this node and update its enabled state + for (auto& batch : packetBatches_) { + auto it = batch.nodeToRangeIndex.find(node); + if (it != batch.nodeToRangeIndex.end()) { + // Found the batch containing this node - update enabled state + batch.setEnabled(node, isEnabled); + return hipSuccess; + } + } + return hipSuccess; // Node not in any batch } // ================================================================================================ @@ -544,26 +626,55 @@ hipError_t GraphExec::EnqueueGraphWithSingleList(hip::Stream* hip_stream) { auto& node = topoOrder_[i]; if (!node->GraphCaptureEnabled()) { - // Node doesn't support capture - execute individually - node->SetStream(hip_stream); - status = node->CreateCommand(node->GetQueue()); - node->EnqueueCommands(hip_stream); + // Node doesn't support capture - execute individually if enabled + if (node->GetEnabled() != 0) { + node->SetStream(hip_stream); + status = node->CreateCommand(node->GetQueue()); + node->EnqueueCommands(hip_stream); + } } else if (i < nodeCaptureStatus_.size() && nodeCaptureStatus_[i]) { - // Node was successfully captured - find which batch it belongs to - // and dispatch the entire batch + // Node was successfully captured - dispatch the batch with enabled nodes only if (batchIndex < packetBatches_.size()) { - // Dispatch this batch - bool batchStatus = hip_stream->vdev()->dispatchAqlPacketBatch( - packetBatches_[batchIndex].packets, packetBatches_[batchIndex].kernelNames, accumulate); - if (!batchStatus) { - status = hipErrorUnknown; - accumulate->release(); - return status; + const auto& batch = packetBatches_[batchIndex]; + // O(1) check: if no disabled nodes, dispatch entire batch directly + // This avoids creating new vectors when all nodes are enabled (common case) + if (batch.disabledNodeCount == 0) { + // Fast path: all nodes enabled, dispatch entire batch + bool batchStatus = hip_stream->vdev()->dispatchAqlPacketBatch( + batch.dispatchPackets, batch.dispatchKernelNames, accumulate); + if (!batchStatus) { + status = hipErrorUnknown; + accumulate->release(); + return status; + } + } else { + // Slow path: some nodes disabled, create filtered vectors + std::vector enabledPackets; + std::vector enabledKernelNames; + for (const auto& range : batch.nodeRanges) { + if (range.enabled) { + // Add packets for this enabled node + for (size_t j = 0; j < range.packetCount; ++j) { + size_t packetIndex = range.startIndex + j; + enabledPackets.push_back(batch.dispatchPackets[packetIndex]); + enabledKernelNames.push_back(batch.dispatchKernelNames[packetIndex]); + } + } + } + // Only dispatch if there are enabled packets + if (!enabledPackets.empty()) { + bool batchStatus = hip_stream->vdev()->dispatchAqlPacketBatch( + enabledPackets, enabledKernelNames, accumulate); + if (!batchStatus) { + status = hipErrorUnknown; + accumulate->release(); + return status; + } + } } // Skip all consecutive captured nodes that belong to this batch - // Use the tracked node count to skip directly instead of parsing one by one - i += packetBatches_[batchIndex].capturedNodeCount - 1; // -1 because loop will increment + i += packetBatches_[batchIndex].nodeRanges.size() - 1; // -1 because loop will increment ++batchIndex; } diff --git a/projects/clr/hipamd/src/hip_graph_internal.hpp b/projects/clr/hipamd/src/hip_graph_internal.hpp index 0d212e38e9..309802da3f 100644 --- a/projects/clr/hipamd/src/hip_graph_internal.hpp +++ b/projects/clr/hipamd/src/hip_graph_internal.hpp @@ -861,6 +861,8 @@ class GraphExec : public amd::ReferenceCountedObject, public Graph { // Capture GPU Packets from graph commands hipError_t CaptureAQLPackets(); hipError_t UpdateAQLPacket(hip::GraphNode* node); + // Handle packetBatches_ updates when nodes are enabled/disabled + hipError_t UpdatePacketBatchesForNodeEnableDisable(hip::GraphNode* node, bool isEnabled); // Kenrel arg manger is for the entire graph. // Child graph also shares the same kernel arg manager object. some apps have 100's of // child graph nodes and each child graph has only one node. @@ -884,15 +886,23 @@ class GraphExec : public amd::ReferenceCountedObject, public Graph { bool hasHiddenHeap_ = false; //!< Hidden heap indicator for Kernel node bool repeatLaunch_ = false; - //! Structure for batch dispatch optimization - packets and kernel names in aligned memory + // PacketBatch structure struct PacketBatch { - std::vector packets; - std::vector kernelNames; - size_t capturedNodeCount; // Number of consecutive captured nodes in this batch - - PacketBatch() : capturedNodeCount(0) {} - PacketBatch(std::vector&& p, std::vector&& k, size_t nodeCount) - : packets(std::move(p)), kernelNames(std::move(k)), capturedNodeCount(nodeCount) {} + // Main dispatch vectors - always ready for batch dispatch + std::vector dispatchPackets; + std::vector dispatchKernelNames; + // Node tracking + struct NodeRange { + size_t startIndex; // Start index in dispatchPackets + size_t packetCount; // Number of packets for this node + bool enabled; // Node enabled state (checked during dispatch) + }; + std::vector nodeRanges; + std::unordered_map nodeToRangeIndex; // O(1) lookup + int disabledNodeCount = 0; // Count of currently disabled nodes + PacketBatch() {} + // O(1) enable/disable operations - just update state + void setEnabled(GraphNode* node, bool enabled); }; //! Batches of accumulated packets and kernel names for batch dispatch optimization diff --git a/projects/clr/rocclr/device/blit.cpp b/projects/clr/rocclr/device/blit.cpp index d58ae48395..749c895c4d 100644 --- a/projects/clr/rocclr/device/blit.cpp +++ b/projects/clr/rocclr/device/blit.cpp @@ -735,6 +735,15 @@ void HostBlitManager::FillBufferInfo::PackInfo(const device::Memory& memory, siz guarantee(fill_size >= pattern_size, "Pattern Size: %u cannot be greater than fill size: %u \n", pattern_size, fill_size); + constexpr bool kDisablePackingOptimization = true; + // Check if packing optimization is disabled + if (kDisablePackingOptimization) { + // Simple case: create a single FillBufferInfo without alignment optimization + FillBufferInfo fill_info(fill_size); + packed_info.push_back(fill_info); + return; + } + // 2. Calculate the next closest dword aligned address for faster processing size_t dst_addr = memory.virtualAddress() + fill_origin; size_t aligned_dst_addr = amd::alignUp(dst_addr, kExtendedSize);