+
    &jR$                         ^ RI t ^ RIt^ RIHt ^ RIHt ^ RIHt ^ RIHtH	t	H
t
HtHt ^ RIHt ^ RIHtHtHt R.t ! R	 R]4      tR# )
    NTensor)constraints)ExponentialFamily)broadcast_allclamp_probslazy_propertylogits_to_probsprobs_to_logits) binary_cross_entropy_with_logits)_Number_sizeNumberContinuousBernoullic                     a a ] tR t^t oRtR]P                  R]P                  /t]P                  t	^ t
RtR#V3R lV 3R llltR$V 3R lltR tR	 tR
 tR t]V3R lR l4       t]V3R lR l4       t]V3R lR l4       t]V3R lR l4       t]V3R lR l4       t]V3R lR l4       t]P6                  ! 4       3R lt]P6                  ! 4       3V3R lR lltR tR tR t R t!]V3R lR  l4       t"R! t#R"t$Vt%V ;t&# )%r   a  
Creates a continuous Bernoulli distribution parameterized by :attr:`probs`
or :attr:`logits` (but not both).

The distribution is supported in [0, 1] and parameterized by 'probs' (in
(0,1)) or 'logits' (real-valued). Note that, unlike the Bernoulli, 'probs'
does not correspond to a probability and 'logits' does not correspond to
log-odds, but the same names are used due to the similarity with the
Bernoulli. See [1] for more details.

Example::

    >>> # xdoctest: +IGNORE_WANT("non-deterministic")
    >>> m = ContinuousBernoulli(torch.tensor([0.3]))
    >>> m.sample()
    tensor([ 0.2538])

Args:
    probs (Number, Tensor): (0,1) valued parameters
    logits (Number, Tensor): real valued parameters whose sigmoid matches 'probs'

