ó
    Eñi3  ã                   ó‚   • S SK r S SKJs  Jr  S SK Jr  S SKJr  S SKJ	r	  S SK
Jr  S SKJrJrJrJr  S/r " S S\	5      rg)	é    N)ÚTensor)Úconstraints)ÚDistribution)ÚGamma)Úbroadcast_allÚlazy_propertyÚlogits_to_probsÚprobs_to_logitsÚNegativeBinomialc                   óâ  ^ • \ rS rSrSr\R                  " S5      \R                  " SS5      \R                  S.r	\R                  r   SS\\-  S	\S-  S
\S-  S\S-  SS4
U 4S jjjrSU 4S jjrS r\S\4S j5       r\S\4S j5       r\S\4S j5       r\S\4S j5       r\S\4S j5       r\S\R4                  4S j5       r\S\4S j5       r\R4                  " 5       4S jrS rSr U =r!$ )r   é   aC  
Creates a Negative Binomial distribution, i.e. distribution
of the number of successful independent and identical Bernoulli trials
before :attr:`total_count` failures are achieved. The probability
of success of each Bernoulli trial is :attr:`probs`.

Args:
    total_count (float or Tensor): non-negative number of negative Bernoulli
        trials to stop, although the distribution is still valid for real
        valued count
    probs (Tensor): Event probabilities of success in the half open interval [0, 1)
    logits (Tensor): Event log-odds for probabilities of success
r   ç        ç      ð?)Útotal_countÚprobsÚlogitsNr   r   r   Úvalidate_argsÚreturnc                 óê  >• US L US L :X  a  [        S5      eUbC  [        X5      u  U l        U l        U R                  R	                  U R                  5      U l        OPUc  [        S5      e[        X5      u  U l        U l        U R                  R	                  U R                  5      U l        Ub  U R                  OU R                  U l        U R                  R                  5       n[        TU ])  XTS9  g )Nz;Either `probs` or `logits` must be specified, but not both.zlogits is unexpectedly None©r   )Ú
ValueErrorr   r   r   Útype_asÚAssertionErrorr   Ú_paramÚsizeÚsuperÚ__init__)Úselfr   r   r   r   Úbatch_shapeÚ	__class__s         €Úb/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/negative_binomial.pyr   ÚNegativeBinomial.__init__+   sä   ø€ ð �TˆM˜v¨˜~Ó.ÜØMóð ð Ñô
 ˜kÓ1ñ	ØÔ à”
à#×/Ñ/×7Ñ7¸¿
¹
ÓCˆDÕà‰~Ü$Ð%BÓCÐCô
 ˜kÓ2ñ	ØÔ à”à#×/Ñ/×7Ñ7¸¿¹ÓDˆDÔà$)Ñ$5�d—j’j¸4¿;¹;ˆŒØ—k‘k×&Ñ&Ó(ˆÜ‰Ñ˜ÐÒBó    c                 óê  >• U R                  [        U5      n[        R                  " U5      nU R                  R                  U5      Ul        SU R                  ;   a1  U R                  R                  U5      Ul        UR                  Ul        SU R                  ;   a1  U R                  R                  U5      Ul	        UR                  Ul        [        [        U]/  USS9  U R                  Ul        U$ )Nr   r   Fr   )Ú_get_checked_instancer   ÚtorchÚSizer   ÚexpandÚ__dict__r   r   r   r   r   Ú_validate_args)r   r   Ú	_instanceÚnewr    s       €r!   r(   ÚNegativeBinomial.expandK   s¾   ø€ Ø×(Ñ(Ô)9¸9ÓEˆÜ—j’j Ó-ˆØ×*Ñ*×1Ñ1°+Ó>ˆŒØ�d—m‘mÓ#ØŸ
