SWDEV-556588 - Handle graph node set params and disabled nodes for AQL packet batching. (#1099)

This commit is contained in:
Jaydeep
2025-10-06 13:26:12 +05:30
committed by GitHub
parent 02883c3d8d
commit 98d6d268a0
4 changed files with 175 additions and 40 deletions
+5
View File
@@ -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);
}
+143 -32
View File
@@ -398,6 +398,30 @@ void GraphExec::GetKernelArgSizeForGraph(std::unordered_map<int, size_t>& 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<uint8_t*> currentBatch;
std::vector<std::string> 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(), &currentBatch,
&currentKernelNames);
if (status != hipSuccess || currentBatch.empty()) {
// Capture packets for this node
std::vector<uint8_t*> nodePackets;
std::vector<std::string> 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<uint8_t*> newPackets;
std::vector<std::string> 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<uint8_t*> enabledPackets;
std::vector<std::string> 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;
}
+18 -8
View File
@@ -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<uint8_t*> packets;
std::vector<std::string> kernelNames;
size_t capturedNodeCount; // Number of consecutive captured nodes in this batch
PacketBatch() : capturedNodeCount(0) {}
PacketBatch(std::vector<uint8_t*>&& p, std::vector<std::string>&& k, size_t nodeCount)
: packets(std::move(p)), kernelNames(std::move(k)), capturedNodeCount(nodeCount) {}
// Main dispatch vectors - always ready for batch dispatch
std::vector<uint8_t*> dispatchPackets;
std::vector<std::string> 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<NodeRange> nodeRanges;
std::unordered_map<GraphNode*, size_t> 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
+9
View File
@@ -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);