mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-19 19:00:47 +00:00
[QNN EP] Fix topological node unit traversal during validation (#17913)
### Description We need to ensure that tensors are first created and validated by their producers. If we don't, then builders that need to modify their outputs may not be able to do so if consumers are processed first (due to caching of tensors). For example, the Tanh builder may need to override its output quant param for 16-bit QDQ. I've encountered a scenario (while working on a partner model) where the override was not being correctly applied due to the graph traversal order. I tried to fix this bug in a previous [PR](https://github.com/microsoft/onnxruntime/pull/17877#discussion_r1353676802), but my fix was incorrect.
This commit is contained in:
parent
dad70ad4e8
commit
4e2a03d5fa
2 changed files with 17 additions and 2 deletions
|
|
@ -371,10 +371,14 @@ Status SimpleOpBuilder::ProcessSigmoidOrTanhOutput(QnnModelWrapper& qnn_model_wr
|
|||
const float scale = output_info.quant_param.scaleOffsetEncoding.scale;
|
||||
|
||||
LOGS(logger, VERBOSE) << "QNN requires that 16-bit quantized " << op_type << " operators use offset/scale values "
|
||||
<< "of <" << offset << ", " << scale << ">. QNN EP will override the original values.";
|
||||
<< "of <" << offset << ", " << scale << ">. QNN EP will override the original values for output "
|
||||
<< output_name;
|
||||
}
|
||||
}
|
||||
|
||||
ORT_RETURN_IF(qnn_model_wrapper.IsQnnTensorWrapperExist(output_name),
|
||||
"QNN EP is unable to override output quantization parameters for ", op_type.c_str(),
|
||||
" operator. Node name: ", node_unit.Name().c_str(), ", output name: ", output_name.c_str());
|
||||
Qnn_TensorType_t tensor_type = qnn_model_wrapper.IsGraphOutput(output_name) ? QNN_TENSOR_TYPE_APP_READ
|
||||
: QNN_TENSOR_TYPE_NATIVE;
|
||||
QnnTensorWrapper output_tensorwrapper(output_name, tensor_type, output_info.qnn_data_type, output_info.quant_param,
|
||||
|
|
|
|||
|
|
@ -242,7 +242,15 @@ QNNExecutionProvider::GetSupportedNodes(const GraphViewer& graph_viewer,
|
|||
for (size_t i = 0; i < node_indices.size(); i++) {
|
||||
gsl::not_null<const onnxruntime::Node*> node(graph_viewer.GetNode(node_indices[i]));
|
||||
|
||||
// Get the node_unit associated with the node. Note that the node may not be the node_unit's target node.
|
||||
const NodeUnit* node_unit = node_unit_map.at(node);
|
||||
|
||||
// Visiting 'nodes' in topological order does not guarantee that 'node_units' are
|
||||
// also visited in topological order. Skip this node if it is not the node_unit's target node
|
||||
// to ensure 'node_units' are visited in topological order.
|
||||
if (node != &node_unit->GetNode()) {
|
||||
continue;
|
||||
}
|
||||
const bool supported = IsNodeSupported(qnn_model_wrapper,
|
||||
*node_unit,
|
||||
node_unit_supported_result,
|
||||
|
|
@ -256,7 +264,10 @@ QNNExecutionProvider::GetSupportedNodes(const GraphViewer& graph_viewer,
|
|||
<< "] name: [" << node_unit->Name()
|
||||
<< "]";
|
||||
if (supported) {
|
||||
supported_nodes.insert(node);
|
||||
// If the node_unit is supported, add all of its nodes to the supported list.
|
||||
for (const auto* node_in_group : node_unit->GetAllNodesInGroup()) {
|
||||
supported_nodes.insert(node_in_group);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue