a
    d2=                  
   @   s  d Z ddlZddlZddlZddlmZ ddlmZ ej	g dg dg dg dg dg d	g d
g dgej
djZeeeZejdej
dZed e	g dg dg dg dgjeddddf< eeeZdd Zdd ZG dd dejZG dd dejZG dd dejZG dd dejZG d d! d!ejZG d"d# d#ejZG d$d% d%ejZG d&d' d'ejZG d(d) d)ejZG d*d+ d+ejZG d,d- d-ejZG d.d/ d/ejZ G d0d1 d1ejZ!G d2d3 d3ejZ"G d4d5 d5ejZ#e$d6krddl%Z%dd7l&m'Z'm(Z( e%)d8d9 Z*e+e%j,d:gZ-e%.d;e*d9 e-\Z/Z0e
e%1e0d<Z2e%3d=e2 e#d>d?4 Z5e'e*Z*e6e*e*g4 Z*e*7d:d@gZ8e5e*e8dAZ9e%3dBe(e9d  e%3dCe(e9d<  dS )Dz
Modified from https://github.com/mlomnitz/DiffJPEG

For images not divisible by 8
https://dsp.stackexchange.com/questions/35339/jpeg-dct-padding/35343#35343
    N)
functional)      
   r      (   3   =   )   r
            :   <   7   )r      r   r   r   9   E   8   )r            r   W   P   >   )   r   %   r   D   m   g   M   )r   #   r   @   Q   h   q   \   )1   r"   N   r   r   y   x   e   )H   r&   _   b   p   d   r   c   Zdtype)   r3   r1   )r   r   r   /   )r      r   B   )r   r   r   r1   )r4   r6   r1   r1      c                 C   s   t | | t |  d  S )z& Differentiable rounding function
       )torchround)x r<   _/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/basicsr/utils/diffjpeg.py
diff_round   s    r>   c                 C   s&   | dk rd|  } nd| d  } | d S )z Calculate factor corresponding to quality

    Args:
        quality(float): Quality for jpeg compression.

    Returns:
        float: Compression factor.
    2   g     @g      i@   g      Y@r<   qualityr<   r<   r=   quality_to_factor    s    	
rC   c                       s(   e Zd ZdZ fddZdd Z  ZS )RGB2YCbCrJpegz! Converts RGB image to YCbCr
    c                    s^   t t|   tjg dg dg dgtjdj}tt	
g d| _tt	|| _d S )N)gA`"?gbX9?gv/?)g4($ſgm3տ      ?)rE   gɑڿgܸдr2   )              `@rG   )superrD   __init__nparrayfloat32Tnn	Parameterr9   tensorshift
from_numpymatrixselfrS   	__class__r<   r=   rI   5   s    zRGB2YCbCrJpeg.__init__c                 C   s4   | dddd}tj|| jdd| j }||jS )z
        Args:
            image(Tensor): batch x 3 x height x width

        Returns:
            Tensor: batch x height x width x 3
        r   r@   r8      dims)permuter9   	tensordotrS   rQ   viewshaperU   imageresultr<   r<   r=   forward<   s    zRGB2YCbCrJpeg.forward__name__
__module____qualname____doc__rI   rb   __classcell__r<   r<   rV   r=   rD   1   s   rD   c                       s(   e Zd ZdZ fddZdd Z  ZS )ChromaSubsamplingz) Chroma subsampling on CbCr channels
    c                    s   t t|   d S N)rH   ri   rI   rU   rV   r<   r=   rI   M   s    zChromaSubsampling.__init__c                 C   s   | dddd }tj|dddddddf ddddd}tj|dddddddf ddddd}| dddd}| dddd}|dddddddf |d|dfS )	z
        Args:
            image(tensor): batch x height x width x 3

        Returns:
            y(tensor): batch x height x width
            cb(tensor): batch x height/2 x width/2
            cr(tensor): batch x height/2 x width/2
        r   r8   rX   r@   N)r@   r@   F)Zkernel_sizeZstrideZcount_include_pad)r[   cloneFZ
avg_pool2d	unsqueezeZsqueeze)rU   r`   Zimage_2cbcrr<   r<   r=   rb   P   s    
00zChromaSubsampling.forwardrc   r<   r<   rV   r=   ri   I   s   ri   c                       s(   e Zd ZdZ fddZdd Z  ZS )BlockSplittingz" Splitting image into patches
    c                    s   t t|   d| _d S )Nr3   )rH   rq   rI   krk   rV   r<   r=   rI   f   s    zBlockSplitting.__init__c                 C   sb   |j dd \}}|j d }|||| j | jd| j}|ddddd}| |d| j| jS )z
        Args:
            image(tensor): batch x height x width

        Returns:
            Tensor:  batch x h*w/64 x h x w
        rX   r8   r   r@   r7   )r^   r]   rr   r[   
contiguous)rU   r`   height_
batch_sizeimage_reshapedimage_transposedr<   r<   r=   rb   j   s
    
