Ë
    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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)Ú_batch_mahalanobisÚ	_batch_mv)Ú_standard_normalÚlazy_property)Ú_sizeÚLowRankMultivariateNormalc                 ó:  — | j                  d«      }| j                  |j                  d«      z  }t        j                  || «      j                  «       }|j                  d||z  «      dd…dd|dz   …fxx   dz  cc<   t        j                  j                  |«      S )zƒ
    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        úm/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributions/lowrank_multivariate_normal.pyÚ_batch_capacitance_trilr      s€   € ð
 	
�‰ˆr‹
€AØ�d‰d�Q—[‘[ “_Ñ$€GÜ�‰�W˜aÓ ×+Ñ+Ó-€AØ‡F�Fˆ2ˆq�1‰uÓ’a™˜A ™E˜�kÓ" aÑ'Ó"Ü�<‰<× Ñ  Ó#Ð#ó    c                 ó¨   — d|j                  dd¬«      j                  «       j                  d«      z  |j                  «       j                  d«      z   S )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                 ó¾   — | j                   |j                  d«      z  }t        ||«      }|j                  d«      |z  j	                  d«      }t        ||«      }||z
  S )a  
    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   (   s]   € ð �d‰d�Q—[‘[ “_Ñ$€GÜ˜' 1Ó%€IØŸ™˜q› A™×*Ñ*¨2Ó.ÐÜ*Ð+;¸YÓGÐØÐ0Ñ0Ð0r    c                   óš  ‡ — e Zd ZdZej
                   ej                  ej                  d«       ej                  ej                  d«      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j0                  «       fdedefd„Zd„ Zd„ Zˆ xZS )r   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_diagTc           	      óî  •— |j                  «       dk  rt        d«      ‚|j                  dd  }|j                  «       dk  rt        d«      ‚|j                  dd |k7  rt        d|d   › d	�«      ‚|j                  dd  |k7  rt        d
|› �«      ‚|j                  d«      }|j                  d«      }	 t	        j
                  |||«      \  }| _        }|d   | _        |d   | _	        | j                  j                  d d }	|| _
        || _        t        ||«      | _        t        ‰
