a
    d^]                     @   s  U d dl mZmZmZmZmZ d dlZd dlmZ d dl	m  m
Z d dlmZ d dlmZmZ d dlmZ d dlmZ dZeed< d	d
gZee ed< g dZee ed< g dZee ed< eeedZeeee f ed< dd dD Zeeef ed< d1eeeejf dddZ eeejdddZ!G dd dej"Z#G dd  d ej"Z$G d!d" d"ej"Z%e&ejd#d$d%Z'G d&d' d'ej"Z(G d(d) d)ej"Z)G d*d+ d+ej"Z*eeed,d-d.Z+G d/d0 d0ej"Z,dS )2    )CallableDictListTupleUnionN)pi)GaussianBlur2dSpatialGradient)cart2pol)create_meshgridg;f?sqrt2gpu?g "?COEFFS_N1_K1)"g#[?x@ٔ+?K'"?COEFFS_N2_K8)r   r   r   gˈfF?COEFFS_N3_K8)xyrhophithetaCOEFFSc                 C   s   i | ]}|d | dqS )zChttps://github.com/manyids2/mkd_pytorch/raw/master/mkd_pytorch/mkd-z-64.pth ).0kr   r   [/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/kornia/feature/mkd.py
<dictcomp>   s   r   )cartpolarconcaturls    )
patch_sizereturnc                 C   s^   t | | dd}|ddddddf }|ddddddf }t||\}}||||d}|S )z1Get cartesian and polar parametrizations of grid.TheightwidthZnormalized_coordinatesr   N   )xyrhophi)r   r
   )r!   kgridr'   r(   r)   r*   Z	grid_dictr   r   r   get_grid_dict   s    r,   )d1d2r"   c                 C   s`   t j| | dgt jd}t| D ]:}t|D ],}|||| | df< |||| | df< q,q |S )z&Get order for doing kronecker product.   )Zdtyper   r&   )torchzerosint64range)r-   r.   
kron_orderijr   r   r   get_kron_order#   s    r7   c                       sH   e Zd ZdZdd fddZejejdddZedd	d
Z	  Z
S )MKDGradientsa  Module, which computes gradients of given patches, stacked as [magnitudes, orientations].

    Given gradients $g_x$, $g_y$ with respect to $x$, $y$ respectively,
      - $\mathbox{mags} = $\sqrt{g_x^2 + g_y^2 + eps}$
      - $\mathbox{oris} = $\mbox{tan}^{-1}(\nicefrac{g_y}{g_x})$.

    Args:
        patch_size: Input patch size in pixels.

    Returns:
        gradients of given patches.

    Shape:
        - Input: (B, 1, patch_size, patch_size)
        - Output: (B, 2, patch_size, patch_size)

    Example:
        >>> patches = torch.rand(23, 1, 32, 32)
        >>> gradient = MKDGradients()
        >>> g = gradient(patches) # 23x2x32x32
    Nr"   c                    s$   t    d| _tdddd| _d S )N:0yE>diffr&   F)modeorder
normalized)super__init__epsr	   gradself	__class__r   r   r@   D   s    
zMKDGradients.__init__r'   r"   c                 C   s   t |tjstdt| t|jdks<td|j | | }|d d d d dd d d d f }|d d d d dd d d d f }tj	t
||| jdd}|S )N&Input type is not a torch.Tensor. Got    -Invalid input shape, we expect Bx1xHxW. Got: r   r&   Zdim)
isinstancer0   Tensor	TypeErrortypelenshape
ValueErrorrB   catr
   rA   )rD   r'   Zgrads_xyZgxgyr(   r   r   r   forwardJ   s    ""zMKDGradients.forwardc                 C   s   | j jS N)rF   __name__rC   r   r   r   __repr__V   s    zMKDGradients.__repr__)rW   
__module____qualname____doc__r@   r0   rM   rU   strrX   __classcell__r   r   rE   r   r8   -   s   r8   c                       sT   e Zd ZdZeeeef dd fddZe	j
e	j
dddZed	d
dZ  ZS )VonMisesKernela  Module, which computes parameters of Von Mises kernel given coefficients, and embeds given patches.

    Args:
        patch_size: Input patch size in pixels.
        coeffs: List of coefficients. Some examples are hardcoded in COEFFS,

    Returns:
        Von Mises embedding of given parametrization.

    Shape:
        - Input: (B, 1, patch_size, patch_size)
        - Output: (B, d, patch_size, patch_size)

    Examples:
        >>> oris = torch.rand(23, 1, 32, 32)
        >>> vm = VonMisesKernel(patch_size=32,
        ...                     coeffs=[0.14343168,
        ...                             0.268285,
        ...                             0.21979234])
        >>> emb = vm(oris) # 23x7x32x32
    N)r!   coeffsr"   c                    s   t    || _t|}| d| t|d }|| _d| d | _t	dd||g}t
