Ë
    óÍ:j½  ã                   ó�   — d dl mZmZ d dlZd dlmZ d dlmZ d dlmZ d dl	m
Z
mZmZmZ d dlmZ d dlmZmZ d	gZ G d
„ d	e«      Zy)é    )ÚOptionalÚUnionN)ÚTensor)Úconstraints)ÚDistribution)Úbroadcast_allÚlazy_propertyÚlogits_to_probsÚprobs_to_logits)Ú binary_cross_entropy_with_logits)Ú_NumberÚNumberÚ	Geometricc            	       óZ  ‡ — e Zd ZdZej
                  ej                  dœZej                  Z		 	 	 dde
eeef      de
eeef      de
e   ddfˆ fd„Zdˆ fd	„	Zedefd
„«       Zedefd„«       Zedefd„«       Zedefd„«       Zedefd„«       Z ej2                  «       fd„Zd„ Zd„ Zˆ xZS )r   a€  
    Creates a Geometric distribution parameterized by :attr:`probs`,
    where :attr:`probs` is the probability of success of Bernoulli trials.

    .. math::

        P(X=k) = (1-p)^{k} p, k = 0, 1, ...

    .. note::
        :func:`torch.distributions.geometric.Geometric` :math:`(k+1)`-th trial is the first success
        hence draws samples in :math:`\{0, 1, \ldots\}`, whereas
        :func:`torch.Tensor.geometric_` `k`-th trial is the first success hence draws samples in :math:`\{1, 2, \ldots\}`.

    Example::

        >>> # xdoctest: +IGNORE_WANT("non-deterministic")
        >>> m = Geometric(torch.tensor([0.3]))
        >>> m.sample()  # underlying Bernoulli has 30% chance 1; 70% chance 0
        tensor([ 2.])

    Args:
        probs (Number, Tensor): the probability of sampling `1`. Must be in range (0, 1]
        logits (Number, Tensor): the log-odds of sampling `1`.
    )ÚprobsÚlogitsNr   r   Úvalidate_argsÚreturnc           
      ó2  •— |d u |d u k(  rt        d«      ‚|�t        |«      \  | _        n|€J ‚t        |«      \  | _        |�|n|}t	        |t
        «      rt        j                  «       }n|€J ‚|j                  «       }t        ‰	| �)  ||¬«       | j                  r{|�x| j                  }|dkD  }|j                  «       sV|j                  |    }t        dt        |«      j                  › dt!        |j"                  «      › dt%        | «      › d|› �«      ‚y y y )Nz;Either `probs` or `logits` must be specified, but not both.©r   r   zExpected parameter probs (z
 of shape z) of distribution z* to be positive but found invalid values:
)Ú
ValueErrorr   r   r   Ú
isinstancer   ÚtorchÚSizeÚsizeÚsuperÚ__init__Ú_validate_argsÚallÚdataÚtypeÚ__name__ÚtupleÚshapeÚrepr)
Úselfr   r   r   Úprobs_or_logitsÚbatch_shapeÚvalueÚvalidÚinvalid_valueÚ	__class__s
            €úr/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torch/distributions/geometric.pyr   zGeometric.__init__2   s:  ø€ ð �TˆM˜v¨˜~Ò.ÜØMóð ð ÐÜ)¨%Ó0‰MˆT�ZàÐ%Ð%Ð%Ü*¨6Ó2‰NˆTŒ[Ø#(Ð#4™%¸&ˆÜ�o¤wÔ/ÜŸ*™*›,‰Kà"Ð.Ð.Ð.Ø)×.Ñ.Ó0ˆKÜ‰Ñ˜°MÐÔBØ×Ò 5Ð#4à—J‘JˆEØ˜A‘IˆEØ—9‘9”;Ø %§
¡
¨E¨6Ñ 2�Ü ðÜ˜U›×,Ñ,Ð-¨Z¼¸e¿k¹kÓ8JÐ7Kð L'Ü'+¨D£z lð 3AØANÀðQóð ð ð	 $5Ðó    c                 ób  •— | j                  t        |«      }t        j                  |«      }d| j                  v r | j
                  j                  |«      |_        d| j                  v r | j                  j                  |«      |_        t        t        |�'  |d¬«       | j                  |_
        |S )Nr   r   Fr   )Ú_get_checked_instancer   r   r   Ú__dict__r   Úexpandr   r   r   r   )r&   r(   Ú	_instanceÚnewr,   s       €r-   r2   zGeometric.expandU   sŽ   ø€ Ø×(Ñ(¬°IÓ>ˆÜ—j‘j Ó-ˆØ�d—m‘mÑ#ØŸ
™
×)Ñ)¨+Ó6ˆCŒIØ�t—}‘}Ñ$ØŸ™×+Ñ+¨KÓ8ˆCŒJÜŒi˜Ñ& {À%Ð&ÔHØ!×0Ñ0ˆÔØˆ
r.   c                 ó&   — d| j                   z  dz
  S ©Ng      ð?©r   ©r&   s    r-   ÚmeanzGeometric.mean`   s   € à�T—Z‘ZÑ #Ñ%Ð%r.   c                 ó@   — t        j                  | j                  «      S ©N)r   Ú
zeros_liker   r8   s    r-   ÚmodezGeometric.moded   s   € ä×Ñ §
¡
Ó+Ð+r.   c                 ó@   — d| j                   z  dz
  | j                   z  S r6   r7   r8   s    r-   ÚvariancezGeometric.varianceh   s   € à�d—j‘jÑ  3Ñ&¨$¯*©*Ñ4Ð4r.   c                 ó0   — t        | j                  d¬«      S ©NT)Ú	is_binary)r   r   r8   s    r-   r   zGeometric.logitsl   s   € ä˜tŸz™z°TÔ:Ð:r.   c                 ó0   — t        | j                  d¬«      S rA   )r
   r   r8   s    r-   r   zGeometric.probsp   s   € ä˜tŸ{™{°dÔ;Ð;r.   c                 óŠ  — | j                  |«      }t        j                  | j                  j                  «      j
                  }t        j                  «       5  t        j                  j                  «       rSt        j                  || j                  j                  | j                  j                  ¬«      }|j                  |¬«      }n+| j                  j                  |«      j                  |d«      }|j                  «       | j                   j                  «       z  j!                  «       cd d d «       S # 1 sw Y   y xY w)N)ÚdtypeÚdevice)Úminé   )Ú_extended_shaper   Úfinfor   rE   ÚtinyÚno_gradÚ_CÚ_get_tracing_stateÚrandrF   Úclampr4   Úuniform_ÚlogÚlog1pÚfloor)r&   Úsample_shaper$   rK   Úus        r-   ÚsamplezGeometric.samplet   sÙ   € Ø×$Ñ$ \Ó2ˆÜ�{‰{˜4Ÿ:™:×+Ñ+Ó,×1Ñ1ˆÜ�]‰]‹_ñ 	=Ü�x‰x×*Ñ*Ô,ä—J‘J˜u¨D¯J©J×,<Ñ,<ÀTÇZÁZ×EVÑEVÔW�Ø—G‘G �GÓ%‘à—J‘J—N‘N 5Ó)×2Ñ2°4¸Ó;�Ø—E‘E“G §
¡
˜{×1Ñ1Ó3Ñ3×:Ñ:Ó<÷	=÷ 	=ò 	=ús   ÁCD9Ä9Ec                 ó(  — | j                   r| j                  |«       t        || j                  «      \  }}|j	                  t
        j                  ¬«      }d||dk(  |dk(  z  <   || j                  «       z  | j                  j                  «       z   S )N)Úmemory_formatr   rH   )	r   Ú_validate_sampler   r   Úcloner   Úcontiguous_formatrS   rR   )r&   r)   r   s      r-   Úlog_probzGeometric.log_prob€   s~   € Ø×ÒØ×!Ñ! %Ô(Ü$ U¨D¯J©JÓ7‰ˆˆuØ—‘¬%×*AÑ*A�ÓBˆØ-.ˆˆu˜‰z˜e q™jÑ)Ñ*Ø˜˜—~‘~Ó'Ñ'¨$¯*©*¯.©.Ó*:Ñ:Ð:r.   c                 ó`   — t        | j                  | j                  d¬«      | j                  z  S )NÚnone)Ú	reduction)r   r   r   r8   s    r-   ÚentropyzGeometric.entropyˆ   s(   € ä,¨T¯[©[¸$¿*¹*ÐPVÔWØ�j‰jñð	
r.   )NNNr;   )r"   Ú
__module__Ú__qualname__Ú__doc__r   Úunit_intervalÚrealÚarg_constraintsÚnonnegative_integerÚsupportr   r   r   r   Úboolr   r2   Úpropertyr9   r=   r?   r	   r   r   r   r   rW   r]   ra   Ú__classcell__)r,   s   @r-   r   r      s*  ø„ ñð2 !,× 9Ñ 9À[×EUÑEUÑV€OØ×-Ñ-€Gð 26Ø26Ø(,ñ	!à˜˜f f˜nÑ-Ñ.ð!ð ˜˜v v˜~Ñ.Ñ/ð!ð   ‘~ð	!ð
 
õ!õF	ð ð&�fò &ó ð&ð ð,�fò ,ó ð,ð ð5˜&ò 5ó ð5ð ð;˜ò ;ó ð;ð ð<�vò <ó ð<ð #- %§*¡*£,ó 
=ò;ö
r.   )Útypingr   r   r   r   Útorch.distributionsr   Ú torch.distributions.distributionr   Útorch.distributions.utilsr   r	   r
   r   Útorch.nn.functionalr   Útorch.typesr   r   Ú__all__r   © r.   r-   ú<module>ru      s>   ðç "ã Ý Ý +Ý 9÷ó õ Aß 'ð ˆ-€ôw
�õ w
r.   