Ë
    f^(h�*  ã                   ó‚   — d dl Z d dlZd dlmZ d dlmZ d dlmZ d dlmZm	Z	 d dl
mZ dgZd„ Zd	„ Zd
„ Z G d„ de«      Zy)é    N)ÚTensor)Úconstraints)ÚDistribution)Ú_standard_normalÚlazy_property)Ú_sizeÚMultivariateNormalc                 ój   — t        j                  | |j                  d«      «      j                  d«      S )a»  
    Performs a batched matrix-vector product, with compatible but different batch shapes.

    This function takes as input `bmat`, containing :math:`n \times n` matrices, and
    `bvec`, containing length :math:`n` vectors.

    Both `bmat` and `bvec` may have any number of leading dimensions, which correspond
    to a batch shape. They are not necessarily assumed to have the same batch shape,
    just ones which can be broadcasted.
    éÿÿÿÿ)ÚtorchÚmatmulÚ	unsqueezeÚsqueeze)ÚbmatÚbvecs     úe/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributions/multivariate_normal.pyÚ	_batch_mvr      s)   € ô �<‰<˜˜dŸn™n¨RÓ0Ó1×9Ñ9¸"Ó=Ð=ó    c                 ó$  — |j                  d«      }|j                  dd }t        |«      }| j                  «       dz
  }||z
  }||z   }|d|z  z   }|j                  d| }	t	        | j                  dd |j                  |d «      D ]  \  }
}|	||
z  |
fz  }	Œ |	|fz  }	|j                  |	«      }t        t        |«      «      t        t        ||d«      «      z   t        t        |dz   |d«      «      z   |gz   }|j                  |«      }| j                  d||«      }|j                  d|j                  d«      |«      }|j                  ddd«      }t        j                  j                  ||d¬«      j                  d«      j                  d«      }|j                  «       }|j                  |j                  dd «      }t        t        |«      «      }t        |«      D ]  }|||z   ||z   gz  }Œ |j                  |«      }|j                  |«      S )	aK  
    Computes the squared Mahalanobis distance :math:`\mathbf{x}^\top\mathbf{M}^{-1}\mathbf{x}`
    for a factored :math:`\mathbf{M} = \mathbf{L}\mathbf{L}^\top`.

    Accepts batches for both bL and bx. They are not necessarily assumed to have the same batch
    shape, but `bL` one should be able to broadcasted to `bx` one.
    r   Né   éþÿÿÿé   r   F©Úupper)ÚsizeÚshapeÚlenÚdimÚzipÚreshapeÚlistÚrangeÚpermuter   ÚlinalgÚsolve_triangularÚpowÚsumÚt)ÚbLÚbxÚnÚbx_batch_shapeÚbx_batch_dimsÚbL_batch_dimsÚouter_batch_dimsÚold_batch_dimsÚnew_batch_dimsÚbx_new_shapeÚsLÚsxÚpermute_dimsÚflat_LÚflat_xÚflat_x_swapÚM_swapÚMÚ
permuted_MÚpermute_inv_dimsÚiÚ
reshaped_Ms                         r   Ú_batch_mahalanobisr?      s0  € ð 	�‰�‹€AØ—X‘X˜c˜r�]€Nô ˜Ó'€MØ—F‘F“H˜q‘L€MØ$ }Ñ4ÐØ%¨Ñ5€NØ%¨¨MÑ(9Ñ9€Nà—8‘8Ð-Ð-Ð.€LÜ�b—h‘h˜s �m R§X¡XÐ.>¸rÐ%BÓCò '‰ˆˆBØ˜˜r™ 2˜Ñ&‰ð'à�Q�DÑ€LØ	�‰�LÓ	!€Bô 	ŒUÐ#Ó$Ó%Ü
ŒuÐ% ~°qÓ9Ó
:ñ	;ä
ŒuÐ%¨Ñ)¨>¸1Ó=Ó
>ñ	?ð Ð
ñ	ð ð 
�‰�LÓ	!€Bà�Z‰Z˜˜A˜qÓ!€FØ�Z‰Z˜˜FŸK™K¨›N¨AÓ.€FØ—.‘.  A qÓ)€Kä�‰×%Ñ% f¨kÀÐ%ÓG×KÑKÈAÓN×RÑRÐSUÓVð ð 	�‰‹
€Að —‘˜2Ÿ8™8 C R˜=Ó)€JÜœEÐ"2Ó3Ó4ÐÜ�=Ó!ò GˆØÐ-°Ñ1°>ÀAÑ3EÐFÑFÑðGà×#Ñ#Ð$4Ó5€JØ×Ñ˜nÓ-Ð-r   c                 óx  — t         j                  j                  t        j                  | d«      «      }t        j                  t        j                  |d«      dd«      }t        j
                  | j                  d   | j                  | j                  ¬«      }t         j                  j                  ||d¬«      }|S )N)r   r   r   r   ©ÚdtypeÚdeviceFr   )
