tomaszki/PythonFileCompressor
1
1import os2import torch3import numpy as np4from torch.utils.cpp_extension import load5 6 7PRECISION = 16 # DO NOT EDIT!8 9 10# Load on-the-fly with ninja.11torchac_dir = os.path.dirname(os.path.realpath(__file__))12backend_dir = os.path.join(torchac_dir, 'backend')13numpyAc_backend = load(14 name="numpyAc_backend",15 sources=[os.path.join(backend_dir, "numpyAc_backend.cpp")],16 verbose=False)17 18def _encode_float_cdf(cdf_float,19 sym,20 needs_normalization=True,21 check_input_bounds=False):22 """Encode symbols `sym` with potentially unnormalized floating point CDF.23 24 Check the README for more details.25 26 :param cdf_float: CDF tensor, float32, on CPU. Shape (N1, ..., Nm, Lp).27 :param sym: The symbols to encode, int16, on CPU. Shape (N1, ..., Nm).28 :param needs_normalization: if True, assume `cdf_float` is un-normalized and29 needs normalization. Otherwise only convert it, without normalizing.30 :param check_input_bounds: if True, ensure inputs have valid values.31 Important: may take significant time. Only enable to check.32 33 :return: byte-string, encoding `sym`.34 """35 if check_input_bounds:36 if cdf_float.min() < 0:37 raise ValueError(f'cdf_float.min() == {cdf_float.min()}, should be >=0.!')38 if cdf_float.max() > 1:39 raise ValueError(f'cdf_float.max() == {cdf_float.max()}, should be <=1.!')40 Lp = cdf_float.shape[-1]41 if sym.max() >= Lp - 1:42 raise ValueError(f'sym.max() == {sym.max()}, should be <=Lp - 1.!')43 cdf_int = _convert_to_int_and_normalize(cdf_float, needs_normalization)44 return _encode_int16_normalized_cdf(cdf_int, sym)45 46 47def _encode_int16_normalized_cdf(cdf_int, sym):48 """Encode symbols `sym` with a normalized integer cdf `cdf_int`.49 50 Check the README for more details.51 52 :param cdf_int: CDF tensor, int16, on CPU. Shape (N1, ..., Nm, Lp).53 :param sym: The symbols to encode, int16, on CPU. Shape (N1, ..., Nm).54 55 :return: byte-string, encoding `sym`56 """57 cdf_int, sym = _check_and_reshape_inputs(cdf_int, sym)58 return numpyAc_backend.encode_cdf( torch.ShortTensor(cdf_int), torch.ShortTensor(sym))59 60 61def _check_and_reshape_inputs(cdf, sym=None):62 """Check device, dtype, and shapes."""63 if sym is not None and sym.dtype != np.int16:64 raise ValueError('Symbols must be int16!')65 if sym is not None:66 if len(cdf.shape) != len(sym.shape) + 1 or cdf.shape[:-1] != sym.shape:67 raise ValueError(f'Invalid shapes of cdf={cdf.shape}, sym={sym.shape}! '68 'The first m elements of cdf.shape must be equal to '69 'sym.shape, and cdf should only have one more dimension.')70 Lp = cdf.shape[-1]71 cdf = cdf.reshape(-1, Lp)72 if sym is None:73 return cdf74 sym = sym.reshape(-1)75 return cdf, sym76 77 78# def _reshape_output(cdf_shape, sym):79# """Reshape single dimension `sym` back to the correct spatial dimensions."""80# spatial_dimensions = cdf_shape[:-1]81# if len(sym) != np.prod(spatial_dimensions):82# raise ValueError()83# return sym.reshape(*spatial_dimensions)84 85 86def _convert_to_int_and_normalize(cdf_float, needs_normalization):87 """Convert floatingpoint CDF to integers. See README for more info.88 89 The idea is the following:90 When we get the cdf here, it is (assumed to be) between 0 and 1, i.e,91 cdf \in [0, 1)92 (note that 1 should not be included.)93 We now want to convert this to int16 but make sure we do not get94 the same value twice, as this would break the arithmetic coder95 (you need a strictly monotonically increasing function).96 So, if needs_normalization==True, we multiply the input CDF97 with 2**16 - (Lp - 1). This means that now,98 cdf \in [0, 2**16 - (Lp - 1)].99 Then, in a final step, we add an arange(Lp), which is just a line with100 slope one. This ensure that for sure, we will get unique, strictly101 monotonically increasing CDFs, which are \in [0, 2**16)102 """103 Lp = cdf_float.shape[-1]104 factor = 2**PRECISION105 new_max_value = factor106 if needs_normalization:107 new_max_value = new_max_value - (Lp - 1)108 cdf_float = cdf_float*(new_max_value)109 cdf_float = np.round(cdf_float)110 cdf = cdf_float.astype(np.int16)111 if needs_normalization:112 r = np.arange(Lp) 113 cdf+=r114 return cdf115 116def pdf_convert_to_cdf_and_normalize(pdf):117 assert pdf.ndim==2118 cdfF = np.cumsum( pdf, axis=1)119 cdfF = cdfF/cdfF[:,-1:]120 cdfF = np.hstack((np.zeros((pdf.shape[0],1)),cdfF))121 return cdfF122 123class arithmeticCoding():124 def __init__(self) -> None:125 self.binfile = None126 self.sysNum = None127 self.byte_stream = None128 129 130 def encode(self,pdf,sym,binfile=None):131 assert pdf.shape[0]==sym.shape[0]132 assert pdf.ndim==2 and sym.ndim==1133 134 self.sysNum = sym.shape[0]135 136 cdfF = pdf_convert_to_cdf_and_normalize(pdf)137 138 # pdf = np.diff(cdfF)139 # print( -np.log2(pdf[range(0,self.sysNum),sym]).sum())140 141 self.byte_stream = _encode_float_cdf(cdfF, sym, check_input_bounds=True)142 real_bits = len(self.byte_stream) * 8143 # # Write to a file.144 if binfile is not None:145 with open(binfile, 'wb') as fout:146 fout.write(self.byte_stream)147 return self.byte_stream,real_bits148 149class arithmeticDeCoding():150 """151 Decoding class152 byte_stream: the bin file stream.153 sysDim: the Number of the possible symbols.154 binfile: bin file path, if it is Not None, 'byte_stream' will read from this file155 and copy to Cpp backend Class 'InCacheString'156 """157 def __init__(self,byte_stream,symDim,binfile=None) -> None:158 if binfile is not None:159 with open(binfile, 'rb') as fin:160 byte_stream = fin.read() 161 self.byte_stream = byte_stream162 self.decoder = numpyAc_backend.decode(self.byte_stream,symDim+1)163 164 def decode(self,pdf):165 cdfF = pdf_convert_to_cdf_and_normalize(pdf)166 pro = _convert_to_int_and_normalize(cdfF,needs_normalization=True)167 pro = pro.squeeze(0).astype(np.uint16).tolist()168 sym_out = self.decoder.decodeAsym(pro)169 return sym_out170 