|d }|ddd}td| d g}t||d |d < t|dd  ||d d < |ddd}| d| | d| | d| d S )Nr_   r&   r/   emb0frangeweights)r?   r@   r!   r0   Ztensorregister_bufferrP   ndonesZarangeZreshaper1   sqrt)rD   r!   r_   Zb_coeffsre   ra   rb   rc   rE   r   r   r@   q   s"    

zVonMisesKernel.__init__rG   c                 C   s   t |tjstdt| t|jdkr:|jd dkrJtd|j tj	tj| j
}|||dddd}| j|| }t|}t|}tj|||gdd}| j| }|S )NrH   rI   r&   rJ   r   rK   )rL   r0   rM   rN   rO   rP   rQ   rR   jitannotatera   torepeatsizerb   cossinrS   rc   )rD   r'   ra   rb   emb1emb2Z	embeddingr   r   r   rU      s    


zVonMisesKernel.forwardr9   c                 C   sT   | j jd d t| j d d t| j d d t| j d d t| j d S )N(patch_size=, zn=zd=zcoeffs=))rF   rW   r\   r!   re   rf   r_   rC   r   r   r   rX      s8    	
zVonMisesKernel.__repr__)rW   rY   rZ   r[   intr   listtupler@   r0   rM   rU   r\   rX   r]   r   r   rE   r   r^   Z   s   r^   c                       sb   e Zd ZdZdeedd fddZejejdd	d
Z	ejejdddZ
edddZ  ZS )EmbedGradientsa<  Module that computes gradient embedding, weighted by sqrt of magnitudes of given patches.

    Args:
        patch_size: Input patch size in pixels.
        relative: absolute or relative gradients.

    Returns:
        Gradient embedding.

    Shape:
        - Input: (B, 2, patch_size, patch_size)
        - Output: (B, 7, patch_size, patch_size)

    Examples:
        >>> grads = torch.rand(23, 2, 32, 32)
        >>> emb_grads = EmbedGradients(patch_size=32,
        ...                            relative=False)
        >>> emb = emb_grads(grads) # 23x7x32x32
    r    FN)r!   relativer"   c                    s   t    || _|| _d| _t|td d| _t||dd}t	|d d d d d d df |d d d d d d df \}}| 
d| d S )	Nr:   r   r!   r_   Tr#   r   r&   r*   )r?   r@   r!   rz   rA   r^   r   kernelr   r
   rd   )rD   r!   rz   r+   _r*   rE   r   r   r@      s    
>zEmbedGradients.__init__)magsr"   c                 C   s   t || j }|S )z@Embed square roots of magnitudes with eps for numerical reasons.)r0   rh   rA   )rD   r~   r   r   r   emb_mags   s    zEmbedGradients.emb_mags)gradsr"   c                 C   s   t |tjstdt| t|jdks<td|j |d d d dd d d d f }|d d dd d d d d f }| jr|| j	
| }| || | }|S )NrH   rI   z-Invalid input shape, we expect Bx2xHxW. Got: r&   )rL   r0   rM   rN   rO   rP   rQ   rR   rz   r*   rk   r|   r   )rD   r   r~   Zorisr(   r   r   r   rU      s      zEmbedGradients.forwardr9   c                 C   s0   | j jd d t| j d d t| j d S )Nrr   rs   rt   z	relative=ru   )rF   rW   r\   r!   rz   rC   r   r   r   rX      s     zEmbedGradients.__repr__)r    F)rW   rY   rZ   r[   rv   boolr@   r0   rM   r   rU   r\   rX   r]   r   r   rE   r   ry      s
   ry   )gridsr"   c                    s  dt t t d t d d | dkr0d}ddg}n| dkrDd	}d
