Team Ai
Apppublic

tomaszki/PythonFileCompressor

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
numpyAc.py170 linesDownload Raw Back to numpyAc
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