Ë
    f^(hß  ã                   óŒ   — d dl Z d dl mZ d dlmZ d dlmZ d dlmZ d dlm	Z	 d dl
mZ dgZd	„ Z G d
„ de«      Z G d„ de	«      Zy)é    N)ÚTensor)ÚFunction)Úonce_differentiable)Úconstraints)ÚExponentialFamily)Ú_sizeÚ	Dirichletc                 ó¨   — |j                  dd«      j                  |«      }t        j                  | ||«      }||| |z  j                  dd«      z
  z  S ©NéÿÿÿÿT)ÚsumÚ	expand_asÚtorchÚ_dirichlet_grad)ÚxÚconcentrationÚgrad_outputÚtotalÚgrads        ú[/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributions/dirichlet.pyÚ_Dirichlet_backwardr      sT   € Ø×Ñ˜b $Ó'×1Ñ1°-Ó@€EÜ× Ñ   M°5Ó9€DØ�; ! k¡/×!6Ñ!6°r¸4Ó!@Ñ@ÑAÐAó    c                   ó6   — e Zd Zed„ «       Zeed„ «       «       Zy)Ú
_Dirichletc                 óT   — t        j                  |«      }| j                  ||«       |S ©N)r   Ú_sample_dirichletÚsave_for_backward)Úctxr   r   s      r   Úforwardz_Dirichlet.forward   s'   € ä×#Ñ# MÓ2ˆØ×Ñ˜a Ô/Øˆr   c                 ó:   — | j                   \  }}t        |||«      S r   )Úsaved_tensorsr   )r   r   r   r   s       r   Úbackwardz_Dirichlet.backward   s#   € ð ×,Ñ,Ñˆˆ=Ü" 1 m°[ÓAÐAr   N)Ú__name__Ú
__module__Ú__qualname__Ústaticmethodr    r   r#   © r   r   r   r      s2   „ Øñó ðð
 ØñBó ó ñBr   r   c                   ó  ‡ — e Zd ZdZd ej
                  ej                  d«      iZej                  Z	dZ
dˆ fd„	Zdˆ fd„	Zddedefd	„Zd
„ Zedefd„«       Zedefd„«       Zedefd„«       Zd„ Zedee   fd„«       Zd„ Zˆ xZS )r	   aÏ  
    Creates a Dirichlet distribution parameterized by concentration :attr:`concentration`.

    Example::

        >>> # xdoctest: +IGNORE_WANT("non-deterministic")
        >>> m = Dirichlet(torch.tensor([0.5, 0.5]))
        >>> m.sample()  # Dirichlet distributed with concentration [0.5, 0.5]
        tensor([ 0.1046,  0.8954])

    Args:
        concentration (Tensor): concentration parameter of the distribution
            (often referred to as alpha)
    r   é   Tc                 ó°   •— |j                  «       dk  rt        d«      ‚|| _        |j                  d d |j                  dd  }}t        ‰| �  |||¬«       y )Nr*   z;`concentration` parameter must be at least one-dimensional.r   ©Úvalidate_args)ÚdimÚ
ValueErrorr   ÚshapeÚsuperÚ__init__)Úselfr   r-   Úbatch_shapeÚevent_shapeÚ	__class__s        €r   r2   zDirichlet.__init__9   sg   ø€ Ø×ÑÓ Ò"ÜØMóð ð +ˆÔØ#0×#6Ñ#6°s¸Ð#;¸]×=PÑ=PÐQSÐQTÐ=U�[ˆÜ‰Ñ˜ kÀÐÕOr   c                 ó  •— | j                  t        |«      }t        j                  |«      }| j                  j                  || j                  z   «      |_        t        t        |�#  || j                  d¬«       | j                  |_	        |S )NFr,   )
Ú_get_checked_instancer	   r   ÚSizer   Úexpandr5   r1   r2   Ú_validate_args)r3   r4   Ú	_instanceÚnewr6   s       €r   r:   zDirichlet.expandB   s}   ø€ Ø×(Ñ(¬°IÓ>ˆÜ—j‘j Ó-ˆØ ×.Ñ.×5Ñ5°kÀD×DTÑDTÑ6TÓUˆÔÜŒi˜Ñ&Ø˜×)Ñ)¸ð 	'ô 	
ð "×0Ñ0ˆÔØˆ
r   Úsample_shapeÚreturnc                 ó„   — | j                  |«      }| j                  j                  |«      }t        j	                  |«      S r   )Ú_extended_shaper   r:   r   Úapply)r3   r>   r0   r   s       r   ÚrsamplezDirichlet.rsampleL   s9   € Ø×$Ñ$ \Ó2ˆØ×*Ñ*×1Ñ1°%Ó8ˆÜ×Ñ Ó.Ð.r   c                 ó\  — | j                   r| j                  |«       t        j                  | j                  dz
  |«      j                  d«      t        j                  | j                  j                  d«      «      z   t        j                  | j                  «      j                  d«      z
  S )Nç      ð?r   )r;   Ú_validate_sampler   Úxlogyr   r   Úlgamma)r3   Úvalues     r   Úlog_probzDirichlet.log_probQ   s†   € Ø×ÒØ×!Ñ! %Ô(ä�K‰K˜×*Ñ*¨SÑ0°%Ó8×<Ñ<¸RÓ@Ü�l‰l˜4×-Ñ-×1Ñ1°"Ó5Ó6ñ7ä�l‰l˜4×-Ñ-Ó.×2Ñ2°2Ó6ñ7ð	
r   c                 óT   — | j                   | j                   j                  dd«      z  S r   )r   r   ©r3   s    r   ÚmeanzDirichlet.meanZ   s&   € à×!Ñ! D×$6Ñ$6×$:Ñ$:¸2¸tÓ$DÑDÐDr   c                 ód  — | j                   dz
  j                  d¬«      }||j                  dd«      z  }| j                   dk  j                  d¬«      }t        j
                  j                  j                  ||   j                  d¬«      |j                  d   «      j                  |«      ||<   |S )Nr*   g        )Úminr   T)r.   )r   Úclampr   Úallr   ÚnnÚ
