microsoft/onnxruntime-extensions

Public

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

CodeCommitsIssuesPull requestsActionsInsightsSecurity
wechi/ort_test

Branches

Tags

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

Clone

HTTPS

Download ZIP

include/custom_op/tensor_api.h

606lines · modecode

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