§
    ŠŠtj:,  ã                   óŠ   — 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dS )é    N)ÚTensor)Úconstraints)ÚDistribution)Ú_standard_normalÚlazy_property)Ú_sizeÚMultivariateNormalc                 óx   — t          j        | |                     d¦  «        ¦  «                             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/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributions/multivariate_normal.pyÚ	_batch_mvr      s0   € õ Œ<˜˜dŸnšn¨RÑ0Ô0Ñ1Ô1×9Ò9¸"Ñ=Ô=Ð=ó    c                 óä  — |                      d¦  «        }|j        dd…         }t          |¦  «        }|                      ¦   «         dz
  }||z
  }||z   }|d|z  z   }|j        d|…         }	t	          | j        dd…         |j        |d…         ¦  «        D ]\  }
}|	||
z  |
fz  }	Œ|	|fz  }	|                     |	¦  «        }t          t          |¦  «        ¦  «        t          t          ||d¦  «        ¦  «        z   t          t          |dz   |d¦  «        ¦  «        z   |gz   }|                     |¦  «        }|                      d||¦  «        }|                     d|                      d¦  «        |¦  «        }|                     ddd¦  «        }t          j
                             ||d¬¦  «                             d¦  «                             d¦  «        }|                     ¦   «         }|                     |j        dd…         ¦  «        }t          t          |¦  «        ¦  «        }t          |¦  «        D ]}|||z   ||z   gz  }Œ|                     |¦  «        }|                     |¦  «        S )	aN  
    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 be 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˜c˜r˜c”]€Nõ ˜Ñ'Ô'€MØ—F’F‘H”H˜q‘L€MØ$ }Ñ4ÐØ%¨Ñ5€NØ%¨¨MÑ(9Ñ9€Nà”8Ð-Ð-Ð-Ô.€LÝ�b”h˜s ˜s”m R¤XÐ.>¸rÐ.AÔ%BÑCÔCð 'ð '‰ˆˆBØ˜˜r™ 2˜Ñ&ˆˆØ�Q�DÑ€LØ	�Š�LÑ	!Ô	!€Bõ 	�UÐ#Ñ$Ô$Ñ%Ô%Ý
�uÐ% ~°qÑ9Ô9Ñ
:Ô
:ñ	;å
�uÐ%¨Ñ)¨>¸1Ñ=Ô=Ñ
>Ô
>ñ	?ð Ð
ñ	ð ð 
�Š�LÑ	!Ô	!€Bà�ZŠZ˜˜A˜qÑ!Ô!€FØ�ZŠZ˜˜FŸKšK¨™NœN¨AÑ.Ô.€FØ—.’.  A qÑ)Ô)€KåŒ×%Ò% f¨kÀÐ%ÑGÔG×KÒKÈAÑNÔN×RÒRÐSUÑVÔVð ð 	�Š‰
Œ
€Að —’˜2œ8 C R Cœ=Ñ)Ô)€JÝ�EÐ"2Ñ3Ô3Ñ4Ô4ÐÝ�=Ñ!Ô!ð Gð GˆØÐ-°Ñ1°>ÀAÑ3EÐFÑFÐÐØ×#Ò#Ð$4Ñ5Ô5€JØ×Ò˜nÑ-Ô-Ð-r   c                 óX  — t           j                             t          j        | d¦  «        ¦  «        }t          j        t          j        |d¦  «        dd¦  «        }t          j        | j        d         | j        | j        ¬¦  «        }t           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¨!¨XÑ6Ô6Ñ	7Ô	7€BÝŒO�EœJ r¨8Ñ4Ô4°b¸"Ñ=Ô=€EÝ	Œ�1”7˜2”; a¤g°a´hÐ	?Ñ	?Ô	?€BÝŒ×%Ò% e¨R°uÐ%Ñ=Ô=€AØ€Hr   c                   ó”  ‡ — e Zd ZdZej        ej        ej        ej        dœZej        Z	dZ
	 	 	 	 ddededz  dedz  dedz  d	edz  d
dfˆ 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_trilTNrO   rP   rQ   rR   Úvalidate_argsÚreturnc                 ó°  •— |                      ¦   «         dk     rt          d¦  «        ‚|d u|d uz   |d uz   dk    rt          d¦  «        ‚|�t|                      ¦   «         dk     rt          d¦  «        ‚t          j        |j        d d…         |j        d d…         ¦  «        }|                     |dz   ¦  «        | _        nú|�t|                      ¦   «         dk     rt          d	¦  «        ‚t          j        |j        d d…         |j        d d…         ¦  «        }|                     |dz   ¦  «        | _        n„|€t          d
¦  «        ‚|                      ¦   «         dk     rt          d¦  «        ‚t          j        |j        d d…         |j        d d…         ¦  «        }|                     |dz   ¦  «        | _	        |                     |dz   ¦  «        | _
        | j
        j        dd …         }t          ¦   «                              |||¬¦  «         |�	|| _        d S |�&t          j                             |¦  «        | _        d S t!          |¦  «        | _        d S )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   ©rS   )r   Ú
