Ë
    ˆ\;j#  ã                   ód   — d dl 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  G d„ d«      Zy)é    N)Ú_C_ops)Úcheck_variable_and_dtypeÚconvert_dtype)ÚVariable)Úin_dynamic_modec                   óº   ‡ — e Zd ZdZdˆ fd„	Zed„ «       Zed„ «       Zed„ «       Zed„ «       Z	dd„Z
dd„Zd	„ Zd
„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zdd„Zdd„Zˆ xZS )ÚDistributiona>  
    The abstract base class for probability distributions. Functions are
    implemented in specific distributions.

    Args:
        batch_shape(Sequence[int], optional):  independent, not identically
            distributed draws, aka a "collection" or "bunch" of distributions.
        event_shape(Sequence[int], optional): the shape of a single
            draw from the distribution; it may be dependent across dimensions.
            For scalar distributions, the event shape is []. For n-dimension
            multivariate distribution, the event shape is [n].
    c                 óª   •— t        |t        «      r|n
t        |«      | _        t        |t        «      r|n
t        |«      | _        t        ‰| �  «        y )N)Ú
isinstanceÚtupleÚ_batch_shapeÚ_event_shapeÚsuperÚ__init__)ÚselfÚbatch_shapeÚevent_shapeÚ	__class__s      €úiG:\00. PROJECTS\API\Inventory\templateJSON\kerjaOCR\Lib\site-packages\paddle/distribution/distribution.pyr   zDistribution.__init__/   sR   ø€ ô ˜+¤uÔ-ñ ä�{Ó#ð 	Ôô ˜+¤uÔ-ñ ä�{Ó#ð 	Ôô 	‰ÑÕó    c                 ó   — | j                   S )zeReturns batch shape of distribution

        Returns:
            Sequence[int]: batch shape
        )r   ©r   s    r   r   zDistribution.batch_shape=   ó   € ð × Ñ Ð r   c                 ó   — | j                   S )zeReturns event shape of distribution

        Returns:
            Sequence[int]: event shape
        )r   r   s    r   r   zDistribution.event_shapeF   r   r   c                 ó   — t         ‚)zMean of distribution©ÚNotImplementedErrorr   s    r   ÚmeanzDistribution.meanO   ó
   € ô "Ð!r   c                 ó   — t         ‚)zVariance of distributionr   r   s    r   ÚvariancezDistribution.varianceT   r   r   c                 ó   — t         ‚)zSampling from the distribution.r   ©r   Úshapes     r   ÚsamplezDistribution.sampleY   ó   € ä!Ð!r   c                 ó   — t         ‚)zreparameterized sampler   r#   s     r   ÚrsamplezDistribution.rsample]   r&   r   c                 ó   — t         ‚)z The entropy of the distribution.r   r   s    r   ÚentropyzDistribution.entropya   r&   r   c                 ó   — t         ‚)z7The KL-divergence between self distributions and other.r   )r   Úothers     r   Úkl_divergencezDistribution.kl_divergencee   r&   r   c                 ó@   — | j                  |«      j                  «       S )z‡Probability density/mass function evaluated at value.

        Args:
            value (Tensor): value which will be evaluated
        )Úlog_probÚexp©r   Úvalues     r   ÚprobzDistribution.probi   s   € ð �}‰}˜UÓ#×'Ñ'Ó)Ð)r   c                 ó   — t         ‚)z&Log probability density/mass function.r   r1   s     r   r/   zDistribution.log_probq   r&   r   c                 ó   — t         ‚)zœProbability density/mass function.

        Note:

            This method will be deprecated in the future, please use `prob`
            instead.
        r   r1   s     r   ÚprobszDistribution.probsu   s
   € ô "Ð!r   c                 óp   — t        |«      t        | j                  «      z   t        | j                  «      z   S )z¥compute shape of the sample

        Args:
            sample_shape (Tensor): sample shape

        Returns:
            Tensor: generated sample data shape
        )r   r   r   )r   Úsample_shapes     r   Ú_extend_shapezDistribution._extend_shape   s7   € ô �,ÓÜ�D×%Ñ%Ó&ñ'ä�D×%Ñ%Ó&ñ'ð	
r   c                 ód   — d}d}|D ]  }t        |t        «      rd}Œd}Œ |r|rt        d«      ‚|S )zá
        Argument validation for distribution args
        Args:
            value (float, list, numpy.ndarray, Tensor)
        Raises
            ValueError: if one argument is Tensor, all arguments should be Tensor
        FTz9if one argument is Tensor, all arguments should be Tensor)r   r   Ú
ValueError)r   ÚargsÚis_variableÚ	is_numberÚargs        r   Ú_validate_argszDistribution._validate_argsŽ   sK   € ð ˆØˆ	ÛˆCÜ˜#œxÔ(Ø"‘à ‘	ð	 ñ ™9ÜØKóð ð Ðr   c           	      ó®  — g }g }d}|D ]Í  }t        |t        t        t        t        j
                  t        f«      s#t        dj                  t        |«      «      «      ‚t	        j                  |«      }|j                  }t        |«      dk7  r4t        |«      dk7  rt        j                  d«       |j                  d«      }||z   }|j!                  |«       ŒÏ |j                  }|D ]b  }t	        j"                  ||«      \  }	}
