§
    ŠŠtj=#  ã                   óˆ   — d dl Z d dl mZ d dlmZ d dlmZ d dlmZ d dlm	Z	m
Z
 d dlmZ d dlmZ d	gZ G d
„ d	e¦  «        ZdS )é    N)ÚTensor)Úconstraints)ÚDistribution)ÚIndependent)ÚComposeTransformÚ	Transform)Ú_sum_rightmost)Ú_sizeÚTransformedDistributionc            	       ó@  ‡ — e Zd ZU dZi Zeeej        f         e	d<   	 dde
deee         z  dedz  ddfˆ fd„Zdˆ fd	„	Z ej        d
¬¦  «        d„ ¦   «         Zedefd„¦   «         Z ej        ¦   «         fd„Z ej        ¦   «         fdedefd„Zd„ Zd„ Zd„ Zd„ Zˆ xZS )r   a±  
    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                 óZ  •— t          |t          ¦  «        r	|g| _        nWt          |t          ¦  «        r0t	          d„ |D ¦   «         ¦  «        st          d¦  «        ‚|| _        nt          d|› �¦  «        ‚|j        |j        z   }t          |j        ¦  «        }t          | j        ¦  «        }t          |¦  «        |j
        j        k     r t          d|j
        j        › d|› d�¦  «        ‚|                     |¦  «        }|                     |¦  «        }||k    r/|d t          |¦  «        |z
  …         }	|                     |	¦  «        }|j
        j        |z
  }
|
dk    rt          ||
¦  «        }|| _        |j        j        |j
        j        z
  }t%          |j        j        ||z   ¦  «        }t          |¦  «        |k     r"t'          dt          |¦  «        › d	|› �¦  «        ‚t          |¦  «        |z
  }|d |…         }||d …         }t)          ¦   «                              |||¬
¦  «         d S )Nc              3   ó@   K  — | ]}t          |t          ¦  «        V — Œd S ©N)Ú
isinstancer   )Ú.0Úts     új/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/torch/distributions/transformed_distribution.pyú	<genexpr>z3TransformedDistribution.__init__.<locals>.<genexpr>?   s,   è è € ÐDÐD°A•z !¥YÑ/Ô/ÐDÐDÐDÐDÐDÐDó    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-   z TransformedDistribution.__init__4   sŒ  ø€ õ �j¥)Ñ,Ô,ð 	àðˆDŒOˆOõ ˜
¥DÑ)Ô)ð 		ÝÐDÐD¸ÐDÑDÔDÑDÔDð Ý ØLñô ð ð )ˆDŒOˆOåØOÀ:ÐOÐOñô ð ð
 'Ô2Ð5FÔ5RÑRˆ
