ó
    Eñi´6  ã                   óØ   • S SK r S SKrS SKrS SKJrJr  S SKJr  S SKJr  S SK	J
r
  S SKJr  S SKJrJrJr  S/r\ R$                  " S	5      rS
\S\S\4S jrS
\S\4S jr " S S\5      rg)é    N)ÚnanÚTensor)Úconstraints)ÚExponentialFamily)Ú_precision_to_scale_tril)Úlazy_property)Ú_NumberÚ_sizeÚNumberÚWisharté   ÚxÚpÚreturnc           	      ó~  • U R                  US-
  S-  5      R                  5       (       d  [        S5      e[        R                  " U R                  S5      [        R                  " XR                  U R                  S9R                  S5      R                  U R                  S-   5      -
  5      R                  S5      $ )Né   r   z/Wrong domain for multivariate digamma function.éÿÿÿÿ©ÚdtypeÚdevice©r   )ÚgtÚallÚAssertionErrorÚtorchÚdigammaÚ	unsqueezeÚaranger   r   ÚdivÚexpandÚshapeÚsum)r   r   s     ÚX/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/wishart.pyÚ
_mvdigammar$      s�   € Ø�4‰4��Q‘˜!‘Ó× Ñ ×"Ñ"ÜÐNÓOÐOÜ�=Š=Ø	�‰�B‹Ü
�,Š,�q§¡°·±Ñ
9×
=Ñ
=¸aÓ
@×
GÑ
GÈÏÉÐRWÉÓ
Xñ	Yó÷ 
�cˆ"ƒgðó    c                 óp   • U R                  [        R                  " U R                  5      R                  S9$ )N)Úmin)Úclampr   Úfinfor   Úeps)r   s    r#   Ú_clamp_above_epsr+      s&   € à�7‰7”u—{’{ 1§7¡7Ó+×/Ñ/ˆ7Ð0Ð0r%   c                   óØ  ^ • \ rS rSrSr\R                  rSrSr	\
S 5       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 jr\R2                  " 5       S4S\S\4S jjrS rS r\
S\\\4   4S j5       r S r!Sr"U =r#$ )r   é!   aW  
Creates a Wishart distribution parameterized by a symmetric positive definite matrix :math:`\Sigma`,
or its Cholesky decomposition :math:`\mathbf{\Sigma} = \mathbf{L}\mathbf{L}^\top`

Example:
    >>> # xdoctest: +SKIP("FIXME: scale_tril must be at least two-dimensional")
    >>> m = Wishart(torch.Tensor([2]), covariance_matrix=torch.eye(2))
    >>> m.sample()  # Wishart distributed with mean=`df * I` and
    >>> # variance(x_ij)=`df` for i != j and variance(x_ij)=`2 * df` for i == j

Args:
    df (float or Tensor): real-valued parameter larger than the (dimension of Square matrix) - 1
    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.
    'torch.distributions.LKJCholesky' is a restricted Wishart distribution.[1]

**References**