ValueErrorr   Úbroadcast_shapesr   ÚexpandrR   rP   ÚAssertionErrorrQ   rO   ÚsuperÚ__init__Ú_unbroadcasted_scale_trilr$   rD   rM   )	ÚselfrO   rP   rQ   rR   rS   Úbatch_shapeÚevent_shapeÚ	__class__s	           €r   r\   zMultivariateNormal.__init__‡   s�  ø€ ð �7Š7‰9Œ9�qŠ=ˆ=ÝÐDÑEÔEÐEØ TÐ)¨jÀÐ.DÑEØ DÐ(ñ
àòð õ Øfñô ð ð Ð!Ø�~Š~ÑÔ !Ò#Ð#Ý ð=ñô ð õ  Ô0°Ô1AÀ#À2À#Ô1FÈÌ	ÐRUÐSUÐRUÌÑWÔWˆKà(×/Ò/°¸hÑ0FÑGÔGˆDŒOˆOØÐ*Ø ×$Ò$Ñ&Ô&¨Ò*Ð*Ý ð=ñô ð õ  Ô0Ø!Ô'¨¨¨Ô,¨c¬i¸¸¸¬nñô ˆKð &7×%=Ò%=¸kÈHÑ>TÑ%UÔ%UˆDÔ"Ð"àÐ'Ý$Ð%LÑMÔMÐMØ×#Ò#Ñ%Ô%¨Ò)Ð)Ý ð=ñô ð õ  Ô0Ø Ô& s¨ sÔ+¨S¬Y°s¸°s¬^ñô ˆKð %5×$;Ò$;¸KÈ(Ñ<RÑ$SÔ$SˆDÔ!Ø—:’:˜k¨EÑ1Ñ2Ô2ˆŒà”h”n R S SÔ)ˆå‰Œ×Ò˜ kÀÐÑOÔOÐOàÐ!Ø-7ˆDÔ*Ð*Ð*ØÐ*Ý-2¬\×-BÒ-BÐCTÑ-UÔ-UˆDÔ*Ð*Ð*å-EÐFVÑ-WÔ-WˆDÔ*Ð*Ð*r   c                 ó\  •— |                       t          |¦  «        }t          j        |¦  «        }|| j        z   }|| j        z   | j        z   }| j                             |¦  «        |_        | j        |_        d| j        v r| j	                             |¦  «        |_	        d| j        v r| j
                             |¦  «        |_
        d| j        v r| j                             |¦  «        |_        t          t          |¦  «                             || j        d¬¦  «         | j        |_        |S )NrP   rR   rQ   FrV   )Ú_get_checked_instancer	   r   ÚSizer`   rO   rY   r]   Ú__dict__rP   rR   rQ   r[   r\   Ú_validate_args)r^   r_   Ú	_instanceÚnewÚ	loc_shapeÚ	cov_shapera   s         €r   rY   zMultivariateNormal.expandÆ   s  ø€ Ø×(Ò(Õ);¸YÑGÔGˆÝ”j Ñ-Ô-ˆØ $Ô"2Ñ2ˆ	Ø $Ô"2Ñ2°TÔ5EÑEˆ	Ø”(—/’/ )Ñ,Ô,ˆŒØ(,Ô(FˆÔ%Ø $¤-Ð/Ð/Ø$(Ô$:×$AÒ$AÀ)Ñ$LÔ$LˆCÔ!Ø˜4œ=Ð(Ð(Ø!œ_×3Ò3°IÑ>Ô>ˆCŒNØ ¤Ð.Ð.Ø#'Ô#8×#?Ò#?À	Ñ#JÔ#JˆCÔ ÝÕ  #Ñ&Ô&×/Ò/Ø˜Ô)¸ð 	0ñ 	
ô 	
ð 	
ð "Ô0ˆÔØˆ
r   c                 ó`   — | j                              | j        | j        z   | j        z   ¦  «        S ©N)r]   rY   Ú_batch_shapeÚ_event_shape©r^   s    r   rR   zMultivariateNormal.scale_trilÙ   s3   € àÔ-×4Ò4ØÔ Ô 1Ñ1°DÔ4EÑEñ
ô 
ð 	
r   c                 óš   — t          j        | j        | j        j        ¦  «                             | j        | j        z   | j        z   ¦  «        S rl   )r   r   r]   ÚmTrY   rm   rn   ro   s    r   rP   z$MultivariateNormal.covariance_matrixß   sE   € åŒ|ØÔ*¨DÔ,JÔ,Mñ
ô 
ç
Š&�Ô" TÔ%6Ñ6¸Ô9JÑJÑ
KÔ
Kð	Lr   c                 ó„   — t          j        | j        ¦  «                             | j        | j        z   | j        z   ¦  «        S rl   )r   Úcholesky_inverser]   rY   rm   rn   ro   s    r   rQ   z#MultivariateNormal.precision_matrixå   s>   € åÔ% dÔ&DÑEÔE×LÒLØÔ Ô 1Ñ1°DÔ4EÑEñ
ô 
ð 	
r   c                 ó   — | j         S rl   ©rO   ro   s    r   ÚmeanzMultivariateNormal.meanë   ó	   € àŒxˆr   c                 ó   — | j         S rl   ru   ro   s    r   ÚmodezMultivariateNormal.modeï   rw   r   c                 óœ   — | j                              d¦  «                             d¦  «                             | j        | j        z   ¦  «        S )Nr   r   )r]   r&   r'   rY   rm   rn   ro   s    r   ÚvariancezMultivariateNormal.varianceó   s@   € ð Ô*×.Ò.¨qÑ1Ô1ßŠS�‰WŒWßŠV�DÔ%¨Ô(9Ñ9Ñ:Ô:ð	
r   Úsample_shapec                 ó²   — |                       |¦  «        }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û   sK   € Ø×$Ò$ \Ñ2Ô2ˆÝ˜u¨D¬H¬NÀ4Ä8Ä?ÐSÑSÔSˆØŒx�) DÔ$BÀCÑHÔHÑHÐHr   c                 ój  — | j         r|                      |¦  «         || j        z
  }t          | j        |¦  «        }| j                             dd¬¦  «                             ¦   «                              d¦  «        }d| j        d         t          j        dt          j
        z  ¦  «        z  |z   z  |z
  S )Nr   r   ©Údim1Údim2g      à¿r   r   )rf   Ú_validate_samplerO   r?   r]   ÚdiagonalÚlogr'   rn   ÚmathÚpi)r^   ÚvalueÚdiffr:   Úhalf_log_dets        r   Úlog_probzMultivariateNormal.log_prob   s¬   € ØÔð 	)Ø×!Ò! %Ñ(Ô(Ð(Ø�t”xÑˆÝ˜tÔ=¸tÑDÔDˆàÔ*×3Ò3¸À"Ð3ÑEÔE×IÒIÑKÔK×OÒOÐPRÑSÔSð 	ð �tÔ(¨Ô+­d¬h°q½4¼7±{Ñ.CÔ.CÑCÀaÑGÑHÈ<ÑWÐWr   c                 ó\  — | j                              dd¬¦  «                             ¦   «                              d¦  «        }d| j        d         z  dt          j        dt
          j        z  ¦  «        z   z  |z   }t          | j        ¦  «        dk    r|S | 	                    | j        ¦  «        S )Nr   r   r‚   g      à?r   g      ð?r   )
