§
    ŠŠtj…  ã                   ól   — d dl Z d dl mZmZ d dlmZ d dlmZ d dlmZm	Z	m
Z
 dgZ G d„ de¦  «        ZdS )é    N)ÚnanÚTensor)Úconstraints)ÚDistribution)Úlazy_propertyÚlogits_to_probsÚprobs_to_logitsÚCategoricalc            	       ó¢  ‡ — e Zd ZdZej        ej        dœZdZ	 	 	 dde	dz  de	dz  de
dz  ddfˆ fd	„Zdˆ fd
„	Zd„ Z ej        dd¬¦  «        d„ ¦   «         Zede	fd„¦   «         Zede	fd„¦   «         Zedej        fd„¦   «         Zede	fd„¦   «         Zede	fd„¦   «         Zede	fd„¦   «         Z ej        ¦   «         fd„Zd„ Zd„ Zdd„Zˆ xZS )r
   aä  
    Creates a categorical distribution parameterized by either :attr:`probs` or
    :attr:`logits` (but not both).

    .. note::
        It is equivalent to the distribution that :func:`torch.multinomial`
        samples from.

    Samples are integers from :math:`\{0, \ldots, K-1\}` where `K` is ``probs.size(-1)``.

    If `probs` is 1-dimensional with length-`K`, each element is the relative probability
    of sampling the class at that index.

    If `probs` is N-dimensional, the first N-1 dimensions are treated as a batch of
    relative probability vectors.

    .. note:: The `probs` argument must be non-negative, finite and have a non-zero sum,
              and it will be normalized to sum to 1 along the last dimension. :attr:`probs`
              will return this normalized value.
              The `logits` argument will be interpreted as unnormalized log probabilities
              and can therefore be any real number. It will likewise be normalized so that
              the resulting probabilities sum to 1 along the last dimension. :attr:`logits`
              will return this normalized value.

    See also: :func:`torch.multinomial`

    Example::

        >>> # xdoctest: +IGNORE_WANT("non-deterministic")
        >>> m = Categorical(torch.tensor([ 0.25, 0.25, 0.25, 0.25 ]))
        >>> m.sample()  # equal probability of 0, 1, 2, 3
        tensor(3)

    Args:
        probs (Tensor): event probabilities
        logits (Tensor): event log probabilities (unnormalized)
    )ÚprobsÚlogitsTNr   r   Úvalidate_argsÚreturnc                 óÔ  •— |d u |d u k    rt          d¦  «        ‚|�G|                     ¦   «         dk     rt          d¦  «        ‚||                     dd¬¦  «        z  | _        nW|€t	          d¦  «        ‚|                     ¦   «         dk     rt          d¦  «        ‚||                     dd¬	¦  «        z
  | _        |�| j        n| j        | _        | j                             ¦   «         d         | _	        | j         
                    ¦   «         dk    r!| j                             ¦   «         d d…         nt          j        ¦   «         }t          ¦   «                              ||¬
