Ë
    ˆ\;j{  ã                   ón   — d dl Z d dlmZ d dlmZ d dlmZ d dlmZ  G d„ dej                  «      Z
d	d„Zy)
é    N)Úcheck_variable_and_dtype)ÚLayerHelper)Úexponential_family)Úin_dynamic_modec                   ór   ‡ — e Zd ZdZˆ fd„Zed„ «       Zed„ «       Zdd„Zd„ Z	d„ Z
d„ Zed	„ «       Zd
„ Zˆ xZS )Ú	Dirichletab  
    Dirichlet distribution with parameter "concentration".

    The Dirichlet distribution is defined over the `(k-1)-simplex` using a
    positive, lenght-k vector concentration(`k > 1`).
    The Dirichlet is identically the Beta distribution when `k = 2`.

    For independent and identically distributed continuous random variable
    :math:`\boldsymbol X \in R_k` , and support
    :math:`\boldsymbol X \in (0,1), ||\boldsymbol X|| = 1` ,
    The probability density function (pdf) is

    .. math::

        f(\boldsymbol X; \boldsymbol \alpha) = \frac{1}{B(\boldsymbol \alpha)} \prod_{i=1}^{k}x_i^{\alpha_i-1}

    where :math:`\boldsymbol \alpha = {\alpha_1,...,\alpha_k}, k \ge 2` is
    parameter, the normalizing constant is the multivariate beta function.

    .. math::

        B(\boldsymbol \alpha) = \frac{\prod_{i=1}^{k} \Gamma(\alpha_i)}{\Gamma(\alpha_0)}

    :math:`\alpha_0=\sum_{i=1}^{k} \alpha_i` is the sum of parameters,
    :math:`\Gamma(\alpha)` is gamma function.

    Args:
        concentration (Tensor): "Concentration" parameter of dirichlet
            distribution, also called :math:`\alpha`. When it's over one
            dimension, the last axis denotes the parameter of distribution,
            ``event_shape=concentration.shape[-1:]`` , axes other than last are
            condsider batch dimensions with ``batch_shape=concentration.shape[:-1]`` .

    Examples:

        .. code-block:: python

            >>> import paddle
            >>> dirichlet = paddle.distribution.Dirichlet(paddle.to_tensor([1., 2., 3.]))
            >>> print(dirichlet.entropy())
            Tensor(shape=[], dtype=float32, place=Place(cpu), stop_gradient=True,
            -1.24434423)

            >>> print(dirichlet.prob(paddle.to_tensor([.3, .5, .6])))
            Tensor(shape=[], dtype=float32, place=Place(cpu), stop_gradient=True,
            10.80000019)
    c                 ó¤   •— |j                  «       dk  rt        d«      ‚|| _        t        ‰| �  |j
                  d d |j
                  dd  «       y )Né   z:`concentration` parameter must be at least one dimensionaléÿÿÿÿ)ÚdimÚ
ValueErrorÚconcentrationÚsuperÚ__init__Úshape)Úselfr   Ú	__class__s     €úfG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/distribution/dirichlet.pyr   zDirichlet.__init__G   sW   ø€ Ø×ÑÓ Ò"ÜØLóð ð +ˆÔÜ‰Ñ˜×,Ñ,¨S¨bÐ1°=×3FÑ3FÀrÀsÐ3KÕLó    c                 óV   — | j                   | j                   j                  dd¬«      z  S )zbMean of Dirichelt distribution.

        Returns:
            Mean value of distribution.
        r   T©Úkeepdim)r   Úsum©r   s    r   ÚmeanzDirichlet.meanP   s+   € ð ×!Ñ! D×$6Ñ$6×$:Ñ$:¸2ÀtÐ$:Ó$LÑLÐLr   c                 ó¤   — | j                   j                  dd¬«      }| j                   || j                   z
  z  |j                  d«      |dz   z  z  S )zjVariance of Dirichlet distribution.

        Returns:
            Variance value of distribution.
        r   Tr   é   r
   )r   r   Úpow)r   Úconcentration0s     r   ÚvariancezDirichlet.varianceY   sZ   € ð ×+Ñ+×/Ñ/°¸DÐ/ÓAˆØ×"Ñ" n°t×7IÑ7IÑ&IÑJØ×Ñ˜qÓ! ^°aÑ%7Ñ8ñ
ð 	
r   c                 ó¢   — t        |t        «      r|n
t        |«      }t        | j                  j	                  | j                  |«      «      «      S )z�Sample from dirichlet distribution.

        Args:
            shape (Sequence[int], optional): Sample shape. Defaults to empty tuple.
        )Ú
