Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -91,7 +91,8 @@ Status GetEpContextFromMainNode(const onnxruntime::Node& main_context_node,
NodeAttrHelper node_helper(main_context_node);
bool is_embed_mode = node_helper.Get(EMBED_MODE, true);
if (is_embed_mode) {
const std::string& context_binary = node_helper.Get(EP_CACHE_CONTEXT, "");
static const std::string empty_context_binary;
const std::string& context_binary = node_helper.Get(EP_CACHE_CONTEXT, empty_context_binary);
return qnn_backend_manager->LoadCachedQnnContextFromBuffer(const_cast<char*>(context_binary.c_str()),
static_cast<uint64_t>(context_binary.length()),
"",
Expand Down
16 changes: 16 additions & 0 deletions onnxruntime/core/providers/qnn/ort_api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,22 @@ const std::string& NodeAttrHelper::Get(const std::string& key, const std::string
return def_val;
}

std::string NodeAttrHelper::Get(const std::string& key, std::string&& def_val) const {
if (auto entry = node_attributes_.find(key); entry != node_attributes_.end()) {
return NODE_ATTR_ITER_VAL(entry).s();
}

return std::move(def_val);
}

std::string NodeAttrHelper::Get(const std::string& key, const char* def_val) const {
if (auto entry = node_attributes_.find(key); entry != node_attributes_.end()) {
return NODE_ATTR_ITER_VAL(entry).s();
}

return def_val;
}

std::vector<std::string> NodeAttrHelper::Get(const std::string& key, const std::vector<std::string>& def_val) const {
if (auto entry = node_attributes_.find(key); entry != node_attributes_.end()) {
std::vector<std::string> res;
Expand Down
3 changes: 3 additions & 0 deletions onnxruntime/core/providers/qnn/ort_api.h
Original file line number Diff line number Diff line change
Expand Up @@ -150,7 +150,10 @@ class NodeAttrHelper {
int64_t Get(const std::string& key, int64_t def_val) const;
std::vector<int64_t> Get(const std::string& key, const std::vector<int64_t>& def_val) const;

// Lvalue defaults may be returned by reference; temporary and literal defaults return owned strings.
const std::string& Get(const std::string& key, const std::string& def_val) const;
std::string Get(const std::string& key, std::string&& def_val) const;
std::string Get(const std::string& key, const char* def_val) const;
std::vector<std::string> Get(const std::string& key, const std::vector<std::string>& def_val) const;

// Convert the i() or ints() of the attribute from int64_t to int32_t
Expand Down
18 changes: 18 additions & 0 deletions onnxruntime/core/providers/shared/utils/utils.cc
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@

#include "utils.h"

#include <utility>

#include "core/common/safeint.h"
#include "core/framework/node_unit.h"
#include "core/framework/tensorprotoutils.h"
Expand Down Expand Up @@ -146,6 +148,22 @@ const std::string& NodeAttrHelper::Get(const std::string& key, const std::string
return def_val;
}

std::string NodeAttrHelper::Get(const std::string& key, std::string&& def_val) const {
if (auto entry = node_attributes_.find(key); entry != node_attributes_.end()) {
return entry->second.s();
}

return std::move(def_val);
}

std::string NodeAttrHelper::Get(const std::string& key, const char* def_val) const {
if (auto entry = node_attributes_.find(key); entry != node_attributes_.end()) {
return entry->second.s();
}

return def_val;
}

std::vector<int32_t> NodeAttrHelper::Get(const std::string& key, const std::vector<int32_t>& def_val) const {
if (auto entry = node_attributes_.find(key); entry != node_attributes_.end()) {
const auto& attr = entry->second;
Expand Down
3 changes: 3 additions & 0 deletions onnxruntime/core/providers/shared/utils/utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,10 @@ class NodeAttrHelper {
int64_t Get(const std::string& key, int64_t def_val) const;
std::vector<int64_t> Get(const std::string& key, const std::vector<int64_t>& def_val) const;

// Lvalue defaults may be returned by reference; temporary and literal defaults return owned strings.
const std::string& Get(const std::string& key, const std::string& def_val) const;
std::string Get(const std::string& key, std::string&& def_val) const;
std::string Get(const std::string& key, const char* def_val) const;
std::vector<std::string> Get(const std::string& key, const std::vector<std::string>& def_val) const;

// Convert the i() or ints() of the attribute from int64_t to int32_t
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -69,7 +69,7 @@ class ResizeOpBuilder : public BaseOpBuilder {
LOGS_DEFAULT(WARNING) << "Antialias attribute is not supported.";
return false;
}
auto& cooridinate = helper.Get("coordinate_transoformation_mode", "half_pixel");
auto cooridinate = helper.Get("coordinate_transoformation_mode", "half_pixel");
if (cooridinate != "align_corners" && cooridinate != "half_pixel") {
LOGS_DEFAULT(WARNING) << "Only support half_pixel and align_corners attributes now.";
return false;
Expand Down
Loading