dg}t| }||d  jd } fdd| D }dd | D }t|t| d}t|t| d}|||d   }	|||d   }
t	|j
|j
}|	d|dddf |
d|dddf  }|S )z<Compute embeddings for cartesian and polar parametrizations.      ?r/   )r*   r)   r'   r(   r   r   r'   r(   r   r   r*   r)   r   r`   c                    s   i | ]\}}|| |  qS r   r   r   r   vZfactorsr   r   r          z,spatial_kernel_embedding.<locals>.<dictcomp>c                 S   s&   i | ]\}}|| d  d  qS )r   )	unsqueezefloatr   r   r   r   r      r   r{   r&   N)r   r   rw   keysrQ   itemsr^   r   Zsqueezer7   rf   index_select)kernel_typer   Zcoeffs_Zparams_r   r!   Zgrids_normedZvm_aZvm_bZemb_aZemb_br4   Zspatial_kernelr   r   r   spatial_kernel_embedding   s$    
0r   c                       s   e Zd ZdZdeeeeedd fdd	Zee	j
d
ddZee	j
e	j
f dddZe	j
e	j
dddZedddZ  ZS )ExplicitSpacialEncodinga  Module that computes explicit cartesian or polar embedding.

    Args:
        kernel_type: Parametrization of kernel ``'polar'`` or ``'cart'``.
        fmap_size: Input feature map size in pixels.
        in_dims: Dimensionality of input feature map.
        do_gmask: Apply gaussian mask.
        do_l2: Apply l2-normalization.

    Returns:
        Explicit cartesian or polar embedding.

    Shape:
        - Input: (B, in_dims, fmap_size, fmap_size)
        - Output: (B, out_dims, fmap_size, fmap_size)

    Example:
        >>> emb_ori = torch.rand(23, 7, 32, 32)
        >>> ese = ExplicitSpacialEncoding(kernel_type='polar',
        ...                               fmap_size=32,
        ...                               in_dims=7,
        ...                               do_gmask=True,
        ...                               do_l2=True)
        >>> desc = ese(emb_ori) # 23x175x32x32
    r   r       TN)r   	fmap_sizein_dimsdo_gmaskdo_l2r"   c           	         s   t    |dvr t| d|| _|| _|| _|| _|| _t|| _	d | _
t| j| j	}| jrz| jdd| _
|| j
 }| d|d |jd | _| j| j | _| j| _|  \}}| d| | d| d S )	N)r   r   z" is not valid, use polar or cart).r   )sigmaembr   rq   idx1)r?   r@   NotImplementedErrorr   r   r   r   r   r,   gridgmaskr   	get_gmaskrd   r   rQ   d_embout_dimsodims	init_kron)	rD   r   r   r   r   r   r   rq   r   rE   r   r   r@   (  s*    


z ExplicitSpacialEncoding.__init__)r   r"   c                 C   s6   | j d | j d   }td|d  |d  }|S )zCompute Gaussian mask.r)   r`   r/   )r   maxr0   exp)rD   r   Znorm_rhor   r   r   r   r   P  s    z!ExplicitSpacialEncoding.get_gmaskr9   c                 C   sN   t | j| j}tjtj| j}t|d|dddf }||dddf fS )z3Initialize helper variables to calculate kronecker.r&   Nr   )	r7   r   r   r0   ri   rj   rM   r   r   )rD   ZkronZ_embrq   r   r   r   r   V  s    z!ExplicitSpacialEncoding.init_kronrG   c                 C   s   t |tjstdt| t|jdk|jd | jkB sTtd| j d|j tj	
tj| j}t|d|}|| j }|jdd}| jrtj|dd}|S )NrH   rI   r&   z!Invalid input shape, we expect BxzxHxW. Got: )r/      rK   )rL   r0   rM   rN   rO   rP   rQ   r   rR   ri   rj   r   r   rq   sumr   F	normalize)rD   r'   r   rp   outputr   r   r   rU   ]  s    
zExplicitSpacialEncoding.forwardc                 C   sx   | j jd d t| j d d t| j d d t| j d d t| j d d t| j d d t| j d	 S )
Nrr   kernel_type=rt   z
fmap_size=in_dims=z	out_dims=z	do_gmask=zdo_l2=ru   )	rF   rW   r\   r   r   r   r   r   r   rC   r   r   r   rX   j  sP    	
z ExplicitSpacialEncoding.__repr__)r   r    r   TT)rW   rY   rZ   r[   r\   rv   r   r@   r   r0   rM   r   r   r   rU   rX   r]   r   r   rE   r   r     s$        (r   c                	       s   e Zd ZdZdeeeeeeejf f df e	e	e	e
dd fddZeeeeejf f dd	d
dZddddZddddZddddZddddZejejdddZedddZ  ZS )	WhiteningaG  Module, performs supervised or unsupervised whitening.

    This is based on the paper "Understanding and Improving Kernel Local Descriptors".
    See :cite:`mukundan2019understanding` for more details.

    Args:
        xform: Variant of whitening to use. None, 'lw', 'pca', 'pcaws', 'pcawt'.
        whitening_model: Dictionary with keys 'mean', 'eigvecs', 'eigvals' holding torch.Tensors.
        in_dims: Dimensionality of input descriptors.
        output_dims: (int) Dimensionality reduction.
        keval: Shrinkage parameter.
        t: Attenuation parameter.

    Returns:
        l2-normalized, whitened descriptors.

    Shape:
        - Input: (B, in_dims, fmap_size, fmap_size)
        - Output: (B, out_dims, fmap_size, fmap_size)

    Examples:
        >>> descs = torch.rand(23, 238)
        >>> whitening_model = {'pca': {'mean': torch.zeros(238),
        ...                            'eigvecs': torch.eye(238),
        ...                            'eigvals': torch.ones(238)}}
        >>> whitening = Whitening(xform='pcawt',
        ...                       whitening_model=whitening_model,
        ...                       in_dims=238,
        ...                       output_dims=128,
        ...                       keval=40,
        ...                       t=0.7)
        >>> wdescs = whitening(descs) # 23x128
       (   ffffff?N)xformwhitening_modelr   output_dimskevaltr"   c                    s   t    || _|| _|| _|| _d| _t||}|| _t	j