zBlockSplitting.forwardrc   r<   r<   rV   r=   rq   b   s   rq   c                       s(   e Zd ZdZ fddZdd Z  ZS )DCT8x8z$ Discrete Cosine Transformation
    c                    s   t t|   tjdtjd}tjtdddD ]X\}}}}t	d| d | tj
 d t	d| d | tj
 d  |||||f< q0td	td gdgd
  }tt| | _ttt||d  | _d S )Nr3   r3   r3   r3   r2   r3   r7   repeatr@   rX   r         ?         ?)rH   rz   rI   rJ   zerosrL   	itertoolsproductrangecospirK   sqrtrN   rO   r9   rR   floatrP   outerscale)rU   rP   r;   yuvalpharV   r<   r=   rI   }   s    N zDCT8x8.__init__c                 C   s0   |d }| j tj|| jdd }||j |S )
        Args:
            image(tensor): batch x height x width

        Returns:
            Tensor: batch x height x width
           r@   rY   )r   r9   r\   rP   r]   r^   r_   r<   r<   r=   rb      s    zDCT8x8.forwardrc   r<   r<   rV   r=   rz   y   s   	rz   c                       s*   e Zd ZdZ fddZdddZ  ZS )	YQuantizeze JPEG Quantization for Y channel

    Args:
        rounding(function): rounding function to use
    c                    s   t t|   || _t| _d S rj   )rH   r   rI   roundingy_tablerU   r   rV   r<   r=   rI      s    zYQuantize.__init__rX   c                 C   sh   t |ttfr"| | j|  }n8|d}| j|ddd||ddd }| | }| |}|S r   r   rX   r3   )
isinstanceintr   r   sizeexpandr]   r   rU   r`   factorbtabler<   r<   r=   rb      s    
"
zYQuantize.forward)rX   rc   r<   r<   rV   r=   r      s   r   c                       s*   e Zd ZdZ fddZdddZ  ZS )	CQuantizezi JPEG Quantization for CbCr channels

    Args:
        rounding(function): rounding function to use
    c                    s   t t|   || _t| _d S rj   )rH   r   rI   r   c_tabler   rV   r<   r=   rI      s    zCQuantize.__init__rX   c                 C   sh   t |ttfr"| | j|  }n8|d}| j|ddd||ddd }| | }| |}|S r   )r   r   r   r   r   r   r]   r   r   r<   r<   r=   rb      s    
"
zCQuantize.forward)rX   rc   r<   r<   rV   r=   r      s   r   c                       s0   e Zd ZdZejf fdd	ZdddZ  ZS )CompressJpegzdFull JPEG compression algorithm

    Args:
        rounding(function): rounding function to use
    c                    sN   t t|   tt t | _tt t	 | _
t|d| _t|d| _d S N)r   )rH   r   rI   rN   Z
SequentialrD   ri   l1rq   rz   l2r   
c_quantizer   
y_quantizer   rV   r<   r=   rI      s
    zCompressJpeg.__init__rX   c           	      C   s   |  |d \}}}|||d}| D ]@}| || }|dv rR| j||d}n| j||d}|||< q(|d |d |d fS )z
        Args:
            image(tensor): batch x 3 x height x width

        Returns:
            dict(tensor): Compressed tensor with batch x h*w/64 x 8 x 8.
           r   ro   rp   ro   rp   r   r   ro   rp   )r   keysr   r   r   )	rU   r`   r   r   ro   rp   
