a
    d)                     @   s  d Z ddlZddlZddlZddlmZmZ ddlmZ ddl	m
Z
 ddlmZmZmZmZmZmZmZmZmZ ddlZddlmZ ddlmZ g d	Zed
ddZdUddZdVddZejdfddZe ej!ej"e dddZ#dd Z$dd Z%dd Z&dWddZ'G d d! d!eZ(d"d# Z)i d$d%d&Z*ej!ej"e+d'd(d)Z,dXej!e-e+e+e+d+d,d-Z.zHdd.l/m0Z1 dd/l2m3Z3 ddd0ejejee+ ee+ edd1d2d3Z0W n\ e4y   dd4l/m5Z1 G d5d6 d6e6Z7ddd0ejejee+ ee+ edd1d7d3Z0Y n0 ee- dd8d9d:Z8dYeee- d;d<d=Z9d>d? Z:dZee- d@dAdBZ;d[ee- d@dCdDZ<d\ee ee- dEdFdGZ=d]eee- dHdIdJZ>d^eee- dHdKdLZ?d_eee- dHdMdNZ@eeedOdPdQZAeddRdSdTZBdS )`z8The testing package contains testing-specific utilities.    N)ABCabstractmethod)deepcopy)product)	AnyIterableListOptionalTupleTypeTypeVarUnioncast)Tensor)	Parameter)tensor_to_gradcheck_varcreate_eye_batchxla_is_availableassert_close)returnc                   C   s   t jddurdS dS )z6Return whether `torch_xla` is available in the system.Z	torch_xlaNTF)	importlibutil	find_spec r   r   `/var/www/html/stable-diffusion-webui/venv/lib/python3.9/site-packages/kornia/testing/__init__.pyr      s    r   c                 C   s$   t j|||dd||| ddS )z3Create a batch of identity matrices of shape Bx3x3.devicedtype   )torcheyeviewexpand)
batch_sizeeye_sizer   r   r   r   r   r      s    r   MbP?c                 C   s*   t | ||}t| |}||| | S )z5Create a batch of random homographies of shape Bx3x3.)r    ZFloatTensorr   Zuniform_)r$   r%   std_valZstdr!   r   r   r   create_random_homography   s    
r(   Tc                 C   s&   t | stt| | ||S )zConvert the input tensor to a valid variable to check the gradient.

    `gradcheck` needs 64-bit floating point and requires gradient.
    )r    Z	is_tensorAssertionErrortypeZrequires_grad_)tensorr   Zrequires_gradr   r   r   r   %   s    
r   )datar   r   r   c                 C   s:   i }|   D ](\}}t|tjr,|||n|||< q|S N)items
isinstancer    r   to)r,   r   r   outkeyvalr   r   r   dict_to/   s    "r4   c                 C   s8   t | | d|d | d |d | d f  S )z+Compute the absolute error between patches..   )r    absmean)xyhwr   r   r   compute_patch_error6   s    r<   c                 C   s"   t | tjstdt|  dS )z.Check whether the supplied object is a tensor.z&Input type is not a torch.Tensor. Got N)r/   r    r   	TypeErrorr*   )objr   r   r   check_is_tensor;   s    r?   c                 C   s8   t g dg dg dgddd}|| dd}|S )z@Create a batch of rectified fundamental matrices of shape Bx3x3.)        r@   r@   )r@   r@   g      )r@   g      ?r@   r      )r    r+   r"   repeat)r$   F_rectZF_repeatr   r   r   #create_rectified_fundamental_matrixA   s    &rD   c                 C   s6   t | }t| d|}t| d|}|ddd| | S )z=Create a batch of random fundamental matrices of shape Bx3x3.rA   r      r   )rD   r(   Zpermute)r$   r'   rC   ZH_leftZH_rightr   r   r    create_random_fundamental_matrixH   s    rF   c                   @   s   e Zd ZejdejdejdiZedd Z	edd Z
edd Zed	d
 Zedd Zedd Zdeeee ee eddddZdS )