[1] The continuous Bernoulli: fixing a pervasive error in variational
autoencoders, Loaiza-Ganem G and Cunningham JP, NeurIPS 2019.
https://arxiv.org/abs/1907.06845
probslogitsTc          
         < V ^8  d   QhRS[ S[,          R,          RS[ S[,          R,          RS[S[S[3,          RS[R,          RR/# )   r   Nr   limsvalidate_argsreturn)r   r   tuplefloatbool)format__classdict__s   "ڀ/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/distributions/continuous_bernoulli.py__annotate__ ContinuousBernoulli.__annotate__7   sc     "C "C%"C $&"C E5L!	"C
 d{"C 
"C    c                  < VR J VR J 8X  d   \        R4      hVe   \        V\        4      p\        V4      w  V n        VeL   V P
                  R,          P                  V P                  4      P                  4       '       g   \        R4      h\        V P                  4      V n        M1Vf   \        R4      h\        V\        4      p\        V4      w  V n
        Ve   V P                  MV P                  V n        V'       d   \        P                  ! 4       pMV P                  P                  4       pW0n        \         SV `E  WdR7       R # )Nz;Either `probs` or `logits` must be specified, but not both.r   z&The parameter probs has invalid valueszlogits is unexpectedly Noner   )
ValueError
isinstancer   r   r   arg_constraintscheckallr   AssertionErrorr   _paramtorchSizesize_limssuper__init__)selfr   r   r   r   	is_scalarbatch_shape	__class__s   &&&&&  r   r0   ContinuousBernoulli.__init__7   s
    TMv~.M  "5'2I)%0MTZ (++G4::4::FJJLL$%MNN$TZZ0DJ~$%BCC"673I*62NT[$)$5djj4;;**,K++**,K
Br!   c                  < V P                  \        V4      pV P                  Vn        \        P                  ! V4      pR V P
                  9   d2   V P                  P                  V4      Vn        VP                  Vn        RV P
                  9   d2   V P                  P                  V4      Vn	        VP                  Vn        \        \        V`/  VRR7       V P                  Vn        V# )r   r   Fr#   )_get_checked_instancer   r.   r+   r,   __dict__r   expandr*   r   r/   r0   _validate_args)r1   r3   	_instancenewr4   s   &&& r   r9   ContinuousBernoulli.expand[   s    (()<iHJJ	jj-dmm#

))+6CICJt}}$++K8CJCJ!30E0R!00
r!   c                :    V P                   P                  ! V/ VB # N)r*   r<   )r1   argskwargss   &*,r   _newContinuousBernoulli._newi   s    {{///r!   c                    \         P                  ! \         P                  ! V P                  V P                  ^ ,          4      \         P
                  ! V P                  V P                  ^,          4      4      # r   )r+   maxler   r.   gtr1   s   &r   _outside_unstable_region,ContinuousBernoulli._outside_unstable_regionl   sG    yyHHTZZA/$**djjQRm1T
 	
r!   c                    \         P                  ! V P                  4       V P                  V P                  ^ ,          \         P
                  ! V P                  4      ,          4      # rE   )r+   whererJ   r   r.   	ones_likerI   s   &r   
_cut_probsContinuousBernoulli._cut_probsq   sC    {{))+JJJJqMEOODJJ77
 	
r!   c           	        V P                  4       p\        P                  ! \        P                  ! VR4      V\        P                  ! V4      4      p\        P                  ! \        P
                  ! VR4      V\        P                  ! V4      4      p\        P                  ! \        P                  ! \        P                  ! V) 4      \        P                  ! V4      ,
          4      4      \        P                  ! \        P                  ! VR4      \        P                  ! RV,          4      \        P                  ! RV,          R,
          4      4      ,
          p\        P                  ! V P                  R,
          ^4      p\        P                  ! R4      RRV,          ,           V,          ,           p\        P                  ! V P                  4       WF4      # )zLcomputes the log normalizing constant as a function of the 'probs' parameter      ?       @      ?g       gUUUUUU?g'}'}@)rO   r+   rM   rG   
zeros_likegerN   logabslog1ppowr   mathrJ   )r1   	cut_probscut_probs_below_halfcut_probs_above_halflog_normxtaylors   &      r   _cont_bern_log_norm'ContinuousBernoulli._cont_bern_log_normx   s9   OO%	${{HHY$i1A1A)1L 
  %{{HHY$i1K 
 99IIekk9*-		)0DDE
KKHHY$KK334IIc00367

 IIdjj3&*#)lQ.>">!!CC{{488:HMMr!   c                    < V ^8  d   QhRS[ /# r   r   r   )r   r   s   "r   r   r       s     I If Ir!   c                   V P                  4       pVR V,          R,
          ,          R\        P                  ! V) 4      \        P                  ! V4      ,
          ,          ,           pV P                  R,
          pRRR\        P
                  ! V^4      ,          ,           V,          ,           p\        P                  ! V P                  4       W$4      # )rS   rT   rR   gUUUUUU?gll?)rO   r+   rY   rW   r   rZ   rM   rJ   )r1   r\   musr`   ra   s   &    r   meanContinuousBernoulli.mean   s    OO%	3?S01CKK
#eii	&::5
 
 JJ	K%))Aq/$AAQFF{{488:CHHr!   c                    < V ^8  d   QhRS[ /# re   r   )r   r   s   "r   r   r       s     ) ) )r!   c                B    \         P                  ! V P                  4      # r?   )r+   sqrtvariancerI   s   &r   stddevContinuousBernoulli.stddev   s    zz$--((r!   c                    < V ^8  d   QhRS[ /# re   r   )r   r   s   "r   r   r       s     J J& Jr!   c                   V P                  4       pWR ,
          ,          \        P                  ! R RV,          ,
          ^4      ,          R \        P                  ! \        P                  ! V) 4      \        P                  ! V4      ,
          ^4      ,          ,           p\        P                  ! V P
                  R,
          ^4      pRRRV,          ,
          V,          ,
          p\        P                  ! V P                  4       W$4      # )rT   rS   rR   gUUUUUU?g?ggjV?)rO   r+   rZ   rY   rW   r   rM   rJ   )r1   r\   varsr`   ra   s   &    r   rm   ContinuousBernoulli.variance   s    OO%	O,uyy#	/!10
 
%))EKK
3eii	6JJANNO IIdjj3&*zMA,==BB{{488:DIIr!   c                    < V ^8  d   QhRS[ /# re   r   )r   r   s   "r   r   r       s     ; ; ;r!   c                0    \        V P                  R R7      # T)	is_binary)r   r   rI   s   &r   r   ContinuousBernoulli.logits   s    tzzT::r!   c                    < V ^8  d   QhRS[ /# re   r   )r   r   s   "r   r   r       s     I Iv Ir!   c                B    \        \        V P                  R R7      4      # rv   )r   r
   r   rI   s   &r   r   ContinuousBernoulli.probs   s    ?4;;$GHHr!   c                4   < V ^8  d   QhRS[ P                  /# re   )r+   r,   )r   r   s   "r   r   r       s     " "UZZ "r!   c                6    V P                   P                  4       # r?   )r*   r-   rI   s   &r   param_shapeContinuousBernoulli.param_shape   s    {{!!r!   c                >   V P                  V4      p\        P                  ! W P                  P                  V P                  P
                  R 7      p\        P                  ! 4       ;_uu_ 4        V P                  V4      uuRRR4       #   + '       g   i     R# ; i)dtypedeviceN)_extended_shaper+   randr   r   r   no_gradicdfr1   sample_shapeshapeus   &&  r   sampleContinuousBernoulli.sample   sX    $$\2JJuJJ$4$4TZZ=N=NO]]__99Q< ___s   /BB	c                &   < V ^8  d   QhRS[ RS[/# )r   r   r   )r   r   )r   r   s   "r   r   r       s      E V r!   c                    V P                  V4      p\        P                  ! W P                  P                  V P                  P
                  R 7      pV P                  V4      # )r   )r   r+   r   r   r   r   r   r   s   &&  r   rsampleContinuousBernoulli.rsample   sD    $$\2JJuJJ$4$4TZZ=N=NOyy|r!   c                    V P                   '       d   V P                  V4       \        V P                  V4      w  r!\	        W!R R7      ) V P                  4       ,           # )none)	reduction)r:   _validate_sampler   r   r   rb   )r1   valuer   s   && r   log_probContinuousBernoulli.log_prob   sQ    !!%(%dkk59-fvNN&&()	
r!   c           
     x   V P                   '       d   V P                  V4       V P                  4       p\        P                  ! W!4      \        P                  ! R V,
          R V,
          4      ,          V,           R ,
          RV,          R ,
          ,          p\        P
                  ! V P                  4       W14      p\        P
                  ! \        P                  ! VR4      \        P                  ! V4      \        P
                  ! \        P                  ! VR 4      \        P                  ! V4      V4      4      # )rT   rS   g        )r:   r   rO   r+   rZ   rM   rJ   rG   rU   rV   rN   )r1   r   r\   cdfsunbounded_cdfss   &&   r   cdfContinuousBernoulli.cdf   s    !!%(OO%	IIi'%))C)OS5[*QQ 9_s"	$
 T%B%B%DdR{{HHUC U#KK,eooe.DnU
 	
r!   c           	     v   V P                  4       p\        P                  ! V P                  4       \        P                  ! V) VR V,          R,
          ,          ,           4      \        P                  ! V) 4      ,
          \        P
                  ! V4      \        P                  ! V) 4      ,
          ,          V4      # )rS   rT   )rO   r+   rM   rJ   rY   rW   )r1   r   r\   s   && r   r   ContinuousBernoulli.icdf   s    OO%	{{))+YJ#	/C2G)HHI++yj)* yy#ekk9*&==	?
 
 	
r!   c                    \         P                  ! V P                  ) 4      p\         P                  ! V P                  4      pV P                  W,
          ,          V P                  4       ,
          V,
          # r?   )r+   rY   r   rW   rh   rb   )r1   
log_probs0
log_probs1s   &  r   entropyContinuousBernoulli.entropy   sU    [[$**-
YYtzz*
II01&&()	
r!   c                0   < V ^8  d   QhRS[ S[,          /# re   )r   r   )r   r   s   "r   r   r       s      v r!   c                    V P                   3# r?   )r   rI   s   &r   _natural_params#ContinuousBernoulli._natural_params   s    ~r!   c                ,   \         P                  ! \         P                  ! WP                  ^ ,          R,
          4      \         P                  ! WP                  ^,          R,
          4      4      p\         P
                  ! W!V P                  ^ ,          R,
          \         P                  ! V4      ,          4      p\         P                  ! \         P                  ! \         P                  P                  V4      4      4      \         P                  ! \         P                  ! V4      4      ,
          pRV,          \         P                  ! V^4      R,          ,           \         P                  ! V^4      R,          ,
          p\         P
                  ! W$V4      # )zLcomputes the log normalizing constant as a function of the natural parameterrR   g      8@g     @)r+   rF   rG   r.   rH   rM   rN   rW   rX   specialexpm1rZ   )r1   r`   out_unst_regcut_nat_paramsr_   ra   s   &&    r   _log_normalizer#ContinuousBernoulli._log_normalizer   s    yyHHQ

1+,ehhq**Q-#:M.N
 djjmc1U__Q5GG
 99IIemm)).9:
IIeii/01 q599Q?T11EIIaOf4LL{{<6::r!   )r.   r*   r   r   )NN)gV-?gx&1?Nr?   )'__name__
__module____qualname____firstlineno____doc__r   unit_intervalrealr&   support_mean_carrier_measurehas_rsampler0   r9   rB   rJ   rO   rb   propertyrh   rn   rm   r	   r   r   r~   r+   r,   r   r   r   r   r   r   r   r   __static_attributes____classdictcell____classcell__)r4   r   s   @@r   r   r      s8    6  9 98[EUEUVO''GK"C "CH0


N( I I ) ) J J ; ; I I " " #(**,   -2JJL  


 


  ; ;r!   )r[   r+   r   torch.distributionsr   torch.distributions.exp_familyr   torch.distributions.utilsr   r   r	   r
   r   torch.nn.functionalr   torch.typesr   r   r   __all__r    r!   r   <module>r      sC       + <  A . . !
!d;+ d;r!   