SWDEV-522841 - Graph nodes must be created/launched on device where they are captured/created (#108)

This commit is contained in:
Godavarthy Surya, Anusha
2025-04-29 22:20:39 +05:30
committed by GitHub
parent eb62fe9f62
commit 2538d7f02b
5 changed files with 55 additions and 53 deletions
+10 -12
View File
@@ -359,9 +359,6 @@ hipError_t GraphExec::Init() {
}
if (DEBUG_CLR_GRAPH_PACKET_CAPTURE) {
if (max_streams_ == 1) {
// Don't wait for other streams to finish.
// Capture stream is to capture AQL packet.
capture_stream_ = hip::getNullStream(false);
// For graph nodes capture AQL packets to dispatch them directly during graph launch.
status = CaptureAQLPackets();
}
@@ -386,8 +383,6 @@ void GraphExec::GetKernelArgSizeForGraph(size_t& kernArgSizeForGraph) {
GraphKernelArgManager* KernelArgManager = GetKernelArgManager();
KernelArgManager->retain();
childNode->SetKernelArgManager(KernelArgManager);
// Set capture stream for child graph
childNode->capture_stream_ = capture_stream_;
if (childNode->GetChildGraph()->max_streams_ == 1) {
childNode->GetKernelArgSizeForGraph(kernArgSizeForGraph);
}
@@ -408,7 +403,7 @@ hipError_t GraphExec::AllocKernelArgForGraphNode() {
}
}
if (node->GraphCaptureEnabled()) {
status = node->CaptureAndFormPacket(capture_stream_, GetKernelArgManager());
status = node->CaptureAndFormPacket(GetKernelArgManager());
} else if (node->GetType() == hipGraphNodeTypeGraph) {
auto childNode = reinterpret_cast<hip::ChildGraphNode*>(node);
if (childNode->GetChildGraph()->max_streams_ == 1) {
@@ -428,9 +423,13 @@ hipError_t GraphExec::CaptureAQLPackets() {
hipError_t status = hipSuccess;
size_t kernArgSizeForGraph = 0;
GetKernelArgSizeForGraph(kernArgSizeForGraph);
auto device = g_devices[ihipGetDevice()]->devices()[0];
// When we support multi device graph lauch we need to allocate the kenel args on respective
// device for each kernel Assume graph has nodes of same device allocate kernel args on the device
// from the first node
auto device = g_devices[topoOrder_[0]->GetDeviceId()]->devices()[0];
// Add a larger initial pool to accomodate for any updates to kernel args
bool bStatus = kernArgManager_->AllocGraphKernargPool(kernArgSizeForGraph + kKernArgChunkSize);
bool bStatus =
kernArgManager_->AllocGraphKernargPool(kernArgSizeForGraph + kKernArgChunkSize, device);
if (bStatus != true) {
return hipErrorMemoryAllocation;
}
@@ -447,7 +446,7 @@ hipError_t GraphExec::CaptureAQLPackets() {
hipError_t GraphExec::UpdateAQLPacket(hip::GraphNode* node) {
hipError_t status = hipSuccess;
if (max_streams_ == 1) {
status = node->CaptureAndFormPacket(capture_stream_, kernArgManager_);
status = node->CaptureAndFormPacket(kernArgManager_);
}
return status;
}
@@ -749,11 +748,10 @@ hipError_t GraphExec::Run(hip::Stream* launch_stream) {
}
// ================================================================================================
bool GraphKernelArgManager::AllocGraphKernargPool(size_t pool_size) {
bool GraphKernelArgManager::AllocGraphKernargPool(size_t pool_size, amd::Device* device) {
bool bStatus = true;
assert(pool_size > 0);
address graph_kernarg_base;
auto device = g_devices[ihipGetDevice()]->devices()[0];
// Current device is stored as part of tls. Save current device to destroy kernelArgs from the
// callback thread.
device_ = device;
@@ -783,7 +781,7 @@ address GraphKernelArgManager::AllocKernArg(size_t size, size_t alignment) {
kernarg_graph_.back().kernarg_pool_offset_ = pool_new_usage;
} else {
// If current chunck is full allocate new chunck with same size as current
bool bStatus = AllocGraphKernargPool(kernarg_graph_.back().kernarg_pool_size_);
bool bStatus = AllocGraphKernargPool(kernarg_graph_.back().kernarg_pool_size_, device_);
if (bStatus == false) {
return nullptr;
} else {