r]   r†   r‡   r'   rn   rˆ   r‰   r   rm   rY   )r^   rŒ   ÚHs      r   ÚentropyzMultivariateNormal.entropy
  sž   € àÔ*×3Ò3¸À"Ð3ÑEÔE×IÒIÑKÔK×OÒOÐPRÑSÔSð 	ð �$Ô# AÔ&Ñ&¨#µ´¸½T¼W¹Ñ0EÔ0EÑ*EÑFÈÑUˆÝˆtÔ Ñ!Ô! QÒ&Ð&ØˆHà—8’8˜DÔ-Ñ.Ô.Ð.r   )NNNNrl   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úreal_vectorÚpositive_definiteÚlower_choleskyÚarg_constraintsÚsupportÚhas_rsampler   Úboolr\   rY   r   rR   rP   rQ   Úpropertyrv   ry   r{   r   rd   r   r€   r�   r�   Ú__classcell__)ra   s   @r   r	   r	   X   sQ  ø€ € € € € ð"ð "ðL Ô&Ø(Ô:Ø'Ô9Ø!Ô0ð	ð €Oð Ô%€GØ€Kð
 ,0Ø*.Ø$(Ø%)ð=Xð =Xàð=Xð " D™=ð=Xð ! 4™-ð	=Xð
 ˜T‘Mð=Xð ˜d‘{ð=Xð 
ð=Xð =Xð =Xð =Xð =Xð =Xð~ð ð ð ð ð ð& ð
˜Fð 
ð 
ð 
ñ „]ð
ð
 ðL 6ð Lð Lð Lñ „]ðLð
 ð
 &ð 
ð 
ð 
ñ „]ð
ð
 ð�fð ð ð ñ „Xðð ð�fð ð ð ñ „Xðð ð
˜&ð 
ð 
ð 
ñ „Xð
ð -7¨E¬J©L¬Lð Ið I Eð I¸Vð Ið Ið Ið Ið
Xð Xð 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¤      sí   ðà €€€à €€€Ø Ð Ð Ð Ð Ð Ø +Ð +Ð +Ð +Ð +Ð +Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø EÐ EÐ EÐ EÐ EÐ EÐ EÐ EØ Ð Ð Ð Ð Ð ð  Ð
 €ð>ð >ð >ð/.ð /.ð /.ðdð ð ðz/ð z/ð z/ð z/ð z/˜ñ z/ô z/ð z/ð z/ð z/r   