ÝÐ.Ô:Ñ;Ô;ˆÝ$ T¤_Ñ5Ô5ˆ	Ýˆz‰?Œ?˜YÔ-Ô7Ò7Ð7ÝØÈIÔL\ÔLfÐÐÐr|ÐÐÐñô ð ð "×/Ò/°
Ñ;Ô;ˆØ'×5Ò5°mÑDÔDÐØÐ,Ò,Ð,Ø2Ø;•#Ð)Ñ*Ô*¨^Ñ;Ð;ô Ðð !2× 8Ò 8Ð9IÑ JÔ JÐØ$-Ô$4Ô$>ÀÑ$OÐ!Ø$ qÒ(Ð(Ý +Ø!Ð#<ñ!ô !Ðð +ˆŒð ÔÔ(¨9Ô+;Ô+EÑEð 	&õ ØÔÔ(ØÐ:Ñ:ñ
ô 
ˆ	õ ˆ}ÑÔ 	Ò)Ð)Ý Ø]­¨MÑ(:Ô(:Ð]Ð]ÐR[Ð]Ð]ñô ð õ �-Ñ Ô  9Ñ,ˆØ# D S DÔ)ˆØ# C D DÔ)ˆÝ‰Œ×Ò˜ kÀÐÑOÔOÐOÐOÐOr   c                 ó  •— |                       t          |¦  «        }t          j        |¦  «        }|| j        z   }t          | j        ¦  «        D ]}|                     |¦  «        }Œ|d t          |¦  «        t          | j	        j        ¦  «        z
  …         }| j	         
                    |¦  «        |_	        | j        |_        t          t          |¦  «                             || j        d¬¦  «         | j        |_        |S )NFr   )Ú_get_checked_instancer   ÚtorchÚSizer!   Úreversedr   r&   r"   r(   r'   r,   r-   Ú_validate_args)r.   r    Ú	_instanceÚnewÚshaper   r3   r7   s          €r   r'   zTransformedDistribution.expandp   sï   ø€ Ø×(Ò(Õ)@À)ÑLÔLˆÝ”j Ñ-Ô-ˆØ˜dÔ.Ñ.ˆÝ˜$œ/Ñ*Ô*ð 	+ð 	+ˆAØ—O’O EÑ*Ô*ˆEˆEØ Ð!O¥3 u¡:¤:µ°D´NÔ4NÑ0OÔ0OÑ#OÐ!OÔPÐØœ×-Ò-Ð.>Ñ?Ô?ˆŒØœˆŒÝÕ% sÑ+Ô+×4Ò4Ø˜Ô)¸ð 	5ñ 	