BaseTester)r&   r&   )gkNuϵ>gh㈵>c                 C   s   t dd S NzImplement a stupid routine.NotImplementedErrorselfr   r   r   r   r   
test_smokeS   s    zBaseTester.test_smokec                 C   s   t dd S rH   rI   rK   r   r   r   test_exceptionW   s    zBaseTester.test_exceptionc                 C   s   t dd S rH   rI   rK   r   r   r   test_cardinality[   s    zBaseTester.test_cardinalityc                 C   s   t dd S rH   rI   rK   r   r   r   test_jit_   s    zBaseTester.test_jitc                 C   s   t dd S rH   rI   )rL   r   r   r   r   test_gradcheckc   s    zBaseTester.test_gradcheckc                 C   s   t dd S rH   rI   rK   r   r   r   test_moduleg   s    zBaseTester.test_moduleNF)actualexpectedrtolatollow_tolerancer   c           
      C   s   t |tr|j}t |tr |j}d|jjv s8d|jjv r@d\}}|du r|du r| j|jd\}}| j|jd\}}	t||t||	 }}|rt	
|n|}|rt	
|n|}t||||dS )a  Asserts that `actual` and `expected` are close.

        Args:
            actual: Actual input.
            expected: Expected input.
            rtol: Relative tolerance.
            atol: Absolute tolerance.
            low_tolerance:
                This parameter allows to reduce tolerance. Half the decimal places.
                Example, 1e-4 -> 1e-2 or 1e-6 -> 1e-3
        xla){Gz?rY   N)r@   r@   rU   rV   )r/   r   r,   r   r*   DTYPE_PRECISIONSgetr   maxmathsqrtr   )
rL   rS   rT   rU   rV   rW   Zactual_rtolZactual_atolZexpected_rtolZexpected_atolr   r   r   r   k   s    

zBaseTester.assert_close)NNF)__name__
__module____qualname__r    float16float32float64r[   r   rM   rN   rO   rP   rQ   rR   r   r	   floatboolr   r   r   r   r   rG   P   s0   





   rG   c                  +   s<      } fdd|D }t| D ]}tt||V  q"dS )z-Create cartesian product of given parameters.c                    s   g | ]} | qS r   r   ).0parameter_namepossible_parametersr   r   
<listcomp>       z3cartesian_product_of_parameters.<locals>.<listcomp>N)keysr   dictzip)rk   Zparameter_namespossible_valuesZparam_combinationr   rj   r   cartesian_product_of_parameters   s    rr   )defaultc                 k   sR   t | tstdt|  | D ](\}}|D ]}t| }|||< |V  q0q$d S )Nzdefault should be a dict not a )r/   ro   r)   r*   r.   r   )rs   rk   ri   rq   vZ	param_setr   r   r   "default_with_one_parameter_changed   s    
ru   )r   r   r   c                 C   s    d| j v rdS |tjkrdS dS )NrX   rY   r&   -C6?)r*   r    rc   r   r   r   r   _get_precision   s
    

rw   rv   )r   device_targettol_valtol_val_defaultr   c                 C   s*   |dvrt d| d|| jv r&|S |S )N)cpuZcudarX   zInvalid device name: .)
ValueErrorr*   )r   rx   ry   rz   r   r   r   _get_precision_by_name   s
    
