microsoft/onnxruntime-extensions

Public

mirrored from https://github.com/microsoft/onnxruntime-extensionsAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
sayanshaw/opencv-build

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

include/custom_op/tensor_api.h

536lines · modecode

1#pragma once
2#include <optional>
3#include <numeric>
4#include <type_traits>
5#include <assert.h>
6
7#include "onnxruntime_f16.h"
8#include "kernel_context.h"
9
10namespace Ort {
11namespace Custom {
12
13template <typename T>
14struct Span {
15 const T* data_ = {};
16 size_t size_ = {};
17 void Assign(const T* data, size_t size) {
18 data_ = data;
19 size_ = size;
20 }
21 size_t size() const { return size_; }
22 T operator[](size_t indice) const {
23 return data_[indice];
24 }
25 const T* data() const { return data_; }
26};
27
28
29#if ORT_API_VERSION >= 16
30
31template <>
32struct Span<MFloat16> {
33 const MFloat16* data_ = {};
34 size_t size_ = {};
35 void Assign(const MFloat16* data, size_t size) {
36 data_ = data;
37 size_ = size;
38 }
39 size_t size() const { return size_; }
40 MFloat16 operator[](size_t indice) const {
41 return data_[indice];
42 }
43 const MFloat16* data() const { return data_; }
44};
45
46template <>
47struct Span<BFloat16> {
48 const BFloat16* data_ = {};
49 size_t size_ = {};
50 void Assign(const BFloat16* data, size_t size) {
51 data_ = data;
52 size_ = size;
53 }
54 size_t size() const { return size_; }
55 BFloat16 operator[](size_t indice) const {
56 return data_[indice];
57 }
58 const BFloat16* data() const { return data_; }
59};
60
61#endif
62
63class ITensorStorage{
64public:
65 virtual const std::vector<int64_t>& Shape() const = 0;
66 virtual const void* DataRaw() const = 0;
67 virtual bool IsInitialized() const = 0;
68 virtual void* Initialize(const std::vector<int64_t>& shape, size_t element_size) = 0;
69 virtual ~ITensorStorage() = default;
70};
71
72
73class IAllocator {
74public:
75 virtual void* Alloc(size_t size) = 0;
76 virtual void Free(void* p) = 0;
77};
78
79
80class OrtEagerTensorStorage : public ITensorStorage {
81public:
82 OrtEagerTensorStorage(const std::vector<int64_t>& shape,
83 void* buffer) : buffer_(buffer), shape_(shape){
84
85 }
86
87 OrtEagerTensorStorage(IAllocator* allocator) : allocator_(allocator){
88 }
89
90 ~OrtEagerTensorStorage() override{
91 if (allocator_ && buffer_)
92 allocator_->Free(buffer_);
93 }
94
95 const std::vector<int64_t>& Shape() const override {
96 if (!IsInitialized())
97 ORTX_CXX_API_THROW("Tensor not initialized", ORT_RUNTIME_EXCEPTION);
98 return *shape_;
99 }
100
101 bool IsInitialized() const override {
102 return shape_.has_value();
103 }
104
105 const void* DataRaw() const override {
106 return buffer_;
107 }
108
109 void* Initialize(const std::vector<int64_t>& shape, size_t element_size) override {
110 if (IsInitialized())
111 return buffer_;
112 assert(allocator_);
113 shape_ = shape;
114 int64_t n_elem = std::accumulate(shape.begin(), shape.end(), 1LL, std::multiplies<int64_t>());
115 auto buffer_size = n_elem * element_size;
116 buffer_ = allocator_->Alloc(buffer_size);
117 return buffer_;
118 }
119
120private:
121 void* buffer_ {};
122 std::optional<std::vector<int64_t>> shape_;
123 // caller need to make sure the allocator is alive
124 IAllocator* allocator_{};
125};
126
127template <typename TT>
128ONNXTensorElementDataType GetOrtDType(){
129 if constexpr (std::is_same<TT, bool>::value)
130 return ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL;
131 else if constexpr (std::is_same<TT, float>::value)
132 return ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT;
133 else if constexpr (std::is_same<TT, double>::value)
134 return ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE;
135 else if constexpr (std::is_same<TT, uint8_t>::value)
136 return ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8;
137 else if constexpr (std::is_same<TT, int8_t>::value)
138 return ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8;
139 else if constexpr (std::is_same<TT, uint16_t>::value)
140 return ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16;
141 else if constexpr (std::is_same<TT, int16_t>::value)
142 return ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16;
143 else if constexpr (std::is_same<TT, uint32_t>::value)
144 return ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32;
145 else if constexpr (std::is_same<TT, int32_t>::value)
146 return ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32;
147 else if constexpr (std::is_same<TT, uint64_t>::value)
148 return ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64;
149 else if constexpr (std::is_same<TT, int64_t>::value)
150 return ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64;
151 else if constexpr (std::is_same<TT, std::string>::value)
152 return ONNX_TENSOR_ELEMENT_DATA_TYPE_STRING;
153 ORTX_CXX_API_THROW("Unexpected type", ORT_RUNTIME_EXCEPTION);
154 return ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT;
155}
156
157class TensorBase : public Arg {
158public:
159 virtual ~TensorBase() = default;
160
161 virtual ONNXTensorElementDataType Type() const = 0;
162 virtual const std::vector<int64_t>& Shape() const = 0;
163 virtual int64_t NumberOfElement() const = 0;
164 virtual const void* DataRaw() const = 0;
165 virtual size_t SizeInBytes() const = 0;
166};
167
168template <typename T>
169class Tensor : public TensorBase {
170 public:
171 using TT = typename std::remove_reference<T>::type;
172 Tensor(std::unique_ptr<ITensorStorage> tensor_storage) : storage_(std::move(tensor_storage)){
173 }
174
175 Tensor(const std::vector<int64_t>& shape, void* buffer) : Tensor(std::make_unique<OrtEagerTensorStorage>(shape, buffer)) {}
176
177 Tensor(IAllocator* allocator) : storage_(std::make_unique<OrtEagerTensorStorage>(allocator)){}
178
179 operator bool() const {
180 return storage_->IsInitialized();
181 }
182
183 ONNXTensorElementDataType Type() const override {
184 return GetOrtDType<T>();
185 }
186
187 const std::vector<int64_t>& Shape() const override {
188 return storage_->Shape();
189 }
190
191 int64_t NumberOfElement() const override {
192 auto& shape = storage_->Shape();
193 return std::accumulate(shape.begin(), shape.end(), 1LL, std::multiplies<int64_t>());
194 }
195
196 std::string Shape2Str() const {
197 if (storage_->IsInitialized()) {
198 std::string shape_str;
199 auto& shape = storage_->Shape();
200 for (const auto& dim : shape) {
201 shape_str.append(std::to_string(dim));
202 shape_str.append(", ");
203 }
204 return shape_str;
205 } else {
206 return "empty";
207 }
208 }
209
210 const TT* Data() const {
211#if ORT_API_VERSION >= 16
212 if constexpr (std::is_same<TT, MFloat16>::value || std::is_same<TT, BFloat16>::value)
213 return reinterpret_cast<const TT*>(storage_->DataRaw());
214 else
215#endif
216 return static_cast<const TT*>(storage_->DataRaw());
217 }
218
219 const void* DataRaw() const override {
220 return storage_->DataRaw();
221 }
222
223 size_t SizeInBytes() const override {
224 return NumberOfElement() * sizeof(TT);
225 }
226
227 TT* Allocate(const std::vector<int64_t>& shape) {
228 // it should be OK to allocate multiple times
229 void* buffer = storage_->Initialize(shape, sizeof(TT));
230#if ORT_API_VERSION >= 16
231 if constexpr (std::is_same<TT, MFloat16>::value || std::is_same<TT, BFloat16>::value)
232 return reinterpret_cast<TT*>(buffer);
233 else
234#endif
235 return static_cast<TT*>(buffer);
236 }
237
238 const Span<T>& AsSpan() {
239#if ORT_API_VERSION >= 16
240 if constexpr (std::is_same<TT, MFloat16>::value || std::is_same<TT, BFloat16>::value) {
241 ORTX_CXX_API_THROW("AsSpan for MFloat16 / BFloat16 not implemented", ORT_RUNTIME_EXCEPTION);
242 }
243 else{
244#endif
245 auto& shape = storage_->Shape();
246 if (shape.size() != 1) {
247 ORTX_CXX_API_THROW("to get a span, shape must be 1-D, actual shape: " + Shape2Str(), ORT_RUNTIME_EXCEPTION);
248 }
249 span_.Assign(Data(), shape[0]);
250 return span_;
251#if ORT_API_VERSION >= 16
252 }
253#endif
254 }
255
256 const T& AsScalar() {
257#if ORT_API_VERSION >= 16
258 if constexpr (std::is_same<TT, MFloat16>::value || std::is_same<TT, BFloat16>::value) {
259 ORTX_CXX_API_THROW("AsScalar for MFloat16 / BFloat16 not implemented", ORT_RUNTIME_EXCEPTION);
260 }
261 else{
262#endif
263 auto& shape = storage_->Shape();
264 if ((shape.size() == 1 && shape[0] != 1) || shape.size() > 1) {
265 ORTX_CXX_API_THROW("to get a scalar, shape must be {1}, actual shape: " + Shape2Str(), ORT_RUNTIME_EXCEPTION);
266 }
267 return *Data();
268#if ORT_API_VERSION >= 16
269 }
270#endif
271 }
272
273 private:
274 std::unique_ptr<ITensorStorage> storage_;
275 Span<T> span_;
276};
277
278template<typename T>
279class IStringTensorStorage{
280public:
281 using strings = std::vector<T>;
282 virtual const std::vector<int64_t>& Shape() const = 0;
283 virtual const void* DataRaw() const = 0;
284 virtual const strings& Data() const = 0;
285 virtual bool IsInitialized() const = 0;
286 virtual void SetStringOutput(const strings& ss, const std::vector<int64_t>& dims) = 0;
287 virtual void SetStringOutput(const std::vector<const char*>& ss, const std::vector<int64_t>& dims) = 0;
288 virtual ~IStringTensorStorage() = default;
289};
290
291template<typename T>
292class EagerStringTensorStorage : public IStringTensorStorage<T>{
293public:
294 using strings = std::vector<T>;
295 EagerStringTensorStorage(const strings& ss) : input_strings_(ss), shape_(std::vector<int64_t>{static_cast<int64_t>(ss.size())}){}
296
297 EagerStringTensorStorage() {}
298
299 const std::vector<int64_t>& Shape() const override {
300 if (!IsInitialized())
301 ORTX_CXX_API_THROW("Tensor not initialized", ORT_RUNTIME_EXCEPTION);
302 return *shape_;
303 }
304
305 const void* DataRaw() const override {
306 if (input_strings_.size() != 1) {
307 ORTX_CXX_API_THROW("DataRaw() only applies to string scalar", ORT_RUNTIME_EXCEPTION);
308 }
309 if constexpr (std::is_same<std::string_view, T>::value)
310 return reinterpret_cast<const void*>(input_strings_[0].data());
311 else
312 return reinterpret_cast<const void*>(input_strings_[0].c_str());
313 }
314
315 bool IsInitialized() const override {
316 return shape_.has_value();
317 }
318
319 void SetStringOutput(const strings& ss, const std::vector<int64_t>& dims) override {
320 if constexpr (std::is_same<std::string_view, T>::value)
321 ORTX_CXX_API_THROW("Set output for string view tensor is not supported", ORT_RUNTIME_EXCEPTION);
322 input_strings_.assign(ss.begin(), ss.end());
323 shape_ = dims;
324 }
325
326 const strings& Data() const override {
327 return input_strings_;
328 }
329
330 void SetStringOutput(const std::vector<const char*>& ss, const std::vector<int64_t>& dims) override {
331 if constexpr (std::is_same<std::string_view, T>::value)
332 ORTX_CXX_API_THROW("Set output for string view tensor is not supported", ORT_RUNTIME_EXCEPTION);
333
334 for (const char* s : ss){
335 input_strings_.push_back(s);
336 }
337 shape_ = dims;
338 }
339
340private:
341 std::vector<T> input_strings_;
342 std::optional<std::vector<int64_t>> shape_;
343};
344
345template <>
346class Tensor<std::string> : public TensorBase {
347 public:
348 using strings = std::vector<std::string>;
349
350 Tensor(std::unique_ptr<IStringTensorStorage<std::string>> storage) : storage_(std::move(storage)) {}
351
352 Tensor(const strings& ss) : storage_(std::make_unique<EagerStringTensorStorage<std::string>>(ss)) {}
353
354 Tensor() : storage_(std::make_unique<EagerStringTensorStorage<std::string>>()) {}
355
356 ONNXTensorElementDataType Type() const override {
357 return GetOrtDType<std::string>();
358 }
359
360 const strings& Data() const {
361 return storage_->Data();
362 }
363
364 const std::vector<int64_t>& Shape() const override {
365 return storage_->Shape();
366 }
367
368 int64_t NumberOfElement() const override {
369 auto& shape = storage_->Shape();
370 return std::accumulate(shape.begin(), shape.end(), 1LL, std::multiplies<int64_t>());
371 }
372
373 std::string Shape2Str() const {
374 if (storage_->IsInitialized()) {
375 std::string shape_str;
376 auto& shape = storage_->Shape();
377 for (const auto& dim : shape) {
378 shape_str.append(std::to_string(dim));
379 shape_str.append(", ");
380 }
381 return shape_str;
382 } else {
383 return "empty";
384 }
385 }
386
387 const void* DataRaw() const override {
388 return storage_->DataRaw();
389 }
390
391 size_t SizeInBytes() const override {
392 auto& ss = storage_->Data();
393 if (ss.size() != 1) {
394 ORTX_CXX_API_THROW("SizeInBytes() only applies to string scalar", ORT_RUNTIME_EXCEPTION);
395 }
396 return ss[0].size();
397 }
398
399 void SetStringOutput(const strings& ss, const std::vector<int64_t>& dims) {
400 storage_->SetStringOutput(ss, dims);
401 }
402 void SetStringOutput(const std::vector<const char*>& ss, const std::vector<int64_t>& dims) {
403 storage_->SetStringOutput(ss, dims);
404 }
405 const Span<std::string>& AsSpan() {
406 ORTX_CXX_API_THROW("span for TensorT of string not implemented", ORT_RUNTIME_EXCEPTION);
407 }
408 const std::string& AsScalar() {
409 auto& ss = storage_->Data();
410 if (ss.size() != 1) {
411 ORTX_CXX_API_THROW("to get a scalar, shape must be {1}, actual shape: " + Shape2Str(), ORT_RUNTIME_EXCEPTION);
412 }
413 return ss[0];
414 }
415
416 private:
417 std::unique_ptr<IStringTensorStorage<std::string>> storage_;
418};
419
420
421template <>
422class Tensor<std::string_view> : public TensorBase {
423 public:
424 using strings = std::vector<std::string_view>;
425
426 Tensor(std::unique_ptr<IStringTensorStorage<std::string_view>> storage) : storage_(std::move(storage)) {}
427
428 Tensor(const strings& ss) : storage_(std::make_unique<EagerStringTensorStorage<std::string_view>>(ss)) {}
429
430 ONNXTensorElementDataType Type() const override {
431 return GetOrtDType<std::string_view>();
432 }
433
434 const strings& Data() const {
435 return storage_->Data();
436 }
437
438 const std::vector<int64_t>& Shape() const override {
439 return storage_->Shape();
440 }
441
442 int64_t NumberOfElement() const override {
443 auto& shape = storage_->Shape();
444 return std::accumulate(shape.begin(), shape.end(), 1LL, std::multiplies<int64_t>());
445 }
446
447 std::string Shape2Str() const {
448 if (storage_->IsInitialized()) {
449 std::string shape_str;
450 auto& shape = storage_->Shape();
451 for (const auto& dim : shape) {
452 shape_str.append(std::to_string(dim));
453 shape_str.append(", ");
454 }
455 return shape_str;
456 } else {
457 return "empty";
458 }
459 }
460
461 const void* DataRaw() const override {
462 return storage_->DataRaw();
463 }
464
465 size_t SizeInBytes() const override {
466 auto& ss = storage_->Data();
467 if (ss.size() != 1) {
468 ORTX_CXX_API_THROW("SizeInBytes() only applies to string scalar", ORT_RUNTIME_EXCEPTION);
469 }
470 return ss[0].size();
471 }
472
473 void SetStringOutput(const strings& ss, const std::vector<int64_t>& dims) {
474 storage_->SetStringOutput(ss, dims);
475 }
476 void SetStringOutput(const std::vector<const char*>& ss, const std::vector<int64_t>& dims) {
477 storage_->SetStringOutput(ss, dims);
478 }
479 const Span<std::string_view>& AsSpan() {
480 ORTX_CXX_API_THROW("span for TensorT of string not implemented", ORT_RUNTIME_EXCEPTION);
481 }
482 const std::string_view& AsScalar() {
483 auto& ss = storage_->Data();
484 if (ss.size() != 1) {
485 ORTX_CXX_API_THROW("to get a scalar, shape must be {1}, actual shape: " + Shape2Str(), ORT_RUNTIME_EXCEPTION);
486 }
487 return ss[0];
488 }
489
490 private:
491 std::unique_ptr<IStringTensorStorage<std::string_view>> storage_;
492};
493
494
495template<typename ...Args>
496class NamedArgumentDict{
497public:
498 using ValueTuple = std::tuple<Args...>;
499
500 NamedArgumentDict(const std::vector<const char*>& keys, const std::tuple<Args...>& args) : entries_(args) {
501 for (const char* key : keys){
502 names_.push_back(key);
503 }
504 }
505
506 template<typename T>
507 T TryToGetAttributeWithDefault(const char* name, const T& default_value) const {
508 return TryToGetAttributeWithDefaultInternal<0>(name, default_value);
509 }
510
511private:
512 template<size_t I, typename T>
513 typename std::enable_if<I == sizeof...(Args), T>::type
514 TryToGetAttributeWithDefaultInternal(const char* name, const T& default_value) const {
515 return default_value;
516 }
517
518 template<size_t I, typename T>
519 typename std::enable_if<I < sizeof...(Args), T>::type
520 TryToGetAttributeWithDefaultInternal(const char* name, const T& default_value) const {
521 if (names_[I] == name){
522 if constexpr (std::is_same<std::tuple_element_t<I, ValueTuple>, T>::value)
523 return std::get<I>(entries_);
524 else
525 throw std::runtime_error("name matched but type is not");
526 }
527 return TryToGetAttributeWithDefaultInternal<I+1>(name, default_value);
528 }
529
530 std::vector<std::string> names_;
531 std::tuple<Args...> entries_;
532
533};
534
535}
536}
537