componentsrr   compr<   r<   r=   rb      s    
zCompressJpeg.forward)rX   	rd   re   rf   rg   r9   r:   rI   rb   rh   r<   r<   rV   r=   r      s   r   c                       s*   e Zd ZdZ fddZdddZ  ZS )YDequantizezDequantize Y channel
    c                    s   t t|   t| _d S rj   )rH   r   rI   r   rk   rV   r<   r=   rI      s    zYDequantize.__init__rX   c                 C   sV   t |ttfr|| j|  }n4|d}| j|ddd||ddd }|| }|S r   )r   r   r   r   r   r   r]   rU   r`   r   outr   r   r<   r<   r=   rb      s    
"zYDequantize.forward)rX   rc   r<   r<   rV   r=   r      s   r   c                       s*   e Zd ZdZ fddZdddZ  ZS )CDequantizezDequantize CbCr channel
    c                    s   t t|   t| _d S rj   )rH   r   rI   r   rk   rV   r<   r=   rI     s    zCDequantize.__init__rX   c                 C   sV   t |ttfr|| j|  }n4|d}| j|ddd||ddd }|| }|S r   )r   r   r   r   r   r   r]   r   r<   r<   r=   rb     s    
"zCDequantize.forward)rX   rc   r<   r<   rV   r=   r     s   r   c                       s(   e Zd ZdZ fddZdd Z  ZS )iDCT8x8z+Inverse discrete Cosine Transformation
    c                    s   t t|   tdtd gdgd  }tt	t
|| | _tjdtjd}tjtddd	D ]X\}}}}td| d | tj d
 td| d | tj d
  |||||f< qntt	| | _d S )Nr~   r@   rX   r   r{   r2   r3   r7   r|   r   )rH   r   rI   rJ   rK   r   rN   rO   r9   rR   r   r   r   r   rL   r   r   r   r   r   rP   )rU   r   rP   r;   r   r   r   rV   r<   r=   rI   -  s     NziDCT8x8.__init__c                 C   s4   || j  }dtj|| jdd d }||j |S )r   r   r@   rY   r   )r   r9   r\   rP   r]   r^   r_   r<   r<   r=   rb   6  s    
ziDCT8x8.forwardrc   r<   r<   rV   r=   r   )  s   	r   c                       s(   e Zd ZdZ fddZdd Z  ZS )BlockMergingzMerge patches into image
    c                    s   t t|   d S rj   )rH   r   rI   rk   rV   r<   r=   rI   H  s    zBlockMerging.__init__c                 C   sL   d}|j d }|||| || ||}|ddddd}| |||S )z
        Args:
            patches(tensor) batch x height*width/64, height x width
            height(int)
            width(int)

        Returns:
            Tensor: batch x height x width
        r3   r   rX   r8   r@   r7   )r^   r]   r[   rt   )rU   Zpatchesru   widthrr   rw   rx   ry   r<   r<   r=   rb   K  s
    

zBlockMerging.forwardrc   r<   r<   rV   r=   r   D  s   r   c                       s(   e Zd ZdZ fddZdd Z  ZS )ChromaUpsamplingzUpsample chroma layers
    c                    s   t t|   d S rj   )rH   r   rI   rk   rV   r<   r=   rI   `  s    zChromaUpsampling.__init__c                 C   s@   ddd}||}||}t j|d|d|dgddS )z
        Args:
            y(tensor): y channel image
            cb(tensor): cb channel
            cr(tensor): cr channel

        Returns:
            Tensor: batch x height x width x 3
        r@   c                 S   sF   | j dd \}}| d} | dd||} | d|| || } | S )NrX   r8   rs   )r^   rn   r}   r]   )r;   rr   ru   r   r<   r<   r=   r}   n  s
    
z(ChromaUpsampling.forward.<locals>.repeatr8   )Zdim)r@   )r9   catrn   )rU   r   ro   rp   r}   r<   r<   r=   rb   c  s    
zChromaUpsampling.forwardrc   r<   r<   rV   r=   r   \  s   r   c                       s(   e Zd ZdZ fddZdd Z  ZS )YCbCr2RGBJpegz%Converts YCbCr image to RGB JPEG
    c                    s^   t t|   tjg dg dg dgtjdj}tt	
