SWDEV-323441 - Phase-II : per thread default stream
Change-Id: I3c796ddaebcf0223d7faf50c425c1674de215f9d
This commit is contained in:
committato da
Sarbojit Sarkar
parent
2b9e39e901
commit
e9961fedd8
+111
-39
@@ -726,17 +726,27 @@ hipError_t capturehipLaunchHostFunc(hipStream_t& stream, hipHostFn_t& fn, void*&
|
||||
return hipSuccess;
|
||||
}
|
||||
|
||||
hipError_t hipStreamIsCapturing(hipStream_t stream, hipStreamCaptureStatus* pCaptureStatus) {
|
||||
HIP_INIT_API(hipStreamIsCapturing, stream, pCaptureStatus);
|
||||
hipError_t hipStreamIsCapturing_common(hipStream_t stream, hipStreamCaptureStatus* pCaptureStatus) {
|
||||
if (pCaptureStatus == nullptr || !hip::isValid(stream)) {
|
||||
HIP_RETURN(hipErrorInvalidValue);
|
||||
return hipErrorInvalidValue;
|
||||
}
|
||||
if (stream == nullptr) {
|
||||
*pCaptureStatus = hipStreamCaptureStatusNone;
|
||||
} else {
|
||||
*pCaptureStatus = reinterpret_cast<hip::Stream*>(stream)->GetCaptureStatus();
|
||||
}
|
||||
HIP_RETURN(hipSuccess);
|
||||
return hipSuccess;
|
||||
}
|
||||
|
||||
hipError_t hipStreamIsCapturing(hipStream_t stream, hipStreamCaptureStatus* pCaptureStatus) {
|
||||
HIP_INIT_API(hipStreamIsCapturing, stream, pCaptureStatus);
|
||||
HIP_RETURN(hipStreamIsCapturing_common(stream, pCaptureStatus));
|
||||
}
|
||||
|
||||
hipError_t hipStreamIsCapturing_spt(hipStream_t stream, hipStreamCaptureStatus* pCaptureStatus) {
|
||||
HIP_INIT_API(hipStreamIsCapturing, stream, pCaptureStatus);
|
||||
PER_THREAD_DEFAULT_STREAM(stream);
|
||||
HIP_RETURN(hipStreamIsCapturing_common(stream, pCaptureStatus));
|
||||
}
|
||||
|
||||
hipError_t hipThreadExchangeStreamCaptureMode(hipStreamCaptureMode* mode) {
|
||||
@@ -754,22 +764,22 @@ hipError_t hipThreadExchangeStreamCaptureMode(hipStreamCaptureMode* mode) {
|
||||
HIP_RETURN_DURATION(hipSuccess);
|
||||
}
|
||||
|
||||
hipError_t hipStreamBeginCapture(hipStream_t stream, hipStreamCaptureMode mode) {
|
||||
HIP_INIT_API(hipStreamBeginCapture, stream, mode);
|
||||
hipError_t hipStreamBeginCapture_common(hipStream_t stream, hipStreamCaptureMode mode) {
|
||||
if (!hip::isValid(stream)) {
|
||||
HIP_RETURN(hipErrorInvalidValue);
|
||||
return hipErrorInvalidValue;
|
||||
}
|
||||
// capture cannot be initiated on legacy stream
|
||||
if (stream == nullptr) {
|
||||
HIP_RETURN(hipErrorStreamCaptureUnsupported);
|
||||
return hipErrorStreamCaptureUnsupported;
|
||||
}
|
||||
if (mode < hipStreamCaptureModeGlobal || mode > hipStreamCaptureModeRelaxed) {
|
||||
HIP_RETURN(hipErrorInvalidValue);
|
||||
if (mode < hipStreamCaptureModeGlobal ||
|
||||
mode > hipStreamCaptureModeRelaxed) {
|
||||
return hipErrorInvalidValue;
|
||||
}
|
||||
hip::Stream* s = reinterpret_cast<hip::Stream*>(stream);
|
||||
// It can be initiated if the stream is not already in capture mode
|
||||
if (s->GetCaptureStatus() == hipStreamCaptureStatusActive) {
|
||||
HIP_RETURN(hipErrorIllegalState);
|
||||
return hipErrorIllegalState;
|
||||
}
|
||||
|
||||
s->SetCaptureGraph(new ihipGraph());
|
||||
@@ -782,35 +792,45 @@ hipError_t hipStreamBeginCapture(hipStream_t stream, hipStreamCaptureMode mode)
|
||||
amd::ScopedLock lock(g_captureStreamsLock);
|
||||
g_captureStreams.push_back(s);
|
||||
}
|
||||
HIP_RETURN_DURATION(hipSuccess);
|
||||
return hipSuccess;
|
||||
}
|
||||
|
||||
hipError_t hipStreamEndCapture(hipStream_t stream, hipGraph_t* pGraph) {
|
||||
HIP_INIT_API(hipStreamEndCapture, stream, pGraph);
|
||||
hipError_t hipStreamBeginCapture(hipStream_t stream, hipStreamCaptureMode mode) {
|
||||
HIP_INIT_API(hipStreamBeginCapture, stream, mode);
|
||||
HIP_RETURN_DURATION(hipStreamBeginCapture_common(stream, mode));
|
||||
}
|
||||
|
||||
hipError_t hipStreamBeginCapture_spt(hipStream_t stream, hipStreamCaptureMode mode) {
|
||||
HIP_INIT_API(hipStreamBeginCapture, stream, mode);
|
||||
PER_THREAD_DEFAULT_STREAM(stream);
|
||||
HIP_RETURN_DURATION(hipStreamBeginCapture_common(stream, mode));
|
||||
}
|
||||
|
||||
hipError_t hipStreamEndCapture_common(hipStream_t stream, hipGraph_t* pGraph) {
|
||||
if (pGraph == nullptr) {
|
||||
HIP_RETURN(hipErrorInvalidValue);
|
||||
return hipErrorInvalidValue;
|
||||
}
|
||||
if (stream == nullptr) {
|
||||
HIP_RETURN(hipErrorIllegalState);
|
||||
return hipErrorIllegalState;
|
||||
}
|
||||
if (!hip::isValid(stream)) {
|
||||
HIP_RETURN(hipErrorContextIsDestroyed);
|
||||
return hipErrorContextIsDestroyed;
|
||||
}
|
||||
hip::Stream* s = reinterpret_cast<hip::Stream*>(stream);
|
||||
// Capture status must be active before endCapture can be initiated
|
||||
if (s->GetCaptureStatus() == hipStreamCaptureStatusNone) {
|
||||
HIP_RETURN(hipErrorIllegalState);
|
||||
return hipErrorIllegalState;
|
||||
}
|
||||
// Capture must be ended on the same stream in which it was initiated
|
||||
if (!s->IsOriginStream()) {
|
||||
HIP_RETURN(hipErrorStreamCaptureUnmatched);
|
||||
return hipErrorStreamCaptureUnmatched;
|
||||
}
|
||||
// If mode is not hipStreamCaptureModeRelaxed, hipStreamEndCapture must be called on the stream
|
||||
// from the same thread
|
||||
const auto& it = std::find(l_captureStreams.begin(), l_captureStreams.end(), s);
|
||||
if (s->GetCaptureMode() != hipStreamCaptureModeRelaxed) {
|
||||
if (it == l_captureStreams.end()) {
|
||||
HIP_RETURN(hipErrorStreamCaptureWrongThread);
|
||||
return hipErrorStreamCaptureWrongThread;
|
||||
}
|
||||
l_captureStreams.erase(it);
|
||||
}
|
||||
@@ -821,7 +841,7 @@ hipError_t hipStreamEndCapture(hipStream_t stream, hipGraph_t* pGraph) {
|
||||
// If capture was invalidated, due to a violation of the rules of stream capture
|
||||
if (s->GetCaptureStatus() == hipStreamCaptureStatusInvalidated) {
|
||||
*pGraph = nullptr;
|
||||
HIP_RETURN(hipErrorStreamCaptureInvalidated);
|
||||
return hipErrorStreamCaptureInvalidated;
|
||||
}
|
||||
// check if all parallel streams have joined
|
||||
// Nodes that are removed from the dependency set via API hipStreamUpdateCaptureDependencies do
|
||||
@@ -838,12 +858,23 @@ hipError_t hipStreamEndCapture(hipStream_t stream, hipGraph_t* pGraph) {
|
||||
}
|
||||
}
|
||||
if (foundInRemovedDep == false) {
|
||||
HIP_RETURN(hipErrorStreamCaptureUnjoined);
|
||||
return hipErrorStreamCaptureUnjoined;
|
||||
}
|
||||
}
|
||||
*pGraph = s->GetCaptureGraph();
|
||||
// end capture on all streams/events part of graph capture
|
||||
HIP_RETURN_DURATION(s->EndCapture());
|
||||
return s->EndCapture();
|
||||
}
|
||||
|
||||
hipError_t hipStreamEndCapture(hipStream_t stream, hipGraph_t* pGraph) {
|
||||
HIP_INIT_API(hipStreamEndCapture, stream, pGraph);
|
||||
HIP_RETURN_DURATION(hipStreamEndCapture_common(stream, pGraph));
|
||||
}
|
||||
|
||||
hipError_t hipStreamEndCapture_spt(hipStream_t stream, hipGraph_t* pGraph) {
|
||||
HIP_INIT_API(hipStreamEndCapture, stream, pGraph);
|
||||
PER_THREAD_DEFAULT_STREAM(stream);
|
||||
HIP_RETURN_DURATION(hipStreamEndCapture_common(stream, pGraph));
|
||||
}
|
||||
|
||||
hipError_t hipGraphCreate(hipGraph_t* pGraph, unsigned int flags) {
|
||||
@@ -1038,12 +1069,22 @@ hipError_t ihipGraphLaunch(hipGraphExec_t graphExec, hipStream_t stream) {
|
||||
return graphExec->Run(stream);
|
||||
}
|
||||
|
||||
hipError_t hipGraphLaunch_common( hipGraphExec_t graphExec, hipStream_t stream) {
|
||||
if (graphExec == nullptr || !hip::isValid(stream) || !hipGraphExec::isGraphExecValid(graphExec)) {
|
||||
return hipErrorInvalidValue;
|
||||
}
|
||||
return ihipGraphLaunch(graphExec, stream);
|
||||
}
|
||||
|
||||
hipError_t hipGraphLaunch(hipGraphExec_t graphExec, hipStream_t stream) {
|
||||
HIP_INIT_API(hipGraphLaunch, graphExec, stream);
|
||||
if (graphExec == nullptr || !hip::isValid(stream) || !hipGraphExec::isGraphExecValid(graphExec)) {
|
||||
HIP_RETURN(hipErrorInvalidValue);
|
||||
}
|
||||
HIP_RETURN_DURATION(ihipGraphLaunch(graphExec, stream));
|
||||
HIP_RETURN_DURATION(hipGraphLaunch_common(graphExec, stream));
|
||||
}
|
||||
|
||||
hipError_t hipGraphLaunch_spt(hipGraphExec_t graphExec, hipStream_t stream) {
|
||||
HIP_INIT_API(hipGraphLaunch, graphExec, stream);
|
||||
PER_THREAD_DEFAULT_STREAM(stream);
|
||||
HIP_RETURN_DURATION(hipGraphLaunch_common(graphExec, stream));
|
||||
}
|
||||
|
||||
hipError_t hipGraphGetNodes(hipGraph_t graph, hipGraphNode_t* nodes, size_t* numNodes) {
|
||||
@@ -1310,37 +1351,47 @@ hipError_t hipGraphExecChildGraphNodeSetParams(hipGraphExec_t hGraphExec, hipGra
|
||||
HIP_RETURN(reinterpret_cast<hipChildGraphNode*>(clonedNode)->SetParams(childGraph));
|
||||
}
|
||||
|
||||
hipError_t hipStreamGetCaptureInfo(hipStream_t stream, hipStreamCaptureStatus* pCaptureStatus,
|
||||
hipError_t hipStreamGetCaptureInfo_common(hipStream_t stream, hipStreamCaptureStatus* pCaptureStatus,
|
||||
unsigned long long* pId) {
|
||||
HIP_INIT_API(hipStreamGetCaptureInfo, stream, pCaptureStatus, pId);
|
||||
if (pCaptureStatus == nullptr || !hip::isValid(stream)) {
|
||||
HIP_RETURN(hipErrorInvalidValue);
|
||||
return hipErrorInvalidValue;
|
||||
}
|
||||
if (stream == nullptr) {
|
||||
HIP_RETURN(hipErrorUnknown);
|
||||
return hipErrorUnknown;
|
||||
}
|
||||
hip::Stream* s = reinterpret_cast<hip::Stream*>(stream);
|
||||
*pCaptureStatus = s->GetCaptureStatus();
|
||||
if (*pCaptureStatus == hipStreamCaptureStatusActive && pId != nullptr) {
|
||||
*pId = s->GetCaptureID();
|
||||
}
|
||||
HIP_RETURN(hipSuccess);
|
||||
return hipSuccess;
|
||||
}
|
||||
|
||||
hipError_t hipStreamGetCaptureInfo_v2(hipStream_t stream, hipStreamCaptureStatus* captureStatus_out,
|
||||
hipError_t hipStreamGetCaptureInfo(hipStream_t stream, hipStreamCaptureStatus* pCaptureStatus,
|
||||
unsigned long long* pId) {
|
||||
HIP_INIT_API(hipStreamGetCaptureInfo, stream, pCaptureStatus, pId);
|
||||
HIP_RETURN(hipStreamGetCaptureInfo_common(stream, pCaptureStatus, pId));
|
||||
}
|
||||
|
||||
hipError_t hipStreamGetCaptureInfo_spt(hipStream_t stream, hipStreamCaptureStatus* pCaptureStatus,
|
||||
unsigned long long* pId) {
|
||||
HIP_INIT_API(hipStreamGetCaptureInfo, stream, pCaptureStatus, pId);
|
||||
PER_THREAD_DEFAULT_STREAM(stream);
|
||||
HIP_RETURN(hipStreamGetCaptureInfo_common(stream, pCaptureStatus, pId));
|
||||
}
|
||||
|
||||
hipError_t hipStreamGetCaptureInfo_v2_common(hipStream_t stream, hipStreamCaptureStatus* captureStatus_out,
|
||||
unsigned long long* id_out, hipGraph_t* graph_out,
|
||||
const hipGraphNode_t** dependencies_out,
|
||||
size_t* numDependencies_out) {
|
||||
HIP_INIT_API(hipStreamGetCaptureInfo_v2, stream, captureStatus_out, id_out, graph_out,
|
||||
dependencies_out, numDependencies_out);
|
||||
if (captureStatus_out == nullptr) {
|
||||
HIP_RETURN(hipErrorInvalidValue);
|
||||
return hipErrorInvalidValue;
|
||||
}
|
||||
if (stream == nullptr) {
|
||||
HIP_RETURN(hipErrorUnknown);
|
||||
return hipErrorUnknown;
|
||||
}
|
||||
if (!hip::isValid(stream)) {
|
||||
HIP_RETURN(hipErrorInvalidValue);
|
||||
return hipErrorInvalidValue;
|
||||
}
|
||||
hip::Stream* s = reinterpret_cast<hip::Stream*>(stream);
|
||||
*captureStatus_out = s->GetCaptureStatus();
|
||||
@@ -1358,7 +1409,28 @@ hipError_t hipStreamGetCaptureInfo_v2(hipStream_t stream, hipStreamCaptureStatus
|
||||
*numDependencies_out = s->GetLastCapturedNodes().size();
|
||||
}
|
||||
}
|
||||
HIP_RETURN(hipSuccess);
|
||||
return hipSuccess;
|
||||
}
|
||||
|
||||
hipError_t hipStreamGetCaptureInfo_v2(hipStream_t stream, hipStreamCaptureStatus* captureStatus_out,
|
||||
unsigned long long* id_out, hipGraph_t* graph_out,
|
||||
const hipGraphNode_t** dependencies_out,
|
||||
size_t* numDependencies_out) {
|
||||
HIP_INIT_API(hipStreamGetCaptureInfo_v2, stream, captureStatus_out, id_out, graph_out,
|
||||
dependencies_out, numDependencies_out);
|
||||
HIP_RETURN(hipStreamGetCaptureInfo_v2_common(stream, captureStatus_out, id_out, graph_out,
|
||||
dependencies_out, numDependencies_out));
|
||||
}
|
||||
|
||||
hipError_t hipStreamGetCaptureInfo_v2_spt(hipStream_t stream, hipStreamCaptureStatus* captureStatus_out,
|
||||
unsigned long long* id_out, hipGraph_t* graph_out,
|
||||
const hipGraphNode_t** dependencies_out,
|
||||
size_t* numDependencies_out) {
|
||||
HIP_INIT_API(hipStreamGetCaptureInfo_v2, stream, captureStatus_out, id_out, graph_out,
|
||||
dependencies_out, numDependencies_out);
|
||||
PER_THREAD_DEFAULT_STREAM(stream);
|
||||
HIP_RETURN(hipStreamGetCaptureInfo_v2_common(stream, captureStatus_out, id_out, graph_out,
|
||||
dependencies_out, numDependencies_out));
|
||||
}
|
||||
|
||||
hipError_t hipStreamUpdateCaptureDependencies(hipStream_t stream, hipGraphNode_t* dependencies,
|
||||
|
||||
Fai riferimento in un nuovo problema
Block a user