™
×)Ñ)¨+Ó6ˆCŒIØŸ™ˆCŒJØ�t—}‘}Ó$ØŸ™×+Ñ+¨KÓ8ˆCŒJØŸ™ˆCŒJÜÔ Ñ-¨kÈÐ-ÑOØ!×0Ñ0ˆÔØˆ
r#   c                 ó:   • U R                   R                  " U0 UD6$ ©N)r   r,   )r   ÚargsÚkwargss      r!   Ú_newÚNegativeBinomial._newY   s   € Ø�{‰{�Š Ð/¨Ñ/Ð/r#   c                 ó\   • U R                   [        R                  " U R                  5      -  $ r/   )r   r&   Úexpr   ©r   s    r!   ÚmeanÚNegativeBinomial.mean\   s    € à×Ñ¤%§)¢)¨D¯K©KÓ"8Ñ8Ð8r#   c                 óŒ   • U R                   S-
  U R                  R                  5       -  R                  5       R	                  SS9$ )Né   r   )Úmin)r   r   r5   ÚfloorÚclampr6   s    r!   ÚmodeÚNegativeBinomial.mode`   s:   € à×!Ñ! AÑ%¨¯©¯©Ó):Ñ:×AÑAÓC×IÑIÈcÐIÐRÐRr#   c                 ó^   • U R                   [        R                  " U R                  * 5      -  $ r/   )r7   r&   Úsigmoidr   r6   s    r!   ÚvarianceÚNegativeBinomial.varianced   s    € à�y‰yœ5Ÿ=š=¨$¯+©+¨Ó6Ñ6Ð6r#   c                 ó*   • [        U R                  SS9$ ©NT)Ú	is_binary)r
   r   r6   s    r!   r   ÚNegativeBinomial.logitsh   s   € ä˜tŸz™z°TÑ:Ð:r#   c                 ó*   • [        U R                  SS9$ rE   )r	   r   r6   s    r!   r   ÚNegativeBinomial.probsl   s   € ä˜tŸ{™{°dÑ;Ð;r#   c                 ó6   • U R                   R                  5       $ r/   )r   r   r6   s    r!   Úparam_shapeÚNegativeBinomial.param_shapep   s   € à�{‰{×ÑÓ!Ð!r#   c                 ój   • [        U R                  [        R                  " U R                  * 5      SS9$ )NF)ÚconcentrationÚrater   )r   r   r&   r5   r   r6   s    r!   Ú_gammaÚNegativeBinomial._gammat   s/   € ô Ø×*Ñ*Ü—’˜DŸK™K˜<Ó(Øñ
ð 	
r#   c                 óÀ   • [         R                  " 5          U R                  R                  US9n[         R                  " U5      sS S S 5        $ ! , (       d  f       g = f)N)Úsample_shape)r&   Úno_gradrP   ÚsampleÚpoisson)r   rS   rO   s      r!   rU   ÚNegativeBinomial.sample}   s8   € Ü�]Š]�_Ø—;‘;×%Ñ%°<Ð%Ð@ˆDÜ—=’= Ó&÷ �_�_ús   –/AÁ
Ac                 óô  • U R                   (       a  U R                  U5        U R                  [        R                  " U R
                  * 5      -  U[        R                  " U R
                  5      -  -   n[        R                  " U R                  U-   5      * [        R                  " SU-   5      -   [        R                  " U R                  5      -   nUR                  U R                  U-   S:H  S5      nX#-
  $ )Nr   r   )	r*   Ú_validate_sampler   ÚFÚ
logsigmoidr   r&   ÚlgammaÚmasked_fill)r   ÚvalueÚlog_unnormalized_probÚlog_normalizations       r!   Úlog_probÚNegativeBinomial.log_prob‚   sØ   € Ø××Ø×!Ñ! %Ô(à $× 0Ñ 0´1·<²<Ø�[‰[ˆLó4
ñ !
à”A—L’L §¡Ó-Ñ-ñ!.Ðô
 �\Š\˜$×*Ñ*¨UÑ2Ó3Ð3Ü�lŠl˜3 ™;Ó'ñ(ä�lŠl˜4×+Ñ+Ó,ñ-ð 	ð .×9Ñ9Ø×Ñ˜uÑ$¨Ñ+¨Só
Ðð %Ñ8Ð8r#   )r   r   r   r   )NNNr/   )"Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   Úgreater_than_eqÚhalf_open_intervalÚrealÚarg_constraintsÚnonnegative_integerÚsupportr   ÚfloatÚboolr   r(   r2   Úpropertyr7   r>   rB   r   r   r   r&   r'   rK   r   rP   rU   ra   Ú__static_attributes__Ú__classcell__)r    s   @r!   r   r      sŠ  ø† ñð  #×2Ò2°1Ó5Ø×/Ò/°°SÓ9Ø×"Ñ"ñ€Oð
 ×-Ñ-€Gð
  $Ø $Ø%)ñCà˜e‘^ðCð ˜‰}ðCð ˜‘ð	Cð
 ˜d‘{ðCð 
÷Cð C÷@ò0ð ð9�fó 9ó ð9ð ðS�fó Só ðSð ð7˜&ó 7ó ð7ð ð;˜ó ;ó ð;ð ð<�vó <ó ð<ð ð"˜UŸZ™Zó "ó ð"ð ð
˜ó 
ó ð
ð #(§*¢*£,ô '÷
9ð 9r#   )r&   Útorch.nn.functionalÚnnÚ
functionalrZ   r   Útorch.distributionsr   Ú torch.distributions.distributionr   Útorch.distributions.gammar   Útorch.distributions.utilsr   r   r	   r
   Ú__all__r   © r#   r!   Ú<module>r|      s>   ðó ß Ð Ý Ý +Ý 9Ý +÷ó ð Ð
€ôB9�|õ B9r#   