isinstanceÚtupleÚ
_dirichletr   ÚexpandÚ_extend_shape)r   r   s     r   ÚsamplezDirichlet.samplee   s?   € ô $ E¬5Ô1‘´u¸U³|ˆÜ˜$×,Ñ,×3Ñ3°D×4FÑ4FÀuÓ4MÓNÓOÐOr   c                 óJ   — t        j                  | j                  |«      «      S )z¶Probability density function(PDF) evaluated at value.

        Args:
            value (Tensor): Value to be evaluated.

        Returns:
            PDF evaluated at value.
        )ÚpaddleÚexpÚlog_prob©r   Úvalues     r   ÚprobzDirichlet.probn   s   € ô �z‰z˜$Ÿ-™-¨Ó.Ó/Ð/r   c                 ó&  — t        j                  |«      | j                  dz
  z  j                  d«      t        j                  | j                  j                  d«      «      z   t        j                  | j                  «      j                  d«      z
  S )zpLog of probability densitiy function.

        Args:
            value (Tensor): Value to be evaluated.
        ç      ð?r   )r)   Úlogr   r   Úlgammar,   s     r   r+   zDirichlet.log_proby   st   € ô �Z‰Z˜Ó $×"4Ñ"4°sÑ":Ñ;×@Ñ@ÀÓDÜ�m‰m˜D×.Ñ.×2Ñ2°2Ó6Ó7ñ8ä�m‰m˜D×.Ñ.Ó/×3Ñ3°BÓ7ñ8ð	
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 )zbEntropy of Dirichlet distribution.

        Returns:
            Entropy of distribution.
        r   r0   )r   r   r   r)   r2   Údigamma)r   r   Úks      r   ÚentropyzDirichlet.entropy…   s±   € ð ×+Ñ+×/Ñ/°Ó3ˆØ×Ñ×$Ñ$ RÑ(ˆä�M‰M˜$×,Ñ,Ó-×1Ñ1°"Ó5Ü�m‰m˜NÓ+ñ,à�>Ñ!¤V§^¡^°NÓ%CÑCñDð ×#Ñ# cÑ)¬V¯^©^¸D×<NÑ<NÓ-OÑOß‰c�"‹gñð	
r   c                 ó   — | j                   fS ©N)r   r   s    r   Ú_natural_parameterszDirichlet._natural_parameters–   s   € à×"Ñ"Ð$Ð$r   c                 óŠ   — |j                  «       j                  d«      t        j                   |j                  d«      «      z
  S )Nr   )r2   r   r)   )r   Úxs     r   Ú_log_normalizerzDirichlet._log_normalizerš   s-   € Ø�x‰x‹z�~‰~˜bÓ!¤F§M¡M°!·%±%¸³)Ó$<Ñ<Ð<r   )© )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úpropertyr   r    r'   r.   r+   r6   r9   r<   Ú__classcell__)r   s   @r   r   r      sg   ø„ ñ.ô`Mð ñMó ðMð ñ	
ó ð	
óPò	0ò

ò
ð" ñ%ó ð%ö=r   r   c                 ó  — t        «       rt        j                  j                  | «      S d}t	        | dg d¢|«       t        |fi t        «       ¤Ž}|j                  | j                  ¬«      }|j                  |d| id|ii ¬«       |S )NÚ	dirichletr   )Úfloat16Úfloat32Úfloat64Úuint16)ÚdtypeÚAlphaÚOut)ÚtypeÚinputsÚoutputsÚattrs)
r   r)   Ú_C_opsrE   r   r   ÚlocalsÚ"create_variable_for_type_inferencerJ   Ú	append_op)r   ÚnameÚop_typeÚhelperÚouts        r   r$   r$   ž   s˜   € ÜÔÜ�}‰}×&Ñ& }Ó5Ð5àˆÜ ØØÚ7Øô		
ô ˜WÑ1¬«Ñ1ˆØ×7Ñ7Ø×%Ñ%ð 8ó 
ˆð 	×ÑØØ˜]Ð+Ø˜C�LØð	 	ô 	
ð ˆ
r   r8   )r)   Úpaddle.base.data_feederr   Úpaddle.base.layer_helperr   Úpaddle.distributionr   Úpaddle.frameworkr   ÚExponentialFamilyr   r$   r=   r   r   Ú<module>r^      s1   ðó Ý <Ý 0Ý 2Ý ,ôE=Ð"×4Ñ4ô E=ôPr   