ô 	
ð 	
ð "Ô0ˆÔØˆ
r   F)Úis_discretec                 óè   — | j         s| j        j        S | j         d         j        }t	          | j        ¦  «        |j        k    r/t          j        |t	          | j        ¦  «        |j        z
  ¦  «        }|S )Néÿÿÿÿ)	r   r(   Úsupportr)   r"   r!   r$   r   Úindependent)r.   rD   s     r   rD   zTransformedDistribution.support   sr   € ð Œð 	*Ø”>Ô)Ð)Ø”/ "Ô%Ô.ˆÝˆtÔÑ Ô  7Ô#4Ò4Ð4Ý!Ô-Ø�˜TÔ-Ñ.Ô.°Ô1BÑBñô ˆGð ˆr   c                 ó   — | j         j        S r   )r(   Úhas_rsample)r.   s    r   rG   z#TransformedDistribution.has_rsample‹   s   € àŒ~Ô)Ð)r   c                 ó¾   — t          j        ¦   «         5  | j                             |¦  «        }| j        D ]} ||¦  «        }Œ|cddd¦  «         S # 1 swxY w Y   dS )a  
        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Úxr1   s       r   rJ   zTransformedDistribution.sample�   s¬   € õ Œ]‰_Œ_ð 	ð 	Ø”×%Ò% lÑ3Ô3ˆAØ!œ_ð !ð !�	Ø�I˜a‘L”L��Øð		ð 	ð 	ð 	ñ 	ô 	ð 	ð 	ð 	ð 	ð 	ð 	øøøð 	ð 	ð 	ð 	ð 	ð 	s   ”1AÁAÁArL   c                 ód   — | j                              |¦  «        }| j        D ]} ||¦  «        }Œ|S )a$  
        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   rK   s       r   rO   zTransformedDistribution.rsampleœ   s>   € ð ŒN×"Ò" <Ñ0Ô0ˆØœð 	ð 	ˆIØ�	˜!‘”ˆAˆAØˆr   c                 óô  — | j         r|                      |¦  «         t          | j        ¦  «        }d}|}t	          | j        ¦  «        D ]i}|                     |¦  «        }||j        j        |j	        j        z
  z  }|t          |                     ||¦  «        ||j        j        z
  ¦  «        z
  }|}Œj|t          | j                             |¦  «        |t          | j        j        ¦  «        z
  ¦  «        z   }|S )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$   rT   Úyr1   rM   s          r   rT   z TransformedDistribution.log_prob¨   s  € ð
 Ôð 	)Ø×!Ò! %Ñ(Ô(Ð(Ý˜Ô(Ñ)Ô)ˆ	Ø#&ˆØˆÝ! $¤/Ñ2Ô2ð 	ð 	ˆIØ—’˜aÑ Ô ˆAØ˜Ô)Ô3°iÔ6HÔ6RÑRÑRˆIØ¥.Ø×.Ò.¨q°!Ñ4Ô4Ø˜IÔ,Ô6Ñ6ñ#ô #ñ ˆHð ˆAˆAà�nØŒN×#Ò# AÑ&Ô&¨	µC¸¼Ô8RÑ4SÔ4SÑ(Sñ
ô 
ñ 
ˆð ˆr   c                 ó~   — d}| j         D ]}||j        z  }Œt          |t          ¦  «        r|dk    r|S ||dz
  z  dz   S )zu
        This conditionally flips ``value -> 1-value`` to ensure :meth:`cdf` is
        monotone increasing.
        é   g      à?)r   Úsignr   Úint)r.   rU   rY   r1   s       r   Ú_monotonize_cdfz'TransformedDistribution._monotonize_cdfÀ   s[   € ð
 ˆØœð 	)ð 	)ˆIØ˜)œ.Ñ(ˆDˆDÝ�d�CÑ Ô ð 	 T¨Q¢Y YØˆLØ�u˜s‘{Ñ# cÑ)Ð)r   c                 óö   — | j         ddd…         D ]}|                     |¦  «        }Œ| j        r| j                             |¦  «         | j                             |¦  «        }|                      |¦  «        }|S )z—
        Computes the cumulative distribution function by inverting the
        transform(s) and computing the score of the base distribution.
        NrC   )r   rR   r=   r(   rQ   Úcdfr[   ©r.   rU   r1   s      r   r]   zTransformedDistribution.cdfÌ   s�   € ð
 œ¨¨¨2¨Ô.ð 	)ð 	)ˆIØ—M’M %Ñ(Ô(ˆEˆEØÔð 	3ØŒN×+Ò+¨EÑ2Ô2Ð2Ø”×"Ò" 5Ñ)Ô)ˆØ×$Ò$ UÑ+Ô+ˆØˆr   c                 óŽ   — |                       |¦  «        }| j                             |¦  «        }| j        D ]} ||¦  «        }Œ|S )z”
        Computes the inverse cumulative distribution function using
        transform(s) and computing the score of the base distribution.
        )r[   r(   Úicdfr   r^   s      r   r`   zTransformedDistribution.icdfÙ   sS   € ð
 ×$Ò$ UÑ+Ô+ˆØ”×#Ò# EÑ*Ô*ˆØœð 	%ð 	%ˆIØ�I˜eÑ$Ô$ˆEˆEØˆr   r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚdictÚstrr   Ú
ConstraintÚ__annotations__r   r   r   Úboolr-   r'   Údependent_propertyrD   ÚpropertyrG   r:   r;   rJ   r
   r   rO   rT   r[   r]   r`   Ú__classcell__)r7   s   @r   r   r      s´  ø€ € € € € € ðð ðB :<€O�T˜#˜{Ô5Ð5Ô6Ð;Ð;Ñ;ð &*ð	:Pð :Pà'ð:Pð   Y¤Ñ/ð:Pð ˜d‘{ð	:Pð
 
ð:Pð :Pð :Pð :Pð :Pð :Pðxð ð ð ð ð ð $€[Ô#°Ð6Ñ6Ô6ðð ñ 7Ô6ðð ð*˜Tð *ð *ð *ñ „Xð*ð #- %¤*¡,¤,ð ð ð ð ð -7¨E¬J©L¬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>ru      sÜ   ðð €€€Ø Ð Ð Ð Ð Ð Ø +Ð +Ð +Ð +Ð +Ð +Ø 9Ð 9Ð 9Ð 9Ð 9Ð 9Ø 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ 4Ð 4Ð 4Ð 4Ð 4Ð 4Ø Ð Ð Ð Ð Ð ð %Ð
%€ðRð Rð Rð Rð R˜lñ Rô Rð Rð Rð Rr   