| �=  |	||¬«       y # t        $ r8}t        d|j                  › d|j                  › d|j                  › �«      |‚d }~ww xY w)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   ©Úvalidate_args)ÚdimÚ
ValueErrorÚshaper   r   Úbroadcast_tensorsr3   ÚRuntimeErrorr2   r4   Ú_unbroadcasted_cov_factorÚ_unbroadcasted_cov_diagr   Ú_capacitance_trilÚsuperÚ__init__)Úselfr2   r3   r4   r7   Úevent_shapeÚloc_Ú	cov_diag_ÚeÚbatch_shapeÚ	__class__s             €r   rA   z"LowRankMultivariateNormal.__init__`   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Ø�j )ó0Ñ,ˆD�$”/ 9ð ˜‘<ˆŒØ! &Ñ)ˆŒØ—h‘h—n‘n S bÐ)ˆà)3ˆÔ&Ø'/ˆÔ$Ü!8¸ÀXÓ!NˆÔÜ‰Ñ˜ kÀÐÕOøô ò 	ÜØ1°#·)±)°¸MÈ*×JZÑJZÐI[Ð[fÐgo×guÑguÐfvÐwóàðûð	ús   Â4 D3 Ä3	E4Ä<3E/Å/E4c                 ó8  •— | j                  t        |«      }t        j                  |«      }|| j                  z   }| j
                  j                  |«      |_        | j                  j                  |«      |_        | j                  j                  || j                  j                  dd  z   «      |_        | j                  |_
        | j                  |_        | j                  |_        t        t        |�;  || j                  d¬«       | j                  |_        |S )Nr   Fr6   )Ú_get_checked_instancer   r   ÚSizerC   r2   Úexpandr4   r3   r:   r=   r>   r?   r@   rA   Ú_validate_args)rB   rG   Ú	_instanceÚnewÚ	loc_shaperH   s        €r   rL   z LowRankMultivariateNormal.expand…   sí   ø€ Ø×(Ñ(Ô)BÀIÓNˆÜ—j‘j Ó-ˆØ $×"2Ñ"2Ñ2ˆ	Ø—(‘(—/‘/ )Ó,ˆŒØ—}‘}×+Ñ+¨IÓ6ˆŒØŸ™×/Ñ/°	¸D¿O¹O×<QÑ<QÐRTÐRUÐ<VÑ0VÓWˆŒØ(,×(FÑ(FˆÔ%Ø&*×&BÑ&BˆÔ#Ø $× 6Ñ 6ˆÔÜÔ'¨Ñ6Ø˜×)Ñ)¸ð 	7ô 	
ð "×0Ñ0ˆÔØˆ
r    Úreturnc                 ó   — | j                   S ©N©r2   ©rB   s    r   ÚmeanzLowRankMultivariateNormal.mean•   ó   € à�x‰xˆr    c                 ó   — | j                   S rS   rT   rU   s    r   ÚmodezLowRankMultivariateNormal.mode™   rW   r    c                 ó¼   — | j                   j                  d«      j                  d«      | j                  z   j	                  | j
                  | j                  z   «      S )Nr"   r   )r=   r+   r'   r>   rL   Ú_batch_shapeÚ_event_shaperU   s    r   Úvariancez"LowRankMultivariateNormal.variance�   sN   € ð ×*Ñ*×.Ñ.¨qÓ1×5Ñ5°bÓ9¸D×<XÑ<XÑXß
‰&�×"Ñ" T×%6Ñ%6Ñ6Ó
7ð	8r    c                 óî  — | j                   d   }| j                  j                  «       j                  d«      }| j                  |z  }t        j                  ||j                  «      j                  «       }|j                  d||z  «      d d …d d |dz   …fxx   dz  cc<   |t
        j                  j                  |«      z  }|j                  | j                  | j                   z   | j                   z   «      S )Nr   r   r   )r\   r>   Úsqrtr   r=   r   r   r   r   r   r   r   rL   r[   )rB   ÚnÚcov_diag_sqrt_unsqueezeÚ
Dinvsqrt_Wr   Ú
scale_trils         r   rc   z$LowRankMultivariateNormal.scale_tril£   sÙ   € ð ×Ñ˜aÑ ˆØ"&×">Ñ">×"CÑ"CÓ"E×"OÑ"OÐPRÓ"SÐØ×3Ñ3Ð6MÑMˆ
Ü�L‰L˜ Z§]¡]Ó3×>Ñ>Ó@ˆØ	�‰ˆr�1�q‘5Óš!™X  A¡˜X˜+Ó&¨!Ñ+Ó&Ø,¬u¯|©|×/DÑ/DÀQÓ/GÑGˆ
Ø× Ñ Ø×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r    c                 ó  — t        j                  | j                  | j                  j                  «      t        j                  | j
                  «      z   }|j                  | j                  | j                  z   | j                  z   «      S rS   )	r   r   r=   r   Ú
diag_embedr>   rL   r[   r\   )rB   Úcovariance_matrixs     r   rf   z+LowRankMultivariateNormal.covariance_matrix´   su   € ä!ŸL™LØ×*Ñ*¨D×,JÑ,J×,MÑ,Mó
ä×Ñ˜T×9Ñ9Ó:ñ;Ðð !×'Ñ'Ø×Ñ × 1Ñ 1Ñ1°D×4EÑ4EÑEó
ð 	
r    c                 ó¢  — | j                   j                  | j                  j                  d«      z  }t        j
                  j                  | j                  |d¬«      }t	        j                  | j                  j                  «       «      |j                  |z  z
  }|j                  | j                  | j                  z   | j                  z   «      S )Nr   F)Úupper)r=   r   r>   r   r   r   Úsolve_triangularr?   re   Ú
