§
    ŠŠtj–  ã                   óp   — d Z ddlZddlZddlmZ ddlmZmZ ddlmZ ddl	m
Z
 dgZ G d„ de¦  «        ZdS )	zÍ
This closely follows the implementation in NumPyro (https://github.com/pyro-ppl/numpyro).

Original copyright notice:

# Copyright: Contributors to the Pyro project.
# SPDX-License-Identifier: Apache-2.0
é    N)ÚTensor)ÚBetaÚconstraints)ÚDistribution)Úbroadcast_allÚLKJCholeskyc            	       óœ   ‡ — e Zd ZdZdej        iZej        Z	 	 dde	de
ez  dedz  ddfˆ fd„Zdˆ fd	„	Z ej        ¦   «         fd
„Zd„ Zˆ xZS )r   a)  
    LKJ distribution for lower Cholesky factor of correlation matrices.
    The distribution is controlled by ``concentration`` parameter :math:`\eta`
    to make the probability of the correlation matrix :math:`M` generated from
    a Cholesky factor proportional to :math:`\det(M)^{\eta - 1}`. Because of that,
    when ``concentration == 1``, we have a uniform distribution over Cholesky
    factors of correlation matrices::

        L ~ LKJCholesky(dim, concentration)
        X = L @ L' ~ LKJCorr(dim, concentration)

    Note that this distribution samples the
    Cholesky factor of correlation matrices and not the correlation matrices
    themselves and thereby differs slightly from the derivations in [1] for
    the `LKJCorr` distribution. For sampling, this uses the Onion method from
    [1] Section 3.

    Example::

        >>> # xdoctest: +IGNORE_WANT("non-deterministic")
        >>> l = LKJCholesky(3, 0.5)
        >>> l.sample()  # l @ l.T is a sample of a correlation 3x3 matrix
        tensor([[ 1.0000,  0.0000,  0.0000],
                [ 0.3516,  0.9361,  0.0000],
                [-0.1899,  0.4748,  0.8593]])

    Args:
        dimension (dim): dimension of the matrices
        concentration (float or Tensor): concentration/shape parameter of the
            distribution (often referred to as eta)

    **References**

    [1] `Generating random correlation matrices based on vines and extended onion method` (2009),
    Daniel Lewandowski, Dorota Kurowicka, Harry Joe.
    Journal of Multivariate Analysis. 100. 10.1016/j.jmva.2009.04.008
    Úconcentrationç      ð?NÚdimÚvalidate_argsÚreturnc                 ód  •— |dk     rt          d|› d�¦  «        ‚|| _        t          |¦  «        \  | _        | j                             ¦   «         }t          j        ||f¦  «        }| j        d| j        dz
  z  z   }t          j        | j        dz
  | j        j        | j        j	        ¬¦  «        }t          j
        |                     d¦  «        |g¦  «        }|dz   }|                     d¦  «        d|z  z
  }	t          ||	¦  «        | _        t          ¦   «                              |||¦  «         d S )	Né   zDExpected dim to be an integer greater than or equal to 2. Found dim=ú.ç      à?é   ©ÚdtypeÚdevice)r   éÿÿÿÿ)Ú
ValueErrorr   r   r
   ÚsizeÚtorchÚSizeÚaranger   r   ÚcatÚ	new_zerosÚ	unsqueezer   Ú_betaÚsuperÚ__init__)Úselfr   r
   r   Úbatch_shapeÚevent_shapeÚmarginal_concÚoffsetÚ
beta_conc1Ú
beta_conc0Ú	__class__s             €ú^/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributions/lkj_cholesky.pyr"   zLKJCholesky.__init__B   s7  ø€ ð �Š7ˆ7ÝØ]ÐWZÐ]Ð]Ð]ñô ð ð ˆŒÝ -¨mÑ <Ô <ÑˆÔ	ØÔ(×-Ò-Ñ/Ô/ˆÝ”j # s Ñ,Ô,ˆàÔ*¨S°D´H¸q±LÑ-AÑAˆÝ”ØŒH�q‰LØÔ$Ô*ØÔ%Ô,ð
ñ 
ô 
ˆõ
 ”˜F×,Ò,¨TÑ2Ô2°FÐ;Ñ<Ô<ˆØ˜c‘\ˆ
Ø"×,Ò,¨RÑ0Ô0°3¸±<Ñ?ˆ
Ý˜* jÑ1Ô1ˆŒ
Ý‰Œ×Ò˜ k°=ÑAÔAÐAÐAÐAó    c                 ó„  •— |                       t          |¦  «        }t          j        |¦  «        }| j        |_        | j                             |¦  «        |_        | j                             || j        fz   ¦  «        |_        t          t          |¦  «         	                    || j
        d¬¦  «         | j        |_        |S )NF)r   )Ú_get_checked_instancer   r   r   r   r
   Úexpandr    r!   r"   r%   Ú_validate_args)r#   r$   Ú	_instanceÚnewr*   s       €r+   r/   zLKJCholesky.expand]   s¬   ø€ Ø×(Ò(­°iÑ@Ô@ˆÝ”j Ñ-Ô-ˆØ”(ˆŒØ Ô.×5Ò5°kÑBÔBˆÔØ”J×%Ò% k°T´X°KÑ&?Ñ@Ô@ˆŒ	Ý�k˜3ÑÔ×(Ò(Ø˜Ô)¸ð 	)ñ 	