functionalÚone_hotÚargmaxr0   Úto)r3   Úconcentrationm1ÚmodeÚmasks       r   rX   zDirichlet.mode^   s¨   € à×-Ñ-°Ñ1×8Ñ8¸SÐ8ÓAˆØ ×!4Ñ!4°R¸Ó!>Ñ>ˆØ×"Ñ" QÑ&×+Ñ+°Ð+Ó3ˆÜ—X‘X×(Ñ(×0Ñ0Ø�‰J×Ñ "ÐÓ% ×'<Ñ'<¸RÑ'@ó
ç
‰"ˆT‹(ð 	ˆT‰
ð ˆr   c                 ó¢   — | j                   j                  dd«      }| j                   || j                   z
  z  |j                  d«      |dz   z  z  S )Nr   Té   r*   )r   r   Úpow)r3   Úcon0s     r   ÚvariancezDirichlet.varianceh   sT   € à×!Ñ!×%Ñ% b¨$Ó/ˆà×ÑØ�d×(Ñ(Ñ(ñ*à�x‰x˜‹{˜d Q™hÑ'ñ)ð	
r   c                 ó¬  — | j                   j                  d«      }| j                   j                  d«      }t        j                  | j                   «      j                  d«      t        j                  |«      z
  ||z
  t        j
                  |«      z  z
  | j                   dz
  t        j
                  | j                   «      z  j                  d«      z
  S )Nr   rE   )r   Úsizer   r   rH   Údigamma)r3   ÚkÚa0s      r   ÚentropyzDirichlet.entropyq   s±   € Ø×Ñ×#Ñ# BÓ'ˆØ×Ñ×#Ñ# BÓ'ˆä�L‰L˜×+Ñ+Ó,×0Ñ0°Ó4Ü�l‰l˜2Óñà�2‰vœŸ™ rÓ*Ñ*ñ+ð ×"Ñ" SÑ(¬E¯M©M¸$×:LÑ:LÓ,MÑM×RÑRÐSUÓVñWð	
r   c                 ó   — | j                   fS r   )r   rL   s    r   Ú_natural_paramszDirichlet._natural_params{   s   € à×"Ñ"Ð$Ð$r   c                 óŠ   — |j                  «       j                  d«      t        j                   |j                  d«      «      z
  S )Nr   )rH   r   r   )r3   r   s     r   Ú_log_normalizerzDirichlet._log_normalizer   s-   € Ø�x‰x‹z�~‰~˜bÓ!¤E§L¡L°·±°r³Ó$;Ñ;Ð;r   r   )r(   )r$   r%   r&   Ú__doc__r   ÚindependentÚpositiveÚarg_constraintsÚsimplexÚsupportÚhas_rsampler2   r:   r   r   rC   rJ   ÚpropertyrM   rX   r^   rd   Útuplerf   rh   Ú__classcell__)r6   s   @r   r	   r	   #   sÞ   ø„ ñð  	Ð0˜×0Ñ0°×1EÑ1EÀqÓIð€Oð ×!Ñ!€GØ€KõPõñ/ Eð /°6ó /ò

ð ðE�fò Eó ðEð ð�fò ó ðð ð
˜&ò 
ó ð
ò
ð ð%  v¡ò %ó ð%ö<r   )r   r   Útorch.autogradr   Útorch.autograd.functionr   Útorch.distributionsr   Útorch.distributions.exp_familyr   Útorch.typesr   Ú__all__r   r   r	   r(   r   r   ú<module>ry      sF   ðã Ý Ý #Ý 7Ý +Ý <Ý ð ˆ-€òBôB�ô Bô]<Ð!õ ]<r   