reciprocalrL   r[   r\   )rB   r   ÚAÚprecision_matrixs       r   rl   z*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                 ó¼  — | j                  |«      }|d d | j                  j                  dd  z   }t        || j                  j
                  | j                  j                  ¬«      }t        || j                  j
                  | j                  j                  ¬«      }| j                  t        | j                  |«      z   | j                  j                  «       |z  z   S )Nr   )ÚdtypeÚdevice)Ú_extended_shaper3   r:   r   r2   ro   rp   r   r=   r>   r_   )rB   rm   r:   ÚW_shapeÚeps_WÚeps_Ds         r   Úrsamplez!LowRankMultivariateNormal.rsampleÎ   s¬   € Ø×$Ñ$ \Ó2ˆØ˜˜�*˜tŸ™×4Ñ4°R°SÐ9Ñ9ˆÜ  °·±·±ÀtÇxÁxÇÁÔWˆÜ  ¨d¯h©h¯n©nÀTÇXÁXÇ_Á_ÔUˆà�H‰HÜ˜×6Ñ6¸Ó>ñ?à×*Ñ*×/Ñ/Ó1°EÑ9ñ:ð	
r    c                 ó†  — | j                   r| j                  |«       || j                  z
  }t        | j                  | j
                  || j                  «      }t        | j                  | j
                  | j                  «      }d| j                  d   t        j                  dt        j                  z  «      z  |z   |z   z  S )Ng      à¿r   r"   )rM   Ú_validate_sampler2   r0   r=   r>   r?   r)   r\   Úmathr&   Úpi)rB   ÚvalueÚdiffÚMÚlog_dets        r   Úlog_probz"LowRankMultivariateNormal.log_probÙ   s¯   € Ø×ÒØ×!Ñ! %Ô(Ø�t—x‘xÑˆÜ&Ø×*Ñ*Ø×(Ñ(ØØ×"Ñ"ó	
ˆô (Ø×*Ñ*Ø×(Ñ(Ø×"Ñ"ó
ˆð
 �t×(Ñ(¨Ñ+¬d¯h©h°q¼4¿7¹7±{Ó.CÑCÀgÑMÐPQÑQÑRÐRr    c                 ó@  — t        | j                  | j                  | j                  «      }d| j                  d   dt        j                  dt
        j                  z  «      z   z  |z   z  }t        | j                  «      dk(  r|S |j                  | j                  «      S )Ng      à?r   g      ð?r"   )r)   r=   r>   r?   r\   rx   r&   ry   Úlenr[   rL   )rB   r}   ÚHs      r   Úentropyz!LowRankMultivariateNormal.entropyê   s‹   € Ü'Ø×*Ñ*Ø×(Ñ(Ø×"Ñ"ó
ˆð
 �4×$Ñ$ QÑ'¨3´·±¸!¼d¿g¹g¹+Ó1FÑ+FÑGÈ'ÑQÑRˆÜˆt× Ñ Ó! QÒ&ØˆHà—8‘8˜D×-Ñ-Ó.Ð.r    rS   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úreal_vectorÚindependentÚrealÚpositiveÚarg_constraintsÚsupportÚhas_rsamplerA   rL   Úpropertyr   rV   rY   r	   r]   rc   rf   rl   r   rK   r
   ru   r~   r‚   Ú__classcell__)rH   s   @r   r   r   6   s6  ø„ ñðD ×&Ñ&Ø-�k×-Ñ-¨k×.>Ñ.>ÀÓBØ+�K×+Ñ+¨K×,@Ñ,@À!ÓDñ€Oð
 ×%Ñ%€GØ€Kõ#PõJð  ð�fò ó ðð ð�fò ó ðð ð8˜&ò 8ó ð8ð
 ð
˜Fò 
ó ð
ð  ð
 6ò 
ó ð
ð ð
 &ò 
ó ð
ð  -7¨E¯J©J«Lñ 	
 Eð 	
¸Vó 	
òSö"
/r    )rx   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ô~/ õ ~/r    