r   r$   ÚcholeskyÚflipÚ	transposeÚeyer   rB   rC   r%   )ÚPÚLfÚL_invÚIdÚLs        r   Ú_precision_to_scale_trilrM   O   s€   € ä	�‰×	Ñ	œuŸz™z¨!¨XÓ6Ó	7€BÜ�O‰OœEŸJ™J r¨8Ó4°b¸"Ó=€EÜ	�‰�1—7‘7˜2‘; a§g¡g°a·h±hÔ	?€BÜ�‰×%Ñ% e¨R°uÐ%Ó=€AØ€Hr   c                   ót  ‡ — e Zd ZdZej
                  ej                  ej                  ej                  dœZej
                  Z	dZ
	 	 	 	 dˆ fd„	Zdˆ fd„	Zedefd„«       Zedefd„«       Zedefd	„«       Zedefd
„«       Zedefd„«       Zedefd„«       Z ej.                  «       fdedefd„Zd„ Zd„ Zˆ xZS )r	   a™  
    Creates a multivariate normal (also called Gaussian) distribution
    parameterized by a mean vector and a covariance matrix.

    The multivariate normal distribution can be parameterized either
    in terms of a positive definite covariance matrix :math:`\mathbf{\Sigma}`
    or a positive definite precision matrix :math:`\mathbf{\Sigma}^{-1}`
    or a lower-triangular matrix :math:`\mathbf{L}` with positive-valued
    diagonal entries, such that
    :math:`\mathbf{\Sigma} = \mathbf{L}\mathbf{L}^\top`. This triangular matrix
    can be obtained via e.g. Cholesky decomposition of the covariance.

    Example:

        >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_LAPACK)
        >>> # xdoctest: +IGNORE_WANT("non-deterministic")
        >>> m = MultivariateNormal(torch.zeros(2), torch.eye(2))
        >>> m.sample()  # normally distributed with mean=`[0,0]` and covariance_matrix=`I`
        tensor([-0.2102, -0.5429])

    Args:
        loc (Tensor): mean of the distribution
        covariance_matrix (Tensor): positive-definite covariance matrix
        precision_matrix (Tensor): positive-definite precision matrix
        scale_tril (Tensor): lower-triangular factor of covariance, with positive-valued diagonal

    Note:
        Only one of :attr:`covariance_matrix` or :attr:`precision_matrix` or
        :attr:`scale_tril` can be specified.

        Using :attr:`scale_tril` will be more efficient: all computations internally
        are based on :attr:`scale_tril`. If :attr:`covariance_matrix` or
        :attr:`precision_matrix` is passed instead, it is only used to compute
        the corresponding lower triangular matrices using a Cholesky decomposition.
    )ÚlocÚcovariance_matrixÚprecision_matrixÚ
