o
    ˜šÍh-G  ã                   @  s&  d dl mZ ddlmZmZmZmZ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r@dd	lmZ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œdndd„Z ddœdod#d$„Z!dpd%d&„Z"d d'œdqd)d*„Z#drd+d,„Z$dpd-d.„Z%dpd/d0„Z&dsd1d2„Z'dd3d4œdtd8d9„Z(dud;d<„Z)dd=œdvd@dA„Z*dpdBdC„Z+dsdDdE„Z,dd=œdvdFdG„Z-dHdIœdwdLdM„Z.dxdNdO„Z/dPdQ„ Z0dsdRdS„Z1dTdUœdydWdX„Z2dzdZd[„Z3dd\œd{d_d`„Z4d ddaœd|ddde„Z5ddœdodfdg„Z6ddddhœd}dkdl„Z7g dm¢Z8dS )~é    )Úannotationsé   )Ú_floating_dtypesÚ_numeric_dtypesÚfloat32Úfloat64Ú	complex64Ú
complex128)Úreshape)ÚArrayé   )Únormalize_axis_tuple)ÚTYPE_CHECKING)ÚLiteralÚOptionalÚSequenceÚTupleÚUnionÚDtype)Ú
NamedTupleNc                   @  ó   e Zd ZU ded< ded< dS )Ú
EighResultr   ZeigenvaluesZeigenvectorsN©Ú__name__Ú
__module__Ú__qualname__Ú__annotations__© r   r   úW/home/www/facesmatcher.com/frenv/lib/python3.10/site-packages/numpy/array_api/linalg.pyr      ó   
 r   c                   @  r   )ÚQRResultr   ÚQÚRNr   r   r   r   r   r       r   r    c                   @  r   )ÚSlogdetResultr   ÚsignZ	logabsdetNr   r   r   r   r   r#   !   r   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)ÚupperÚxr   r(   ÚboolÚreturnc               C  s:   | j tvr	tdƒ‚tj | j¡}|rt |¡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   r0   ,   s   

r0   éÿÿÿÿ©ÚaxisÚx1Úx2r6   Úintc               C  sr   | j tvs
|j tvrtdƒ‚| j|jkrtdƒ‚| jdkr!tdƒ‚| j| dkr,t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 3r5   )r,   r   r-   ÚshapeÚ
ValueErrorÚndimr   r2   r.   Úcrossr1   )r7   r8   r6   r   r   r   r>   <   s   
r>   c                C  ó&   | j tv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   r2   r.   r/   Údetr1   ©r)   r   r   r   r@   N   s   
r@   )ÚoffsetrB   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.
    éþÿÿÿr4   )rB   Úaxis1Úaxis2)r   r2   r.   Údiagonalr1   )r)   rB   r   r   r   rF   [   s   rF   c                C  ó,   | j tv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   r2   r.   r/   Úeighr1   rA   r   r   r   rI   f   ó   
rI   c                C  r?   )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   r2   r.   r/   Úeigvalshr1   rA   r   r   r   rK   v   ó   
rK   c                C  r?   )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   r2   r.   r/   Úinvr1   rA   r   r   r   rM   ƒ   rL   rM   c                C  s2   | j tvs
|j tv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   r2   r.   Úmatmulr1   ©r7   r8   r   r   r   rN   ’   s   rN   Zfro)ÚkeepdimsÚordrP   rQ   ú2Optional[Union[int, float, Literal['fro', 'nuc']]]c               C  s.   | j tv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)rC   r4   ©r6   rP   rQ   )	r,   r   r-   r   r2   r.   r/   Únormr1   )r)   rP   rQ   r   r   r   Úmatrix_norm¦   s   
rV   Únc                C  s(   | j tv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   r2   r.   r/   Úmatrix_powerr1   )r)   rW   r   r   r   rX   ´   s   
rX   )ÚrtolrY   úOptional[Union[float, Array]]c               C  sª   | j dk rtj d¡‚tjj| jdd�}|du r0|jddd�t| jd	d… ƒ t |j	¡j
 }nt|tƒr8|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_uvNr4   T)r6   rP   rC   .r5   )r=   r.   r/   ZLinAlgErrorÚsvdr1   Úmaxr;   Úfinfor,   ÚepsÚ
isinstancer   ÚasarrayZnewaxisr2   Zcount_nonzero)r)   rY   r'   Ztolr   r   r   Úmatrix_rankÃ   s   
0
"rb   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_transposer4   rC   )r=   r<   r   r2   r.   Zswapaxesr1   rA   r   r   r   Úmatrix_transposeÛ   s   
rc   c                C  sN   | j tvs
|j tvrtdƒ‚| jdks|jdkrt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-   r=   r<   r   r2   r.   Úouterr1   rO   r   r   r   rd   á   s
   rd   c               C  sR   | j tvr	tdƒ‚|du 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 pinvNrC   )Zrcond)r,   r   r-   r]   r;   r.   r^   r_   r   r2   r/   Úpinvr1   )r)   rY   r   r   r   re   ó   s
   
 re   Zreduced©Úmoderg   úLiteral['reduced', 'complete']c               C  ó0   | j tv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 qrrf   )r,   r   r-   r    rH   r   r2   r.   r/   Úqrr1   )r)   rg   r   r   r   rj     ó   
