ó
    Eñi=#  ã                   ó€   • S SK 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
  S SKJr  S SKJr  S	/r " S
 S	\5      rg)é    N)ÚTensor)Úconstraints)ÚDistribution)ÚIndependent)ÚComposeTransformÚ	Transform)Ú_sum_rightmost)Ú_sizeÚTransformedDistributionc            	       óR  ^ • \ rS rSr% Sr0 r\\\R                  4   \
S'    SS\S\\\   -  S\S-  SS4U 4S	 jjjrSU 4S
 jjr\R"                  " SS9S 5       r\S\4S j5       r\R,                  " 5       4S jr\R,                  " 5       4S\S\4S jjrS rS rS rS rSrU =r $ )r   é   aI  
Extension of the Distribution class, which applies a sequence of Transforms
to a base distribution.  Let f be the composition of transforms applied::

    X ~ BaseDistribution
    Y = f(X) ~ TransformedDistribution(BaseDistribution, f)
    log p(Y) = log p(X) + log |det (dX/dY)|

Note that the ``.event_shape`` of a :class:`TransformedDistribution` is the
maximum shape of its base distribution and its transforms, since transforms
can introduce correlations among events.

An example for the usage of :class:`TransformedDistribution` would be::

    # Building a Logistic Distribution
    # X ~ Uniform(0, 1)
    # f = a + b * logit(X)
    # Y ~ f(X) ~ Logistic(a, b)
    base_distribution = Uniform(0, 1)
    transforms = [SigmoidTransform().inv, AffineTransform(loc=a, scale=b)]
    logistic = TransformedDistribution(base_distribution, transforms)

