ó
    Eñià  ã                   óŒ   • S SK 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	  S SK
Jr  S/rS	 r " S
 S\5      r " S S\	5      rg)é    N)ÚTensor)ÚFunction)Úonce_differentiable)Úconstraints)ÚExponentialFamily)Ú_sizeÚ	Dirichletc                 ó¤   • UR                  SS5      R                  U5      n[        R                  " XU5      nXBX-  R                  SS5      -
  -  $ ©NéÿÿÿÿT)ÚsumÚ	expand_asÚtorchÚ_dirichlet_grad)ÚxÚconcentrationÚgrad_outputÚtotalÚgrads        ÚZ/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/dirichlet.pyÚ_Dirichlet_backwardr      sN   € Ø×Ñ˜b $Ó'×1Ñ1°-Ó@€EÜ× Ò  °5Ó9€DØ !¡/×!6Ñ!6°r¸4Ó!@Ñ@ÑAÐAó    c                   ó>   • \ rS rSr\S 5       r\\S 5       5       rSrg)Ú
_Dirichleté   c                 óT   • [         R                  " U5      nU R                  X!5        U$ ©N)r   Ú_sample_dirichletÚsave_for_backward)Úctxr   r   s      r   ÚforwardÚ_Dirichlet.forward   s'   € ô ×#Ò# MÓ2ˆØ×Ñ˜aÔ/Øˆr   c                 ó6   • U R                   u  p#[        X#U5      $ r   )Úsaved_tensorsr   )r    r   r   r   s       r   ÚbackwardÚ_Dirichlet.backward   s   € ð ×,Ñ,ÑˆÜ" 1°[ÓAÐAr   © N)	Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ústaticmethodr!   r   r%   Ú__static_attributes__r'   r   r   r   r      s2   † Øñó ðð
 ØñBó ó óBr   r   c                   ó@  ^ • \ rS rSrSrS\R                  " \R                  S5      0r\R                  r
Sr SS\S\S-  SS4U 4S	 jjjrSU 4S
 jjrSS\S\4S jjrS r\S\4S j5       r\S\4S j5       r\S\4S j5       rS r\S\\   4S j5       rS rSrU =r$ )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   é   TNÚvalidate_argsÚreturnc                 ó¦   >• UR                  5       S:  a  [        S5      eXl        UR                  S S UR                  SS  pC[        TU ]  X4US9  g )Nr0   z;`concentration` parameter must be at least one-dimensional.r   ©r1   )ÚdimÚ