scale_trilTc                 óú  •— |j                  «       dk  rt        d«      ‚|d u|d uz   |d uz   dk7  rt        d«      ‚|�h|j                  «       dk  rt        d«      ‚t        j                  |j                  d d |j                  d d «      }|j                  |dz   «      | _        nÑ|�h|j                  «       dk  rt        d	«      ‚t        j                  |j                  d d |j                  d d «      }|j                  |dz   «      | _        ng|j                  «       dk  rt        d
«      ‚t        j                  |j                  d d |j                  d d «      }|j                  |dz   «      | _        |j                  |dz   «      | _	        | j                  j                  dd  }t        ‰| �-  |||¬«       |�|| _        y |�%t        j                  j                  |«      | _        y t        |«      | _        y )Nr   z%loc must be at least one-dimensional.zTExactly one of covariance_matrix or precision_matrix or scale_tril may be specified.r   zZscale_tril matrix must be at least two-dimensional, with optional leading batch dimensionsr   r   )r   r   zZcovariance_matrix must be at least two-dimensional, with optional leading batch dimensionszYprecision_matrix must be at least two-dimensional, with optional leading batch dimensions)r   ©Úvalidate_args)r   Ú
ValueErrorr   Úbroadcast_shapesr   ÚexpandrR   rP   rQ   rO   ÚsuperÚ__init__Ú_unbroadcasted_scale_trilr$   rD   rM   )	ÚselfrO   rP   rQ   rR   rU   Úbatch_shapeÚevent_shapeÚ	__class__s	           €r   rZ   zMultivariateNormal.__init__†   s  ø€ ð �7‰7‹9�qŠ=ÜÐDÓEÐEØ TÐ)¨jÀÐ.DÑEØ DÐ(ñ
àòô Øfóð ð Ð!Ø�~‰~Ó !Ò#Ü ð=óð ô  ×0Ñ0°×1AÑ1AÀ#À2Ð1FÈÏ	É	ÐRUÐSUÈÓWˆKØ(×/Ñ/°¸hÑ0FÓGˆD�OØÐ*Ø ×$Ñ$Ó&¨Ò*Ü ð=óð ô  ×0Ñ0Ø!×'Ñ'¨¨Ð,¨c¯i©i¸¸¨nóˆKð &7×%=Ñ%=¸kÈHÑ>TÓ%UˆDÕ"à×#Ñ#Ó%¨Ò)Ü ð=óð ô  ×0Ñ0Ø ×&Ñ& s¨Ð+¨S¯Y©Y°s¸¨^óˆKð %5×$;Ñ$;¸KÈ(Ñ<RÓ$SˆDÔ!Ø—:‘:˜k¨EÑ1Ó2ˆŒà—h‘h—n‘n R SÐ)ˆÜ‰Ñ˜ kÀÐÔOàÐ!Ø-7ˆDÕ*ØÐ*Ü-2¯\©\×-BÑ-BÐCTÓ-UˆDÕ*ä-EÐFVÓ-WˆDÕ*r   c                 óŒ  •— | j                  t        |«      }t        j                  |«      }|| j                  z   }|| j                  z   | j                  z   }| j
                  j                  |«      |_        | j                  |_        d| j                  v r | j                  j                  |«      |_	        d| j                  v r | j                  j                  |«      |_
        d| j                  v r | j                  j                  |«      |_        t        t        |�7  || j                  d¬«       | j                  |_        |S )NrP   rR   rQ   FrT   )Ú_get_checked_instancer	   r   ÚSizer^   rO   rX   r[   Ú__dict__rP   rR   rQ   rY   rZ   Ú_validate_args)r\   r]   Ú	_instanceÚnewÚ	loc_shapeÚ	cov_shaper_   s         €r   rX   zMultivariateNormal.expand¿   s  ø€ Ø×(Ñ(Ô);¸YÓGˆÜ—j‘j Ó-ˆØ $×"2Ñ"2Ñ2ˆ	Ø $×"2Ñ"2Ñ2°T×5EÑ5EÑEˆ	Ø—(‘(—/‘/ )Ó,ˆŒØ(,×(FÑ(FˆÔ%Ø $§-¡-Ñ/Ø$(×$:Ñ$:×$AÑ$AÀ)Ó$LˆCÔ!Ø˜4Ÿ=™=Ñ(Ø!Ÿ_™_×3Ñ3°IÓ>ˆCŒNØ §¡Ñ.Ø#'×#8Ñ#8×#?Ñ#?À	Ó#JˆCÔ ÜÔ  #Ñ/Ø˜×)Ñ)¸ð 	0ô 	
ð "×0Ñ0ˆÔØˆ
r   Úreturnc                 ó€   — | j                   j                  | j                  | j                  z   | j                  z   «      S ©N)r[   rX   Ú_batch_shapeÚ_event_shape©r\   s    r   rR   zMultivariateNormal.scale_trilÒ   s:   € à×-Ñ-×4Ñ4Ø×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r   c                 óÐ   — t        j                  | j                  | j                  j                  «      j	                  | j
                  | j                  z   | j                  z   «      S rk   )r   r   r[   ÚmTrX   rl   rm   rn   s    r   rP   z$MultivariateNormal.covariance_matrixØ   sQ   € ä�|‰|Ø×*Ñ*¨D×,JÑ,J×,MÑ,Mó
ç
‰&�×"Ñ" T×%6Ñ%6Ñ6¸×9JÑ9JÑJÓ
Kð	Lr   c                 ó¦   — t        j                  | j                  «      j                  | j                  | j
                  z   | j
                  z   «      S rk   )r   Úcholesky_inverser[   rX   rl   rm   rn   s    r   rQ   z#MultivariateNormal.precision_matrixÞ   sE   € ä×%Ñ% d×&DÑ&DÓE×LÑLØ×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r   c                 ó   — | j                   S rk   ©rO   rn   s    r   ÚmeanzMultivariateNormal.meanä   ó   € à�x‰xˆr   c                 ó   — | j                   S rk   rt   rn   s    r   ÚmodezMultivariateNormal.modeè   rv   r   c                 ó¢   — | j                   j                  d«      j                  d«      j                  | j                  | j
                  z   «      S )Nr   r   )r[   r&   r'   rX   rl   rm   rn   s    r   ÚvariancezMultivariateNormal.varianceì   sA   € ð ×*Ñ*×.Ñ.¨qÓ1ß‰S�‹Wß‰V�D×%Ñ%¨×(9Ñ(9Ñ9Ó:ð	
r   Úsample_shapec                 óÖ   — | j                  |«      }t        || j                  j                  | j                  j                  ¬«      }| j                  t        | j                  |«      z   S )NrA   )Ú_extended_shaper   rO   rB   rC   r   r[   )r\   r{   r   Úepss       r   ÚrsamplezMultivariateNormal.rsampleô   sL   € Ø×$Ñ$ \Ó2ˆÜ˜u¨D¯H©H¯N©NÀ4Ç8Á8Ç?Á?ÔSˆØ�x‰xœ) D×$BÑ$BÀCÓHÑHÐHr   c                 óx  — | j                   r| j                  |«       || j                  z
  }t        | j                  |«      }| j                  j                  dd¬«      j                  «       j                  d«      }d| j                  d   t        j                  dt        j                  z  «      z  |z   z  |z
  S )Nr   r   ©Údim1Údim2g      à¿r   r   )rd   Ú_validate_samplerO   r?   r[   ÚdiagonalÚlogr'   rm   ÚmathÚpi)r\   ÚvalueÚdiffr:   Úhalf_log_dets        r   Úlog_probzMultivariateNormal.log_probù   s¤   € Ø×ÒØ×!Ñ! %Ô(Ø�t—x‘xÑˆÜ˜t×=Ñ=¸tÓDˆà×*Ñ*×3Ñ3¸À"Ð3ÓE×IÑIÓK×OÑOÐPRÓSð 	ð �t×(Ñ(¨Ñ+¬d¯h©h°q¼4¿7¹7±{Ó.CÑCÀaÑGÑHÈ<ÑWÐWr   c                 ó^  — | j                   j                  dd¬«      j                  «       j                  d«      }d| j                  d   z  dt        j                  dt
        j                  z  «      z   z  |z   }t        | j                  «      dk(  r|S |j                  | j                  «      S )Nr   r   r�   g      à?r   g      ð?r   )