r~   )r   )_get_default_tolerancerZ   )rS   rT   rU   rV   kwargsr   c                K   sd   |d u rH|d u rHt t t| |\}}W d    n1 s>0    Y  t| |f||ddd|S )NFT)rU   rV   Zcheck_strideZ	equal_nan)
contextlibsuppress	Exceptionr   _assert_close)rS   rT   rU   rV   r   r   r   r   r      s    ,r   )assert_allclosec                   @   s   e Zd ZdS )
UsageErrorN)r`   ra   rb   r   r   r   r   r      s   r   c             
   K   sR   zt | |f||d|W S  tyL } ztt||W Y d }~n
d }~0 0 d S )NrZ   )r   r}   r   str)rS   rT   rU   rV   r   errorr   r   r   r      s    )shaper   c                 C   s   t |  d|d kr2d}| jt| d d  }n
d}| j}t|t|D ]D}|| }| s`qJt|}|| |krJt|  d| d| j qJd S )N*r   r   z shape should be must be [z]. Got )KORNIA_CHECK_IS_TENSORr   lenrange	isnumericintr=   )r8   r   Z	start_idxZx_shape_to_checkiZdim_Zdimr   r   r   KORNIA_CHECK_SHAPE   s    r   	conditionmsgc                 C   s   | st |  d| d S )Nz not true.
)r   r   r   r   r   KORNIA_CHECK   s    r   c                 C   s
   t || S r-   )r   )Z	maybe_objtypr   r   r   KORNIA_UNWRAP   s    r   )r   c                 C   s&   t | |s"tdt|  d| d S )NzInvalid type: .
)r/   r=   r*   )r8   r   r   r   r   r   KORNIA_CHECK_TYPE  s    
r   c                 C   s&   t | ts"tdt|  d| d S )NzNot a Tensor type. Got: r   )r/   r   r=   r*   r8   r   r   r   r   r     s    
r   tensorsr   c                    sT   t t tot dkd t fdd D sPtddd  D  d| d S )	Nr   z)Expected a list with at least one elementc                 3   s   | ]} d  j |j kV  qdS )r   Nr   rh   r8   r   r   r   	<genexpr>  rm   z,KORNIA_CHECK_SAME_DEVICES.<locals>.<genexpr>z"Not same device for tensors. Got: c                 S   s   g | ]
}|j qS r   r   r   r   r   r   rl     rm   z-KORNIA_CHECK_SAME_DEVICES.<locals>.<listcomp>r   )r   r/   listr   allr   r   r   r   r   KORNIA_CHECK_SAME_DEVICES  s    r   r   c                 C   s8   t | jdk s| jd dkr4tdt|  d| d S )NrA   zNot a color tensor. Got: r   r   r   r=   r*   r   r   r   r   KORNIA_CHECK_IS_COLOR  s    r   c                 C   sF   t | jdk s*t | jdkrB| jd dkrBtdt|  d| d S )NrE   rA   r   r   zNot a gray tensor. Got: r   r   r   r   r   r   KORNIA_CHECK_IS_GRAY  s    *r   c                 C   s8   t | jdk s| jd dvr4tdt|  d| d S )NrA   r   )r   rA   z"Not an color or gray tensor. Got: r   r   r   r   r   r   KORNIA_CHECK_IS_COLOR_OR_GRAY  s    r   )desc1desc2dmc                 C   sP   | d|  dkr(| d| dksLd|j d| j d|j }t|d S )Nr   r   zdistance matrix shape zG is not
                      consistent with descriptors shape: desc1 z
                      desc2 )sizer   r=   )r   r   r   messager   r   r   KORNIA_CHECK_DM_DESC"  s    (
r   )lafr   c                 C   s   t | g d dS )z\Auxiliary function, which verifies that input.

    Args:
        laf: [BxNx2x3] shape.
    )BN23N)r   )r   r   r   r   KORNIA_CHECK_LAF*  s    r   )NN)r&   )r&   )rv   )N)N)N)N)N)N)N)C__doc__r   r   r^   abcr   r   copyr   	itertoolsr   typingr   r   r   r	   r
   r   r   r   r   r    r   Ztorch.nnr   __all__rg   r   r   r(   re   r   ro   r   r   r4   r<   r?   rD   rF   rG   rr   ru   rf   rw   r   r~   Ztorch.testingr   r   Ztorch.testing._corer   ImportErrorr   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   r   <module>   s   ,



B		 
