U
    ¸mœdÄE  ã                   @  sð  d dl mZ ddlmZmZ ddlmZ ddlmZ ddl	m
Z
 d dlmZ erldd	lmZmZmZmZmZ d d
lmZ d dlZd dlZG dd„ deƒZG dd„ deƒZG dd„ deƒZG dd„ deƒZddœddddœdd„Zddœdddddœdd „Zddd!œd"d#„Zd d$œdddd%œd&d'„Zddd!œd(d)„Zddd!œd*d+„Z ddd!œd,d-„Z!dddd.œd/d0„Z"dd1d2œddd3dd4œd5d6„Z#dddd7œd8d9„Z$dd:œdd;dd<œd=d>„Z%ddd!œd?d@„Z&dddd.œdAdB„Z'dd:œdd;dd<œdCdD„Z(dEdFœddGddHœdIdJ„Z)ddd!œdKdL„Z*dMdN„ Z+dddd.œdOdP„Z,dQdRœddddSœdTdU„Z-ddVd!œdWdX„Z.ddYœdddZdd[œd\d]„Z/d d$œdddd%œd^d_„Z0ddœdddddœd`da„Z1ddddbœddcdddddeœdfdg„Z2dd d#d'd)d+d-d0d6d9d>d@dBdDdJdLdPdUdXd]d_dadggZ3dS )hé    )Úannotationsé   )Ú_floating_dtypesÚ_numeric_dtypes)Úreshape)ÚArrayé   )Únormalize_axis_tuple)ÚTYPE_CHECKING)ÚLiteralÚOptionalÚSequenceÚTupleÚUnion)Ú
NamedTupleNc                   @  s   e Zd ZU ded< ded< dS )Ú
EighResultr   ZeigenvaluesZeigenvectorsN©Ú__name__Ú
__module__Ú__qualname__Ú__annotations__© r   r   úO/home/sam/Atlas/atlas_env/lib/python3.8/site-packages/numpy/array_api/linalg.pyr      s   
r   c                   @  s   e Zd ZU ded< ded< dS )ÚQRResultr   ÚQÚRNr   r   r   r   r   r      s   
r   c                   @  s   e Zd ZU ded< ded< dS )ÚSlogdetResultr   ÚsignZ	logabsdetNr   r   r   r   r   r      s   
r   c                   @  s&   e Zd ZU ded< ded< ded< dS )Ú	SVDResultr   ÚUÚSZVhNr   r   r   r   r   r      s   
r   F)Úupperr   Úbool)Úxr!   Úreturnc               C  s:   | j tkrtdƒ‚tj | j¡}|r0t |¡j	S t |¡S )zŽ
    Array API compatible wrapper for :py:func:`np.linalg.cholesky <numpy.linalg.cholesky>`.

    See its docstring for more information.
    z2Only floating-point dtypes are allowed in cholesky)
Údtyper   Ú	TypeErrorÚnpÚlinalgÚcholeskyÚ_arrayr   Ú_newZmT)r#   r!   ÚLr   r   r   r)   %   s    
r)   éÿÿÿÿ©ÚaxisÚint)Úx1Úx2r/   r$   c               C  sr   | j tks|j tkrtdƒ‚| j|jkr0tdƒ‚| jdkrBtdƒ‚| j| dkrXtdƒ‚t tj	| j
|j
|d�¡S )zz
    Array API compatible wrapper for :py:func:`np.cross <numpy.cross>`.

    See its docstring for more information.
    z(Only numeric dtypes are allowed in crossz"x1 and x2 must have the same shaper   z/cross() requires arrays of dimension at least 1é   zcross() dimension must equal 3r.   )r%   r   r&   ÚshapeÚ
ValueErrorÚndimr   r+   r'   Úcrossr*   )r1   r2   r/   r   r   r   r7   5   s    
r7   )r#   r$   c                C  s&   | j tkrtdƒ‚t tj | j¡¡S )z„
    Array API compatible wrapper for :py:func:`np.linalg.det <numpy.linalg.det>`.

    See its docstring for more information.
    z-Only floating-point dtypes are allowed in det)	r%   r   r&   r   r+   r'   r(   Údetr*   ©r#   r   r   r   r8   G   s    
