* \brief A simple encapsulation class for n-dimensional tensor. */
| 460 | * \brief A simple encapsulation class for n-dimensional tensor. |
| 461 | */ |
| 462 | struct TensorND { |
| 463 | TensorLayout layout; |
| 464 | |
| 465 | TensorND() : m_ref_ptr(RefPtr((void*)nullptr)) {} |
| 466 | |
| 467 | TensorND(void* raw_ptr_, const TensorLayout& layout_) |
| 468 | : layout(layout_), m_ref_ptr(raw_ptr_) {} |
| 469 | |
| 470 | TensorND(const TensorLayout& layout_, const RefPtr& ref_ptr) |
| 471 | : layout(layout_), m_ref_ptr(ref_ptr) {} |
| 472 | |
| 473 | MGE_WIN_DECLSPEC_FUC void reset_ptr(void* ptr, size_t offset = 0); |
| 474 | |
| 475 | void* raw_ptr() const { return m_ref_ptr.get_ptr(); } |
| 476 | |
| 477 | const RefPtr get_ref_ptr() const { return m_ref_ptr; } |
| 478 | |
| 479 | RefPtr& get_ref_ptr() { return m_ref_ptr; } |
| 480 | |
| 481 | //! get typed pointer; type check is performed |
| 482 | template <typename T> |
| 483 | T* ptr() const { |
| 484 | layout.dtype.assert_is_ctype<T>(); |
| 485 | return static_cast<T*>(m_ref_ptr.get_ptr()); |
| 486 | } |
| 487 | |
| 488 | //! get typed pointer of compatible type |
| 489 | template <typename T> |
| 490 | T* compatible_ptr() const { |
| 491 | layout.dtype.assert_is_compatible_ctype<T>(); |
| 492 | return reinterpret_cast<T*>(m_ref_ptr.get_ptr()); |
| 493 | } |
| 494 | |
| 495 | private: |
| 496 | RefPtr m_ref_ptr; |
| 497 | }; |
| 498 | |
| 499 | #if MEGDNN_CC_HOST |
| 500 | using TensorFormat = TensorLayout::Format; |
no outgoing calls