U
    ÉZjDB  ã                   @  sØ  d dl mZ ddlmZmZ ddlmZ d dlmZ erTd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,d-„Zdd.d/œddd0dd1œd2d3„Zdddd4œd5d6„Z dd7œdd8dd9œd:d;„Z!dddœd<d=„Z"dddd+œd>d?„Z#dd7œdd8dd9œd@dA„Z$dBdCœddDddEœdFdG„Z%dddœdHdI„Z&dJdK„ Z'dddd+œdLdM„Z(dNdOœddddPœdQdR„Z)ddSdœdTdU„Z*dVdWœdddXddYœdZd[„Z+d d!œdddd"œd\d]„Z,ddœdddddœd^d_„Z-dddVd`œddaddbddcœddde„Z.ddd d$d&d(d*d-d3d6d;d=d?dAdGdIdMdRdUd[d]d_degZ/dS )fé    )Úannotationsé   )Ú_floating_dtypesÚ_numeric_dtypes)ÚArray)Ú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   úY/var/www/html/TRUCKING_PROJECT/venv/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'   )r.   r/   r,   r   r   r   r4   2   s    
r4   )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   r5   D   s    
r5   )Úoffset)r    r7   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*   ©r7   Zaxis1Zaxis2)r   r(   r$   Údiagonalr'   ©r    r7   r   r   r   r:   Q   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'   r6   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'   r6   r   r   r   r>   l   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.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'   r6   r   r   r   r?   y   s    
r?   )r.   r/   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'   ©r.   r/   r   r   r   r@   ˆ   s    r@   Zfro)ÚkeepdimsÚordz4Optional[Union[int, float, Literal[('fro', 'nuc')]]])r    rB   rC   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)r8   r*   ©r,   rB   rC   )	r"   r   r#   r   r(   r$   r%   Únormr'   )r    rB   rC   r   r   r   Úmatrix_normœ   s    
rG   )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    rH   r   r   r   rI   ª   s    
rI   )ÚrtolzOptional[Union[float, Array]])r    rJ   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.
    é   zA1-dimensional array given. Array must be at least two-dimensionalF©Z
compute_uvNr*   T)r,   rB   r8   .r+   )r3   r$   r%   ZLinAlgErrorÚsvdr'   Úmaxr1   Úfinfor"   ÚepsÚ
isinstancer   ÚasarrayZnewaxisr(   Zcount_nonzero)r    rJ   r   Ztolr   r   r   Úmatrix_rank¹   s    
0
"rS   c                C  s(   | j dk rtdƒ‚t t | jdd¡¡S )NrK   z5x must be at least 2-dimensional for matrix_transposer*   r8   )r3   r2   r   r(   r$   Zswapaxesr'   r6   r   r   r   Úmatrix_transposeÑ   s    
rT   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#   r3   r2   r   r(   r$   Úouterr'   rA   r   r   r   rU   ×   s
    rU   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 pinvNr8   )Zrcond)r"   r   r#   rN   r1   r$   rO   rP   r   r(   r%   Úpinvr'   )r    rJ   r   r   r   rV   é   s
    
 rV   Zreduced©Úmodez Literal[('reduced', 'complete')])r    rX   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 qrrW   )r"   r   r#   r   r<   r   r(   r$   r%   Úqrr'   )r    rX   r   r   r   rY   ú   s    
rY   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'   r6   r   r   r   rZ   	  s    
rZ   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 )
NrK   )Ú
_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^   r_   r`   ra   r%   rb   r3   Zsolve1ÚsolveZastype)ÚaÚbr[   r\   r]   r^   r_   r`   ra   rb   Ú_ÚwrapÚtZresult_tZgufuncrc   rd   Úrr   r   r   Ú_solve!  s    $
rm   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(   rm   r'   rA   r   r   r   rf   <  s    rf   T©Úfull_matrices)r    ro   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 svdrn   )r"   r   r#   r   r<   r   r(   r$   r%   rM   r'   )r    ro   r   r   r   rM   I  s    
rM   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 svdvalsFrL   )	r"   r   r#   r   r(   r$   r%   rM   r'   r6   r   r   r   ÚsvdvalsZ  s    
rp   rK   ©Úaxesz/Union[int, Tuple[Sequence[int], Sequence[int]]])r.   r/   rr   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 tensordotrq   )r"   r   r#   r   r(   r$   Ú	tensordotr'   )r.   r/   rr   r   r   r   rs   b  s    rs   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 tracer8   r*   r9   )	r"   r   r#   r   r(   r$   rR   Útracer'   r;   r   r   r   rt   k  s    
rt   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#   rN   r3   Útupler1   r2   r$   Zbroadcast_arraysr'   Zmoveaxisr   r(   )	r.   r/   r,   r3   Zx1_shapeZx2_shapeZx1_Zx2_Úresr   r   r   Úvecdotx  s    rw   rE   z%Optional[Union[int, Tuple[int, ...]]]zOptional[Union[int, float]])r    r,   rB   rC   r!   c                 s´   | j tkrtdƒ‚| j‰ ˆdkr.ˆ  ¡ ‰ d‰nltˆtƒršt‡fdd„tˆ jƒD ƒƒ}ˆ| }t	 
ˆ |¡ t	 ‡ fdd„ˆD ƒ¡f‡ fdd„|D ƒ˜¡‰ d‰t t	jjˆ ˆ||d�¡S )	rD   z.Only floating-point dtypes are allowed in normNr   c                 3  s   | ]}|ˆ kr|V  qd S )Nr   ©Ú.0Úir+   r   r   Ú	<genexpr>   s      zvector_norm.<locals>.<genexpr>c                   s   g | ]}ˆ j | ‘qS r   )r1   rx   )rg   r   r   Ú
<listcomp>¢  s     zvector_norm.<locals>.<listcomp>rE   )r"   r   r#   r'   ÚflattenrQ   ru   Úranger3   r$   Z	transposeZreshapeÚprodr   r(   r%   rF   )r    r,   rB   rC   ÚrestZnewshaper   )rg   r,   r   Úvector_normŽ  s    

:r�   )0Ú
__future__r   Z_dtypesr   r   Z_array_objectr   Útypingr   Z_typingr   r	   r
   r   r   r   Znumpy.linalgÚnumpyr$   r   r   r   r   r&   r4   r5   r:   r=   r>   r?   r@   rG   rI   rS   rT   rU   rV   rY   rZ   rm   rf   rM   rp   rs   rt   rw   r�   Ú__all__r   r   r   r   Ú<module>   sJ   	 