SWDEV-522841 - Graph nodes must be created/launched on device where they are captured/created (#108)
This commit is contained in:
committed by
GitHub
parent
eb62fe9f62
commit
2538d7f02b
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user