r8   )Úoffset)r#   r:   r$   c               C  s   t  tj| j|ddd�¡S )z€
    Array API compatible wrapper for :py:func:`np.diagonal <numpy.diagonal>`.

    See its docstring for more information.
    éþÿÿÿr-   ©r:   Zaxis1Zaxis2)r   r+   r'   Údiagonalr*   ©r#   r:   r   r   r   r=   T   s    r=   c                C  s,   | j tkrtdƒ‚tttjtj 	| j
¡ƒŽ S )z†
    Array API compatible wrapper for :py:func:`np.linalg.eigh <numpy.linalg.eigh>`.

    See its docstring for more information.
    z.Only floating-point dtypes are allowed in eigh)r%   r   r&   r   Úmapr   r+   r'   r(   Úeighr*   r9   r   r   r   r@   _   s    
r@   c                C  s&   | j tkrtdƒ‚t tj | j¡¡S )zŽ
    Array API compatible wrapper for :py:func:`np.linalg.eigvalsh <numpy.linalg.eigvalsh>`.

    See its docstring for more information.
    z2Only floating-point dtypes are allowed in eigvalsh)	r%   r   r&   r   r+   r'   r(   Úeigvalshr*   r9   r   r   r   rA   o   s    
rA   c                C  s&   | j tkrtdƒ‚t tj | j¡¡S )z„
    Array API compatible wrapper for :py:func:`np.linalg.inv <numpy.linalg.inv>`.

    See its docstring for more information.
    z-Only floating-point dtypes are allowed in inv)	r%   r   r&   r   r+   r'   r(   Úinvr*   r9   r   r   r   rB   |   s    
rB   )r1   r2   r$   c                C  s2   | j tks|j tkrtdƒ‚t t | j|j¡¡S )z|
    Array API compatible wrapper for :py:func:`np.matmul <numpy.matmul>`.

    See its docstring for more information.
    z)Only numeric dtypes are allowed in matmul)r%   r   r&   r   r+   r'   Úmatmulr*   ©r1   r2   r   r   r   rC   ‹   s    rC   Zfro)ÚkeepdimsÚordz4Optional[Union[int, float, Literal[('fro', 'nuc')]]])r#   rE   rF   r$   c               C  s.   | j tkrtdƒ‚t tjj| jd||d�¡S )ú†
    Array API compatible wrapper for :py:func:`np.linalg.norm <numpy.linalg.norm>`.

    See its docstring for more information.
    z5Only floating-point dtypes are allowed in matrix_norm)r;   r-   ©r/   rE   rF   )	r%   r   r&   r   r+   r'   r(   Únormr*   )r#   rE   rF   r   r   r   Úmatrix_normŸ   s    
rJ   )r#   Únr$   c                C  s(   | j tkrtdƒ‚t tj | j|¡¡S )zˆ
    Array API compatible wrapper for :py:func:`np.matrix_power <numpy.matrix_power>`.

    See its docstring for more information.
    zMOnly floating-point dtypes are allowed for the first argument of matrix_power)	r%   r   r&   r   r+   r'   r(   Úmatrix_powerr*   )r#   rK   r   r   r   rL   ­   s    
rL   )ÚrtolzOptional[Union[float, Array]])r#   rM   r$   c               C  sª   | j dk rtj d¡‚tjj| jdd�}|dkr`|jddd�t| jd	d… ƒ t |j	¡j
 }n2t|tƒrp|j}|jddd�t |¡d
tjf  }t tj||kdd�¡S )z†
    Array API compatible wrapper for :py:func:`np.matrix_rank <numpy.matrix_rank>`.

    See its docstring for more information.
    r   zA1-dimensional array given. Array must be at least two-dimensionalF©Z
