| 627 | } |
| 628 | |
| 629 | void GraphLoaderOSS::OprLoadContextImpl::load_tensor_value( |
| 630 | HostTensorND* dest, const TensorLayout& layout, const fbs::Tensor* tensor) { |
| 631 | auto&& loader = m_loader->m_cur_load_config->tensor_value_loader; |
| 632 | auto&& file = m_loader->m_file; |
| 633 | auto begin_pos = file->tell(); |
| 634 | file->skip(tensor->offset()); |
| 635 | if (loader) { |
| 636 | // call custom loader |
| 637 | void* dest_ptr = nullptr; |
| 638 | if (dest) { |
| 639 | dest->dtype(layout.dtype).resize(layout); |
| 640 | dest_ptr = dest->raw_ptr(); |
| 641 | } |
| 642 | loader(dest_ptr, layout, *file); |
| 643 | } else { |
| 644 | if (dest) { |
| 645 | file->read_into_tensor(*dest, layout); |
| 646 | } else { |
| 647 | file->skip(layout.span().high_byte); |
| 648 | } |
| 649 | } |
| 650 | mgb_throw_if( |
| 651 | file->tell() < begin_pos, SerializationError, |
| 652 | "Custom tensor value loader accessed out of range data before " |
| 653 | "start of data blob"); |
| 654 | auto data_size = tensor->data_size(); |
| 655 | auto consumed_size = file->tell() - begin_pos; |
| 656 | mgb_throw_if( |
| 657 | consumed_size > data_size, SerializationError, |
| 658 | "Custom tensor value loader consumed more data than " |
| 659 | "available: consumed %zu, has %u", |
| 660 | consumed_size, data_size); |
| 661 | if (consumed_size < data_size) { |
| 662 | mgb_log_warn( |
| 663 | "Tensor value loader consumed less data than available: " |
| 664 | "consumed %zu bytes, has %u bytes", |
| 665 | consumed_size, data_size); |
| 666 | file->skip(data_size - consumed_size); |
| 667 | } |
| 668 | } |
| 669 | |
| 670 | std::shared_ptr<HostTensorND> GraphLoaderOSS::OprLoadContextImpl::load_tensor() { |
| 671 | mgb_assert( |