| 14 | }; |
| 15 | |
| 16 | void lite::decompressed_tensor_value_loader( |
| 17 | void* ptr_, const mgb::TensorLayout& layout, |
| 18 | mgb::serialization::InputFile& fin) { |
| 19 | uint8_t compress_flag; |
| 20 | fin.read(&compress_flag, sizeof(compress_flag)); |
| 21 | size_t num_weights = layout.total_nr_elems(); |
| 22 | switch (CompressionMethod(compress_flag)) { |
| 23 | case CompressionMethod::NO_COMPRESSION: { |
| 24 | mgb::serialization::GraphLoadConfig::default_tensor_value_loader( |
| 25 | ptr_, layout, fin); |
| 26 | break; |
| 27 | } |
| 28 | case CompressionMethod::FLOAT32_STRIDE_FLOAT32_BASE_UINT8_WEIGHTS: { |
| 29 | if (ptr_) { |
| 30 | float stride, base; |
| 31 | std::vector<uint8_t> weights(num_weights); |
| 32 | fin.read(&stride, sizeof(stride)); |
| 33 | fin.read(&base, sizeof(base)); |
| 34 | fin.read(weights.data(), num_weights * sizeof(uint8_t)); |
| 35 | auto* ptr = static_cast<float*>(ptr_); |
| 36 | for (size_t i = 0; i < num_weights; ++i) |
| 37 | ptr[i] = stride * weights[i] + base; |
| 38 | } else { |
| 39 | fin.skip(sizeof(float) * 2 + num_weights * sizeof(uint8_t)); |
| 40 | } |
| 41 | break; |
| 42 | } |
| 43 | case CompressionMethod::FLOAT32_STRIDE_FLOAT32_BASE_UINT16_WEIGHTS: { |
| 44 | if (ptr_) { |
| 45 | float stride, base; |
| 46 | std::vector<uint16_t> weights(num_weights); |
| 47 | fin.read(&stride, sizeof(stride)); |
| 48 | fin.read(&base, sizeof(base)); |
| 49 | fin.read(weights.data(), num_weights * sizeof(uint16_t)); |
| 50 | auto* ptr = static_cast<float*>(ptr_); |
| 51 | for (size_t i = 0; i < num_weights; ++i) |
| 52 | ptr[i] = stride * weights[i] + base; |
| 53 | } else { |
| 54 | fin.skip(sizeof(float) * 2 + num_weights * sizeof(uint16_t)); |
| 55 | } |
| 56 | break; |
| 57 | } |
| 58 | default: |
| 59 | LITE_THROW("Unexpected compression method"); |
| 60 | } |
| 61 | } |
| 62 | |
| 63 | LTensorLayout lite::to_impl_layout(const Layout& layout) { |
| 64 | mgb::TensorLayout mge_layout; |
nothing calls this directly
no test coverage detected