g d| _tt	|| _d S )N)r~   rF   g;On?)rX   gX Sֿg!3)rX   g'1Z?r   r2   )r         `r   )rH   r   rI   rJ   rK   rL   rM   rN   rO   r9   rP   rQ   rR   rS   rT   rV   r<   r=   rI   ~  s    $zYCbCr2RGBJpeg.__init__c                 C   s0   t j|| j | jdd}||jddddS )z
        Args:
            image(tensor): batch x height x width x 3

        Returns:
            Tensor: batch x 3 x height x width
        rX   rY   r   r8   r@   )r9   r\   rQ   rS   r]   r^   r[   r_   r<   r<   r=   rb     s    zYCbCr2RGBJpeg.forwardrc   r<   r<   rV   r=   r   z  s   r   c                       s0   e Zd ZdZejf fdd	ZdddZ  ZS )DeCompressJpegzfFull JPEG decompression algorithm

    Args:
        rounding(function): rounding function to use
    c                    sB   t t|   t | _t | _t | _t	 | _
t | _t | _d S rj   )rH   r   rI   r   c_dequantizer   y_dequantizer   idctr   mergingr   chromar   colorsr   rV   r<   r=   rI     s    zDeCompressJpeg.__init__rX   c                 C   s   |||d}|  D ]r}|dv rN| j|| |d}	t|d t|d  }
}n| j|| |d}	|| }
}| |	}	| |	|
|||< q| |d |d |d }| |}t	dt
| tt||}|d S )	z
        Args:
            compressed(dict(tensor)): batch x h*w/64 x 8 x 8
            imgh(int)
            imgw(int)
            factor(float)

        Returns:
            Tensor: batch x 3 x height x width
        r   r   r   r@   r   ro   rp   r   )r   r   r   r   r   r   r   r   r9   minZ	ones_likemaxZ
zeros_like)rU   r   ro   rp   ZimghZimgwr   r   rr   r   ru   r   r`   r<   r<   r=   rb     s    


$zDeCompressJpeg.forward)rX   r   r<   r<   rV   r=   r     s   	r   c                       s*   e Zd ZdZd fdd	Zdd Z  ZS )DiffJPEGzThis JPEG algorithm result is slightly different from cv2.
    DiffJPEG supports batch processing.

    Args:
        differentiable(bool): If True, uses custom differentiable rounding function, if False, uses standard torch.round
    Tc                    s:   t t|   |rt}ntj}t|d| _t|d| _	d S r   )
rH   r   rI   r>   r9   r:   r   compressr   
decompress)rU   differentiabler   rV   r<   r=   rI     s    zDiffJPEG.__init__c                 C   s   |}t |ttfrt|}n$t|dD ]}t|| ||< q*| dd \}}d\}}|d dkrtd|d  }|d dkrd|d  }tj|d|d|fddd}| j||d\}	}
}| j	|	|
||| || |d}|ddddd|d|f }|S )	z
        Args:
            x (Tensor): Input image, bchw, rgb, [0, 1]
            quality(float): Quality factor for jpeg compression scheme.
        r   N)r   r   r   Zconstant)modevaluer   )
r   r   r   rC   r   r   rm   padr   r   )rU   r;   rB   r   ihwZh_padZw_padr   ro   rp   Z	recoveredr<   r<   r=   rb     s     
 zDiffJPEG.forward)Trc   r<   r<   rV   r=   r     s   
r   __main__)
img2tensor
tensor2imgztest.pngg     o@   z.jpgrX   zcv2_JPEG_20.pngF)r   r   rA   zpt_JPEG_20.pngzpt_JPEG_40.png):rg   r   numpyrJ   r9   Ztorch.nnrN   r   rm   rK   rL   rM   r   rO   rR   emptyr   fillr>   rC   ModulerD   ri   rq   rz   r   r   r   r   r   r   r   r   r   r   r   rd   Zcv2Zbasicsr.utilsr   r   ZimreadZimg_gtr   ZIMWRITE_JPEG_QUALITYZencode_paramZimencoderv   ZencimgZimdecodeZimg_lqZimwriteZcudaZjpegerstackZ
new_tensorrB   r   r<   r<   r<   r=   <module>   sb   
4'0-