compute_uvNr-   T)r/   rE   r;   .r.   )r6   r'   r(   ZLinAlgErrorÚsvdr*   Úmaxr4   Úfinfor%   ÚepsÚ
isinstancer   ÚasarrayZnewaxisr+   Zcount_nonzero)r#   rM   r    Ztolr   r   r   Úmatrix_rank¼   s    
0
"rU   c                C  s(   | j dk rtdƒ‚t t | jdd¡¡S )Nr   z5x must be at least 2-dimensional for matrix_transposer-   r;   )r6   r5   r   r+   r'   Zswapaxesr*   r9   r   r   r   Úmatrix_transposeÔ   s    
rV   c                C  sN   | j tks|j tkrtdƒ‚| jdks0|jdkr8tdƒ‚t t | j	|j	¡¡S )zz
    Array API compatible wrapper for :py:func:`np.outer <numpy.outer>`.

    See its docstring for more information.
    z(Only numeric dtypes are allowed in outerr   z/The input arrays to outer must be 1-dimensional)
r%   r   r&   r6   r5   r   r+   r'   Úouterr*   rD   r   r   r   rW   Ú   s
    rW   c               C  sR   | j tkrtdƒ‚|dkr:t| jdd… ƒt | j ¡j }t 	tj
j| j|d�¡S )z†
    Array API compatible wrapper for :py:func:`np.linalg.pinv <numpy.linalg.pinv>`.

    See its docstring for more information.
    z.Only floating-point dtypes are allowed in pinvNr;   )Zrcond)r%   r   r&   rP   r4   r'   rQ   rR   r   r+   r(   Úpinvr*   )r#   rM   r   r   r   rX   ì   s
    
 rX   Zreduced©Úmodez Literal[('reduced', 'complete')])r#   rZ   r$   c               C  s0   | j tkrtdƒ‚tttjtjj	| j
|d�ƒŽ S )z‚
    Array API compatible wrapper for :py:func:`np.linalg.qr <numpy.linalg.qr>`.

    See its docstring for more information.
    z,Only floating-point dtypes are allowed in qrrY   )r%   r   r&   r   r?   r   r+   r'   r(   Úqrr*   )r#   rZ   r   r   r   r[   ý   s    
r[   c                C  s,   | j tkrtdƒ‚tttjtj 	| j
¡ƒŽ S )zŒ
    Array API compatible wrapper for :py:func:`np.linalg.slogdet <numpy.linalg.slogdet>`.

    See its docstring for more information.
    z1Only floating-point dtypes are allowed in slogdet)r%   r   r&   r   r?   r   r+   r'   r(   Úslogdetr*   r9   r   r   r   r\     s    
r\   c                 C  s¸   ddl m}m}m}m}m}m}m} ddlm	}	 || ƒ\} }
|| ƒ || ƒ ||ƒ\}}|| |ƒ\}}|j
dkrx|	j}n|	j}||ƒrŠdnd}||ƒ}|| |||d�}||j|dd	�ƒS )
Nr   )Ú
_makearrayÚ_assert_stacked_2dÚ_assert_stacked_squareÚ_commonTypeÚisComplexTypeÚget_linalg_error_extobjÚ_raise_linalgerror_singular)Ú_umath_linalgr   zDD->Dzdd->d)Ú	signatureÚextobjF)Úcopy)Zlinalg.linalgr]   r^   r_   r`   ra   rb   rc   r(   rd   r6   Zsolve1ÚsolveZastype)ÚaÚbr]   r^   r_   r`   ra   rb   rc   rd   Ú_ÚwrapÚtZresult_tZgufuncre   rf   Úrr   r   r   Ú_solve$  s    $
ro   c                C  s0   | j tks|j tkrtdƒ‚t t| j|jƒ¡S )zˆ
    Array API compatible wrapper for :py:func:`np.linalg.solve <numpy.linalg.solve>`.

    See its docstring for more information.
    z/Only floating-point dtypes are allowed in solve)r%   r   r&   r   r+   ro   r*   rD   r   r   r   rh   ?  s    rh   T©Úfull_matrices)r#   rq   r$   c               C  s0   | j tkrtdƒ‚tttjtjj	| j
