MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / decompressed_tensor_value_loader

Method decompressed_tensor_value_loader

lite/src/mge/common.cpp:16–61  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

14};
15
16void 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
63LTensorLayout lite::to_impl_layout(const Layout& layout) {
64 mgb::TensorLayout mge_layout;

Callers

nothing calls this directly

Calls 5

CompressionMethodEnum · 0.85
readMethod · 0.45
total_nr_elemsMethod · 0.45
dataMethod · 0.45
skipMethod · 0.45

Tested by

no test coverage detected