SWDEV-556588 - Handle graph node set params and disabled nodes for AQL packet batching. (#1099)
This commit is contained in:
@@ -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);
|
||||
}
|
||||
|
||||
|
||||
@@ -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(), ¤tBatch,
|
||||
¤tKernelNames);
|
||||
|
||||
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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user