+
    &j                         ^ RI t ^ RI Ht ^ RIHt ^ RIHt ^ RIHt ^ RIH	t	 ^ RI
Ht R.tR	 t ! R
 R]4      t ! R R]	4      tR# )    NTensor)Function)once_differentiable)constraints)ExponentialFamily)_size	Dirichletc                     VP                  RR4      P                  V4      p\        P                  ! WV4      pWBW,          P                  RR4      ,
          ,          #    T)sum	expand_astorch_dirichlet_grad)xconcentrationgrad_outputtotalgrads   &&&  u/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/distributions/dirichlet.py_Dirichlet_backwardr      sN    b$'11-@E  59D!/!6!6r4!@@AA    c                   J   a  ] tR t^t o ]R 4       t]]R 4       4       tRtV t	R# )
_Dirichletc                T    \         P                  ! V4      pV P                  W!4       V# N)r   _sample_dirichletsave_for_backward)ctxr   r   s   && r   forward_Dirichlet.forward   s'     ##M2a/r   c                6    V P                   w  r#\        W#V4      # 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____classdictcell__)__classdict__s   @r   r   r      s5      
 B  Br   r   c                   H  a a ] tR t^&t oRtR]P                  ! ]P                  ^4      /t]P                  t
RtRV3R lV 3R llltRV 3R lltRV3R lR lltR	 t]V3R
 lR l4       t]V3R lR l4       t]V3R lR l4       tR t]V3R lR l4       tR tRtVtV ;t# )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                8   < V ^8  d   QhRS[ RS[R,          RR/# )   r   validate_argsNreturn)r   bool)formatr0   s   "r   __annotate__Dirichlet.__annotate__=   s2     P PP d{P 
	Pr   c                   < VP                  4       ^8  d   \        R4      hWn        VP                  RR VP                  RR rC\        SV `  W4VR7       R# )r   z;`concentration` parameter must be at least one-dimensional.Nr4   r   )dim
ValueErrorr   shapesuper__init__)selfr   r4   batch_shapeevent_shape	__class__s   &&&  r   r@   Dirichlet.__init__=   sc    
 "M  +#0#6#6s#;]=P=PQSQT=U[Or   c                   < V P                  \        V4      p\        P                  ! V4      pV P                  P                  WP                  ,           4      Vn        \        \        V`#  WP                  R R7       V P                  Vn	        V# )Fr;   )
_get_checked_instancer
   r   Sizer   expandrC   r?   r@   _validate_args)rA   rB   	_instancenewrD   s   &&& r   rI   Dirichlet.expandK   sy    ((I>jj- ..55kDTDT6TUi&)) 	' 	
 "00
r   c                &   < V ^8  d   QhRS[ RS[/# )r3   sample_shaper5   )r	   r   )r7   r0   s   "r   r8   r9   U   s     / /E /6 /r   c                    V P                  V4      pV P                  P                  V4      p\        P	                  V4      # r   )_extended_shaper   rI   r   apply)rA   rO   r>   r   s   &&  r   rsampleDirichlet.rsampleU   s9    $$\2**11%8..r   c                   V P                   '       d   V P                  V4       \        P                  ! V P                  R ,
          V4      P                  R4      \        P                  ! V P                  P                  R4      4      ,           \        P                  ! V P                  4      P                  R4      ,
          # )      ?r   )rJ   _validate_sampler   xlogyr   r   lgamma)rA   values   &&r   log_probDirichlet.log_probZ   s    !!%(KK**S0%8<<R@ll4--11"567ll4--.22267	
r   c                    < V ^8  d   QhRS[ /# r3   r5   r   )r7   r0   s   "r   r8   r9   d   s     E Ef Er   c                \    V P                   V P                   P                  RR4      ,          # r   )r   r   rA   s   &r   meanDirichlet.meanc   s&    !!D$6$6$:$:2t$DDDr   c                    < V ^8  d   QhRS[ /# r^   r   )r7   r0   s   "r   r8   r9   h   s      f r   c                ~   V P                   ^,
          P                  RR7      pWP                  RR4      ,          pV P                   ^8  P                  RR7      p\        P
                  P                  P                  W#,          P                  RR7      VP                  R,          4      P                  V4      W#&   V# )r   g        )minT)r<   r   )r   clampr   allr   nn
functionalone_hotargmaxr>   to)rA   concentrationm1modemasks   &   r   rn   Dirichlet.modeg   s    --188S8A!4!4R!>>""Q&+++3XX((00J"%'<'<R'@

"T( 	
 r   c                    < V ^8  d   QhRS[ /# r^   r   )r7   r0   s   "r   r8   r9   r   s     
 
& 
r   c                    V P                   P                  RR4      pV P                   WP                   ,
          ,          VP                  ^4      V^,           ,          ,          # r   )r   r   pow)rA   con0s   & r   varianceDirichlet.varianceq   sR    !!%%b$/(((*xx{dQh')	
r   c                   V P                   P                  R4      pV P                   P                  R4      p\        P                  ! V P                   4      P                  R4      \        P                  ! V4      ,
          W,
          \        P
                  ! V4      ,          ,
          V P                   R,
          \        P
                  ! V P                   4      ,          P                  R4      ,
          # )r   rV   r   )r   sizer   r   rY   digamma)rA   ka0s   &  r   entropyDirichlet.entropyz   s    ##B'##B'LL++,004ll2vr**+ ""S(EMM$:L:L,MMRRSUVW	
r   c                0   < V ^8  d   QhRS[ S[,          /# r^   )tupler   )r7   r0   s   "r   r8   r9      s     % %v %r   c                    V P                   3# r   r   r`   s   &r   _natural_paramsDirichlet._natural_params   s    ""$$r   c                    VP                  4       P                  R4      \        P                   ! VP                  R4      4      ,
          # )r   r   )rY   r   r   )rA   r   s   &&r   _log_normalizerDirichlet._log_normalizer   s-    xxz~~b!ELLr$;;;r   r   r   )r(   )r)   r*   r+   r,   __doc__r   independentpositivearg_constraintssimplexsupporthas_rsampler@   rI   rS   r[   propertyra   rn   ru   r|   r   r   r.   r/   __classcell__)rD   r0   s   @@r   r
   r
   &   s     " 	001E1EqIO !!GKP P/ /

 E E   
 

 % %< <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   