SWDEV-374375 - moving nullptr check into hipGraphNode::isValidNode

Change-Id: Ic9264521d6f0b36ffc948ca9feeaa70eb6e98eb6


[ROCm/clr commit: 03200953e5]
This commit is contained in:
Jiabao Xie
2023-01-09 16:17:19 -05:00
parent 507c87da5f
commit 3ab54974dc
2 changed files with 31 additions and 30 deletions
+30 -29
View File
@@ -1068,7 +1068,8 @@ hipError_t hipGraphAddMemcpyNode1D(hipGraphNode_t* pGraphNode, hipGraph_t graph,
hipError_t hipGraphMemcpyNodeSetParams1D(hipGraphNode_t node, void* dst, const void* src,
size_t count, hipMemcpyKind kind) {
HIP_INIT_API(hipGraphMemcpyNodeSetParams1D, node, dst, src, count, kind);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || dst == nullptr || src == nullptr || count == 0 || src == dst) {
if (!hipGraphNode::isNodeValid(node) || dst == nullptr || src == nullptr || count == 0 ||
src == dst) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1079,8 +1080,8 @@ hipError_t hipGraphExecMemcpyNodeSetParams1D(hipGraphExec_t hGraphExec, hipGraph
void* dst, const void* src, size_t count,
hipMemcpyKind kind) {
HIP_INIT_API(hipGraphExecMemcpyNodeSetParams1D, hGraphExec, node, dst, src, count, kind);
if (hGraphExec == nullptr || node == nullptr || !hipGraphNode::isNodeValid(node) || dst == nullptr || src == nullptr || count == 0 ||
src == dst) {
if (hGraphExec == nullptr || !hipGraphNode::isNodeValid(node) || dst == nullptr ||
src == nullptr || count == 0 || src == dst) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1272,7 +1273,7 @@ hipError_t hipGraphGetRootNodes(hipGraph_t graph, hipGraphNode_t* pRootNodes,
hipError_t hipGraphKernelNodeGetParams(hipGraphNode_t node, hipKernelNodeParams* pNodeParams) {
HIP_INIT_API(hipGraphKernelNodeGetParams, node, pNodeParams);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
reinterpret_cast<hipGraphKernelNode*>(node)->GetParams(pNodeParams);
@@ -1282,7 +1283,7 @@ hipError_t hipGraphKernelNodeGetParams(hipGraphNode_t node, hipKernelNodeParams*
hipError_t hipGraphKernelNodeSetParams(hipGraphNode_t node,
const hipKernelNodeParams* pNodeParams) {
HIP_INIT_API(hipGraphKernelNodeSetParams, node, pNodeParams);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr || pNodeParams->func == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNodeParams == nullptr || pNodeParams->func == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
HIP_RETURN(reinterpret_cast<hipGraphKernelNode*>(node)->SetParams(pNodeParams));
@@ -1290,7 +1291,7 @@ hipError_t hipGraphKernelNodeSetParams(hipGraphNode_t node,
hipError_t hipGraphMemcpyNodeGetParams(hipGraphNode_t node, hipMemcpy3DParms* pNodeParams) {
HIP_INIT_API(hipGraphMemcpyNodeGetParams, node, pNodeParams);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
reinterpret_cast<hipGraphMemcpyNode*>(node)->GetParams(pNodeParams);
@@ -1325,7 +1326,7 @@ hipError_t hipGraphKernelNodeGetAttribute(hipGraphNode_t hNode, hipKernelNodeAtt
hipError_t hipGraphMemcpyNodeSetParams(hipGraphNode_t node, const hipMemcpy3DParms* pNodeParams) {
HIP_INIT_API(hipGraphMemcpyNodeSetParams, node, pNodeParams);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
HIP_RETURN(reinterpret_cast<hipGraphMemcpyNode*>(node)->SetParams(pNodeParams));
@@ -1334,7 +1335,7 @@ hipError_t hipGraphMemcpyNodeSetParams(hipGraphNode_t node, const hipMemcpy3DPar
hipError_t hipGraphExecMemcpyNodeSetParams(hipGraphExec_t hGraphExec, hipGraphNode_t node,
hipMemcpy3DParms* pNodeParams) {
HIP_INIT_API(hipGraphExecMemcpyNodeSetParams, hGraphExec, node, pNodeParams);
if (hGraphExec == nullptr || node == nullptr || !hipGraphNode::isNodeValid(node)) {
if (hGraphExec == nullptr || !hipGraphNode::isNodeValid(node)) {
HIP_RETURN(hipErrorInvalidValue);
}
if (ihipMemcpy3D_validate(pNodeParams) != hipSuccess) {
@@ -1354,7 +1355,7 @@ hipError_t hipGraphExecMemcpyNodeSetParams(hipGraphExec_t hGraphExec, hipGraphNo
hipError_t hipGraphMemsetNodeGetParams(hipGraphNode_t node, hipMemsetParams* pNodeParams) {
HIP_INIT_API(hipGraphMemsetNodeGetParams, node, pNodeParams);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
reinterpret_cast<hipGraphMemsetNode*>(node)->GetParams(pNodeParams);
@@ -1363,7 +1364,7 @@ hipError_t hipGraphMemsetNodeGetParams(hipGraphNode_t node, hipMemsetParams* pNo
hipError_t hipGraphMemsetNodeSetParams(hipGraphNode_t node, const hipMemsetParams* pNodeParams) {
HIP_INIT_API(hipGraphMemsetNodeSetParams, node, pNodeParams);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
if (pNodeParams->height > 1 && pNodeParams->pitch < (pNodeParams->width * pNodeParams->elementSize)) {
@@ -1375,7 +1376,7 @@ hipError_t hipGraphMemsetNodeSetParams(hipGraphNode_t node, const hipMemsetParam
hipError_t hipGraphExecMemsetNodeSetParams(hipGraphExec_t hGraphExec, hipGraphNode_t node,
const hipMemsetParams* pNodeParams) {
HIP_INIT_API(hipGraphExecMemsetNodeSetParams, hGraphExec, node, pNodeParams);
if (hGraphExec == nullptr || node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr ||
if (hGraphExec == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr ||
pNodeParams->dst == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1426,7 +1427,7 @@ hipError_t hipGraphAddDependencies(hipGraph_t graph, const hipGraphNode_t* from,
hipError_t hipGraphExecKernelNodeSetParams(hipGraphExec_t hGraphExec, hipGraphNode_t node,
const hipKernelNodeParams* pNodeParams) {
HIP_INIT_API(hipGraphExecKernelNodeSetParams, hGraphExec, node, pNodeParams);
if (hGraphExec == nullptr || node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr ||
if (hGraphExec == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr ||
pNodeParams->func == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1439,7 +1440,7 @@ hipError_t hipGraphExecKernelNodeSetParams(hipGraphExec_t hGraphExec, hipGraphNo
hipError_t hipGraphChildGraphNodeGetGraph(hipGraphNode_t node, hipGraph_t* pGraph) {
HIP_INIT_API(hipGraphChildGraphNodeGetGraph, node, pGraph);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pGraph == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pGraph == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
*pGraph = reinterpret_cast<hipGraphNode*>(node)->GetChildGraph();
@@ -1452,7 +1453,7 @@ hipError_t hipGraphChildGraphNodeGetGraph(hipGraphNode_t node, hipGraph_t* pGrap
hipError_t hipGraphExecChildGraphNodeSetParams(hipGraphExec_t hGraphExec, hipGraphNode_t node,
hipGraph_t childGraph) {
HIP_INIT_API(hipGraphExecChildGraphNodeSetParams, hGraphExec, node, childGraph);
if (hGraphExec == nullptr || node == nullptr || !hipGraphNode::isNodeValid(node) || childGraph == nullptr ||
if (hGraphExec == nullptr || !hipGraphNode::isNodeValid(node) || childGraph == nullptr ||
!ihipGraph::isGraphValid(childGraph)) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1662,7 +1663,7 @@ hipError_t hipGraphGetEdges(hipGraph_t graph, hipGraphNode_t* from, hipGraphNode
hipError_t hipGraphNodeGetDependencies(hipGraphNode_t node, hipGraphNode_t* pDependencies,
size_t* pNumDependencies) {
HIP_INIT_API(hipGraphNodeGetDependencies, node, pDependencies, pNumDependencies);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNumDependencies == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNumDependencies == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
const std::vector<hipGraphNode_t>& dependencies = node->GetDependencies();
@@ -1690,7 +1691,7 @@ hipError_t hipGraphNodeGetDependencies(hipGraphNode_t node, hipGraphNode_t* pDep
hipError_t hipGraphNodeGetDependentNodes(hipGraphNode_t node, hipGraphNode_t* pDependentNodes,
size_t* pNumDependentNodes) {
HIP_INIT_API(hipGraphNodeGetDependentNodes, node, pDependentNodes, pNumDependentNodes);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNumDependentNodes == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNumDependentNodes == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
const std::vector<hipGraphNode_t>& dependents = node->GetEdges();
@@ -1717,7 +1718,7 @@ hipError_t hipGraphNodeGetDependentNodes(hipGraphNode_t node, hipGraphNode_t* pD
hipError_t hipGraphNodeGetType(hipGraphNode_t node, hipGraphNodeType* pType) {
HIP_INIT_API(hipGraphNodeGetType, node, pType);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pType == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pType == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
*pType = node->GetType();
@@ -1726,7 +1727,7 @@ hipError_t hipGraphNodeGetType(hipGraphNode_t node, hipGraphNodeType* pType) {
hipError_t hipGraphDestroyNode(hipGraphNode_t node) {
HIP_INIT_API(hipGraphDestroyNode, node);
if (node == nullptr || !hipGraphNode::isNodeValid(node)) {
if (!hipGraphNode::isNodeValid(node)) {
HIP_RETURN(hipErrorInvalidValue);
}
node->GetParentGraph()->RemoveNode(node);
@@ -1795,7 +1796,7 @@ hipError_t hipGraphMemcpyNodeSetParamsFromSymbol(hipGraphNode_t node, void* dst,
if (symbol == nullptr) {
HIP_RETURN(hipErrorInvalidSymbol);
}
if (node == nullptr || !hipGraphNode::isNodeValid(node) || dst == nullptr || count == 0 || symbol == dst) {
if (!hipGraphNode::isNodeValid(node) || dst == nullptr || count == 0 || symbol == dst) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1811,7 +1812,7 @@ hipError_t hipGraphExecMemcpyNodeSetParamsFromSymbol(hipGraphExec_t hGraphExec,
if (symbol == nullptr) {
HIP_RETURN(hipErrorInvalidSymbol);
}
if (hGraphExec == nullptr || node == nullptr || !hipGraphNode::isNodeValid(node) || dst == nullptr || count == 0 || symbol == dst) {
if (hGraphExec == nullptr || !hipGraphNode::isNodeValid(node) || dst == nullptr || count == 0 || symbol == dst) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1855,7 +1856,7 @@ hipError_t hipGraphMemcpyNodeSetParamsToSymbol(hipGraphNode_t node, const void*
if (symbol == nullptr) {
HIP_RETURN(hipErrorInvalidSymbol);
}
if (node == nullptr || !hipGraphNode::isNodeValid(node) || src == nullptr || count == 0 || symbol == src) {
if (!hipGraphNode::isNodeValid(node) || src == nullptr || count == 0 || symbol == src) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1873,7 +1874,7 @@ hipError_t hipGraphExecMemcpyNodeSetParamsToSymbol(hipGraphExec_t hGraphExec, hi
if (symbol == nullptr) {
HIP_RETURN(hipErrorInvalidSymbol);
}
if (hGraphExec == nullptr || src == nullptr || node == nullptr || !hipGraphNode::isNodeValid(node) || count == 0 || src == symbol) {
if (hGraphExec == nullptr || src == nullptr || !hipGraphNode::isNodeValid(node) || count == 0 || src == symbol) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -1901,7 +1902,7 @@ hipError_t hipGraphAddEventRecordNode(hipGraphNode_t* pGraphNode, hipGraph_t gra
hipError_t hipGraphEventRecordNodeGetEvent(hipGraphNode_t node, hipEvent_t* event_out) {
HIP_INIT_API(hipGraphEventRecordNodeGetEvent, node, event_out);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || event_out == nullptr || node->GetType() != hipGraphNodeTypeEventRecord) {
if (!hipGraphNode::isNodeValid(node) || event_out == nullptr || node->GetType() != hipGraphNodeTypeEventRecord) {
HIP_RETURN(hipErrorInvalidValue);
}
reinterpret_cast<hipGraphEventRecordNode*>(node)->GetParams(event_out);
@@ -1910,7 +1911,7 @@ hipError_t hipGraphEventRecordNodeGetEvent(hipGraphNode_t node, hipEvent_t* even
hipError_t hipGraphEventRecordNodeSetEvent(hipGraphNode_t node, hipEvent_t event) {
HIP_INIT_API(hipGraphEventRecordNodeSetEvent, node, event);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || event == nullptr || node->GetType() != hipGraphNodeTypeEventRecord) {
if (!hipGraphNode::isNodeValid(node) || event == nullptr || node->GetType() != hipGraphNodeTypeEventRecord) {
HIP_RETURN(hipErrorInvalidValue);
}
HIP_RETURN(reinterpret_cast<hipGraphEventRecordNode*>(node)->SetParams(event));
@@ -1945,7 +1946,7 @@ hipError_t hipGraphAddEventWaitNode(hipGraphNode_t* pGraphNode, hipGraph_t graph
hipError_t hipGraphEventWaitNodeGetEvent(hipGraphNode_t node, hipEvent_t* event_out) {
HIP_INIT_API(hipGraphEventWaitNodeGetEvent, node, event_out);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || event_out == nullptr || node->GetType() != hipGraphNodeTypeWaitEvent) {
if (!hipGraphNode::isNodeValid(node) || event_out == nullptr || node->GetType() != hipGraphNodeTypeWaitEvent) {
HIP_RETURN(hipErrorInvalidValue);
}
reinterpret_cast<hipGraphEventWaitNode*>(node)->GetParams(event_out);
@@ -1954,7 +1955,7 @@ hipError_t hipGraphEventWaitNodeGetEvent(hipGraphNode_t node, hipEvent_t* event_
hipError_t hipGraphEventWaitNodeSetEvent(hipGraphNode_t node, hipEvent_t event) {
HIP_INIT_API(hipGraphEventWaitNodeSetEvent, node, event);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || event == nullptr || node->GetType() != hipGraphNodeTypeWaitEvent) {
if (!hipGraphNode::isNodeValid(node) || event == nullptr || node->GetType() != hipGraphNodeTypeWaitEvent) {
HIP_RETURN(hipErrorInvalidValue);
}
HIP_RETURN(reinterpret_cast<hipGraphEventWaitNode*>(node)->SetParams(event));
@@ -1990,7 +1991,7 @@ hipError_t hipGraphAddHostNode(hipGraphNode_t* pGraphNode, hipGraph_t graph,
hipError_t hipGraphHostNodeGetParams(hipGraphNode_t node, hipHostNodeParams* pNodeParams) {
HIP_INIT_API(hipGraphHostNodeGetParams, node, pNodeParams);
if (node == nullptr || !hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
if (!hipGraphNode::isNodeValid(node) || pNodeParams == nullptr) {
HIP_RETURN(hipErrorInvalidValue);
}
reinterpret_cast<hipGraphHostNode*>(node)->GetParams(pNodeParams);
@@ -1999,7 +2000,7 @@ hipError_t hipGraphHostNodeGetParams(hipGraphNode_t node, hipHostNodeParams* pNo
hipError_t hipGraphHostNodeSetParams(hipGraphNode_t node, const hipHostNodeParams* pNodeParams) {
HIP_INIT_API(hipGraphHostNodeSetParams, node, pNodeParams);
if (pNodeParams == nullptr || pNodeParams->fn == nullptr || node == nullptr ||
if (pNodeParams == nullptr || pNodeParams->fn == nullptr ||
!hipGraphNode::isNodeValid(node)) {
HIP_RETURN(hipErrorInvalidValue);
}
@@ -2010,7 +2011,7 @@ hipError_t hipGraphExecHostNodeSetParams(hipGraphExec_t hGraphExec, hipGraphNode
const hipHostNodeParams* pNodeParams) {
HIP_INIT_API(hipGraphExecHostNodeSetParams, hGraphExec, node, pNodeParams);
if (hGraphExec == nullptr || pNodeParams == nullptr || pNodeParams->fn == nullptr ||
node == nullptr || !hipGraphNode::isNodeValid(node)) {
!hipGraphNode::isNodeValid(node)) {
HIP_RETURN(hipErrorInvalidValue);
}
hipGraphNode_t clonedNode = hGraphExec->GetClonedNode(node);
@@ -214,7 +214,7 @@ struct hipGraphNode : public hipGraphNodeDOTAttribute {
// check node validity
static bool isNodeValid(hipGraphNode* pGraphNode) {
amd::ScopedLock lock(nodeSetLock_);
if (nodeSet_.find(pGraphNode) == nodeSet_.end()) {
if (pGraphNode == nullptr || nodeSet_.find(pGraphNode) == nodeSet_.end()) {
return false;
}
return true;