ó
    Eñi7,  ã                   ó‚   • S SK r S SKrS SKJr  S SKJr  S SKJr  S SKJrJ	r	  S SK
Jr  S/rS rS	 rS
 r " S S\5      rg)é    N)ÚTensor)Úconstraints)ÚDistribution)Ú_standard_normalÚlazy_property)Ú_sizeÚMultivariateNormalc                 ój   • [         R                  " XR                  S5      5      R                  S5      $ )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     Úd/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/multivariate_normal.pyÚ	_batch_mvr      s'   € ô �<Š<˜Ÿn™n¨RÓ0Ó1×9Ñ9¸"Ó=Ð=ó    c                 ó  • UR                  S5      nUR                  SS n[        U5      nU R                  5       S-
  nXE-
  nXe-   nUSU-  -   nUR                  SU n	[	        U R                  SS UR                  US 5       H  u  p«X›U
-  U
4-  n	M     X’4-  n	UR                  U	5      n[        [        U5      5      [        [        XhS5      5      -   [        [        US-   US5      5      -   U/-   nUR                  U5      nU R                  SX"5      nUR                  SUR                  S5      U5      nUR                  SSS5      n[        R                  R                  XßSS9R                  S5      R                  S5      nUR                  5       nUR                  UR                  SS 5      n[        [        U5      5      n[        U5       H  nUUU-   UU-   /-  nM     UR                  U5      nUR                  U5      $ )	a7  
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?      s  € ð 	�‰�‹€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‰ˆØ˜r™ 2˜Ñ&Šñ Dà�DÑ€LØ	�‰�LÓ	!€Bô 	ŒUÐ#Ó$Ó%Ü
ŒuÐ%°qÓ9Ó
:ñ	;ä
ŒuÐ%¨Ñ)¨>¸1Ó=Ó
>ñ	?ð Ð
ñ	ð ð 
�‰�LÓ	!€Bà�Z‰Z˜˜AÓ!€FØ�Z‰Z˜˜FŸK™K¨›N¨AÓ.€FØ—.‘.  A qÓ)€Kä�‰×%Ñ% fÀÐ%ÐG×KÑKÈAÓN×RÑRÐSUÓVð ð 	�‰‹
€Að —‘˜2Ÿ8™8 C R˜=Ó)€JÜœEÐ"2Ó3Ó4ÐÜ�=Ö!ˆØÐ-°Ñ1°>ÀAÑ3EÐFÑFÒñ "à×#Ñ#Ð$4Ó5€JØ×Ñ˜nÓ-Ð-r   c                 ór  • [         R                  R                  [         R                  " U S5      5      n[         R                  " [         R                  " US5      SS5      n[         R
                  " U R                  S   U R                  U R                  S9n[         R                  R                  X#SS9nU$ )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°uÐ%Ð=€AØ€Hr   c                   óÈ  ^ • \ rS rSrSr\R                  \R                  \R                  \R                  S.r	\R                  r
Sr    SS\S\S-  S\S-  S	\S-  S
\S-  SS4U 4S jjjrSU 4S jjr\S\4S j5       r\S\4S j5       r\S\4S j5       r\S\4S j5       r\S\4S j5       r\S\4S j5       r\R2                  " 5       4S\S\4S jjrS rS rSrU =r$ )r	   éX   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_trilTNrP   rQ   rR   rS   Úvalidate_argsÚreturnc                 ó$  >• UR                  5       S:  a  [        S5      eUS LUS L-   US L-   S:w  a  [        S5      eUbj  UR                  5       S:  a  [        S5      e[        R                  " UR                  S S UR                  S S 5      nUR                  US-   5      U l        OäUbj  UR                  5       S:  a  [        S	5      e[        R                  " UR                  S S UR                  S S 5      nUR                  US-   5      U l        OwUc  [        S
5      eUR                  5       S:  a  [        S5      e[        R                  " UR                  S S UR                  S S 5      nUR                  US-   5      U l	        UR                  US-   5      U l
        U R                  R                  SS  n[        TU ]1  XgUS9  Ub  X@l        g Ub%  [        R                  R                  U5      U l        g [!        U5      U l        g )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 dimensionsz%precision_matrix is unexpectedly NonezYprecision_matrix must be at least two-dimensional, with optional leading batch dimensions)r   ©rT   )r   Ú
ValueErrorr   Úbroadcast_shapesr   ÚexpandrS   rQ   ÚAssertionErrorrR   rP   ÚsuperÚ__init__Ú_unbroadcasted_scale_trilr$   rD   rM   )	ÚselfrP   rQ   rR   rS   rT   Úbatch_shapeÚevent_shapeÚ	__class__s	           €r   r]   Ú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Õ"àÑ'Ü$Ð%LÓMÐMØ×#Ñ#Ó%¨Ó)Ü ð=óð ô  ×0Ò0Ø ×&Ñ& s¨Ð+¨S¯Y©Y°s¸¨^óˆKð %5×$;Ñ$;¸KÈ(Ñ<RÓ$SˆDÔ!Ø—:‘:˜k¨EÑ1Ó2ˆŒà—h‘h—n‘n R SÐ)ˆä‰Ñ˜ÀÐÑOàÑ!Ø-7Õ*ØÑ*Ü-2¯\©\×-BÑ-BÐCTÓ-UˆDÕ*ä-EÐFVÓ-WˆDÕ*r   c                 óŽ  >• U R                  [        U5      n[        R                  " U5      nXR                  -   nXR                  -   U R                  -   nU R
                  R                  U5      Ul        U R                  Ul        SU R                  ;   a   U R                  R                  U5      Ul	        SU R                  ;   a   U R                  R                  U5      Ul
        SU R                  ;   a   U R                  R                  U5      Ul        [        [        U]7  XR                  SS9  U R                  Ul        U$ )NrQ   rS   rR   FrW   )Ú_get_checked_instancer	   r   ÚSizera   rP   rZ   r^   Ú__dict__rQ   rS   rR   r\   r]   Ú_validate_args)r_   r`   Ú	_instanceÚnewÚ	loc_shapeÚ	cov_shaperb   s         €r   rZ   Ú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   c                 ó€   • U R                   R                  U R                  U R                  -   U R                  -   5      $ ©N)r^   rZ   Ú_batch_shapeÚ_event_shape©r_   s    r   rS   ÚMultivariateNormal.scale_trilÙ   s:   € à×-Ñ-×4Ñ4Ø×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r   c                 óÒ   • [         R                  " U R                  U R                  R                  5      R	                  U R
                  U R                  -   U R                  -   5      $ ro   )r   r   r^   ÚmTrZ   rp   rq   rr   s    r   rQ   Ú$MultivariateNormal.covariance_matrixß   sQ   € ä�|Š|Ø×*Ñ*¨D×,JÑ,J×,MÑ,Mó
ç
‰&�×"Ñ" T×%6Ñ%6Ñ6¸×9JÑ9JÑJÓ
Kð	Lr   c                 ó¨   • [         R                  " U R                  5      R                  U R                  U R
                  -   U R
                  -   5      $ ro   )r   Úcholesky_inverser^   rZ   rp   rq   rr   s    r   rR   Ú#MultivariateNormal.precision_matrixå   sE   € ä×%Ò% d×&DÑ&DÓE×LÑLØ×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r   c                 ó   • U R                   $ ro   ©rP   rr   s    r   ÚmeanÚMultivariateNormal.meanë   ó   € à�x‰xˆr   c                 ó   • U R                   $ ro   r{   rr   s    r   ÚmodeÚMultivariateNormal.modeï   r~   r   c                 ó¢   • U R                   R                  S5      R                  S5      R                  U R                  U R
                  -   5      $ )Nr   r   )r^   r&   r'   rZ   rp   rq   rr   s    r   ÚvarianceÚMultivariateNormal.varianceó   sA   € ð ×*Ñ*×.Ñ.¨qÓ1ß‰S�‹Wß‰V�D×%Ñ%¨×(9Ñ(9Ñ9Ó:ð	
r   Úsample_shapec                 óÎ   • U R                  U5      n[        X R                  R                  U R                  R                  S9nU R                  [        U R                  U5      -   $ )NrA   )Ú_extended_shaper   rP   rB   rC   r   r^   )r_   r…   r   Úepss       r   ÚrsampleÚMultivariateNormal.rsampleû   sJ   € Ø×$Ñ$ \Ó2ˆÜ˜u¯H©H¯N©NÀ4Ç8Á8Ç?Á?ÑSˆØ�x‰xœ) D×$BÑ$BÀCÓHÑHÐHr   c                 ó|  • U R                   (       a  U R                  U5        XR                  -
  n[        U R                  U5      nU R                  R                  SSS9R                  5       R                  S5      nSU R                  S   [        R                  " S[        R                  -  5      -  U-   -  U-
  $ )Nr   r   ©Údim1Údim2g      à¿r   r   )rh   Ú_validate_samplerP   r?   r^   ÚdiagonalÚlogr'   rq   ÚmathÚpi)r_   ÚvalueÚdiffr:   Úhalf_log_dets        r   Úlog_probÚMultivariateNormal.log_prob   s¢   € Ø××Ø×!Ñ! %Ô(Ø—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                 ó\  • U R                   R                  SSS9R                  5       R                  S5      nSU R                  S   -  S[
        R                  " S[
        R                  -  5      -   -  U-   n[        U R                  5      S:X  a  U$ UR                  U R                  5      $ )Nr   r   rŒ   g      à?r   g      ð?r   )
r^   r�   r‘   r'   rq   r’   r“   r   rp   rZ   )r_   r–   ÚHs      r   ÚentropyÚ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   )r^   rQ   rP   rR   rS   )NNNNro   ) Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Úreal_vectorÚpositive_definiteÚlower_choleskyÚarg_constraintsÚsupportÚhas_rsampler   Úboolr]   rZ   r   rS   rQ   rR   Úpropertyr|   r€   rƒ   r   rf   r   r‰   r—   r›   Ú__static_attributes__Ú__classcell__)rb   s   @r   r	   r	   X   sŽ  ø† ñ"ðL ×&Ñ&Ø(×:Ñ:Ø'×9Ñ9Ø!×0Ñ0ñ	€Oð ×%Ñ%€GØ€Kð
 ,0Ø*.Ø$(Ø%)ñ=Xàð=Xð " D™=ð=Xð ! 4™-ð	=Xð
 ˜T‘Mð=Xð ˜d‘{ð=Xð 
÷=Xð =X÷~ð& ð
˜Fó 
ó ð
ð
 ðL 6ó Ló ðLð
 ð
 &ó 
ó ð
ð
 ð�fó ó ðð ð�fó ó ðð ð
˜&ó 
ó ð
ð -2¯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ôz/˜õ z/r   