[1] Wang, Z., Wu, Y. and Chu, H., 2018. `On equivalence of the LKJ distribution and the restricted Wishart distribution`.
[2] Sawyer, S., 2007. `Wishart Distributions and Inverse-Wishart Sampling`.
[3] Anderson, T. W., 2003. `An Introduction to Multivariate Statistical Analysis (3rd ed.)`.
[4] Odell, P. L. & Feiveson, A. H., 1966. `A Numerical Procedure to Generate a SampleCovariance Matrix`. JASA, 61(313):199-203.
[5] Ku, Y.-C. & Bloomfield, P., 2010. `Generating Random Wishart Matrices with Fractional Degrees of Freedom in OX`.
Tr   c                 ó¬   • [         R                  [         R                  [         R                  [         R                  " U R                  S   S-
  5      S.$ )Nr   r   )Úcovariance_matrixÚprecision_matrixÚ
scale_trilÚdf)r   Úpositive_definiteÚlower_choleskyÚgreater_thanÚevent_shape©Úselfs    r#   Úarg_constraintsÚWishart.arg_constraintsG   sG   € ô "-×!>Ñ!>Ü +× =Ñ =Ü%×4Ñ4Ü×*Ò*¨4×+;Ñ+;¸BÑ+?À!Ñ+CÓDñ	
ð 	
r%   Nr2   r/   r0   r1   Úvalidate_argsr   c           	      óH  >• US LUS L-   US L-   S:w  a  [        S5      e[        S X#U4 5       5      nUR                  5       S:  a  [        S5      e[	        U[
        5      (       aR  [        R                  " UR                  S S 5      n[        R                  " XR                  UR                  S9U l        OD[        R                  " UR                  S S UR                  5      nUR                  U5      U l        UR                  SS  nU R                  R                  US   S-
  5      R!                  5       (       a  [        S	U S
US   S-
   S35      eUb  UR                  US-   5      U l        O9Ub  UR                  US-   5      U l        OUb  UR                  US-   5      U l        U R                  R)                  US   5      R!                  5       (       a  [*        R,                  " SSS9  [.        T
U ]a  XxUS9  [3        [5        U R6                  5      5       V	s/ s H  o™S-   * PM
     sn	U l        Ub  X@l        O8Ub%  [        R<                  R?                  U5      U l        O[A        U5      U l        [        RB                  RD                  RG                  U R                  RI                  S5      [        RJ                  " U RL                  S   U R:                  R                  U R:                  R                  S9R                  US-   5      -
  S9U l'        g s  sn	f )Nr   zTExactly one of covariance_matrix or precision_matrix or scale_tril may be specified.c              3   ó0   #   • U  H  nUc  M  Uv •  M     g 7f©N© )Ú.0r   s     r#   Ú	<genexpr>Ú#Wishart.__init__.<locals>.<genexpr>a   s   é € ð 
âF�Ø÷ ‰AÚFùs   ‚�	r   zSscale_tril must be at least two-dimensional, with optional leading batch dimensionséþÿÿÿr   r   zValue of df=z( expected to be greater than ndim - 1 = Ú.)r   r   z]Low df values detected. Singular samples are highly likely to occur for ndim - 1 < df < ndim.©Ú
stacklevel©r;   r   ©r2   )(r   ÚnextÚdimÚ
ValueErrorÚ
isinstancer	   r   ÚSizer!   Útensorr   r   r2   Úbroadcast_shapesr    ÚleÚanyr1   r/   r0   ÚltÚwarningsÚwarnÚsuperÚ__init__ÚrangeÚlenÚ_batch_shapeÚ_batch_dimsÚ_unbroadcasted_scale_trilÚlinalgÚcholeskyr   ÚdistributionsÚchi2ÚChi2r   r   Ú_event_shapeÚ
_dist_chi2)r8   r2   r/   r0   r1   r;   ÚparamÚbatch_shaper6   r   Ú	__class__s             €r#   rV   ÚWishart.__init__P   sà  ø€ ð  dÐ*Ø Ð%ñ'à tÐ+ñ-ð ó	ô
 !Øfóð ô ñ 
à'¸:ÑFó
ó 
ˆð �9‰9‹;˜‹?ÜØeóð ô �bœ'×"Ñ"ÜŸ*š* U§[¡[°°"Ð%5Ó6ˆKÜ—l’l 2¯[©[ÀÇÁÑNˆD�Gä×0Ò0°·±¸S¸bÐ1AÀ2Ç8Á8ÓLˆKØ—i‘i Ó,ˆDŒGØ—k‘k " #Ð&ˆà�7‰7�:‰:�k "‘o¨Ñ)Ó*×.Ñ.×0Ñ0ÜØ˜r˜dÐ"JÈ;ÐWYÉ?Ð]^ÑK^ÐJ_Ð_`Ðaóð ð Ñ!à#Ÿl™l¨;¸Ñ+AÓBˆD�OØÑ*à%*§\¡\°+ÀÑ2HÓ%IˆDÕ"ØÑ)à$)§L¡L°¸xÑ1GÓ$HˆDÔ!à�7‰7�:‰:�k "‘oÓ&×*Ñ*×,Ñ,Ü�MŠMØoØòô 	‰Ñ˜ÀÐÑOÜ.3´C¸×8IÑ8IÓ4JÔ.KÓLÒ.K¨ !™e›HÑ.KÑLˆÔàÑ!Ø-7Õ*ØÑ*Ü-2¯\©\×-BÑ-BÐCTÓ-UˆDÕ*ä-EÐFVÓ-WˆDÔ*ô  ×-Ñ-×2Ñ2×7Ñ7à—‘×!Ñ! "Ó%Ü—,’,Ø×%Ñ% bÑ)Ø×8Ñ8×>Ñ>Ø×9Ñ9×@Ñ@ñ÷ ‘&˜ uÑ,Ó-ñ.ð 8ð 	
ˆ�ùò Ms   È"Lc           	      ó   >• U R                  [        U5      n[        R                  " U5      nXR                  -   nU R
                  R                  U5      Ul        U R                  R                  U5      Ul        [        [        U5      5       Vs/ s H  oUS-   * PM
     sn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        [        R                  R                   R#                  UR                  R%                  S5      [        R&                  " U R                  S   UR
                  R(                  UR
                  R*                  S9R                  US-   5      -
  S9Ul        [.        [        U]c  XR                  S	S
9  U R2                  Ul        U$ s  snf )Nr   r/   r1   r0   r   r   r   rH   FrG   )Ú_get_checked_instancer   r   rM   r6   r[   r    r2   rW   rX   rZ   Ú__dict__r/   r1   r0   r^   r_   r`   r   r   r   r   rb   rU   rV   Ú_validate_args)r8   rd   Ú	_instanceÚnewÚ	cov_shaper   re   s         €r#   r    ÚWishart.expand    s«  ø€ Ø×(Ñ(¬°)Ó<ˆÜ—j’j Ó-ˆØ×"2Ñ"2Ñ2ˆ	Ø(,×(FÑ(F×(MÑ(MÈiÓ(XˆÔ%Ø—‘—‘ Ó,ˆŒä-2´3°{Ó3CÔ-DÓEÒ-D¨ ™U›8Ñ-DÑEˆŒà $§-¡-Ó/Ø$(×$:Ñ$:×$AÑ$AÀ)Ó$LˆCÔ!Ø˜4Ÿ=™=Ó(Ø!Ÿ_™_×3Ñ3°IÓ>ˆCŒNØ §¡Ó.Ø#'×#8Ñ#8×#?Ñ#?À	Ó#JˆCÔ ô ×,Ñ,×1Ñ1×6Ñ6à—‘× Ñ  Ó$Ü—,’,Ø×$Ñ$ RÑ(Ø×7Ñ7×=Ñ=Ø×8Ñ8×?Ñ?ñ÷ ‘&˜ uÑ,Ó-ñ.ð 7ð 	
ˆŒô 	Œg�sÑ$ [×2BÑ2BÐRWÐ$ÑXØ!×0Ñ0ˆÔØˆ
ùò/ Fs   ÂHc                 óf   • U R                   R                  U R                  U R                  -   5      $ r>   )r[   r    rY   ra   r7   s    r#   r1   ÚWishart.scale_trilÀ   s/   € à×-Ñ-×4Ñ4Ø×Ñ × 1Ñ 1Ñ1ó
ð 	
r%   c                 ó    • U R                   U R                   R                  SS5      -  R                  U R                  U R                  -   5      $ )NrC   r   )r[   Ú	transposer    rY   ra   r7   s    r#   r/   ÚWishart.covariance_matrixÆ   sH   € ð ×*Ñ*Ø×,Ñ,×6Ñ6°r¸2Ó>ñ?ç
‰&�×"Ñ" T×%6Ñ%6Ñ6Ó
7ð	8r%   c                 ó$  • [         R                  " U R                  S   U R                  R                  U R                  R
                  S9n[         R                  " XR                  5      R                  U R                  U R                  -   5      $ )Nr   )r   r   )	r   Úeyera   r[   r   r   Úcholesky_solver    rY   )r8   Úidentitys     r#   r0   ÚWishart.precision_matrixÍ   sv   € ä—9’9Ø×Ñ˜bÑ!Ø×1Ñ1×8Ñ8Ø×0Ñ0×6Ñ6ñ
ˆô
 ×#Ò# H×.LÑ.LÓM×TÑTØ×Ñ × 1Ñ 1Ñ1ó
ð 	
r%   c                 ól   • U R                   R                  U R                  S-   5      U R                  -  $ )N©r   r   )r2   ÚviewrY   r/   r7   s    r#   ÚmeanÚWishart.meanØ   s+   € à�w‰w�|‰|˜D×-Ñ-°Ñ6Ó7¸$×:PÑ:PÑPÐPr%   c                 óÀ   • U R                   U R                  R                  S   -
  S-
  n[        XS:*  '   UR	                  U R
                  S-   5      U R                  -  $ )Nr   r   r   rz   )r2   r/   r!   r   r{   rY   )r8   Úfactors     r#   ÚmodeÚWishart.modeÜ   sW   € à—‘˜4×1Ñ1×7Ñ7¸Ñ;Ñ;¸aÑ?ˆÜ!ˆ˜‰{ÑØ�{‰{˜4×,Ñ,¨vÑ5Ó6¸×9OÑ9OÑOÐOr%   c                 óÞ   • U R                   nUR                  SSS9nU R                  R                  U R                  S-   5      UR                  S5      [        R                  " SX"5      -   -  $ )NrC   r   ©Údim1Údim2rz   r   z...i,...j->...ij)r/   Údiagonalr2   r{   rY   Úpowr   Úeinsum)r8   ÚVÚdiag_Vs      r#   ÚvarianceÚWishart.varianceâ   s`   € à×"Ñ"ˆØ—‘ ¨"�Ð-ˆØ�w‰w�|‰|˜D×-Ñ-°Ñ6Ó7Ø�E‰E�!‹H”u—|’|Ð$6¸ÓGÑGñ
ð 	
r%   c                 óÞ  • U R                   S   n[        U R                  R                  U5      R	                  5       5      R                  SSS9n[        R                  " X"SS9u  pE[        R                  " [        R                  " U5      U R                  -   [        X"S-
  -  S-  5      4-   UR                  UR                  S9USXE4'   U R                  U-  nXfR                  SS5      -  $ )	Nr   rC   rƒ   )Úoffsetr   r   r   .)ra   r+   rb   ÚrsampleÚsqrtÚ
diag_embedr   Útril_indicesÚrandnrM   rY   Úintr   r   r[   rr   )r8   Úsample_shaper   ÚnoiseÚiÚjÚchols          r#   Ú_bartlett_samplingÚWishart._bartlett_samplingê   sÙ   € Ø×Ñ˜bÑ!ˆô !Ø�O‰O×#Ñ# LÓ1×6Ñ6Ó8ó
ç
‰*˜" 2ˆ*Ð
&ð 	ô ×!Ò! !¨rÑ2‰ˆÜ Ÿ;š;Ü�JŠJ�|Ó$ t×'8Ñ'8Ñ8¼CÀÈÁUÁÈaÁÓ<PÐ;RÑRØ—+‘+Ø—<‘<ñ
ˆˆc�1ˆiÑð
 ×-Ñ-°Ñ5ˆØ—n‘n R¨Ó,Ñ,Ð,r%   r•   c                 ó&  • Uc'  [         R                  R                  5       (       a  SOSn[         R                  " U5      nU R	                  U5      nU R
                  R                  U5      nU R                  (       a  UR                  U R                  5      n[         R                  R                  5       (       a†  [        U5       Hu  nU R	                  U5      n[         R                  " XFU5      nU R
                  R                  U5      ) nU R                  (       d  MZ  UR                  U R                  5      nMw     U$ UR                  5       (       aº  [        R                  " SSS9  [        U5       H–  nU R	                  XD   R                  5      nXcU'   U R
                  R                  U5      ) nU R                  (       a  UR                  U R                  5      nXtUR!                  5       '   UR                  5       (       a  M•    U$    U$ )aà  
.. warning::
    In some cases, sampling algorithm based on Bartlett decomposition may return singular matrix samples.
    Several tries to correct singular samples are performed by default, but it may end up returning
    singular matrix samples. Singular samples may return `-inf` values in `.log_prob()`.
    In those cases, the user should validate the samples and either fix the value of `df`
    or adjust `max_try_correction` value for argument in `.rsample` accordingly.
é   é
   zSingular sample detected.r   rE   )r   Ú_CÚ_get_tracing_staterM   rš   ÚsupportÚcheckrY   ÚamaxrZ   rW   ÚwhererQ   rS   rT   r!   Úclone)r8   r•   Úmax_try_correctionÚsampleÚis_singularÚ_Ú
sample_newÚis_singular_news           r#   r�   ÚWishart.rsampleû   s¬  € ð Ñ%Ü&+§h¡h×&AÑ&A×&CÑ&C¡ÈÐä—z’z ,Ó/ˆØ×(Ñ(¨Ó6ˆð —l‘l×(Ñ(¨Ó0ˆØ××Ø%×*Ñ*¨4×+;Ñ+;Ó<ˆKä�8‰8×&Ñ&×(Ñ(äÐ-Ö.�Ø!×4Ñ4°\ÓB�
ÜŸš [¸fÓE�à#Ÿ|™|×1Ñ1°&Ó9Ð9�Ø×$×$Ñ$Ø"-×"2Ñ"2°4×3CÑ3CÓ"D’Kñ /ð2 ˆð �‰× Ñ Ü—’Ð9ÀaÒHäÐ1Ö2�AØ!%×!8Ñ!8¸Ñ9Q×9WÑ9WÓ!X�JØ*4˜;Ñ'à'+§|¡|×'9Ñ'9¸*Ó'EÐ&E�OØ×(×(Ø*9×*>Ñ*>¸t×?OÑ?OÓ*P˜Ø7F × 1Ñ 1Ó 3Ñ4à&Ÿ?™?×,Ó,Øàˆñ 3ð ˆr%   c                 ó&  • U R                   (       a  U R                  U5        U R                  nU R                  S   nU* U[        -  S-  U R
                  R                  SSS9R                  5       R                  S5      -   -  [        R                  " US-  US9-
  X#-
  S-
  S-  [        R                  R                  U5      R                  -  -   [        R                  " XR
                  5      R                  SSS9R                  SS9S-  -
  $ )Nr   r   rC   rƒ   ©r   r   )rJ   )rj   Ú_validate_sampler2   ra   Ú_log_2r[   r†   Úlogr"   r   Úmvlgammar\   ÚslogdetÚ	logabsdetrv   )r8   ÚvalueÚnur   s       r#   Úlog_probÚWishart.log_prob/  sÿ   € Ø××Ø×!Ñ! %Ô(Ø�W‰WˆØ×Ñ˜bÑ!ˆàˆCà”F‘
˜Q‘Ø×0Ñ0×9Ñ9¸rÈÐ9ÐKß‘“ß‘�R“ññô �nŠn˜R !™V qÑ)ñ*ð ‰v˜‰z˜QÑ¤§¡×!5Ñ!5°eÓ!<×!FÑ!FÑFñGô ×"Ò" 5×*HÑ*HÓIß‰X˜2 BˆXÐ'ß‰S�RˆSˆ[Øññð	
r%   c                 ó@  • U R                   nU R                  S   nUS-   U[        -  S-  U R                  R	                  SSS9R                  5       R                  S5      -   -  [        R                  " US-  US9-   X-
  S-
  S-  [        US-  US9-  -
  X-  S-  -   $ )Nr   r   r   rC   rƒ   r®   )
r2   ra   r°   r[   r†   r±   r"   r   r²   r$   ©r8   r¶   r   s      r#   ÚentropyÚWishart.entropyD  s³   € Ø�W‰WˆØ×Ñ˜bÑ!ˆà�‰Uà”F‘
˜Q‘Ø×0Ñ0×9Ñ9¸rÈÐ9ÐKß‘“ß‘�R“ññô �nŠn˜R !™V qÑ)ñ*ð ‰v˜‰z˜QÑ¤¨B°©F°aÑ!8Ñ8ñ9ð ‰f�q‰jñ	ð	
r%   c                 ól   • U R                   nU R                  S   nU R                  * S-  X-
  S-
  S-  4$ )Nr   r   r   )r2   ra   r0   rº   s      r#   Ú_natural_paramsÚWishart._natural_paramsT  s?   € à�W‰WˆØ×Ñ˜bÑ!ˆØ×%Ñ%Ð%¨Ñ)¨B©F°Q©J¸!Ñ+;Ð;Ð;r%   c                 óà   • U R                   S   nX#S-   S-  -   [        R                  R                  SU-  5      R                  * [
        U-  -   -  [        R                  " X#S-   S-  -   US9-   $ )Nr   r   r   rC   r®   )ra   r   r\   r³   r´   r°   r²   )r8   r   Úyr   s       r#   Ú_log_normalizerÚWishart._log_normalizer[  sn   € Ø×Ñ˜bÑ!ˆØ˜‘U˜a‘K‘Ü�\‰\×!Ñ! " q¡&Ó)×3Ñ3Ð3´f¸q±jÑ@ñ
ä�NŠN˜1 A¡¨™{™?¨aÑ0ñ1ð 	1r%   )rZ   rb   r[   r/   r2   r0   r1   )NNNNr>   )$Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   r3   r¡   Úhas_rsampleÚ_mean_carrier_measureÚpropertyr9   r   r   ÚboolrV   r    r   r1   r/   r0   r|   r€   r‹   r   rM   rš   r
   r�   r·   r»   Útupler¾   rÂ   Ú__static_attributes__Ú__classcell__)re   s   @r#   r   r   !   s»  ø† ñðB ×+Ñ+€GØ€KØÐàñ
ó ð
ð ,0Ø*.Ø$(Ø%)ñN
à�V‰OðN
ð " D™=ðN
ð ! 4™-ð	N
ð
 ˜T‘MðN
ð ˜d‘{ðN
ð 
÷N
ð N
÷`ð@ ð
˜Fó 
ó ð
ð
 ð8 6ó 8ó ð8ð ð
 &ó 
ó ð
ð ðQ�fó Qó ðQð ðP�fó Pó ðPð
 ð
˜&ó 
ó ð
ð /4¯jªj«lô -ð$ %*§J¢J£LÀTñ2Ø!ð2à	õ2òh
ò*
ð  ð<  v¨v ~Ñ!6ó <ó ð<÷1ð 1r%   )ÚmathrS   r   r   r   Útorch.distributionsr   Útorch.distributions.exp_familyr   Ú'torch.distributions.multivariate_normalr   Útorch.distributions.utilsr   Útorch.typesr	   r
   r   Ú__all__r±   r°   r”   r$   r+   r   r?   r%   r#   Ú<module>r×      su   ðã Û ã ß Ý +Ý <Ý LÝ 3ß .Ñ .ð ˆ+€à	�Š�!‹€ð�&ð ˜Sð  Vô ð1˜ð 1 6ô 1ô
~1Ðõ ~1r%   