rj   c                C  rG   )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#   rH   r   r2   r.   r/   Úslogdetr1   rA   r   r   r   rl     rJ   rl   c                 C  s¸   ddl m}m}m}m}m}m}m} ddlm	}	 || ƒ\} }
|| ƒ || ƒ ||ƒ\}}|| |ƒ\}}|j
dkr<|	j}n|	j}||ƒrEd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.linalgrm   rn   ro   rp   rq   rr   rs   r/   rt   r=   Zsolve1ÚsolveZastype)ÚaÚbrm   rn   ro   rp   rq   rr   rs   rt   Ú_ÚwrapÚtZresult_tZgufuncru   rv   Úrr   r   r   Ú_solve+  s   $
r   c                C  s0   | j tvs
|j tv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   r2   r   r1   rO   r   r   r   rx   F  s   rx   T©Úfull_matricesr�   c               C  ri   )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 svdr€   )r,   r   r-   r%   rH   r   r2   r.   r/   r\   r1   )r)   r�   r   r   r   r\   S  rk   r\   úUnion[Array, Tuple[Array, ...]]c                C  s*   | j tvr	tdƒ‚t tjj| jdd�¡S )Nz1Only floating-point dtypes are allowed in svdvalsFr[   )	r,   r   r-   r   r2   r.   r/   r\   r1   rA   r   r   r   Úsvdvalsd  s   
rƒ   ©Úaxesr…   ú/Union[int, Tuple[Sequence[int], Sequence[int]]]c               C  s6   | j tvs
|j tvrtdƒ‚t tj| j|j|d�¡S )Nz,Only numeric dtypes are allowed in tensordotr„   )r,   r   r-   r   r2   r.   Ú	tensordotr1   )r7   r8   r…   r   r   r   r‡   l  s   r‡   )rB   r,   r,   úOptional[Dtype]c               C  sZ   | j tvr	tdƒ‚|du r| j tkrt}n| j tkrt}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 traceNrC   r4   )rB   rD   rE   r,   )r,   r   r-   r   r   r   r	   r   r2   r.   ra   Útracer1   )r)   rB   r,   r   r   r   r‰   u  s   


"r‰   c         	      C  sÊ   | j tvs
|j tvrtdƒ‚t| j|jƒ}d|| j  t| jƒ }d||j  t|jƒ }|| || kr9t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 axisr4   .).N).r   r   )r,   r   r-   r]   r=   Útupler;   r<   r.   Zbroadcast_arraysr1   Zmoveaxisr   r2   )	r7   r8   r6   r=   Zx1_shapeZx2_shapeZx1_Zx2_Úresr   r   r   ÚvecdotŠ  s   rŒ   rT   ú%Optional[Union[int, Tuple[int, ...]]]úOptional[Union[int, float]]c         
        s  | j tvr	tdƒ‚| j‰ |du rˆ  ¡ ‰ d}nBt|tƒrWt|| jƒ‰t‡fdd„t	ˆ jƒD ƒƒ}|| }t
 ˆ |¡ t
j‡ fdd„|D ƒtd�g‡ fdd„|D ƒ¢R ¡‰ d}n|}t t
jjˆ ||d	�¡}|r‹t| jƒ}t|du rvt	| jƒn|| jƒ}|D ]}	d
||	< q}t|t|ƒƒ}|S )rS   z.Only floating-point dtypes are allowed in normNr   c                 3  s   � | ]	}|ˆ v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   )r;   r�   )ry   r   r   Ú
<listcomp>»  s    zvector_norm.<locals>.<listcomp>)r,   )r6   rQ   r   )r,   r   r-   r1   Zravelr`   rŠ   r   r=   Úranger.   Z	transposer
   Úprodr9   r   r2   r/   rU   Úlistr;   )
r)   r6   rP   rQ   Z_axisÚrestZnewshaper‹   r;   r‘   r   )ry   r’   r   Úvector_norm   s.   

0ÿ

r™   )r0   r>   r@   rF   rI   rK   rM   rN   rV   rX   rb   rc   rd   re   rj   rl   rx   r\   rƒ   r‡   r‰   rŒ   r™   )r)   r   r(   r*   r+   r   )r7   r   r8   r   r6   r9   r+   r   )r)   r   r+   r   )r)   r   rB   r9   r+   r   )r)   r   r+   r   )r7   r   r8   r   r+   r   )r)   r   rP   r*   rQ   rR   r+   r   )r)   r   rW   r9   r+   r   )r)   r   rY   rZ   r+   r   )r)   r   rg   rh   r+   r    )r)   r   r+   r#   )r)   r   r�   r*   r+   r%   )r)   r   r+   r‚   )r7   r   r8   r   r…   r†   r+   r   )r)   r   rB   r9   r,   rˆ   r+   r   )
r)   r   r6   r�   rP   r*   rQ   rŽ   r+   r   )9Ú
__future__r   Z_dtypesr   r   r   r   r   r	   Z_manipulation_functionsr
   Z_array_objectr   Zcore.numericr   Útypingr   Z_typingr   r   r   r   r   r   r   Znumpy.linalgÚnumpyr.   r   r    r#   r%   r0   r>   r@   rF   rI   rK   rM   rN   rV   rX   rb   rc   rd   re   rj   rl   r   rx   r\   rƒ   r‡   r‰   rŒ   r™   Ú__all__r   r   r   r   Ú<module>   sP      










	-