For more examples, please look at the implementations of
:class:`~torch.distributions.gumbel.Gumbel`,
:class:`~torch.distributions.half_cauchy.HalfCauchy`,
:class:`~torch.distributions.half_normal.HalfNormal`,
:class:`~torch.distributions.log_normal.LogNormal`,
:class:`~torch.distributions.pareto.Pareto`,
:class:`~torch.distributions.weibull.Weibull`,
:class:`~torch.distributions.relaxed_bernoulli.RelaxedBernoulli` and
:class:`~torch.distributions.relaxed_categorical.RelaxedOneHotCategorical`
Úarg_constraintsNÚbase_distributionÚ
transformsÚvalidate_argsÚreturnc                 ó  >• [        U[        5      (       a	  U/U l        OL[        U[        5      (       a)  [	        S U 5       5      (       d  [        S5      eX l        O[        SU 35      eUR                  UR                  -   n[        UR                  5      n[        U R                  5      n[        U5      UR                  R                  :  a&  [        SUR                  R                   SU S35      eUR                  U5      nUR                  U5      nXH:w  a"  US [        U5      U-
   n	UR                  U	5      nUR                  R                  U-
  n
U
S:”  a  [        X5      nXl        UR"                  R                  UR                  R                  -
  n[%        UR"                  R                  X[-   5      n[        U5      U:  a  ['        S[        U5       S	U 35      e[        U5      U-
  nUS U nX}S  n[(        TU ]U  XïUS
9  g )Nc              3   óB   #   • U  H  n[        U[        5      v •  M     g 7f©N)Ú
isinstancer   )Ú.0Úts     Úi/home/mande/repo/quber/.venv/lib/python3.13/site-packages/torch/distributions/transformed_distribution.pyÚ	<genexpr>Ú3TransformedDistribution.__init__.<locals>.<genexpr>?   s   é € ÐDº°A”z !¤Y×/Ð/ºùs   ‚z6transforms must be a Transform or a list of Transformsz0transforms must be a Transform or list, but was z9base_distribution needs to have shape with size at least z
, but got Ú.r   zforward_shape length z must be >= event_dim ©r   )r   r   r   ÚlistÚallÚ
ValueErrorÚbatch_shapeÚevent_shapeÚlenr   ÚdomainÚ	event_dimÚforward_shapeÚinverse_shapeÚexpandr   Ú	base_distÚcodomainÚmaxÚAssertionErrorÚsuperÚ__init__)Úselfr   r   r   Ú
base_shapeÚbase_event_dimÚ	transformr&   Úexpanded_base_shapeÚbase_batch_shapeÚreinterpreted_batch_ndimsÚtransform_change_in_event_dimr%   Úcutr!   r"   Ú	__class__s                   €r   r.   Ú TransformedDistribution.__init__4   s.  ø€ ô �j¤)×,Ñ,àðˆD�Oô ˜
¤D×)Ñ)ÜÑD¹ÓD×DÑDÜ ØLóð ð )�OäØBÀ:À,ÐOóð ð
 '×2Ñ2Ð5F×5RÑ5RÑRˆ
ÜÐ.×:Ñ:Ó;ˆÜ$ T§_¡_Ó5ˆ	Üˆz‹?˜Y×-Ñ-×7Ñ7Ó7ÜØKÈI×L\ÑL\×LfÑLfÐKgÐgqÐr|Ðq}Ð}~Ðóð ð "×/Ñ/°
Ó;ˆØ'×5Ñ5°mÓDÐØÓ,Ø2Ø;”#Ð)Ó*¨^Ñ;ð Ðð !2× 8Ñ 8Ð9IÓ JÐØ$-×$4Ñ$4×$>Ñ$>ÀÑ$OÐ!Ø$ qÓ(Ü +Ø!ó!Ðð +Œð ×Ñ×(Ñ(¨9×+;Ñ+;×+EÑ+EÑEð 	&ô Ø×Ñ×(Ñ(ØÑ:ó
ˆ	ô ˆ}Ó 	Ó)Ü Ø'¬¨MÓ(:Ð';Ð;QÐR[ÐQ\Ð]óð ô �-Ó  9Ñ,ˆØ# D SÐ)ˆØ# DÐ)ˆÜ‰Ñ˜ÀÐÒOó    c                 óî  >• U R                  [        U5      n[        R                  " U5      nXR                  -   n[        U R                  5       H  nUR                  U5      nM     US [        U5      [        U R                  R                  5      -
   nU R                  R                  U5      Ul	        U R                  Ul        [        [        U]3  XR                  SS9  U R                  Ul        U$ )NFr   )Ú_get_checked_instancer   ÚtorchÚSizer"   Úreversedr   r'   r#   r)   r(   r-   r.   Ú_validate_args)r/   r!   Ú	_instanceÚnewÚshaper   r4   r8   s          €r   r(   ÚTransformedDistribution.expandp   sÐ   ø€ Ø×(Ñ(Ô)@À)ÓLˆÜ—j’j Ó-ˆØ×.Ñ.Ñ.ˆÜ˜$Ÿ/™/Ö*ˆAØ—O‘O EÓ*ŠEñ +à Ð!O¤3 u£:´°D·N±N×4NÑ4NÓ0OÑ#OÐPÐØŸ™×-Ñ-Ð.>Ó?ˆŒØŸ™ˆŒÜÔ% sÑ4Ø×)Ñ)¸ð 	5ñ 	
ð "×0Ñ0ˆÔØˆ
r:   F)Úis_discretec                 ó:  • U R                   (       d  U R                  R                  $ U R                   S   R                  n[	        U R
                  5      UR                  :”  a7  [        R                  " U[	        U R
                  5      UR                  -
  5      nU$ )Néÿÿÿÿ)	r   r)   Úsupportr*   r#   r"   r%   r   Úindependent)r/   rH   s     r   rH   ÚTransformedDistribution.support   sz   € ð ��Ø—>‘>×)Ñ)Ð)Ø—/‘/ "Ñ%×.Ñ.ˆÜˆt×ÑÓ  7×#4Ñ#4Ó4Ü!×-Ò-Øœ˜T×-Ñ-Ó.°×1BÑ1BÑBóˆGð ˆr:   c                 ó.   • U R                   R                  $ r   )r)   Úhas_rsample)r/   s    r   rL   Ú#TransformedDistribution.has_rsample‹   s   € à�~‰~×)Ñ)Ð)r:   c                 óÒ   • [         R                  " 5          U R                  R                  U5      nU R                   H  nU" U5      nM     UsSSS5        $ ! , (       d  f       g= f)zÜ
Generates a sample_shape shaped sample or sample_shape shaped batch of
samples if the distribution parameters are batched. Samples first from
base distribution and applies `transform()` for every transform in the
list.
N)r=   Úno_gradr)   Úsampler   ©r/   Úsample_shapeÚxr2   s       r   rP   ÚTransformedDistribution.sample�   sE   € ô �]Š]�_Ø—‘×%Ñ% lÓ3ˆAØ!Ÿ_œ_�	Ù˜a“L’ñ -à÷	 �_�_ús   –8AÁ
A&rR   c                 ór   • U R                   R                  U5      nU R                   H  nU" U5      nM     U$ )zü
Generates a sample_shape shaped reparameterized sample or sample_shape
shaped batch of reparameterized samples if the distribution parameters
are batched. Samples first from base distribution and applies
`transform()` for every transform in the list.
)r)   Úrsampler   rQ   s       r   rV   ÚTransformedDistribution.rsampleœ   s4   € ð �N‰N×"Ñ" <Ó0ˆØŸœˆIÙ˜!“ŠAñ )àˆr:   c                 ó0  • U R                   (       a  U R                  U5        [        U R                  5      nSnUn[	        U R
                  5       Hy  nUR                  U5      nX%R                  R                  UR                  R                  -
  -  nU[        UR                  Xd5      X%R                  R                  -
  5      -
  nUnM{     U[        U R                  R                  U5      U[        U R                  R                  5      -
  5      -   nU$ )z�
Scores the sample by inverting the transform(s) and computing the score
using the score of the base distribution and the log abs det jacobian.
g        )r@   Ú_validate_sampler#   r"   r?   r   Úinvr$   r%   r*   r	   Úlog_abs_det_jacobianr)   Úlog_prob)r/   Úvaluer%   r\   Úyr2   rS   s          r   r\   Ú TransformedDistribution.log_prob¨   sõ   € ð
 ××Ø×!Ñ! %Ô(Ü˜×(Ñ(Ó)ˆ	Ø#&ˆØˆÜ! $§/¡/Ö2ˆIØ—‘˜aÓ ˆAØ×)Ñ)×3Ñ3°i×6HÑ6H×6RÑ6RÑRÑRˆIØ¤.Ø×.Ñ.¨qÓ4Ø×,Ñ,×6Ñ6Ñ6ó#ñ ˆHð ŠAñ 3ð œnØ�N‰N×#Ñ# AÓ&¨	´C¸¿¹×8RÑ8RÓ4SÑ(Só
ñ 
ˆð ˆr:   c                 ó–   • SnU R                    H  nX#R                  -  nM     [        U[        5      (       a  US:X  a  U$ X!S-
  -  S-   $ )z]
This conditionally flips ``value -> 1-value`` to ensure :meth:`cdf` is
monotone increasing.
é   g      à?)r   Úsignr   Úint)r/   r]   rb   r2   s       r   Ú_monotonize_cdfÚ'TransformedDistribution._monotonize_cdfÀ   sM   € ð
 ˆØŸœˆIØŸ.™.Ñ(ŠDñ )ä�dœC× Ñ  T¨Q£YØˆLØ˜s‘{Ñ# cÑ)Ð)r:   c                 ó
  • U R                   SSS2    H  nUR                  U5      nM     U R                  (       a  U R                  R	                  U5        U R                  R                  U5      nU R                  U5      nU$ )z
Computes the cumulative distribution function by inverting the
transform(s) and computing the score of the base distribution.
NrG   )r   rZ   r@   r)   rY   Úcdfrd   ©r/   r]   r2   s      r   rg   ÚTransformedDistribution.cdfÌ   sm   € ð
 Ÿ™©¨2¨Ô.ˆIØ—M‘M %Ó(ŠEñ /à××Ø�N‰N×+Ñ+¨EÔ2Ø—‘×"Ñ" 5Ó)ˆØ×$Ñ$ UÓ+ˆØˆr:   c                 ó”   • U R                  U5      nU R                  R                  U5      nU R                   H  nU" U5      nM     U$ )z|