ValueErrorr   ÚshapeÚsuperÚ__init__)Úselfr   r1   Úbatch_shapeÚevent_shapeÚ	__class__s        €r   r9   ÚDirichlet.__init__=   sc   ø€ ð
 ×ÑÓ Ó"ÜØMóð ð +ÔØ#0×#6Ñ#6°s¸Ð#;¸]×=PÑ=PÐQSÐQTÐ=U�[ä‰Ñ˜ÀÐÒOr   c                 ó  >• U R                  [        U5      n[        R                  " U5      nU R                  R                  XR                  -   5      Ul        [        [        U]#  XR                  SS9  U R                  Ul	        U$ )NFr4   )
Ú_get_checked_instancer	   r   ÚSizer   Úexpandr<   r8   r9   Ú_validate_args)r:   r;   Ú	_instanceÚnewr=   s       €r   rB   ÚDirichlet.expandK   sy   ø€ Ø×(Ñ(¬°IÓ>ˆÜ—j’j Ó-ˆØ ×.Ñ.×5Ñ5°k×DTÑDTÑ6TÓUˆÔÜŒi˜Ñ&Ø×)Ñ)¸ð 	'ñ 	
ð "×0Ñ0ˆÔØˆ
r   Úsample_shapec                 ó„   • U R                  U5      nU R                  R                  U5      n[        R	                  U5      $ r   )Ú_extended_shaper   rB   r   Úapply)r:   rG   r7   r   s       r   ÚrsampleÚDirichlet.rsampleU   s9   € Ø×$Ñ$ \Ó2ˆØ×*Ñ*×1Ñ1°%Ó8ˆÜ×Ñ Ó.Ð.r   c                 ól  • U R                   (       a  U R                  U5        [        R                  " U R                  S-
  U5      R                  S5      [        R                  " U R                  R                  S5      5      -   [        R                  " U R                  5      R                  S5      -
  $ )Nç      ð?r   )rC   Ú_validate_sampler   Úxlogyr   r   Úlgamma)r:   Úvalues     r   Úlog_probÚDirichlet.log_probZ   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   • U R                   U R                   R                  SS5      -  $ r   )r   r   ©r:   s    r   ÚmeanÚDirichlet.meanc   s&   € à×!Ñ! D×$6Ñ$6×$:Ñ$:¸2¸tÓ$DÑDÐDr   c                 óL  • U R                   S-
  R                  SS9nXR                  SS5      -  nU R                   S:  R                  SS9n[        R
                  R                  R                  X#   R                  SS9UR                  S   5      R                  U5      X#'   U$ )Nr0   g        )Úminr   T)r5   )r   Úclampr   Úallr   ÚnnÚ
functionalÚone_hotÚargmaxr7   Úto)r:   Úconcentrationm1ÚmodeÚmasks       r   rc   ÚDirichlet.modeg   s¢   € à×-Ñ-°Ñ1×8Ñ8¸SÐ8ÐAˆØ×!4Ñ!4°R¸Ó!>Ñ>ˆØ×"Ñ" QÑ&×+Ñ+°Ð+Ð3ˆÜ—X‘X×(Ñ(×0Ñ0Ø‰J×Ñ "ÐÐ% ×'<Ñ'<¸RÑ'@ó
ç
‰"ˆT‹(ð 	‰
ð ˆr   c                 ó    • U R                   R                  SS5      nU R                   XR                   -
  -  UR                  S5      US-   -  -  $ )Nr   Té   r0   )r   r   Úpow)r:   Úcon0s     r   ÚvarianceÚDirichlet.varianceq   sR   € à×!Ñ!×%Ñ% b¨$Ó/ˆà×ÑØ×(Ñ(Ñ(ñ*à�x‰x˜‹{˜d Q™hÑ'ñ)ð	
r   c                 ó²  • U R                   R                  S5      nU R                   R                  S5      n[        R                  " U R                   5      R                  S5      [        R                  " U5      -
  X-
  [        R
                  " U5      -  -
  U R                   S-
  [        R
                  " U R                   5      -  R                  S5      -
  $ )Nr   rN   )r   Úsizer   r   rQ   Údigamma)r:   ÚkÚa0s      r   ÚentropyÚDirichlet.entropyz   s¯   € Ø×Ñ×#Ñ# BÓ'ˆØ×Ñ×#Ñ# BÓ'ˆä�LŠL˜×+Ñ+Ó,×0Ñ0°Ó4Ü�lŠl˜2Óñà‰vœŸš rÓ*Ñ*ñ+ð ×"Ñ" SÑ(¬E¯MªM¸$×:LÑ:LÓ,MÑM×RÑRÐSUÓVñWð	
r   c                 ó   • U R                   4$ r   ©r   rV   s    r   Ú_natural_paramsÚDirichlet._natural_params„   s   € à×"Ñ"Ð$Ð$r   c                 óŒ   • UR                  5       R                  S5      [        R                   " UR                  S5      5      -
  $ )Nr   )rQ   r   r   )r:   r   s     r   Ú_log_normalizerÚDirichlet._log_normalizer‰   s-   € Ø�x‰x‹z�~‰~˜bÓ!¤E§L¢L°·±°r³Ó$;Ñ;Ð;r   rt   r   )r'   )r(   r)   r*   r+   Ú__doc__r   ÚindependentÚpositiveÚarg_constraintsÚsimplexÚsupportÚhas_rsampler   Úboolr9   rB   r   rK   rS   ÚpropertyrW   rc   rj   rq   Útupleru   rx   r-   Ú__classcell__)r=   s   @r   r	   r	   &   s  ø† ñð" 	˜×0Ò0°×1EÑ1EÀqÓIð€Oð ×!Ñ!€GØ€Kð
 &*ñPàðPð ˜d‘{ðPð 
÷	Pð 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>r‹      sH   ðó Ý Ý #Ý 7Ý +Ý <Ý ð ˆ-€òBôB�ô Bô d<Ð!õ d<r   