r[   r…   r†   r'   rm   r‡   rˆ   r   rl   rX   )r\   r‹   ÚHs      r   ÚentropyzMultivariateNormal.entropy  s™   € à×*Ñ*×3Ñ3¸À"Ð3ÓE×IÑIÓK×OÑOÐPRÓSð 	ð �$×#Ñ# AÑ&Ñ&¨#´·±¸¼T¿W¹W¹Ó0EÑ*EÑFÈÑUˆÜˆt× Ñ Ó! QÒ&ØˆHà—8‘8˜D×-Ñ-Ó.Ð.r   )NNNNrk   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úreal_vectorÚpositive_definiteÚlower_choleskyÚarg_constraintsÚsupportÚhas_rsamplerZ   rX   r   r   rR   rP   rQ   Úpropertyru   rx   rz   r   rb   r   r   rŒ   r�   Ú__classcell__)r_   s   @r   r	   r	   X   s5  ø„ ñ"ðJ ×&Ñ&Ø(×:Ñ:Ø'×9Ñ9Ø!×0Ñ0ñ	€Oð ×%Ñ%€GØ€Kð
 ØØØõ7Xõrð& ð
˜Fò 
ó ð
ð
 ðL 6ò Ló ðLð
 ð
 &ò 
ó ð
ð
 ð�fò ó ðð ð�fò ó ðð ð
˜&ò 
ó ð
ð -7¨E¯J©J«Lñ I Eð I¸Vó Iò
Xö/r   )r‡   r   r   Útorch.distributionsr   Ú torch.distributions.distributionr   Útorch.distributions.utilsr   r   Útorch.typesr   Ú__all__r   r?   rM   r	   © r   r   ú<module>r¢      sB   ðã ã Ý Ý +Ý 9ß EÝ ð  Ð
 €ò>ò/.òdôs/˜õ s/r   