Computes the inverse cumulative distribution function using
transform(s) and computing the score of the base distribution.
)rd   r)   Úicdfr   rh   s      r   rk   ÚTransformedDistribution.icdfÙ   sE   € ð
 ×$Ñ$ UÓ+ˆØ—‘×#Ñ# EÓ*ˆØŸœˆIÙ˜eÓ$ŠEñ )àˆr:   )r)   r   r   )!Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__r   ÚdictÚstrr   Ú
ConstraintÚ__annotations__r   r   r   Úboolr.   r(   Údependent_propertyrH   ÚpropertyrL   r=   r>   rP   r
   r   rV   r\   rd   rg   rk   Ú__static_attributes__Ú__classcell__)r8   s   @r   r   r      sô   ø‡ ñðB :<€O�T˜#˜{×5Ñ5Ð5Ñ6Ó;ð &*ñ	:Pà'ð:Pð   Y¡Ñ/ð:Pð ˜d‘{ð	:Pð
 
÷:Pð :P÷xð ×#Ò#°Ñ6ñó 7ðð ð*˜Tó *ó ð*ð #(§*¢*£,ô ð -2¯JªJ«Lñ 
 Eð 
¸Võ 
òò0
*ò÷	ð 	r:   )r=   r   Útorch.distributionsr   Ú torch.distributions.distributionr   Útorch.distributions.independentr   Útorch.distributions.transformsr   r   Útorch.distributions.utilsr	   Útorch.typesr
   Ú__all__r   © r:   r   Ú<module>rƒ      s7   ðó Ý Ý +Ý 9Ý 7ß FÝ 4Ý ð %Ð
%€ôR˜lõ Rr:   