+
    &j                         ^ RI t ^ RI Ht ^ RIH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HtHt R	R
.t ! R R	]4      t ! R R
]4      tR# )    NTensor)constraints)Distribution)TransformedDistribution)SigmoidTransform)broadcast_allclamp_probslazy_propertylogits_to_probsprobs_to_logits)_Number_sizeNumberLogitRelaxedBernoulliRelaxedBernoullic                   8  a a ] tR t^t oRtR]P                  R]P                  /t]P                  t	RV3R lV 3R lllt
RV 3R ll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]P&                  ! 4       3V3R lR lltR tRtVtV ;t# )r   aO  
Creates a LogitRelaxedBernoulli distribution parameterized by :attr:`probs`
or :attr:`logits` (but not both), which is the logit of a RelaxedBernoulli
distribution.

Samples are logits of values in (0, 1). See [1] for more details.

Args:
    temperature (Tensor): relaxation temperature
    probs (Number, Tensor): the probability of sampling `1`
    logits (Number, Tensor): the log-odds of sampling `1`

[1] The Concrete Distribution: A Continuous Relaxation of Discrete Random
Variables (Maddison et al., 2017)

[2] Categorical Reparametrization with Gumbel-Softmax
(Jang et al., 2017)
probslogitsc          
         < V ^8  d   QhRS[ RS[ S[,          R,          RS[ S[,          R,          RS[R,          RR/#    temperaturer   Nr   validate_argsreturnr   r   bool)format__classdict__s   "}/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/distributions/relaxed_bernoulli.py__annotate__"LogitRelaxedBernoulli.__annotate__.   sZ     C CC %C $&	C
 d{C 
C    c                  < Wn         VR J VR J 8X  d   \        R4      hVe$   \        V\        4      p\	        V4      w  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\        SV `5  WdR7       R # )Nz;Either `probs` or `logits` must be specified, but not both.zlogits is unexpectedly Noner   )r   
ValueError
isinstancer   r	   r   AssertionErrorr   _paramtorchSizesizesuper__init__)selfr   r   r   r   	is_scalarbatch_shape	__class__s   &&&&&  r    r.   LogitRelaxedBernoulli.__init__.   s     'TMv~.M  "5'2I)%0MTZ~$%BCC"673I*62NT[$)$5djj4;;**,K++**,KBr#   c                  < V P                  \        V4      p\        P                  ! V4      pV P                  Vn        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-   r.   _validate_argsr/   r1   	_instancenewr2   s   &&& r    r7   LogitRelaxedBernoulli.expandK   s    (()>	Jjj-**dmm#

))+6CICJt}}$++K8CJCJ#S2;e2T!00
r#   c                :    V P                   P                  ! V/ VB # N)r)   r;   )r/   argskwargss   &*,r    _newLogitRelaxedBernoulli._newY   s    {{///r#   c                    < V ^8  d   QhRS[ /# r   r   r   )r   r   s   "r    r!   r"   ]   s     ; ; ;r#   c                0    \        V P                  R R7      # T)	is_binary)r   r   r/   s   &r    r   LogitRelaxedBernoulli.logits\   s    tzzT::r#   c                    < V ^8  d   QhRS[ /# rD   r   )r   r   s   "r    r!   r"   a   s     < <v <r#   c                0    \        V P                  R R7      # rF   )r   r   rH   s   &r    r   LogitRelaxedBernoulli.probs`   s    t{{d;;r#   c                4   < V ^8  d   QhRS[ P                  /# rD   )r*   r+   )r   r   s   "r    r!   r"   e   s     " "UZZ "r#   c                6    V P                   P                  4       # r>   )r)   r,   rH   s   &r    param_shape!LogitRelaxedBernoulli.param_shaped   s    {{!!r#   c                &   < V ^8  d   QhRS[ RS[/# )r   sample_shaper   )r   r   )r   r   s   "r    r!   r"   h   s      E V r#   c                   V P                  V4      p\        V P                  P                  V4      4      p\        \        P
                  ! W#P                  VP                  R 7      4      pVP                  4       V) P                  4       ,
          VP                  4       ,           V) P                  4       ,
          V P                  ,          # ))dtypedevice)_extended_shaper
   r   r7   r*   randrT   rU   loglog1pr   )r/   rR   shaper   uniformss   &&   r    rsampleLogitRelaxedBernoulli.rsampleh   s    $$\2DJJ--e45JJuKKE
 LLNxi..00599;>5&AQQ 	r#   c                P   V P                   '       d   V P                  V4       \        V P                  V4      w  r!W!P	                  V P
                  4      ,
          pV P
                  P                  4       V,           ^VP                  4       P                  4       ,          ,
          # )r   )	r8   _validate_sampler	   r   mulr   rX   exprY   )r/   valuer   diffs   &&  r    log_probLogitRelaxedBernoulli.log_probr   sx    !!%(%dkk59		$"2"233##%,q488:3C3C3E/EEEr#   )r)   r   r   r   NNNr>   )__name__
__module____qualname____firstlineno____doc__r   unit_intervalrealarg_constraintssupportr.   r7   rA   r   r   r   propertyrO   r*   r+   r\   rd   __static_attributes____classdictcell____classcell__r2   r   s   @@r    r   r      s     (  9 98[EUEUVOGC C:0 ; ; < < " " -2JJL  F Fr#   c                     a a ] tR t^zt oRtR]P                  R]P                  /t]P                  t	Rt
RV3R lV 3R llltRV 3R ll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tRtVtV ;t# )r   a  
Creates a RelaxedBernoulli distribution, parametrized by
:attr:`temperature`, and either :attr:`probs` or :attr:`logits`
(but not both). This is a relaxed version of the `Bernoulli` distribution,
so the values are in (0, 1), and has reparametrizable samples.

Example::

    >>> # xdoctest: +IGNORE_WANT("non-deterministic")
    >>> m = RelaxedBernoulli(torch.tensor([2.2]),
    ...                      torch.tensor([0.1, 0.2, 0.3, 0.99]))
    >>> m.sample()
    tensor([ 0.2951,  0.3442,  0.8918,  0.9021])

Args:
    temperature (Tensor): relaxation temperature
    probs (Number, Tensor): the probability of sampling `1`
    logits (Number, Tensor): the log-odds of sampling `1`
r   r   Tc          
         < V ^8  d   QhRS[ RS[ S[,          R,          RS[ S[,          R,          RS[R,          RR/# r   r   )r   r   s   "r    r!   RelaxedBernoulli.__annotate__   sZ     U UU %U $&	U
 d{U 
Ur#   c                T   < \        WV4      p\        SV `	  V\        4       VR 7       R# )r%   N)r   r-   r.   r   )r/   r   r   r   r   	base_distr2   s   &&&&& r    r.   RelaxedBernoulli.__init__   s)     *+fE	$4$6mTr#   c                P   < V P                  \        V4      p\        SV `  WR 7      # ))r:   )r5   r   r-   r7   r9   s   &&& r    r7   RelaxedBernoulli.expand   s'    (()99Ew~k~99r#   c                    < V ^8  d   QhRS[ /# rD   r   )r   r   s   "r    r!   rw      s     * *V *r#   c                .    V P                   P                  # r>   )ry   r   rH   s   &r    r   RelaxedBernoulli.temperature   s    ~~)))r#   c                    < V ^8  d   QhRS[ /# rD   r   )r   r   s   "r    r!   rw      s     % % %r#   c                .    V P                   P                  # r>   )ry   r   rH   s   &r    r   RelaxedBernoulli.logits   s    ~~$$$r#   c                    < V ^8  d   QhRS[ /# rD   r   )r   r   s   "r    r!   rw      s     $ $v $r#   c                .    V P                   P                  # r>   )ry   r   rH   s   &r    r   RelaxedBernoulli.probs   s    ~~###r#   c                &   < V ^8  d   Qh/ S[ ;R&   # )r   ry   )r   )r   r   s   "r    r!   rw   z   s     4 %$5 r#    rf   r>   )rg   rh   ri   rj   rk   r   rl   rm   rn   ro   has_rsampler.   r7   rp   r   r   r   __annotate_func__rq   rr   rs   rt   s   @@r    r   r   z   s     (  9 98[EUEUVO''GKU U: * * % % $ $g  r#   )r*   r   torch.distributionsr    torch.distributions.distributionr   ,torch.distributions.transformed_distributionr   torch.distributions.transformsr   torch.distributions.utilsr	   r
   r   r   r   torch.typesr   r   r   __all__r   r   r   r#   r    <module>r      sV      + 9 P ;  / . #$6
7aFL aFH4$. 4$r#   