t$        j&                  j)                  |¬«      }t%        j*                  |	|«       |j!                  |«       Œd t        |«      S )z¤
        Argument convert args to Tensor

        Args:
            value (float, list, numpy.ndarray, Tensor)
        Returns:
            Tensor of args.
        g        z\Type of input args must be float, list, tuple, numpy.ndarray or Tensor, but received type {}Úfloat32Úfloat64zadata type of argument only support float32 and float64, your argument will be convert to float32.©Údtype)r   ÚfloatÚlistr   ÚnpÚndarrayr   Ú	TypeErrorÚformatÚtypeÚarrayrE   ÚstrÚwarningsÚwarnÚastypeÚappendÚbroadcast_arraysÚpaddleÚtensorÚcreate_tensorÚassign)r   r<   Ú
numpy_argsÚvariable_argsÚtmpr?   Úarg_npÚ	arg_dtyperE   Úarg_broadcastedÚ_Úarg_variables               r   Ú
_to_tensorzDistribution._to_tensor¥   s,  € ð ˆ
ØˆØˆãˆCÜ˜c¤E¬4´¼¿
¹
ÄHÐ#MÔNÜØr×yÑyÜ˜S›	óóð ô —X‘X˜c“]ˆFØŸ™ˆIÜ�9‹~ Ò*Ü�y“> YÒ.ô —M‘MØ{ôð  Ÿ™ yÓ1�à˜‘,ˆCØ×Ñ˜fÕ%ð) ð, —	‘	ˆÛˆCÜ!#×!4Ñ!4°S¸#Ó!>ÑˆO˜QÜ!Ÿ=™=×6Ñ6¸UÐ6ÓCˆLÜ�M‰M˜/¨<Ô8Ø× Ñ  Õ.ð	 ô �]Ó#Ð#r   c                 ó¦  — t        «       rg|j                  |j                  k7  rLt        |j                  «      dv r5t        j                  d«       t        j                  ||j                  «      S |S t        |dddgd«       |j                  |j                  k7  r6t        j                  d«       t        j                  ||j                  ¬«      S |S )a³  
        Log_prob and probs methods have input ``value``, if value's dtype is different from param,
        convert value's dtype to be consistent with param's dtype.

        Args:
            param (Tensor): low and high in Uniform class, loc and scale in Normal class.
            value (Tensor): The input tensor.

        Returns:
            value (Tensor): Change value's dtype if value's dtype is different from param.
        )rB   rC   ztdtype of input 'value' needs to be the same as parameters of distribution class. dtype of 'value' will be converted.r2   rB   rC   r/   rD   )	r   rE   r   rO   rP   r   Úcastr   rT   )r   Úparamr2   s      r   Ú_check_values_dtype_in_probsz)Distribution._check_values_dtype_in_probsÑ   s¶   € ô ÔØ�{‰{˜eŸk™kÒ)¬m¸E¿K¹KÓ.Hð Mñ /ô —‘ð Kôô —{‘{ 5¨%¯+©+Ó6Ð6ØˆLä Ø�7˜Y¨	Ð2°Jô	
ð �;‰;˜%Ÿ+™+Ò%Ü�M‰Mð Gôô —;‘;˜u¨E¯K©KÔ8Ð8Øˆr   c                 óˆ   — |r,t        j                  |«      t        j                  | «      z
  S t        j                  |«      S )a  
        Converts probabilities into logits. For the binary, probs denotes the
        probability of occurrence of the event indexed by `1`. For the
        multi-dimensional, values of last axis denote the probabilities of
        occurrence of each of the events.
        )rT   ÚlogÚlog1p)r   r6   Ú	is_binarys      r   Ú_probs_to_logitszDistribution._probs_to_logitsò   s=   € ñ ô �Z‰Z˜Ó¤§¡¨u¨fÓ!5Ñ5ð	
ô —‘˜EÓ"ð	
r   c                 ó®   — |r)t         j                  j                  j                  |«      S t         j                  j                  j	                  |d¬«      S )zê
        Converts logits into probabilities. For the binary, each value denotes
        log odds, whereas for the multi-dimensional case, the values along the
        last dimension denote the log probabilities of the events.
        éÿÿÿÿ)Úaxis)rT   ÚnnÚ
functionalÚsigmoidÚsoftmax)r   Úlogitsrh   s      r   Ú_logits_to_probszDistribution._logits_to_probsÿ   sJ   € ñ ô �I‰I× Ñ ×(Ñ(¨Ó0ð	
ô —‘×%Ñ%×-Ñ-¨f¸2Ð-Ó>ð	
r   )© rs   )rs   )F)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   Úpropertyr   r   r   r!   r%   r(   r*   r-   r3   r/   r6   r9   r@   r`   rd   ri   rr   Ú__classcell__)r   s   @r   r	   r	   !   s    ø„ ñõð ñ!ó ð!ð ñ!ó ð!ð ñ"ó ð"ð ñ"ó ð"ó"ó"ò"ò"ò*ò"ò"ò
òò.*$òXóB
÷

r   r	   )rO   ÚnumpyrH   rT   r   Úpaddle.base.data_feederr   r   Úpaddle.base.frameworkr   Úpaddle.frameworkr   r	   rs   r   r   Ú<module>r~      s(   ðó, ã ã Ý ß KÝ *Ý ,÷h
ò h
r   