t|dd| _t	j
t|d d d |f dd| _t	j
t|d | dd| _|d ur| | d S )Nr   T)Zrequires_grad)r?   r@   r   r   r   r   pvalminr   nn	Parameterr0   r1   meanZeyeevecsrg   evalsload_whitening_parameters)rD   r   r   r   r   r   r   rE   r   r   r@     s    	

&zWhitening.__init__)r   r"   c                 C   s   | j dkrdnd}|| }|d | j_|d d d d | jf | j_|d d | j | j_| j| j| j| j	d}|| j    d S )Nlwpcar   ZeigvecsZeigvals)r   r   Zpcawspcawt)
r   r   datar   r   r   _modify_pca
_modify_lw_modify_pcaws_modify_pcawt)rD   r   algoZwh_modelZmodificationsr   r   r   r     s    z#Whitening.load_whitening_parametersr9   c                 C   s
   d| _ dS )zModify powerlaw parameter.g      ?N)r   rC   r   r   r   r     s    zWhitening._modify_pcac                 C   s   dS )zNo modification required.Nr   rC   r   r   r   r     s    zWhitening._modify_lwc                 C   s>   | j | j }d| | j  | }| jtt|d | j_dS )zShrinkage for eigenvalues.r&         N)r   r   r   r0   diagpowr   )rD   alphar   r   r   r   r     s    zWhitening._modify_pcawsc                 C   s,   d| j  }| jtt| j| | j_dS )zAttenuation for eigenvalues.r   N)r   r   r0   r   r   r   r   )rD   mr   r   r   r     s    
zWhitening._modify_pcawtrG   c                 C   s|   t |tjstdt| t|jdks<td|j || j }|| j	 }t
|tt|| j }tj|ddS )NrH   r/   z)Invalid input shape, we expect NxD. Got: r&   rK   )rL   r0   rM   rN   rO   rP   rQ   rR   r   r   signr   absr   r   r   rD   r'   r   r   r   rU     s    

zWhitening.forwardc                 C   sB   | j jd d t| j d d t| j d d t| j d S )Nrr   zxform=rt   r   output_dims=ru   )rF   rW   r\   r   r   r   rC   r   r   r   rX     s,    	
zWhitening.__repr__)r   r   r   )rW   rY   rZ   r[   r\   r   r   r0   rM   rv   r   r@   r   r   r   r   r   rU   rX   r]   r   r   rE   r   r     s(   '   "
r   c                       sT   e Zd ZdZdeeeeedd fd	d
ZejejdddZ	edddZ
  ZS )MKDDescriptora  Module that computes Multiple Kernel local descriptors.

    This is based on the paper "Understanding and Improving Kernel Local Descriptors".
    See :cite:`mukundan2019understanding` for more details.

    Args:
        patch_size: Input patch size in pixels.
        kernel_type: Parametrization of kernel ``'concat'``, ``'cart'``, ``'polar'``.
        whitening: Whitening transform to apply ``None``, ``'lw'``, ``'pca'``, ``'pcawt'``, ``'pcaws'``.
        training_set: Set that model was trained on ``'liberty'``, ``'notredame'``, ``'yosemite'``.
        output_dims: Dimensionality reduction.

    Returns:
        Explicit cartesian or polar embedding.

    Shape:
        - Input: :math:`(B, in_{dims}, fmap_{size}, fmap_{size})`.
        - Output: :math:`(B, out_{dims}, fmap_{size}, fmap_{size})`,

    Examples:
        >>> patches = torch.rand(23, 1, 32, 32)
        >>> mkd = MKDDescriptor(patch_size=32,
        ...                     kernel_type='concat',
        ...                     whitening='pcawt',
        ...                     training_set='liberty',
        ...                     output_dims=128)
        >>> desc = mkd(patches) # 23x128
    r    r   r   libertyr   Nr!   r   	whiteningtraining_setr   r"   c                    sD  t    || _|| _|| _|| _d|d  | _td| j| jfd| _t	 | _