ô 	
ð 	
ð "Ô0ˆÔØˆ
r,   c                 ó~  — | j                              |¦  «                             d¦  «        }t          j        |                      |¦  «        |j        |j        ¬¦  «                             d¦  «        }|| 	                    dd¬¦  «        z  }|ddd d …f          
                    d¦  «         t          j        |¦  «        |z  }t          j        |j        ¦  «        j        }t          j        dt          j        |d	z  d¬
¦  «        z
  |¬¦  «                             ¦   «         }|t          j        |¦  «        z  }|S )Nr   r   T)r   Úkeepdim.r   g        r   r   ©r   )Úmin)r    Úsampler   r   ÚrandnÚ_extended_shaper   r   ÚtrilÚnormÚfill_ÚsqrtÚfinfoÚtinyÚclampÚsumÚ
diag_embed)r#   Úsample_shapeÚyÚu_normalÚu_hypersphereÚwÚepsÚ
diag_elemss           r+   r7   zLKJCholesky.samplei   s  € ð ŒJ×Ò˜lÑ+Ô+×5Ò5°bÑ9Ô9ˆÝ”;Ø× Ò  Ñ.Ô.°a´gÀaÄhð
ñ 
ô 
ç
Š$ˆr‰(Œ(ð 	ð ! 8§=¢=°RÀ =Ñ#FÔ#FÑFˆà�c˜1˜a˜a˜a�iÔ ×&Ò& sÑ+Ô+Ð+ÝŒJ�q‰MŒM˜MÑ)ˆåŒk˜!œ'Ñ"Ô"Ô'ˆÝ”[ ¥U¤Y¨q°!©t¸Ð%<Ñ%<Ô%<Ñ!<À#ÐFÑFÔF×KÒKÑMÔMˆ
Ø	�UÔ˜jÑ)Ô)Ñ)ˆØˆr,   c                 óh  — | j         r|                      |¦  «         |                     dd¬¦  «        ddd …f         }t          j        d| j        dz   | j        j        ¬¦  «        }d| j        dz
                       d¦  «        z  | j        z   |z
  }t          j	        || 
                    ¦   «         z  d¬¦  «        }| j        dz
  }| j        d	|z  z   }t          j        |¦  «        |z  }t          j        |d	z
  |¦  «        }d	|z  t          j
        t          j        ¦  «        z  }	|	|z   |z
  }
||
z
  S )
Nr   éþÿÿÿ)Údim1Údim2.r   r   )r   r5   r   )r0   Ú_validate_sampleÚdiagonalr   r   r   r
   r   r   rA   ÚlogÚlgammaÚmvlgammaÚmathÚpi)r#   ÚvaluerI   ÚorderÚunnormalized_log_pdfÚdm1ÚalphaÚdenominatorÚ	numeratorÚpi_constantÚnormalize_terms              r+   Úlog_probzLKJCholesky.log_prob~   s2  € ð Ôð 	)Ø×!Ò! %Ñ(Ô(Ð(Ø—^’^¨°"�^Ñ5Ô5°c¸1¸2¸2°gÔ>ˆ
Ý”˜Q ¤¨1¡°TÔ5GÔ5NÐOÑOÔOˆØ�TÔ'¨!Ñ+×6Ò6°rÑ:Ô:Ñ:¸T¼XÑEÈÑMˆÝ$œy¨°·²Ñ1AÔ1AÑ)AÀrÐJÑJÔJÐàŒh˜‰lˆØÔ" S¨3¡YÑ.ˆÝ”l 5Ñ)Ô)¨CÑ/ˆÝ”N 5¨3¡;°Ñ4Ô4ˆ	ð ˜C‘i¥$¤(­4¬7Ñ"3Ô"3Ñ3ˆØ$ yÑ0°;Ñ>ˆØ# nÑ4Ð4r,   )r   N)N)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚpositiveÚarg_constraintsÚcorr_choleskyÚsupportÚintr   ÚfloatÚboolr"   r/   r   r   r7   r^   Ú__classcell__)r*   s   @r+   r   r      só   ø€ € € € € ð$ð $ðN '¨Ô(<Ð=€OØÔ'€Gð
 ),Ø%)ð	Bð BàðBð  ‘~ðBð ˜d‘{ð	Bð
 
ðBð Bð Bð Bð Bð Bð6
ð 
ð 
ð 
ð 
ð 
ð #- %¤*¡,¤,ð ð ð ð ð*5ð 5ð 5ð 5ð 5ð 5ð 5r,   )rb   rS   r   r   Útorch.distributionsr   r   Ú torch.distributions.distributionr   Útorch.distributions.utilsr   Ú__all__r   © r,   r+   ú<module>rp      s¸   ððð ð €€€à €€€Ø Ð Ð Ð Ð Ð Ø 1Ð 1Ð 1Ð 1Ð 1Ð 1Ð 1Ð 1Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø 3Ð 3Ð 3Ð 3Ð 3Ð 3ð ˆ/€ðA5ð A5ð A5ð A5ð A5�,ñ A5ô A5ð A5ð A5ð A5r,   