ó
    EñiÂ'  ã                   ó’   • 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Jr  S SKJr  S/rS	 rS
 rS r " S S\5      rg)é    N)ÚTensor)Úconstraints)ÚDistribution)Ú_batch_mahalanobisÚ	_batch_mv)Ú_standard_normalÚlazy_property)Ú_sizeÚLowRankMultivariateNormalc                 ó8  • U R                  S5      nU R                  UR                  S5      -  n[        R                  " X05      R                  5       nUR                  SX"-  5      SS2SSUS-   24==   S-  ss'   [        R                  R                  U5      $ )zw
Computes Cholesky of :math:`I + W.T @ inv(D) @ W` for a batch of matrices :math:`W`
and a batch of vectors :math:`D`.
éÿÿÿÿéþÿÿÿNé   )	ÚsizeÚmTÚ	unsqueezeÚtorchÚmatmulÚ
contiguousÚviewÚlinalgÚcholesky)ÚWÚDÚmÚWt_DinvÚKs        Úl/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/lowrank_multivariate_normal.pyÚ_batch_capacitance_trilr      s|   € ð
 	
�‰ˆr‹
€AØ�d‰d�Q—[‘[ “_Ñ$€GÜ�Š�WÓ ×+Ñ+Ó-€AØ‡F�Fˆ2ˆq‰uÓ’a™˜A ™E˜�kÓ" aÑ'Ó"Ü�<‰<× Ñ  Ó#Ð#ó    c                 ó¢   • SUR                  SSS9R                  5       R                  S5      -  UR                  5       R                  S5      -   $ )z³
Uses "matrix determinant lemma"::
    log|W @ W.T + D| = log|C| + log|D|,
where :math:`C` is the capacitance matrix :math:`I + W.T @ inv(D) @ W`, to compute
the log determinant.
é   r   r   )Údim1Údim2)ÚdiagonalÚlogÚsum)r   r   Úcapacitance_trils      r   Ú_batch_lowrank_logdetr)      sP   € ð Ð×(Ñ(¨b°rÐ(Ð:×>Ñ>Ó@×DÑDÀRÓHÑHÈ1Ï5É5Ë7Ï;É;Ø
óLñ ð r    c                 ó¸   • U R                   UR                  S5      -  n[        XB5      nUR                  S5      U-  R	                  S5      n[        X55      nXg-
  $ )zÿ
Uses "Woodbury matrix identity"::
    inv(W @ W.T + D) = inv(D) - inv(D) @ W @ inv(C) @ W.T @ inv(D),
where :math:`C` is the capacitance matrix :math:`I + W.T @ inv(D) @ W`, to compute the squared
Mahalanobis distance :math:`x.T @ inv(W @ W.T + D) @ x`.
r   r"   r   )r   r   r   Úpowr'   r   )r   r   Úxr(   r   Ú	Wt_Dinv_xÚmahalanobis_term1Úmahalanobis_term2s           r   Ú_batch_lowrank_mahalanobisr0   (   sV   € ð �d‰d�Q—[‘[ “_Ñ$€GÜ˜'Ó%€IØŸ™˜q› A™×*Ñ*¨2Ó.ÐÜ*Ð+;ÓGÐØÑ0Ð0r    c                   óÚ  ^ • \ rS rSrSr\R                  \R                  " \R                  S5      \R                  " \R                  S5      S.r
\R                  rSr 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\R4                  " 5       4S\S\4S jjrS rS rSrU =r $ )r   é6   a³  
Creates a multivariate normal distribution with covariance matrix having a low-rank form
parameterized by :attr:`cov_factor` and :attr:`cov_diag`::

    covariance_matrix = cov_factor @ cov_factor.T + cov_diag

Example:
    >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_LAPACK)
    >>> # xdoctest: +IGNORE_WANT("non-deterministic")
    >>> m = LowRankMultivariateNormal(
    ...     torch.zeros(2), torch.tensor([[1.0], [0.0]]), torch.ones(2)
    ... )
    >>> m.sample()  # normally distributed with mean=`[0,0]`, cov_factor=`[[1],[0]]`, cov_diag=`[1,1]`
    tensor([-0.2102, -0.5429])

Args:
    loc (Tensor): mean of the distribution with shape `batch_shape + event_shape`
    cov_factor (Tensor): factor part of low-rank form of covariance matrix with shape
        `batch_shape + event_shape + (rank,)`
    cov_diag (Tensor): diagonal part of low-rank form of covariance matrix with shape
        `batch_shape + event_shape`

Note:
    The computation for determinant and inverse of covariance matrix is avoided when
    `cov_factor.shape[1] << cov_factor.shape[0]` thanks to `Woodbury matrix identity
    <https://en.wikipedia.org/wiki/Woodbury_matrix_identity>`_ and
    `matrix determinant lemma <https://en.wikipedia.org/wiki/Matrix_determinant_lemma>`_.
    Thanks to these formulas, we just need to compute the determinant and inverse of
    the small size "capacitance" matrix::

        capacitance = I + cov_factor.T @ inv(cov_diag) @ cov_factor
r"   r   )ÚlocÚ
cov_factorÚcov_diagTNr3   r4   r5   Úvalidate_argsÚreturnc           	      óè  >• UR                  5       S:  a  [        S5      eUR                  SS  nUR                  5       S:  a  [        S5      eUR                  SS U:w  a  [        SUS    S	35      eUR                  SS  U:w  a  [        S
U 35      eUR                  S5      nUR                  S5      n [        R
                  " XbU5      u  o`l        nUS   U l        US   U l	        U R                  R                  S S n	X l
        X0l        [        X#5      U l        [        T
U ]=  X•US9  g ! [         a8  n[        SUR                   SUR                   SUR                   35      UeS nAff = f)Nr   z%loc must be at least one-dimensional.r   r"   zScov_factor must be at least two-dimensional, with optional leading batch dimensionsr   z2cov_factor must be a batch of matrices with shape r   z x mz/cov_diag must be a batch of vectors with shape zIncompatible batch shapes: loc z, cov_factor z, cov_diag ).r   ©r6   )ÚdimÚ
ValueErrorÚshaper   r   Úbroadcast_tensorsr4   ÚRuntimeErrorr3   r5   Ú_unbroadcasted_cov_factorÚ_unbroadcasted_cov_diagr   Ú_capacitance_trilÚsuperÚ__init__)Úselfr3   r4   r5   r6   Úevent_shapeÚloc_Ú	cov_diag_ÚeÚbatch_shapeÚ	__class__s             €r   rC   Ú"LowRankMultivariateNormal.__init__a   s–  ø€ ð �7‰7‹9�q‹=ÜÐDÓEÐEØ—i‘i  �nˆØ�>‰>Ó˜aÓÜð9óð ð ×Ñ˜B˜rÐ" kÓ1ÜØDÀ[ÐQRÁ^ÐDTÐTXÐYóð ð �>‰>˜"˜#Ð +Ó-ÜØAÀ+ÀÐOóð ð �}‰}˜RÓ ˆØ×&Ñ& rÓ*ˆ	ð	Ü/4×/FÒ/FØ )ó0Ñ,ˆD”/ 9ð ˜‘<ˆŒØ! &Ñ)ˆŒØ—h‘h—n‘n S bÐ)ˆà)3Ô&Ø'/Ô$Ü!8¸Ó!NˆÔä‰Ñ˜ÀÐÒOøô ó 	ÜØ1°#·)±)°¸MÈ*×JZÑJZÐI[Ð[fÐgo×guÑguÐfvÐwóàðûð	ús   Â8D/ Ä/
E1Ä93E,Å,E1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 R                  R                  X@R                  R                  SS  -   5      Ul        U R                  Ul
        U R                  Ul        U R                  Ul        [        [        U];  XR                  SS9  U R                  Ul        U$ )Nr   Fr9   )Ú_get_checked_instancer   r   ÚSizerE   r3   Úexpandr5   r4   r<   r?   r@   rA   rB   rC   Ú_validate_args)rD   rI   Ú	_instanceÚnewÚ	loc_shaperJ   s        €r   rO   Ú LowRankMultivariateNormal.expand�   sç   ø€ Ø×(Ñ(Ô)BÀIÓNˆÜ—j’j Ó-ˆØ×"2Ñ"2Ñ2ˆ	Ø—(‘(—/‘/ )Ó,ˆŒØ—}‘}×+Ñ+¨IÓ6ˆŒØŸ™×/Ñ/°	¿O¹O×<QÑ<QÐRTÐRUÐ<VÑ0VÓWˆŒØ(,×(FÑ(FˆÔ%Ø&*×&BÑ&BˆÔ#Ø $× 6Ñ 6ˆÔÜÔ'¨Ñ6Ø×)Ñ)¸ð 	7ñ 	
ð "×0Ñ0ˆÔØˆ
r    c                 ó   • U R                   $ ©N©r3   ©rD   s    r   ÚmeanÚLowRankMultivariateNormal.mean�   ó   € à�x‰xˆr    c                 ó   • U R                   $ rV   rW   rX   s    r   ÚmodeÚLowRankMultivariateNormal.mode¡   r[   r    c                 ó¼   • U R                   R                  S5      R                  S5      U R                  -   R	                  U R
                  U R                  -   5      $ )Nr"   r   )r?   r+   r'   r@   rO   Ú_batch_shapeÚ_event_shaperX   s    r   ÚvarianceÚ"LowRankMultivariateNormal.variance¥   sN   € ð ×*Ñ*×.Ñ.¨qÓ1×5Ñ5°bÓ9¸D×<XÑ<XÑXß
‰&�×"Ñ" T×%6Ñ%6Ñ6Ó
7ð	8r    c                 óì  • U R                   S   nU R                  R                  5       R                  S5      nU R                  U-  n[
        R                  " X3R                  5      R                  5       nUR                  SX-  5      S S 2S S US-   24==   S-  ss'   U[
        R                  R                  U5      -  nUR                  U R                  U R                   -   U R                   -   5      $ )Nr   r   r   )ra   r@   Úsqrtr   r?   r   r   r   r   r   r   r   rO   r`   )rD   ÚnÚcov_diag_sqrt_unsqueezeÚ
Dinvsqrt_Wr   Ú
scale_trils         r   ri   Ú$LowRankMultivariateNormal.scale_tril«   sÕ   € ð ×Ñ˜aÑ ˆØ"&×">Ñ">×"CÑ"CÓ"E×"OÑ"OÐPRÓ"SÐØ×3Ñ3Ð6MÑMˆ
Ü�LŠL˜§]¡]Ó3×>Ñ>Ó@ˆØ	�‰ˆr�1‘5Óš!™X  A¡˜X˜+Ó&¨!Ñ+Ó&Ø,¬u¯|©|×/DÑ/DÀQÓ/GÑGˆ
Ø× Ñ Ø×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r    c                 ó  • [         R                  " U R                  U R                  R                  5      [         R                  " U R
                  5      -   nUR                  U R                  U R                  -   U R                  -   5      $ rV   )	r   r   r?   r   Ú
diag_embedr@   rO   r`   ra   )rD   Úcovariance_matrixs     r   rm   Ú+LowRankMultivariateNormal.covariance_matrix¼   su   € ä!ŸLšLØ×*Ñ*¨D×,JÑ,J×,MÑ,Mó
ä×Ò˜T×9Ñ9Ó:ñ;Ðð !×'Ñ'Ø×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r    c                 óž  • U R                   R                  U R                  R                  S5      -  n[        R
                  R                  U R                  USS9n[        R                  " U R                  R                  5       5      UR                  U-  -
  nUR                  U R                  U R                  -   U R                  -   5      $ )Nr   F)Úupper)r?   r   r@   r   r   r   Úsolve_triangularrA   rl   Ú
reciprocalrO   r`   ra   )rD   r   ÚAÚprecision_matrixs       r   rt   Ú*LowRankMultivariateNormal.precision_matrixÅ   s¹   € ð ×*Ñ*×-Ñ-Ø×*Ñ*×4Ñ4°RÓ8ñ9ð 	ô �L‰L×)Ñ)¨$×*@Ñ*@À'ÐQVÐ)ÐWˆä×Ò˜T×9Ñ9×DÑDÓFÓGÈ!Ï$É$ÐQRÉ(ÑRð 	ð  ×&Ñ&Ø×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r    Úsample_shapec                 ó¬  • U R                  U5      nUS S U R                  R                  SS  -   n[        X0R                  R
                  U R                  R                  S9n[        X R                  R
                  U R                  R                  S9nU R                  [        U R                  U5      -   U R                  R                  5       U-  -   $ )Nr   )ÚdtypeÚdevice)Ú_extended_shaper4   r<   r   r3   rx   ry   r   r?   r@   re   )rD   rv   r<   ÚW_shapeÚeps_WÚeps_Ds         r   ÚrsampleÚ!LowRankMultivariateNormal.rsampleÖ   s¨   € Ø×$Ñ$ \Ó2ˆØ˜˜�*˜tŸ™×4Ñ4°R°SÐ9Ñ9ˆÜ  ·±·±ÀtÇxÁxÇÁÑWˆÜ  ¯h©h¯n©nÀTÇXÁXÇ_Á_ÑUˆà�H‰HÜ˜×6Ñ6¸Ó>ñ?à×*Ñ*×/Ñ/Ó1°EÑ9ñ:ð	
r    c                 ó�  • U R                   (       a  U R                  U5        XR                  -
  n[        U R                  U R
                  UU R                  5      n[        U R                  U R
                  U R                  5      nSU R                  S   [        R                  " S[        R                  -  5      -  U-   U-   -  $ )Ng      à¿r   r"   )rP   Ú_validate_sampler3   r0   r?   r@   rA   r)   ra   Úmathr&   Úpi)rD   ÚvalueÚdiffÚMÚlog_dets        r   Úlog_probÚ"LowRankMultivariateNormal.log_probá   s­   € Ø××Ø×!Ñ! %Ô(Ø—x‘xÑˆÜ&Ø×*Ñ*Ø×(Ñ(ØØ×"Ñ"ó	
ˆô (Ø×*Ñ*Ø×(Ñ(Ø×"Ñ"ó
ˆð
 �t×(Ñ(¨Ñ+¬d¯hªh°q¼4¿7¹7±{Ó.CÑCÀgÑMÐPQÑQÑRÐRr    c                 óD  • [        U R                  U R                  U R                  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      $ )Ng      à?r   g      ð?r"   )r)   r?   r@   rA   ra   r‚   r&   rƒ   Úlenr`   rO   )rD   r‡   ÚHs      r   ÚentropyÚ!LowRankMultivariateNormal.entropyò   s‹   € Ü'Ø×*Ñ*Ø×(Ñ(Ø×"Ñ"ó
ˆð
 �4×$Ñ$ QÑ'¨3´·²¸!¼d¿g¹g¹+Ó1FÑ+FÑGÈ'ÑQÑRˆÜˆt× Ñ Ó! QÓ&ØˆHà—8‘8˜D×-Ñ-Ó.Ð.r    )rA   r@   r?   r5   r4   r3   rV   )!Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Úreal_vectorÚindependentÚrealÚpositiveÚarg_constraintsÚsupportÚhas_rsampler   ÚboolrC   rO   ÚpropertyrY   r]   r	   rb   ri   rm   rt   r   rN   r
   r~   rˆ   r�   Ú__static_attributes__Ú__classcell__)rJ   s   @r   r   r   6   sy  ø† ñðF ×&Ñ&Ø!×-Ò-¨k×.>Ñ.>ÀÓBØ×+Ò+¨K×,@Ñ,@À!ÓDñ€Oð
 ×%Ñ%€GØ€Kð &*ñ*Pàð*Pð ð*Pð ð	*Pð
 ˜d‘{ð*Pð 
÷*Pð *P÷Xð  ð�fó ó ðð ð�fó ó ðð ð8˜&ó 8ó ð8ð
 ð
˜Fó 
ó ð
ð  ð
 6ó 
ó ð
ð ð
 &ó 
ó ð
ð  -2¯JªJ«Lñ 	
 Eð 	
¸Võ 	
òS÷"
/ð 
/r    )r‚   r   r   Útorch.distributionsr   Ú torch.distributions.distributionr   Ú'torch.distributions.multivariate_normalr   r   Útorch.distributions.utilsr   r	   Útorch.typesr
   Ú__all__r   r)   r0   r   © r    r   Ú<module>r¦      sD   ðã ã Ý Ý +Ý 9ß Qß EÝ ð 'Ð
'€ò	$ò	ò1ôF/ õ F/r    