|d�ƒŽ S )z„
    Array API compatible wrapper for :py:func:`np.linalg.svd <numpy.linalg.svd>`.

    See its docstring for more information.
    z-Only floating-point dtypes are allowed in svdrp   )r%   r   r&   r   r?   r   r+   r'   r(   rO   r*   )r#   rq   r   r   r   rO   L  s    
rO   zUnion[Array, Tuple[Array, ...]]c                C  s*   | j tkrtdƒ‚t tjj| jdd�¡S )Nz1Only floating-point dtypes are allowed in svdvalsFrN   )	r%   r   r&   r   r+   r'   r(   rO   r*   r9   r   r   r   Úsvdvals]  s    
rr   ©Úaxesz/Union[int, Tuple[Sequence[int], Sequence[int]]])r1   r2   rt   r$   c               C  s6   | j tks|j tkrtdƒ‚t tj| j|j|d�¡S )Nz,Only numeric dtypes are allowed in tensordotrs   )r%   r   r&   r   r+   r'   Ú	tensordotr*   )r1   r2   rt   r   r   r   ru   e  s    ru   c            
   C  s2   | j tkrtdƒ‚t t tj| j|ddd�¡¡S )zz
    Array API compatible wrapper for :py:func:`np.trace <numpy.trace>`.

    See its docstring for more information.
    z(Only numeric dtypes are allowed in tracer;   r-   r<   )	r%   r   r&   r   r+   r'   rT   Útracer*   r>   r   r   r   rv   n  s    
rv   c         	      C  sÊ   | j tks|j tkrtdƒ‚t| j|jƒ}d|| j  t| jƒ }d||j  t|jƒ }|| || krrtdƒ‚t 	| j
|j
¡\}}t ||d¡}t ||d¡}|dd d d …f |d  }t |d ¡S )Nz)Only numeric dtypes are allowed in vecdot)r   z6x1 and x2 must have the same size along the given axisr-   .).N).r   r   )r%   r   r&   rP   r6   Útupler4   r5   r'   Zbroadcast_arraysr*   Zmoveaxisr   r+   )	r1   r2   r/   r6   Zx1_shapeZx2_shapeZx1_Zx2_Úresr   r   r   Úvecdot{  s    ry   rH   z%Optional[Union[int, Tuple[int, ...]]]zOptional[Union[int, float]])r#   r/   rE   rF   r$   c         
        s  | j tkrtdƒ‚| j‰ |dkr.ˆ  ¡ ‰ d}n‚t|tƒr¬t|| jƒ‰t‡fdd„t	ˆ jƒD ƒƒ}|| }t
 ˆ |¡ t
j‡ fdd„|D ƒtd�f‡ fdd„|D ƒ˜¡‰ d}n|}t t
jjˆ ||d	�¡}|�rt| jƒ}t|dkrìt	| jƒn|| jƒ}|D ]}	d
||	< qút|t|ƒƒ}|S )rG   z.Only floating-point dtypes are allowed in normNr   c                 3  s   | ]}|ˆ kr|V  qd S )Nr   ©Ú.0Úi)Únormalized_axisr   r   Ú	<genexpr>©  s      zvector_norm.<locals>.<genexpr>c                   s   g | ]}ˆ j | ‘qS r   )r4   rz   )ri   r   r   Ú
<listcomp>¬  s     zvector_norm.<locals>.<listcomp>)r%   )r/   rF   r   )r%   r   r&   r*   ZravelrS   rw   r	   r6   Úranger'   Z	transposer   Úprodr0   r   r+   r(   rI   Úlistr4   )
r#   r/   rE   rF   Z_axisÚrestZnewshaperx   r4   r|   r   )ri   r}   r   Úvector_norm‘  s.    

.ÿ

r„   )4Ú
__future__r   Z_dtypesr   r   Z_manipulation_functionsr   Z_array_objectr   Zcore.numericr	   Útypingr
   Z_typingr   r   r   r   r   r   Znumpy.linalgÚnumpyr'   r   r   r   r   r)   r7   r8   r=   r@   rA   rB   rC   rJ   rL   rU   rV   rW   rX   r[   r\   ro   rh   rO   rr   ru   rv   ry   r„   Ú__all__r   r   r   r   Ú<module>   sN   	 -