d}d}| jdkrh||gn| jg| _d| _|d	|d
i}i | _| jD ]H}	t|||	 d}
t|	||
jjd}t|
|| j|	< |  j|j7  _qt|| j| _| jd ur8dd }tjjt| j |d}|| }t||| j| jd| _| j| _|   d S )Nffffff?@      r   	replicater   r   r   r   TFr!   rz   r   r   r   c                 S   s   | S rV   r   Zstoragelocr   r   r   <lambda>F  r   z(MKDDescriptor.__init__.<locals>.<lambda>Zmap_locationr   r   )r?   r@   r!   r   r   r   r   r   	smoothingr8   	gradientsparametrizationsr   featsry   r   r|   rf   r   
Sequentialr   r   r0   hubload_state_dict_from_urlr   r   whitening_layereval)rD   r!   r   r   r   r   Zpolar_sZcart_sZrelative_orientationsparametrizationZgradient_embeddingZspatial_encodingstorage_fcnwhitening_modelsr   rE   r   r   r@     s>    


zMKDDescriptor.__init__)patchesr"   c                 C   s   t |tjstdt| t|jdks<td|j | |}| 	|}g }| j
D ]*}| j| |j || j| | qZtj|dd}tj|dd}| jd ur| |}|S )NrH   rI   rJ   r&   rK   )rL   r0   rM   rN   rO   rP   rQ   rR   r   r   r   r   rk   ZdeviceappendrS   r   r   r   r   )rD   r   gfeaturesr   r(   r   r   r   rU   O  s    




zMKDDescriptor.forwardr9   c                 C   sf   | j jd d t| j d d t| j d d t| j d d t| j d d t| j d S )	Nrr   rs   rt   r   z
whitening=ztraining_set=r   ru   )rF   rW   r\   r!   r   r   r   r   rC   r   r   r   rX   j  sD    	
zMKDDescriptor.__repr__)r    r   r   r   r   )rW   rY   rZ   r[   rv   r\   r@   r0   rM   rU   rX   r]   r   r   rE   r   r      s         1r   )r   r   r"   c                 C   s(   dd }t jjt|  |d}|| }|S )Nc                 S   s   | S rV   r   r   r   r   r   r     r   z&load_whitening_model.<locals>.<lambda>r   )r0   r   r   r   )r   r   r   r   r   r   r   r   load_whitening_model  s    r   c                       sF   e Zd ZdZdeeeeedd fd	d
ZejejdddZ	  Z
S )SimpleKDz+Example to write custom Kernel Descriptors.r    r   r   r   r   Nr   c                    s   t    |dk}d|d  }|| _td||fd}t }	t||d}
t|||
jjd}t	|t
|||j|d}t||	|
||| _d S )	Nr   r   r   r   r   r   r   r   )r?   r@   r!   r   r8   ry   r   r|   rf   r   r   r   r   r   r   )rD   r!   r   r   r   r   rz   r   r   r   ZoriZeseZwhrE   r   r   r@     s    
zSimpleKD.__init__rG   c                 C   s
   |  |S rV   )r   r   r   r   r   rU     s    zSimpleKD.forward)r    r   r   r   r   )rW   rY   rZ   r[   rv   r\   r@   r0   rM   rU   r]   r   r   rE   r   r     s        r   )r    )-typingr   r   r   r   r   r0   Ztorch.nnr   Ztorch.nn.functionalZ
functionalr   Zkornia.constantsr   Zkornia.filtersr   r	   Zkornia.geometry.conversionsr
   Zkornia.utilsr   r   r   __annotations__r   r   r   r   r\   r   rv   rM   r,   r7   Moduler8   r^   ry   dictr   r   r   r   r   r   r   r   r   r   <module>   s6    

-SAv} 