Mage-Fans
/

File size: 2,456 Bytes
8ad2fa0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
// Copyright (c) Microsoft Corporation.
// Licensed under the MIT License.

#pragma once
#include "rans.h"
#include <memory>

#include <pybind11/numpy.h>
#include <pybind11/pybind11.h>
#include <pybind11/stl.h>

namespace py = pybind11;

// the classes in this file only perform the type conversion
// from python type (numpy) to C++ type (vector)
class RansEncoder {
public:
    RansEncoder();

    RansEncoder(const RansEncoder&) = delete;
    RansEncoder(RansEncoder&&) = delete;
    RansEncoder& operator=(const RansEncoder&) = delete;
    RansEncoder& operator=(RansEncoder&&) = delete;

    void encode_y(const py::array_t<int16_t>& symbols, const int cdf_group_index);
    void encode_z(const py::array_t<int8_t>& symbols, const int cdf_group_index,
                  const int start_offset, const int per_channel_size);
    void flush();
    py::array_t<uint8_t> get_encoded_stream();
    void reset();
    int add_cdf(const py::array_t<int32_t>& cdfs, const py::array_t<int32_t>& cdfs_sizes,
                const py::array_t<int32_t>& offsets);
    void empty_cdf_buffer();
    void set_use_two_encoders(bool b);
    bool get_use_two_encoders();

private:
    std::shared_ptr<RansEncoderLib> m_encoder0;
    std::shared_ptr<RansEncoderLib> m_encoder1;
    bool m_use_two_encoders{ false };
};

class RansDecoder {
public:
    RansDecoder();

    RansDecoder(const RansDecoder&) = delete;
    RansDecoder(RansDecoder&&) = delete;
    RansDecoder& operator=(const RansDecoder&) = delete;
    RansDecoder& operator=(RansDecoder&&) = delete;

    void set_stream(const py::array_t<uint8_t>&);

    void decode_y(const py::array_t<uint8_t>& indexes, const int cdf_group_index);
    py::array_t<int8_t> decode_and_get_y(const py::array_t<uint8_t>& indexes, const int cdf_group_index);
    void decode_z(const int total_size, const int cdf_group_index, const int start_offset,
                  const int per_channel_size);
    py::array_t<int8_t> get_decoded_tensor();
    int add_cdf(const py::array_t<int32_t>& cdfs, const py::array_t<int32_t>& cdfs_sizes,
                const py::array_t<int32_t>& offsets);
    void empty_cdf_buffer();
    void set_use_two_decoders(bool b);
    bool get_use_two_decoders();

private:
    std::shared_ptr<RansDecoderLib> m_decoder0;
    std::shared_ptr<RansDecoderLib> m_decoder1;
    bool m_use_two_decoders{ false };
};

std::vector<uint32_t> pmf_to_quantized_cdf(const std::vector<float>& pmf, int precision);