¦  «         d S )Nz;Either `probs` or `logits` must be specified, but not both.é   z3`probs` parameter must be at least one-dimensional.éÿÿÿÿT)Úkeepdimzlogits is unexpectedly Nonez4`logits` parameter must be at least one-dimensional.)Údimr   ©r   )Ú
ValueErrorr   Úsumr   ÚAssertionErrorÚ	logsumexpr   Ú_paramÚsizeÚ_num_eventsÚ
ndimensionÚtorchÚSizeÚsuperÚ__init__)Úselfr   r   r   Úbatch_shapeÚ	__class__s        €ú]/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributions/categorical.pyr!   zCategorical.__init__8   s`  ø€ ð �TˆM˜v¨˜~Ò.Ð.ÝØMñô ð ð ÐØ�yŠy‰{Œ{˜QŠˆÝ Ð!VÑWÔWÐWà §¢¨2°t Ñ!<Ô!<Ñ<ˆDŒJˆJàˆ~Ý$Ð%BÑCÔCÐCØ�zŠz‰|Œ|˜aÒÐÝ Ð!WÑXÔXÐXð ! 6×#3Ò#3¸ÀDÐ#3Ñ#IÔ#IÑIˆDŒKØ$)Ð$5�d”j�j¸4¼;ˆŒØœ;×+Ò+Ñ-Ô-¨bÔ1ˆÔà'+¤{×'=Ò'=Ñ'?Ô'?À!Ò'CÐ'CˆDŒK×ÒÑÔ˜s ˜sÔ#Ð#ÍÌÉÌð 	õ 	‰Œ×Ò˜°MÐÑBÔBÐBÐBÐBó    c                 óô  •— |                       t          |¦  «        }t          j        |¦  «        }|t          j        | j        f¦  «        z   }d| j        v r+| j                             |¦  «        |_        |j        |_        d| j        v r+| j	                             |¦  «        |_	        |j	        |_        | j        |_        t          t          |¦  «                             |d¬¦  «         | j        |_        |S )Nr   r   Fr   )Ú_get_checked_instancer
   r   r   r   Ú__dict__r   Úexpandr   r   r    r!   Ú_validate_args)r"   r#   Ú	_instanceÚnewÚparam_shaper$   s        €r%   r*   zCategorical.expandW   sØ   ø€ Ø×(Ò(­°iÑ@Ô@ˆÝ”j Ñ-Ô-ˆØ!¥E¤J°Ô0@Ð/BÑ$CÔ$CÑCˆØ�d”mÐ#Ð#Øœ
×)Ò)¨+Ñ6Ô6ˆCŒIØœˆCŒJØ�t”}Ð$Ð$Øœ×+Ò+¨KÑ8Ô8ˆCŒJØœˆCŒJØÔ*ˆŒÝ�k˜3ÑÔ×(Ò(¨ÀEÐ(ÑJÔJÐJØ!Ô0ˆÔØˆ
r&   c                 ó&   —  | j         j        |i |¤ŽS ©N)r   r-   )r"   ÚargsÚkwargss      r%   Ú_newzCategorical._newf   s   € ØˆtŒ{Œ Ð/¨Ð/Ð/Ð/r&   r   )Úis_discreteÚ	event_dimc                 ó<   — t          j        d| j        dz
  ¦  «        S )Nr   r   )r   Úinteger_intervalr   ©r"   s    r%   ÚsupportzCategorical.supporti   s   € õ Ô+¨A¨tÔ/?À!Ñ/CÑDÔDÐDr&   c                 ó*   — t          | j        ¦  «        S r0   )r	   r   r8   s    r%   r   zCategorical.logitsn   s   € å˜tœzÑ*Ô*Ð*r&   c                 ó*   — t          | j        ¦  «        S r0   )r   r   r8   s    r%   r   zCategorical.probsr   s   € å˜tœ{Ñ+Ô+Ð+r&   c                 ó4   — | j                              ¦   «         S r0   )r   r   r8   s    r%   r.   zCategorical.param_shapev   s   € àŒ{×ÒÑ!Ô!Ð!r&   c                 óˆ   — t          j        |                      ¦   «         t          | j        j        | j        j        ¬¦  «        S ©N©ÚdtypeÚdevice©r   ÚfullÚ_extended_shaper   r   r@   rA   r8   s    r%   ÚmeanzCategorical.meanz   ó=   € åŒzØ× Ò Ñ"Ô"ÝØ”*Ô"Ø”:Ô$ð	
ñ 
ô 
ð 	
r&   c                 ó8   — | j                              d¬¦  «        S )Nr   )r   )r   Úargmaxr8   s    r%   ÚmodezCategorical.modeƒ   s   € àŒz× Ò  RÐ Ñ(Ô(Ð(r&   c                 óˆ   — t          j        |                      ¦   «         t          | j        j        | j        j        ¬¦  «        S r>   rB   r8   s    r%   ÚvariancezCategorical.variance‡   rF   r&   c                 óH  — t          |t          j        ¦  «        st          j        |¦  «        }| j                             d| j        ¦  «        }t          j        ||                     ¦   «         d¦  «        j        }|                     |  	                    |¦  «        ¦  «        S )Nr   T)
Ú
isinstancer   r   r   Úreshaper   ÚmultinomialÚnumelÚTrD   )r"   Úsample_shapeÚprobs_2dÚ
samples_2ds       r%   ÚsamplezCategorical.sample�   s„   € Ý˜,­¬
Ñ3Ô3ð 	4Ý œ: lÑ3Ô3ˆLØ”:×%Ò% b¨$Ô*:Ñ;Ô;ˆÝÔ& x°×1CÒ1CÑ1EÔ1EÀtÑLÔLÔNˆ
Ø×!Ò! $×"6Ò"6°|Ñ"DÔ"DÑEÔEÐEr&   c                 ó,  — | j         r|                      |¦  «         |                     ¦   «                              d¦  «        }t	          j        || j        ¦  «        \  }}|dd d…f         }|                     d|¦  «                             d¦  «        S )Nr   .r   )	r+   Ú_validate_sampleÚlongÚ	unsqueezer   Úbroadcast_tensorsr   ÚgatherÚsqueeze)r"   ÚvalueÚlog_pmfs      r%   Úlog_probzCategorical.log_prob—   s‡   € ØÔð 	)Ø×!Ò! %Ñ(Ô(Ð(Ø—
’
‘”×&Ò& rÑ*Ô*ˆÝÔ0°¸¼ÑDÔD‰ˆˆwØ�c˜2˜A˜2�g”ˆØ�~Š~˜b %Ñ(Ô(×0Ò0°Ñ4Ô4Ð4r&   c                 ó¾   — t          j        | j        j        ¦  «        j        }t          j        | j        |¬¦  «        }|| j        z  }|                     d¦  «         S )N)Úminr   )r   Úfinfor   r@   ra   Úclampr   r   )r"   Úmin_realr   Úp_log_ps       r%   ÚentropyzCategorical.entropyŸ   sN   € Ý”;˜tœ{Ô0Ñ1Ô1Ô5ˆÝ”˜Tœ[¨hÐ7Ñ7Ô7ˆØ˜4œ:Ñ%ˆØ—’˜B‘”ÐÐr&   c                 ó  — | j         }t          j        |t          j        | j        j        ¬¦  «        }|                     ddt          | j        ¦  «        z  z   ¦  «        }|r| 	                    d| j        z   ¦  «        }|S )Nr?   )r   )r   )
r   r   ÚarangerX   r   rA   ÚviewÚlenÚ_batch_shaper*   )r"   r*   Ú
num_eventsÚvaluess       r%   Úenumerate_supportzCategorical.enumerate_support¥   ss   € ØÔ%ˆ
Ý”˜jµ´
À4Ä;ÔCUÐVÑVÔVˆØ—’˜U T­C°Ô0AÑ,BÔ,BÑ%BÑBÑCÔCˆØð 	>Ø—]’] 5¨4Ô+<Ñ#<Ñ=Ô=ˆFØˆr&   )NNNr0   )T)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚsimplexÚreal_vectorÚarg_constraintsÚhas_enumerate_supportr   Úboolr!   r*   r3   Údependent_propertyr9   r   r   r   Úpropertyr   r   r.   rE   rI   rK   rU   r_   rf   rn   Ú__classcell__)r$   s   @r%   r
   r
      sR  ø€ € € € € ð$ð $ðN !,Ô 3¸{Ô?VÐWÐW€OØ Ðð  $Ø $Ø%)ð	Cð Cà˜‰}ðCð ˜‘ðCð ˜d‘{ð	Cð
 
ðCð Cð Cð Cð Cð Cð>ð ð ð ð ð ð0ð 0ð 0ð $€[Ô#°ÀÐBÑBÔBðEð Eñ CÔBðEð ð+˜ð +ð +ð +ñ „]ð+ð ð,�vð ,ð ,ð ,ñ „]ð,ð ð"˜UœZð "ð "ð "ñ „Xð"ð ð
�fð 
ð 
ð 
ñ „Xð
ð ð)�fð )ð )ð )ñ „Xð)ð ð
˜&ð 
ð 
ð 
ñ „Xð
ð #- %¤*¡,¤,ð Fð Fð Fð Fð5ð 5ð 5ð ð  ð  ðð ð ð ð ð ð ð r&   )r   r   r   Útorch.distributionsr   Ú torch.distributions.distributionr   Útorch.distributions.utilsr   r   r	   Ú__all__r
   © r&   r%   ú<module>r€      s±   ðð €€€Ø Ð Ð Ð Ð Ð Ð Ð Ø +Ð +Ð +Ð +Ð +Ð +Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø UÐ UÐ UÐ UÐ UÐ UÐ UÐ UÐ UÐ Uð ˆ/€ð^ð ^ð ^ð ^ð ^�,ñ ^ô ^ð ^ð ^ð ^r&   