Ë
    *êñi6´  ã                   óz  — d dl Z d dlZd dlZd dlZd dlmZ d dlZd dlmc m	Z
 d dlmZ d dlmZ d dlmZ d dlmZmZmZmZmZ d dlmZmZ d dlmZ g d	¢Z G d
„ d«      Z G d„ de«      Z G d„ de«      Z eg «      Z G d„ de«      Z G d„ de«      Z  G d„ de«      Z! G d„ de«      Z"d„ Z# G d„ de«      Z$ G d„ de«      Z% G d„ de«      Z& G d„ d e«      Z' G d!„ d"e«      Z( G d#„ d$e«      Z) G d%„ d&e«      Z* G d'„ d(e«      Z+ G d)„ d*e«      Z, G d+„ d,e«      Z- G d-„ d.e«      Z. G d/„ d0e«      Z/ G d1„ d2e«      Z0y)3é    N)ÚSequence)ÚTensor)Úconstraints)ÚDistribution)Ú_sum_rightmostÚbroadcast_allÚlazy_propertyÚtril_matrix_to_vecÚvec_to_tril_matrix)ÚpadÚsoftplus)Ú_Number)ÚAbsTransformÚAffineTransformÚCatTransformÚComposeTransformÚCorrCholeskyTransformÚCumulativeDistributionTransformÚExpTransformÚIndependentTransformÚLowerCholeskyTransformÚPositiveDefiniteTransformÚPowerTransformÚReshapeTransformÚSigmoidTransformÚSoftplusTransformÚTanhTransformÚSoftmaxTransformÚStackTransformÚStickBreakingTransformÚ	TransformÚidentity_transformc                   óø   ‡ — e Zd ZU dZdZej                  ed<   ej                  ed<   ddeddfˆ fd„Z	d	„ Z
edefd
„«       Zedd„«       Zedefd„«       Zdd„Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zˆ xZS )r!   aï  
    Abstract class for invertable transformations with computable log
    det jacobians. They are primarily used in
    :class:`torch.distributions.TransformedDistribution`.

    Caching is useful for transforms whose inverses are either expensive or
    numerically unstable. Note that care must be taken with memoized values
    since the autograd graph may be reversed. For example while the following
    works with or without caching::

        y = t(x)
        t.log_abs_det_jacobian(x, y).backward()  # x will receive gradients.

    However the following will error when caching due to dependency reversal::

        y = t(x)
        z = t.inv(y)
        grad(z.sum(), [y])  # error because z is x

    Derived classes should implement one or both of :meth:`_call` or
    :meth:`_inverse`. Derived classes that set `bijective=True` should also
    implement :meth:`log_abs_det_jacobian`.

    Args:
        cache_size (int): Size of cache. If zero, no caching is done. If one,
            the latest single value is cached. Only 0 and 1 are supported.

    Attributes:
        domain (:class:`~torch.distributions.constraints.Constraint`):
            The constraint representing valid inputs to this transform.
        codomain (:class:`~torch.distributions.constraints.Constraint`):
            The constraint representing valid outputs to this transform
            which are inputs to the inverse transform.
        bijective (bool): Whether this transform is bijective. A transform
            ``t`` is bijective iff ``t.inv(t(x)) == x`` and
            ``t(t.inv(y)) == y`` for every ``x`` in the domain and ``y`` in
            the codomain. Transforms that are not bijective should at least
            maintain the weaker pseudoinverse properties
            ``t(t.inv(t(x)) == t(x)`` and ``t.inv(t(t.inv(y))) == t.inv(y)``.
        sign (int or Tensor): For bijective univariate transforms, this
            should be +1 or -1 depending on whether transform is monotone
            increasing or decreasing.
    FÚdomainÚcodomainÚ
cache_sizeÚreturnNc                 óz   •— || _         d | _        |dk(  rn|dk(  rd| _        nt        d«      ‚t        ‰| �  «        y )Nr   é   )NNzcache_size must be 0 or 1)Ú_cache_sizeÚ_invÚ_cached_x_yÚ
ValueErrorÚsuperÚ__init__)Úselfr&   Ú	__class__s     €ú`/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/torch/distributions/transforms.pyr/   zTransform.__init__a   sB   ø€ Ø%ˆÔØ=AˆŒ	Ø˜Š?ØØ˜1Š_Ø)ˆDÕäÐ8Ó9Ð9Ü‰ÑÕó    c                 óD   — | j                   j                  «       }d |d<   |S )Nr+   )Ú__dict__Úcopy)r0   Ústates     r2   Ú__getstate__zTransform.__getstate__l   s"   € Ø—‘×"Ñ"Ó$ˆØˆˆf‰Øˆr3   c                 óž   — | j                   j                  | j                  j                  k(  r| j                   j                  S t        d«      ‚)Nz:Please use either .domain.event_dim or .codomain.event_dim)r$   Ú	event_dimr%   r-   ©r0   s    r2   r:   zTransform.event_dimq   s:   € à�;‰;× Ñ  D§M¡M×$;Ñ$;Ò;Ø—;‘;×(Ñ(Ð(ÜÐUÓVÐVr3   c                 ó�   — d}| j                   �| j                  «       }|€%t        | «      }t        j                  |«      | _         |S )z{
        Returns the inverse :class:`Transform` of this transform.
        This should satisfy ``t.inv.inv is t``.
        N)r+   Ú_InverseTransformÚweakrefÚref)r0   Úinvs     r2   r@   zTransform.invw   sB   € ð ˆØ�9‰9Ð Ø—)‘)“+ˆCØˆ;Ü# DÓ)ˆCÜŸ™ CÓ(ˆDŒIØˆ
r3   c                 ó   — t         ‚)z˜
        Returns the sign of the determinant of the Jacobian, if applicable.
        In general this only makes sense for bijective transforms.
        ©ÚNotImplementedErrorr;   s    r2   ÚsignzTransform.sign…   s
   € ô "Ð!r3   c                 óÀ   — | j                   |k(  r| S t        | «      j                  t        j                  u r t        | «      |¬«      S t	        t        | «      › d�«      ‚)N©r&   z.with_cache is not implemented)r*   Útyper/   r!   rC   ©r0   r&   s     r2   Ú
with_cachezTransform.with_cache�   sU   € Ø×Ñ˜zÒ)ØˆKÜ�‹:×Ñ¤)×"4Ñ"4Ñ4Ø”4˜“:¨Ô4Ð4Ü!¤T¨$£Z LÐ0NÐ"OÓPÐPr3   c                 ó
   — | |u S ©N© ©r0   Úothers     r2   Ú__eq__zTransform.__eq__”   s   € Ø�uˆ}Ðr3   c                 ó&   — | j                  |«       S rK   )rO   rM   s     r2   Ú__ne__zTransform.__ne__—   s   € à—;‘;˜uÓ%Ð%Ð%r3   c                 ó¤   — | j                   dk(  r| j                  |«      S | j                  \  }}||u r|S | j                  |«      }||f| _        |S )z2
        Computes the transform `x => y`.
        r   )r*   Ú_callr,   )r0   ÚxÚx_oldÚy_oldÚys        r2   Ú__call__zTransform.__call__›   sY   € ð ×Ñ˜qÒ Ø—:‘:˜a“=Ð Ø×'Ñ'‰ˆˆuØ�‰:ØˆLØ�J‰J�q‹MˆØ˜a˜4ˆÔØˆr3   c                 ó¤   — | j                   dk(  r| j                  |«      S | j                  \  }}||u r|S | j                  |«      }||f| _        |S )z1
        Inverts the transform `y => x`.
        r   )r*   Ú_inverser,   )r0   rW   rU   rV   rT   s        r2   Ú	_inv_callzTransform._inv_call¨   s[   € ð ×Ñ˜qÒ Ø—=‘= Ó#Ð#Ø×'Ñ'‰ˆˆuØ�‰:ØˆLØ�M‰M˜!ÓˆØ˜a˜4ˆÔØˆr3   c                 ó   — t         ‚)zD
        Abstract method to compute forward transformation.
        rB   ©r0   rT   s     r2   rS   zTransform._callµ   ó
   € ô "Ð!r3   c                 ó   — t         ‚)zD
        Abstract method to compute inverse transformation.
        rB   ©r0   rW   s     r2   rZ   zTransform._inverse»   r^   r3   c                 ó   — t         ‚)zU
        Computes the log det jacobian `log |dy/dx|` given input and output.
        rB   ©r0   rT   rW   s      r2   Úlog_abs_det_jacobianzTransform.log_abs_det_jacobianÁ   r^   r3   c                 ó4   — | j                   j                  dz   S )Nz())r1   Ú__name__r;   s    r2   Ú__repr__zTransform.__repr__Ç   s   € Ø�~‰~×&Ñ&¨Ñ-Ð-r3   c                 ó   — |S )z{
        Infers the shape of the forward computation, given the input shape.
        Defaults to preserving shape.
        rL   ©r0   Úshapes     r2   Úforward_shapezTransform.forward_shapeÊ   ó	   € ð
 ˆr3   c                 ó   — |S )z}
        Infers the shapes of the inverse computation, given the output shape.
        Defaults to preserving shape.
        rL   rh   s     r2   Úinverse_shapezTransform.inverse_shapeÑ   rk   r3   ©r   )r'   r!   ©r)   )re   Ú
__module__Ú__qualname__Ú__doc__Ú	bijectiver   Ú
ConstraintÚ__annotations__Úintr/   r8   Úpropertyr:   r@   rD   rI   rO   rQ   rX   r[   rS   rZ   rc   rf   rj   rm   Ú__classcell__©r1   s   @r2   r!   r!   0   sÅ   ø… ñ*ðX €IØ×"Ñ"Ó"Ø×$Ñ$Ó$ñ	 3ð 	¨tõ 	òð
 ðW˜3ò Wó ðWð
 òó ðð ð"�cò "ó ð"óQòò&òòò"ò"ò"ò.òör3   r!   c                   óþ   ‡ — e Zd ZdZdeddfˆ fd„Z ej                  d¬«      d„ «       Z ej                  d¬«      d	„ «       Z	e
defd
„«       Ze
defd„«       Ze
defd„«       Zdd„Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zˆ xZS )r=   z|
    Inverts a single :class:`Transform`.
    This class is private; please instead use the ``Transform.inv`` property.
    Ú	transformr'   Nc                 óH   •— t         ‰| �  |j                  ¬«       || _        y ©NrF   )r.   r/   r*   r+   )r0   r{   r1   s     €r2   r/   z_InverseTransform.__init__ß   s    ø€ Ü‰Ñ I×$9Ñ$9ÐÔ:Ø(ˆ�	r3   F©Úis_discretec                 ó\   — | j                   €t        d«      ‚| j                   j                  S ©Nú_inv must not be None)r+   ÚAssertionErrorr%   r;   s    r2   r$   z_InverseTransform.domainã   s*   € ð �9‰9ÐÜ Ð!8Ó9Ð9Ø�y‰y×!Ñ!Ð!r3   c                 ó\   — | j                   €t        d«      ‚| j                   j                  S r�   )r+   rƒ   r$   r;   s    r2   r%   z_InverseTransform.codomainê   s*   € ð �9‰9ÐÜ Ð!8Ó9Ð9Ø�y‰y×ÑÐr3   c                 ó\   — | j                   €t        d«      ‚| j                   j                  S r�   )r+   rƒ   rs   r;   s    r2   rs   z_InverseTransform.bijectiveñ   s(   € à�9‰9ÐÜ Ð!8Ó9Ð9Ø�y‰y×"Ñ"Ð"r3   c                 ó\   — | j                   €t        d«      ‚| j                   j                  S r�   )r+   rƒ   rD   r;   s    r2   rD   z_InverseTransform.sign÷   s&   € à�9‰9ÐÜ Ð!8Ó9Ð9Ø�y‰y�~‰~Ðr3   c                 ó   — | j                   S rK   )r+   r;   s    r2   r@   z_InverseTransform.invý   s   € à�y‰yÐr3   c                 óz   — | j                   €t        d«      ‚| j                  j                  |«      j                  S r�   )r+   rƒ   r@   rI   rH   s     r2   rI   z_InverseTransform.with_cache  s3   € Ø�9‰9ÐÜ Ð!8Ó9Ð9Ø�x‰x×"Ñ" :Ó.×2Ñ2Ð2r3   c                 ó„   — t        |t        «      sy| j                  €t        d«      ‚| j                  |j                  k(  S )NFr‚   )Ú
isinstancer=   r+   rƒ   rM   s     r2   rO   z_InverseTransform.__eq__  s9   € Ü˜%Ô!2Ô3ØØ�9‰9ÐÜ Ð!8Ó9Ð9Ø�y‰y˜EŸJ™JÑ&Ð&r3   c                 ó`   — | j                   j                  › dt        | j                  «      › d�S )Nú(ú))r1   re   Úreprr+   r;   s    r2   rf   z_InverseTransform.__repr__  s)   € Ø—.‘.×)Ñ)Ð*¨!¬D°·±«OÐ+<¸AÐ>Ð>r3   c                 óf   — | j                   €t        d«      ‚| j                   j                  |«      S r�   )r+   rƒ   r[   r]   s     r2   rX   z_InverseTransform.__call__  s-   € Ø�9‰9ÐÜ Ð!8Ó9Ð9Ø�y‰y×"Ñ" 1Ó%Ð%r3   c                 ój   — | j                   €t        d«      ‚| j                   j                  ||«       S r�   )r+   rƒ   rc   rb   s      r2   rc   z&_InverseTransform.log_abs_det_jacobian  s2   € Ø�9‰9ÐÜ Ð!8Ó9Ð9Ø—	‘	×.Ñ.¨q°!Ó4Ð4Ð4r3   c                 ó8   — | j                   j                  |«      S rK   )r+   rm   rh   s     r2   rj   z_InverseTransform.forward_shape  ó   € Ø�y‰y×&Ñ& uÓ-Ð-r3   c                 ó8   — | j                   j                  |«      S rK   )r+   rj   rh   s     r2   rm   z_InverseTransform.inverse_shape  r’   r3   ro   )re   rp   rq   rr   r!   r/   r   Údependent_propertyr$   r%   rw   Úboolrs   rv   rD   r@   rI   rO   rf   rX   rc   rj   rm   rx   ry   s   @r2   r=   r=   Ù   sÑ   ø„ ñð
) )ð )°õ )ð $€[×#Ñ#°Ô6ñ"ó 7ð"ð
 $€[×#Ñ#°Ô6ñ ó 7ð ð
 ð#˜4ò #ó ð#ð
 ð�cò ó ðð
 ð�Yò ó ðó3ò
'ò?ò&ò
5ò
.ö.r3   r=   c                   ó
  ‡ — e Zd ZdZddee   deddfˆ fd„Zd„ Z e	j                  d¬	«      d
„ «       Z e	j                  d¬	«      d„ «       Zedefd„«       Zedefd„«       Zedefd„«       Zdd„Zd„ Zd„ Zd„ Zd„ Zd„ Zˆ xZS )r   ab  
    Composes multiple transforms in a chain.
    The transforms being composed are responsible for caching.

    Args:
        parts (list of :class:`Transform`): A list of transforms to compose.
        cache_size (int): Size of cache. If zero, no caching is done. If one,
            the latest single value is cached. Only 0 and 1 are supported.
    Úpartsr&   r'   Nc                 ó~   •— |r|D �cg c]  }|j                  |«      ‘Œ }}t        ‰| �	  |¬«       || _        y c c}w r}   )rI   r.   r/   r—   )r0   r—   r&   Úpartr1   s       €r2   r/   zComposeTransform.__init__,  s?   ø€ ÙØ=BÖC°T�T—_‘_ ZÕ0ÐCˆEÐCÜ‰Ñ JÐÔ/Øˆ�
ùò Ds   ˆ:c                 óV   — t        |t        «      sy| j                  |j                  k(  S ©NF)rŠ   r   r—   rM   s     r2   rO   zComposeTransform.__eq__2  s#   € Ü˜%Ô!1Ô2ØØ�z‰z˜UŸ[™[Ñ(Ð(r3   Fr~   c                 óB  — | j                   st        j                  S | j                   d   j                  }| j                   d   j                  j
                  }t        | j                   «      D ]R  }||j                  j
                  |j                  j
                  z
  z  }t        ||j                  j
                  «      }ŒT ||j
                  k  rt        d|› d|j
                  › �«      ‚||j
                  kD  r#t        j                  |||j
                  z
  «      }|S )Nr   éÿÿÿÿú
event_dim z must be >= domain.event_dim )
r—   r   Úrealr$   r%   r:   ÚreversedÚmaxrƒ   Úindependent)r0   r$   r:   r™   s       r2   r$   zComposeTransform.domain7  sû   € ð �zŠzÜ×#Ñ#Ð#Ø—‘˜A‘×%Ñ%ˆà—J‘J˜r‘N×+Ñ+×5Ñ5ˆ	Ü˜TŸZ™ZÓ(ò 	>ˆDØ˜Ÿ™×.Ñ.°·±×1HÑ1HÑHÑHˆIÜ˜I t§{¡{×'<Ñ'<Ó=‰Ið	>ð �v×'Ñ'Ò'Ü Ø˜Y˜KÐ'DÀV×EUÑEUÐDVÐWóð ð �v×'Ñ'Ò'Ü ×,Ñ,¨V°YÀ×AQÑAQÑ5QÓRˆFØˆr3   c                 ó0  — | j                   st        j                  S | j                   d   j                  }| j                   d   j                  j
                  }| j                   D ]R  }||j                  j
                  |j                  j
                  z
  z  }t        ||j                  j
                  «      }ŒT ||j
                  k  rt        d|› d|j
                  › �«      ‚||j
                  kD  r#t        j                  |||j
                  z
  «      }|S )Nr�   r   rž   z must be >= codomain.event_dim )	r—   r   rŸ   r%   r$   r:   r¡   rƒ   r¢   )r0   r%   r:   r™   s       r2   r%   zComposeTransform.codomainJ  sø   € ð �zŠzÜ×#Ñ#Ð#Ø—:‘:˜b‘>×*Ñ*ˆà—J‘J˜q‘M×(Ñ(×2Ñ2ˆ	Ø—J‘Jò 	@ˆDØ˜Ÿ™×0Ñ0°4·;±;×3HÑ3HÑHÑHˆIÜ˜I t§}¡}×'>Ñ'>Ó?‰Ið	@ð �x×)Ñ)Ò)Ü Ø˜Y˜KÐ'FÀx×GYÑGYÐFZÐ[óð ð �x×)Ñ)Ò)Ü"×.Ñ.¨x¸ÀX×EWÑEWÑ9WÓXˆHØˆr3   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrK   ©rs   )Ú.0Úps     r2   ú	<genexpr>z-ComposeTransform.bijective.<locals>.<genexpr>_  s   è ø€ Ò3 1�1—;•;Ñ3ùó   ‚)Úallr—   r;   s    r2   rs   zComposeTransform.bijective]  s   € äÑ3¨¯
©
Ô3Ó3Ð3r3   c                 óJ   — d}| j                   D ]  }||j                  z  }Œ |S ©Nr)   )r—   rD   )r0   rD   r¨   s      r2   rD   zComposeTransform.signa  s,   € àˆØ—‘ò 	!ˆAØ˜!Ÿ&™&‘=‰Dð	!àˆr3   c                 ó$  — d }| j                   �| j                  «       }|€jt        t        | j                  «      D �cg c]  }|j                  ‘Œ c}«      }t        j                  |«      | _         t        j                  | «      |_         |S c c}w rK   )r+   r   r    r—   r@   r>   r?   )r0   r@   r¨   s      r2   r@   zComposeTransform.invh  sn   € àˆØ�9‰9Ð Ø—)‘)“+ˆCØˆ;Ü"´8¸D¿J¹JÓ3GÖ#H¨a A§E£EÒ#HÓIˆCÜŸ™ CÓ(ˆDŒIÜ—{‘{ 4Ó(ˆCŒHØˆ
ùò $Is   ½Bc                 óR   — | j                   |k(  r| S t        | j                  |¬«      S r}   )r*   r   r—   rH   s     r2   rI   zComposeTransform.with_caches  s&   € Ø×Ñ˜zÒ)ØˆKÜ §
¡
°zÔBÐBr3   c                 ó8   — | j                   D ]
  } ||«      }Œ |S rK   )r—   )r0   rT   r™   s      r2   rX   zComposeTransform.__call__x  s#   € Ø—J‘Jò 	ˆDÙ�Q“‰Að	àˆr3   c           	      óp  — | j                   st        j                  |«      S |g}| j                   d d D ]  }|j                   ||d   «      «       Œ |j                  |«       g }| j                  j
                  }t        | j                   |d d |dd  «      D ]x  \  }}}|j                  t        |j                  ||«      ||j                  j
                  z
  «      «       ||j                  j
                  |j                  j
                  z
  z  }Œz t        j                  t        j                  |«      S )Nr�   r)   )r—   ÚtorchÚ
zeros_likeÚappendr$   r:   Úzipr   rc   r%   Ú	functoolsÚreduceÚoperatorÚadd)r0   rT   rW   Úxsr™   Útermsr:   s          r2   rc   z%ComposeTransform.log_abs_det_jacobian}  s  € Ø�zŠzÜ×#Ñ# AÓ&Ð&ð ˆSˆØ—J‘J˜s �Oò 	$ˆDØ�I‰I‘d˜2˜b™6“lÕ#ð	$à
�	‰	�!ŒàˆØ—K‘K×)Ñ)ˆ	Ü˜dŸj™j¨"¨S¨b¨'°2°a°b°6Ó:ò 	I‰JˆD�!�QØ�L‰LÜØ×-Ñ-¨a°Ó3°YÀÇÁ×AVÑAVÑ5Vóôð
 ˜Ÿ™×0Ñ0°4·;±;×3HÑ3HÑHÑH‰Ið	Iô ×Ñ¤§¡¨eÓ4Ð4r3   c                 óJ   — | j                   D ]  }|j                  |«      }Œ |S rK   )r—   rj   ©r0   ri   r™   s      r2   rj   zComposeTransform.forward_shape’  s*   € Ø—J‘Jò 	.ˆDØ×&Ñ& uÓ-‰Eð	.àˆr3   c                 ó\   — t        | j                  «      D ]  }|j                  |«      }Œ |S rK   )r    r—   rm   r½   s      r2   rm   zComposeTransform.inverse_shape—  s/   € Ü˜TŸZ™ZÓ(ò 	.ˆDØ×&Ñ& uÓ-‰Eð	.àˆr3   c                 óÀ   — | j                   j                  dz   }|dj                  | j                  D �cg c]  }|j	                  «       ‘Œ c}«      z  }|dz  }|S c c}w )Nz(
    z,
    z
))r1   re   Újoinr—   rf   )r0   Ú
fmt_stringr¨   s      r2   rf   zComposeTransform.__repr__œ  sT   € Ø—^‘^×,Ñ,¨yÑ8ˆ
Ø�i—n‘n¸D¿J¹JÖ%G°q a§j¡j¥lÒ%GÓHÑHˆ
Ø�eÑˆ
ØÐùò &Hs   ´A
rn   ro   )re   rp   rq   rr   Úlistr!   rv   r/   rO   r   r”   r$   r%   r	   r•   rs   rD   rw   r@   rI   rX   rc   rj   rm   rf   rx   ry   s   @r2   r   r   !  sÝ   ø„ ññ˜d 9™oð ¸3ð Àtõ ò)ð
 $€[×#Ñ#°Ô6ñó 7ðð" $€[×#Ñ#°Ô6ñó 7ðð" ð4˜4ò 4ó ð4ð ð�cò ó ðð ð�Yò ó ðóCò
ò
5ò*ò
ö
r3   r   c            	       óô   ‡ — e Zd ZdZ	 ddedededdfˆ fd„Zdd„Z ej                  d	¬
«      d„ «       Z
 ej                  d	¬
«      d„ «       Zedefd„«       Zedefd„«       Zd„ Zd„ Zd„ Zd„ Zd„ Zd„ Zˆ xZS )r   a  
    Wrapper around another transform to treat
    ``reinterpreted_batch_ndims``-many extra of the right most dimensions as
    dependent. This has no effect on the forward or backward transforms, but
    does sum out ``reinterpreted_batch_ndims``-many of the rightmost dimensions
    in :meth:`log_abs_det_jacobian`.

    Args:
        base_transform (:class:`Transform`): A base transform.
        reinterpreted_batch_ndims (int): The number of extra rightmost
            dimensions to treat as dependent.
    Úbase_transformÚreinterpreted_batch_ndimsr&   r'   Nc                 ó`   •— t         ‰| �  |¬«       |j                  |«      | _        || _        y r}   )r.   r/   rI   rÄ   rÅ   )r0   rÄ   rÅ   r&   r1   s       €r2   r/   zIndependentTransform.__init__´  s0   ø€ ô 	‰Ñ JÐÔ/Ø,×7Ñ7¸
ÓCˆÔØ)BˆÕ&r3   c                 óh   — | j                   |k(  r| S t        | j                  | j                  |¬«      S r}   )r*   r   rÄ   rÅ   rH   s     r2   rI   zIndependentTransform.with_cache¾  s5   € Ø×Ñ˜zÒ)ØˆKÜ#Ø×Ñ ×!?Ñ!?ÈJô
ð 	
r3   Fr~   c                 ój   — t        j                  | j                  j                  | j                  «      S rK   )r   r¢   rÄ   r$   rÅ   r;   s    r2   r$   zIndependentTransform.domainÅ  s.   € ô ×&Ñ&Ø×Ñ×&Ñ&¨×(FÑ(Fó
ð 	
r3   c                 ój   — t        j                  | j                  j                  | j                  «      S rK   )r   r¢   rÄ   r%   rÅ   r;   s    r2   r%   zIndependentTransform.codomainÌ  s.   € ô ×&Ñ&Ø×Ñ×(Ñ(¨$×*HÑ*Hó
ð 	
r3   c                 ó.   — | j                   j                  S rK   )rÄ   rs   r;   s    r2   rs   zIndependentTransform.bijectiveÓ  s   € à×"Ñ"×,Ñ,Ð,r3   c                 ó.   — | j                   j                  S rK   )rÄ   rD   r;   s    r2   rD   zIndependentTransform.sign×  s   € à×"Ñ"×'Ñ'Ð'r3   c                 óˆ   — |j                  «       | j                  j                  k  rt        d«      ‚| j	                  |«      S ©NúToo few dimensions on input)Údimr$   r:   r-   rÄ   r]   s     r2   rS   zIndependentTransform._callÛ  s7   € Ø�5‰5‹7�T—[‘[×*Ñ*Ò*ÜÐ:Ó;Ð;Ø×"Ñ" 1Ó%Ð%r3   c                 óœ   — |j                  «       | j                  j                  k  rt        d«      ‚| j                  j                  |«      S rÍ   )rÏ   r%   r:   r-   rÄ   r@   r`   s     r2   rZ   zIndependentTransform._inverseà  s=   € Ø�5‰5‹7�T—]‘]×,Ñ,Ò,ÜÐ:Ó;Ð;Ø×"Ñ"×&Ñ& qÓ)Ð)r3   c                 ój   — | j                   j                  ||«      }t        || j                  «      }|S rK   )rÄ   rc   r   rÅ   )r0   rT   rW   Úresults       r2   rc   z)IndependentTransform.log_abs_det_jacobianå  s1   € Ø×$Ñ$×9Ñ9¸!¸QÓ?ˆÜ ¨×(FÑ(FÓGˆØˆr3   c                 óz   — | j                   j                  › dt        | j                  «      › d| j                  › d�S )NrŒ   z, r�   )r1   re   rŽ   rÄ   rÅ   r;   s    r2   rf   zIndependentTransform.__repr__ê  s:   € Ø—.‘.×)Ñ)Ð*¨!¬D°×1DÑ1DÓ,EÐ+FÀbÈ×IgÑIgÐHhÐhiÐjÐjr3   c                 ó8   — | j                   j                  |«      S rK   )rÄ   rj   rh   s     r2   rj   z"IndependentTransform.forward_shapeí  ó   € Ø×"Ñ"×0Ñ0°Ó7Ð7r3   c                 ó8   — | j                   j                  |«      S rK   )rÄ   rm   rh   s     r2   rm   z"IndependentTransform.inverse_shapeð  rÕ   r3   rn   ro   )re   rp   rq   rr   r!   rv   r/   rI   r   r”   r$   r%   rw   r•   rs   rD   rS   rZ   rc   rf   rj   rm   rx   ry   s   @r2   r   r   ¦  sÙ   ø„ ñð" ñ	Cà!ðCð $'ðCð ð	Cð
 
õCó
ð $€[×#Ñ#°Ô6ñ
ó 7ð
ð
 $€[×#Ñ#°Ô6ñ
ó 7ð
ð
 ð-˜4ò -ó ð-ð ð(�cò (ó ð(ò&ò
*ò
ò
kò8ö8r3   r   c            	       óÒ   ‡ — e Zd ZdZdZ	 ddej                  dej                  deddfˆ fd„Ze	j                  d	„ «       Ze	j                  d
„ «       Zdd„Zd„ Zd„ Zd„ Zd„ Zd„ Zˆ xZS )r   aó  
    Unit Jacobian transform to reshape the rightmost part of a tensor.

    Note that ``in_shape`` and ``out_shape`` must have the same number of
    elements, just as for :meth:`torch.Tensor.reshape`.

    Arguments:
        in_shape (torch.Size): The input event shape.
        out_shape (torch.Size): The output event shape.
        cache_size (int): Size of cache. If zero, no caching is done. If one,
            the latest single value is cached. Only 0 and 1 are supported. (Default 0.)
    TÚin_shapeÚ	out_shaper&   r'   Nc                 ó  •— t        j                  |«      | _        t        j                  |«      | _        | j                  j	                  «       | j                  j	                  «       k7  rt        d«      ‚t        ‰| �  |¬«       y )Nz6in_shape, out_shape have different numbers of elementsrF   )r²   ÚSizerØ   rÙ   Únumelr-   r.   r/   )r0   rØ   rÙ   r&   r1   s       €r2   r/   zReshapeTransform.__init__  sc   ø€ ô Ÿ
™
 8Ó,ˆŒÜŸ™ IÓ.ˆŒØ�=‰=×ÑÓ  D§N¡N×$8Ñ$8Ó$:Ò:ÜÐUÓVÐVÜ‰Ñ JÐÕ/r3   c                 óp   — t        j                  t         j                  t        | j                  «      «      S rK   )r   r¢   rŸ   ÚlenrØ   r;   s    r2   r$   zReshapeTransform.domain  s&   € ô ×&Ñ&¤{×'7Ñ'7¼¸T¿]¹]Ó9KÓLÐLr3   c                 óp   — t        j                  t         j                  t        | j                  «      «      S rK   )r   r¢   rŸ   rÞ   rÙ   r;   s    r2   r%   zReshapeTransform.codomain  s&   € ô ×&Ñ&¤{×'7Ñ'7¼¸T¿^¹^Ó9LÓMÐMr3   c                 óh   — | j                   |k(  r| S t        | j                  | j                  |¬«      S r}   )r*   r   rØ   rÙ   rH   s     r2   rI   zReshapeTransform.with_cache  s,   € Ø×Ñ˜zÒ)ØˆKÜ §¡¨t¯~©~È*ÔUÐUr3   c                 ó¤   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  || j
                  z   «      S rK   )ri   rÏ   rÞ   rØ   ÚreshaperÙ   )r0   rT   Úbatch_shapes      r2   rS   zReshapeTransform._call  s?   € Ø—g‘gÐ< §¡£¬#¨d¯m©mÓ*<Ñ <Ð=ˆØ�y‰y˜ t§~¡~Ñ5Ó6Ð6r3   c                 ó¤   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  || j
                  z   «      S rK   )ri   rÏ   rÞ   rÙ   râ   rØ   )r0   rW   rã   s      r2   rZ   zReshapeTransform._inverse#  s?   € Ø—g‘gÐ= §¡£¬#¨d¯n©nÓ*=Ñ =Ð>ˆØ�y‰y˜ t§}¡}Ñ4Ó5Ð5r3   c                 óŠ   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  |«      S rK   )ri   rÏ   rÞ   rØ   Ú	new_zeros)r0   rT   rW   rã   s       r2   rc   z%ReshapeTransform.log_abs_det_jacobian'  s6   € Ø—g‘gÐ< §¡£¬#¨d¯m©mÓ*<Ñ <Ð=ˆØ�{‰{˜;Ó'Ð'r3   c                 ó   — t        |«      t        | j                  «      k  rt        d«      ‚t        |«      t        | j                  «      z
  }||d  | j                  k7  rt        d||d  › d| j                  › �«      ‚|d | | j                  z   S ©NrÎ   zShape mismatch: expected z	 but got )rÞ   rØ   r-   rÙ   ©r0   ri   Úcuts      r2   rj   zReshapeTransform.forward_shape+  sŠ   € Üˆu‹:œ˜DŸM™MÓ*Ò*ÜÐ:Ó;Ð;Ü�%‹jœ3˜tŸ}™}Ó-Ñ-ˆØ��ˆ;˜$Ÿ-™-Ò'ÜØ+¨E°#°$¨K¨=¸	À$Ç-Á-ÀÐQóð ð �T�cˆ{˜TŸ^™^Ñ+Ð+r3   c                 ó   — t        |«      t        | j                  «      k  rt        d«      ‚t        |«      t        | j                  «      z
  }||d  | j                  k7  rt        d||d  › d| j                  › �«      ‚|d | | j                  z   S rè   )rÞ   rÙ   r-   rØ   ré   s      r2   rm   zReshapeTransform.inverse_shape5  s‹   € Üˆu‹:œ˜DŸN™NÓ+Ò+ÜÐ:Ó;Ð;Ü�%‹jœ3˜tŸ~™~Ó.Ñ.ˆØ��ˆ;˜$Ÿ.™.Ò(ÜØ+¨E°#°$¨K¨=¸	À$Ç.Á.ÐAQÐRóð ð �T�cˆ{˜TŸ]™]Ñ*Ð*r3   rn   ro   )re   rp   rq   rr   rs   r²   rÛ   rv   r/   r   r”   r$   r%   rI   rS   rZ   rc   rj   rm   rx   ry   s   @r2   r   r   ô  sž   ø„ ñð €Ið ñ	
0à—*‘*ð
0ð —:‘:ð
0ð ð	
0ð
 
õ
0ð ×#Ñ#ñMó $ðMð ×#Ñ#ñNó $ðNóVò
7ò6ò(ò,ö+r3   r   c                   ó`   — e Zd ZdZej
                  Zej                  ZdZ	dZ
d„ Zd„ Zd„ Zd„ Zy)	r   z8
    Transform via the mapping :math:`y = \exp(x)`.
    Tr)   c                 ó"   — t        |t        «      S rK   )rŠ   r   rM   s     r2   rO   zExpTransform.__eq__J  ó   € Ü˜%¤Ó.Ð.r3   c                 ó"   — |j                  «       S rK   )Úexpr]   s     r2   rS   zExpTransform._callM  ó   € Ø�u‰u‹wˆr3   c                 ó"   — |j                  «       S rK   ©Úlogr`   s     r2   rZ   zExpTransform._inverseP  rñ   r3   c                 ó   — |S rK   rL   rb   s      r2   rc   z!ExpTransform.log_abs_det_jacobianS  ó   € Øˆr3   N©re   rp   rq   rr   r   rŸ   r$   Úpositiver%   rs   rD   rO   rS   rZ   rc   rL   r3   r2   r   r   @  s=   „ ñð ×Ñ€FØ×#Ñ#€HØ€IØ€Dò/òòór3   r   c                   ó¨   ‡ — e Zd ZdZej
                  Zej
                  ZdZdde	de
ddfˆ fd„Zdd„Zede
fd	„«       Zd
„ Zd„ Zd„ Zd„ Zd„ Zd„ Zˆ xZS )r   zD
    Transform via the mapping :math:`y = x^{\text{exponent}}`.
    TÚexponentr&   r'   Nc                 óJ   •— t         ‰| �  |¬«       t        |«      \  | _        y r}   )r.   r/   r   rú   )r0   rú   r&   r1   s      €r2   r/   zPowerTransform.__init__`  s"   ø€ Ü‰Ñ JÐÔ/Ü(¨Ó2Ñˆ�r3   c                 óR   — | j                   |k(  r| S t        | j                  |¬«      S r}   )r*   r   rú   rH   s     r2   rI   zPowerTransform.with_cached  s&   € Ø×Ñ˜zÒ)ØˆKÜ˜dŸm™m¸
ÔCÐCr3   c                 ó6   — | j                   j                  «       S rK   )rú   rD   r;   s    r2   rD   zPowerTransform.signi  s   € à�}‰}×!Ñ!Ó#Ð#r3   c                 ó¦   — t        |t        «      sy| j                  j                  |j                  «      j	                  «       j                  «       S r›   )rŠ   r   rú   Úeqr«   ÚitemrM   s     r2   rO   zPowerTransform.__eq__m  s:   € Ü˜%¤Ô0ØØ�}‰}×Ñ §¡Ó/×3Ñ3Ó5×:Ñ:Ó<Ð<r3   c                 ó8   — |j                  | j                  «      S rK   ©Úpowrú   r]   s     r2   rS   zPowerTransform._callr  s   € Ø�u‰u�T—]‘]Ó#Ð#r3   c                 ó>   — |j                  d| j                  z  «      S r­   r  r`   s     r2   rZ   zPowerTransform._inverseu  s   € Ø�u‰u�Q˜Ÿ™Ñ&Ó'Ð'r3   c                 ó^   — | j                   |z  |z  j                  «       j                  «       S rK   )rú   Úabsrô   rb   s      r2   rc   z#PowerTransform.log_abs_det_jacobianx  s(   € Ø—‘ Ñ! AÑ%×*Ñ*Ó,×0Ñ0Ó2Ð2r3   c                 óX   — t        j                  |t        | j                  dd«      «      S ©Nri   rL   ©r²   Úbroadcast_shapesÚgetattrrú   rh   s     r2   rj   zPowerTransform.forward_shape{  ó"   € Ü×%Ñ% e¬W°T·]±]ÀGÈRÓ-PÓQÐQr3   c                 óX   — t        j                  |t        | j                  dd«      «      S r  r	  rh   s     r2   rm   zPowerTransform.inverse_shape~  r  r3   rn   ro   )re   rp   rq   rr   r   rø   r$   r%   rs   r   rv   r/   rI   r	   rD   rO   rS   rZ   rc   rj   rm   rx   ry   s   @r2   r   r   W  s€   ø„ ñð ×!Ñ!€FØ×#Ñ#€HØ€Iñ3 ð 3°Sð 3Àõ 3óDð
 ð$�cò $ó ð$ò=ò
$ò(ò3òRöRr3   r   c                 óÄ   — t        j                  | j                  «      }t        j                  t        j                  | «      |j
                  d|j                  z
  ¬«      S ©Nç      ð?©Úminr¡   )r²   ÚfinfoÚdtypeÚclampÚsigmoidÚtinyÚeps)rT   r  s     r2   Ú_clipped_sigmoidr  ‚  s<   € Ü�K‰K˜Ÿ™Ó €EÜ�;‰;”u—}‘} QÓ'¨U¯Z©Z¸SÀ5Ç9Á9¹_ÔMÐMr3   c                   ó`   — e Zd ZdZej
                  Zej                  ZdZ	dZ
d„ Zd„ Zd„ Zd„ Zy)	r   zg
    Transform via the mapping :math:`y = \frac{1}{1 + \exp(-x)}` and :math:`x = \text{logit}(y)`.
    Tr)   c                 ó"   — t        |t        «      S rK   )rŠ   r   rM   s     r2   rO   zSigmoidTransform.__eq__‘  ó   € Ü˜%Ô!1Ó2Ð2r3   c                 ó   — t        |«      S rK   )r  r]   s     r2   rS   zSigmoidTransform._call”  s   € Ü Ó"Ð"r3   c                 óØ   — t        j                  |j                  «      }|j                  |j                  d|j
                  z
  ¬«      }|j                  «       | j                  «       z
  S r  )r²   r  r  r  r  r  rô   Úlog1p)r0   rW   r  s      r2   rZ   zSigmoidTransform._inverse—  sK   € Ü—‘˜AŸG™GÓ$ˆØ�G‰G˜Ÿ
™
¨¨e¯i©i©ˆGÓ8ˆØ�u‰u‹w˜1˜"Ÿ™›Ñ%Ð%r3   c                 ó\   — t        j                  | «       t        j                  |«      z
  S rK   )ÚFr   rb   s      r2   rc   z%SigmoidTransform.log_abs_det_jacobianœ  s!   € Ü—
‘
˜A˜2“ˆ¤§¡¨A£Ñ.Ð.r3   N)re   rp   rq   rr   r   rŸ   r$   Úunit_intervalr%   rs   rD   rO   rS   rZ   rc   rL   r3   r2   r   r   ‡  s=   „ ñð ×Ñ€FØ×(Ñ(€HØ€IØ€Dò3ò#ò&ó
/r3   r   c                   ó`   — e Zd ZdZej
                  Zej                  ZdZ	dZ
d„ Zd„ Zd„ Zd„ Zy)	r   zž
    Transform via the mapping :math:`\text{Softplus}(x) = \log(1 + \exp(x))`.
    The implementation reverts to the linear function when :math:`x > 20`.
    Tr)   c                 ó"   — t        |t        «      S rK   )rŠ   r   rM   s     r2   rO   zSoftplusTransform.__eq__«  s   € Ü˜%Ô!2Ó3Ð3r3   c                 ó   — t        |«      S rK   ©r   r]   s     r2   rS   zSoftplusTransform._call®  s   € Ü˜‹{Ðr3   c                 ób   — | j                  «       j                  «       j                  «       |z   S rK   )Úexpm1Únegrô   r`   s     r2   rZ   zSoftplusTransform._inverse±  s'   € Ø��z‰z‹|×ÑÓ!×%Ñ%Ó'¨!Ñ+Ð+r3   c                 ó   — t        | «       S rK   r&  rb   s      r2   rc   z&SoftplusTransform.log_abs_det_jacobian´  s   € Ü˜!˜“ˆ}Ðr3   Nr÷   rL   r3   r2   r   r      s=   „ ñð
 ×Ñ€FØ×#Ñ#€HØ€IØ€Dò4òò,ór3   r   c                   ón   — e Zd ZdZej
                  Z ej                  dd«      ZdZ	dZ
d„ Zd„ Zd„ Zd	„ Zy
)r   aé  
    Transform via the mapping :math:`y = \tanh(x)`.

    It is equivalent to

    .. code-block:: python

        ComposeTransform(
            [
                AffineTransform(0.0, 2.0),
                SigmoidTransform(),
                AffineTransform(-1.0, 2.0),
            ]
        )

    However this might not be numerically stable, thus it is recommended to use `TanhTransform`
    instead.

    Note that one should use `cache_size=1` when it comes to `NaN/Inf` values.

    g      ð¿r  Tr)   c                 ó"   — t        |t        «      S rK   )rŠ   r   rM   s     r2   rO   zTanhTransform.__eq__Ô  s   € Ü˜%¤Ó/Ð/r3   c                 ó"   — |j                  «       S rK   )Útanhr]   s     r2   rS   zTanhTransform._call×  s   € Ø�v‰v‹xˆr3   c                 ó,   — t        j                  |«      S rK   )r²   Úatanhr`   s     r2   rZ   zTanhTransform._inverseÚ  s   € ô �{‰{˜1‹~Ðr3   c                 óV   — dt        j                  d«      |z
  t        d|z  «      z
  z  S )Nç       @g       À)Úmathrô   r   rb   s      r2   rc   z"TanhTransform.log_abs_det_jacobianß  s*   € ð ”d—h‘h˜s“m aÑ'¬(°4¸!±8Ó*<Ñ<Ñ=Ð=r3   N)re   rp   rq   rr   r   rŸ   r$   Úintervalr%   rs   rD   rO   rS   rZ   rc   rL   r3   r2   r   r   ¸  sF   „ ñð, ×Ñ€FØ#ˆ{×#Ñ# D¨#Ó.€HØ€IØ€Dò0òòó
>r3   r   c                   óR   — e Zd ZdZej
                  Zej                  Zd„ Z	d„ Z
d„ Zy)r   z*Transform via the mapping :math:`y = |x|`.c                 ó"   — t        |t        «      S rK   )rŠ   r   rM   s     r2   rO   zAbsTransform.__eq__ë  rî   r3   c                 ó"   — |j                  «       S rK   )r  r]   s     r2   rS   zAbsTransform._callî  rñ   r3   c                 ó   — |S rK   rL   r`   s     r2   rZ   zAbsTransform._inverseñ  rö   r3   N)re   rp   rq   rr   r   rŸ   r$   rø   r%   rO   rS   rZ   rL   r3   r2   r   r   å  s*   „ Ù5à×Ñ€FØ×#Ñ#€Hò/òór3   r   c                   ó  ‡ — e Zd ZdZdZ	 	 ddeez  deez  dededdf
ˆ fd	„Ze	defd
„«       Z
 ej                  d¬«      d„ «       Z ej                  d¬«      d„ «       Zdd„Zd„ Ze	deez  fd„«       Zd„ Zd„ Zd„ Zd„ Zd„ Zˆ xZS )r   a¤  
    Transform via the pointwise affine mapping :math:`y = \text{loc} + \text{scale} \times x`.

    Args:
        loc (Tensor or float): Location parameter.
        scale (Tensor or float): Scale parameter.
        event_dim (int): Optional size of `event_shape`. This should be zero
            for univariate random variables, 1 for distributions over vectors,
            2 for distributions over matrices, etc.
    TÚlocÚscaler:   r&   r'   Nc                 óP   •— t         ‰| �  |¬«       || _        || _        || _        y r}   )r.   r/   r:  r;  Ú
_event_dim)r0   r:  r;  r:   r&   r1   s        €r2   r/   zAffineTransform.__init__  s*   ø€ ô 	‰Ñ JÐÔ/ØˆŒØˆŒ
Ø#ˆ�r3   c                 ó   — | j                   S rK   )r=  r;   s    r2   r:   zAffineTransform.event_dim  s   € à�‰Ðr3   Fr~   c                 óœ   — | j                   dk(  rt        j                  S t        j                  t        j                  | j                   «      S ©Nr   ©r:   r   rŸ   r¢   r;   s    r2   r$   zAffineTransform.domain  ó9   € ð �>‰>˜QÒÜ×#Ñ#Ð#Ü×&Ñ&¤{×'7Ñ'7¸¿¹ÓHÐHr3   c                 óœ   — | j                   dk(  rt        j                  S t        j                  t        j                  | j                   «      S r@  rA  r;   s    r2   r%   zAffineTransform.codomain  rB  r3   c                 ó~   — | j                   |k(  r| S t        | j                  | j                  | j                  |¬«      S r}   )r*   r   r:  r;  r:   rH   s     r2   rI   zAffineTransform.with_cache!  s7   € Ø×Ñ˜zÒ)ØˆKÜØ�H‰H�d—j‘j $§.¡.¸Zô
ð 	
r3   c                 ó8  — t        |t        «      syt        | j                  t        «      r4t        |j                  t        «      r| j                  |j                  k7  r7y| j                  |j                  k(  j	                  «       j                  «       syt        | j                  t        «      r5t        |j                  t        «      r| j                  |j                  k7  ryy| j                  |j                  k(  j	                  «       j                  «       syy)NFT)rŠ   r   r:  r   r«   r   r;  rM   s     r2   rO   zAffineTransform.__eq__(  s¿   € Ü˜%¤Ô1Øä�d—h‘h¤Ô(¬Z¸¿	¹	Ä7Ô-KØ�x‰x˜5Ÿ9™9Ò$Øà—H‘H §	¡	Ñ)×.Ñ.Ó0×5Ñ5Ô7Øä�d—j‘j¤'Ô*¬z¸%¿+¹+ÄwÔ/OØ�z‰z˜UŸ[™[Ò(Øð
 ð —J‘J %§+¡+Ñ-×2Ñ2Ó4×9Ñ9Ô;Øàr3   c                 óÖ   — t        | j                  t        «      r6t        | j                  «      dkD  rdS t        | j                  «      dk  rdS dS | j                  j	                  «       S )Nr   r)   r�   )rŠ   r;  r   ÚfloatrD   r;   s    r2   rD   zAffineTransform.sign<  sR   € ä�d—j‘j¤'Ô*Ü˜dŸj™jÓ)¨AÒ-�1ÐU¼¸t¿z¹zÓ9JÈQÒ9N°2ÐUÐTUÐUØ�z‰z�‰Ó Ð r3   c                 ó:   — | j                   | j                  |z  z   S rK   ©r:  r;  r]   s     r2   rS   zAffineTransform._callB  s   € Ø�x‰x˜$Ÿ*™* q™.Ñ(Ð(r3   c                 ó:   — || j                   z
  | j                  z  S rK   rI  r`   s     r2   rZ   zAffineTransform._inverseE  s   € Ø�D—H‘H‘ §
¡
Ñ*Ð*r3   c                 óÚ  — |j                   }| j                  }t        |t        «      r3t	        j
                  |t        j                  t        |«      «      «      }n#t	        j                  |«      j                  «       }| j                  rQ|j                  «       d | j                    dz   }|j                  |«      j                  d«      }|d | j                    }|j                  |«      S )N)r�   r�   )ri   r;  rŠ   r   r²   Ú	full_liker3  rô   r  r:   ÚsizeÚviewÚsumÚexpand)r0   rT   rW   ri   r;  rÒ   Úresult_sizes          r2   rc   z$AffineTransform.log_abs_det_jacobianH  s²   € Ø—‘ˆØ—
‘
ˆÜ�eœWÔ%Ü—_‘_ Q¬¯©´°U³Ó(<Ó=‰Fä—Y‘Y˜uÓ%×)Ñ)Ó+ˆFØ�>Š>Ø Ÿ+™+›-Ð(9¨4¯>©>¨/Ð:¸UÑBˆKØ—[‘[ Ó-×1Ñ1°"Ó5ˆFØÐ+˜TŸ^™^˜OÐ,ˆEØ�}‰}˜UÓ#Ð#r3   c           	      ó„   — t        j                  |t        | j                  dd«      t        | j                  dd«      «      S r  ©r²   r
  r  r:  r;  rh   s     r2   rj   zAffineTransform.forward_shapeU  ó7   € Ü×%Ñ%Ø”7˜4Ÿ8™8 W¨bÓ1´7¸4¿:¹:ÀwÐPRÓ3Só
ð 	
r3   c           	      ó„   — t        j                  |t        | j                  dd«      t        | j                  dd«      «      S r  rS  rh   s     r2   rm   zAffineTransform.inverse_shapeZ  rT  r3   ©r   r   ro   )re   rp   rq   rr   rs   r   rG  rv   r/   rw   r:   r   r”   r$   r%   rI   rO   rD   rS   rZ   rc   rj   rm   rx   ry   s   @r2   r   r   õ  s÷   ø„ ñ	ð €Ið Øñ
$à�e‰^ð
$ð ˜‰~ð
$ð ð	
$ð
 ð
$ð 
õ
$ð ð˜3ò ó ðð $€[×#Ñ#°Ô6ñIó 7ðIð
 $€[×#Ñ#°Ô6ñIó 7ðIó

òð( ð!�f˜s‘lò !ó ð!ò
)ò+ò$ò
ö

r3   r   c                   ód   — e Zd ZdZej
                  Zej                  ZdZ	d„ Z
d„ Zd	d„Zd„ Zd„ Zy)
r   a°  
    Transforms an unconstrained real vector :math:`x` with length :math:`D*(D-1)/2` into the
    Cholesky factor of a D-dimension correlation matrix. This Cholesky factor is a lower
    triangular matrix with positive diagonals and unit Euclidean norm for each row.
    The transform is processed as follows:

        1. First we convert x into a lower triangular matrix in row order.
        2. For each row :math:`X_i` of the lower triangular part, we apply a *signed* version of
           class :class:`StickBreakingTransform` to transform :math:`X_i` into a
           unit Euclidean length vector using the following steps:
           - Scales into the interval :math:`(-1, 1)` domain: :math:`r_i = \tanh(X_i)`.
           - Transforms into an unsigned domain: :math:`z_i = r_i^2`.
           - Applies :math:`s_i = StickBreakingTransform(z_i)`.
           - Transforms back into signed domain: :math:`y_i = sign(r_i) * \sqrt{s_i}`.
    Tc                 óÈ  — t        j                  |«      }t        j                  |j                  «      j                  }|j                  d|z   d|z
  ¬«      }t        |d¬«      }|dz  }d|z
  j                  «       j                  d«      }|t        j                  |j                  d   |j                  |j                  ¬«      z   }|t        |dd d…f   ddgd¬	«      z  }|S )
Nr�   r)   r  ©Údiagé   )r  Údevice.r   ©Úvalue)r²   r.  r  r  r  r  r   ÚsqrtÚcumprodÚeyeri   r\  r   )r0   rT   r  ÚrÚzÚz1m_cumprod_sqrtrW   s          r2   rS   zCorrCholeskyTransform._callu  sÄ   € Ü�J‰J�q‹MˆÜ�k‰k˜!Ÿ'™'Ó"×&Ñ&ˆØ�G‰G˜˜S™ a¨#¡gˆGÓ.ˆÜ˜q rÔ*ˆð
 ˆq‰DˆØ ™EŸ<™<›>×1Ñ1°"Ó5Ðà”—	‘	˜!Ÿ'™' "™+¨Q¯W©W¸Q¿X¹XÔFÑFˆØ”Ð$ S¨#¨2¨# XÑ.°°A°¸aÔ@Ñ@ˆØˆr3   c                 ó,  — dt        j                  ||z  d¬«      z
  }t        |dd d…f   ddgd¬«      }t        |d¬«      }t        |d¬«      }||j	                  «       z  }|j                  «       |j                  «       j                  «       z
  dz  }|S )	Nr)   r�   ©rÏ   .r   r]  rY  r[  )r²   Úcumsumr   r
   r_  r  r)  )r0   rW   Úy_cumsumÚy_cumsum_shiftedÚy_vecÚy_cumsum_vecÚtrT   s           r2   rZ   zCorrCholeskyTransform._inverse…  s�   € ð ”u—|‘| A¨¡E¨rÔ2Ñ2ˆÜ˜x¨¨S¨b¨S¨Ñ1°A°q°6ÀÔCÐÜ" 1¨2Ô.ˆÜ)Ð*:ÀÔDˆØ�\×'Ñ'Ó)Ñ)ˆà�W‰W‹Y˜Ÿ™›Ÿ™›Ñ(¨AÑ-ˆØˆr3   Nc                 ó  — d||z  j                  d¬«      z
  }t        |d¬«      }d|j                  «       j                  d«      z  }d|t	        d|z  «      z   t        j                  d«      z
  j                  d¬«      z  }||z   S )Nr)   r�   rf  éþÿÿÿrY  ç      à?r2  )rg  r
   rô   rO  r   r3  )r0   rT   rW   ÚintermediatesÚ
y1m_cumsumÚy1m_cumsum_trilÚstick_breaking_logdetÚtanh_logdets           r2   rc   z*CorrCholeskyTransform.log_abs_det_jacobian‘  sˆ   € ð ˜!˜a™%Ÿ™¨B˜Ó/Ñ/ˆ
ô -¨Z¸bÔAˆØ # ×&;Ñ&;Ó&=×&AÑ&AÀ"Ó&EÑ EÐØ˜A¤¨¨a©Ó 0Ñ0´4·8±8¸C³=Ñ@×EÑEÈ"ÐEÓMÑMˆØ$ {Ñ2Ð2r3   c                 ó²   — t        |«      dk  rt        d«      ‚|d   }t        dd|z  z   dz  dz   «      }||dz
  z  dz  |k7  rt        d«      ‚|d d ||fz   S )Nr)   rÎ   r�   g      Ð?r[  ro  z.Input is not a flattened lower-diagonal number)rÞ   r-   Úround)r0   ri   ÚNÚDs       r2   rj   z#CorrCholeskyTransform.forward_shapeŸ  st   € äˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�"‰IˆÜ�4˜!˜a™%‘< CÑ'¨#Ñ-Ó.ˆØ��A‘‰;˜!Ñ˜qÒ ÜÐMÓNÐNØ�S�bˆz˜Q ˜FÑ"Ð"r3   c                 ó’   — t        |«      dk  rt        d«      ‚|d   |d   k7  rt        d«      ‚|d   }||dz
  z  dz  }|d d |fz   S )Nr[  rÎ   rn  r�   zInput is not squarer)   ©rÞ   r-   )r0   ri   rx  rw  s       r2   rm   z#CorrCholeskyTransform.inverse_shape©  sc   € äˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�‰9˜˜b™	Ò!ÜÐ2Ó3Ð3Ø�"‰IˆØ��Q‘‰K˜1ÑˆØ�S�bˆz˜Q˜DÑ Ð r3   rK   )re   rp   rq   rr   r   Úreal_vectorr$   Úcorr_choleskyr%   rs   rS   rZ   rc   rj   rm   rL   r3   r2   r   r   `  s=   „ ñð  ×$Ñ$€FØ×(Ñ(€HØ€Iòò 
ó3ò#ó!r3   r   c                   ó^   — e Zd ZdZej
                  Zej                  Zd„ Z	d„ Z
d„ Zd„ Zd„ Zy)r   a<  
    Transform from unconstrained space to the simplex via :math:`y = \exp(x)` then
    normalizing.

    This is not bijective and cannot be used for HMC. However this acts mostly
    coordinate-wise (except for the final normalization), and thus is
    appropriate for coordinate-wise optimization algorithms.
    c                 ó"   — t        |t        «      S rK   )rŠ   r   rM   s     r2   rO   zSoftmaxTransform.__eq__Á  r  r3   c                 ó|   — |}||j                  dd«      d   z
  j                  «       }||j                  dd«      z  S )Nr�   Tr   )r¡   rð   rO  )r0   rT   ÚlogprobsÚprobss       r2   rS   zSoftmaxTransform._callÄ  s@   € ØˆØ˜HŸL™L¨¨TÓ2°1Ñ5Ñ5×:Ñ:Ó<ˆØ�u—y‘y  TÓ*Ñ*Ð*r3   c                 ó&   — |}|j                  «       S rK   ró   )r0   rW   r�  s      r2   rZ   zSoftmaxTransform._inverseÉ  s   € ØˆØ�y‰y‹{Ðr3   c                 ó8   — t        |«      dk  rt        d«      ‚|S ©Nr)   rÎ   rz  rh   s     r2   rj   zSoftmaxTransform.forward_shapeÍ  ó   € Üˆu‹:˜Š>ÜÐ:Ó;Ð;Øˆr3   c                 ó8   — t        |«      dk  rt        d«      ‚|S r„  rz  rh   s     r2   rm   zSoftmaxTransform.inverse_shapeÒ  r…  r3   N)re   rp   rq   rr   r   r{  r$   Úsimplexr%   rO   rS   rZ   rj   rm   rL   r3   r2   r   r   ´  s8   „ ñð ×$Ñ$€FØ×"Ñ"€Hò3ò+ò
òó
r3   r   c                   óh   — e Zd ZdZej
                  Zej                  ZdZ	d„ Z
d„ Zd„ Zd„ Zd„ Zd„ Zy	)
r    a  
    Transform from unconstrained space to the simplex of one additional
    dimension via a stick-breaking process.

    This transform arises as an iterated sigmoid transform in a stick-breaking
    construction of the `Dirichlet` distribution: the first logit is
    transformed via sigmoid to the first probability and the probability of
    everything else, and then the process recurses.

    This is bijective and appropriate for use in HMC; however it mixes
    coordinates together and is less appropriate for optimization.
    Tc                 ó"   — t        |t        «      S rK   )rŠ   r    rM   s     r2   rO   zStickBreakingTransform.__eq__ê  ó   € Ü˜%Ô!7Ó8Ð8r3   c                 ó(  — |j                   d   dz   |j                  |j                   d   «      j                  d«      z
  }t        ||j	                  «       z
  «      }d|z
  j                  d«      }t        |ddgd¬«      t        |ddgd¬«      z  }|S )Nr�   r)   r   r]  )ri   Únew_onesrg  r  rô   r`  r   )r0   rT   Úoffsetrc  Ú	z_cumprodrW   s         r2   rS   zStickBreakingTransform._callí  s„   € Ø—‘˜‘˜q‘ 1§:¡:¨a¯g©g°b©kÓ#:×#AÑ#AÀ"Ó#EÑEˆÜ˜Q §¡£Ñ-Ó.ˆØ˜‘U—O‘O BÓ'ˆ	Ü��A�q�6 Ô#¤c¨)°a¸°VÀ1Ô&EÑEˆØˆr3   c                 óš  — |dd d…f   }|j                   d   |j                  |j                   d   «      j                  d«      z
  }d|j                  d«      z
  }t        j                  |t        j
                  |j                  «      j                  ¬«      }|j                  «       |j                  «       z
  |j                  «       z   }|S )N.r�   r)   )r  )	ri   rŒ  rg  r²   r  r  r  r  rô   )r0   rW   Úy_cropr�  ÚsfrT   s         r2   rZ   zStickBreakingTransform._inverseô  s    € Ø�3˜˜˜�8‘ˆØ—‘˜‘˜qŸz™z¨&¯,©,°rÑ*:Ó;×BÑBÀ2ÓFÑFˆØ�—‘˜rÓ"Ñ"ˆô �[‰[˜¤§¡¨Q¯W©WÓ!5×!:Ñ!:Ô;ˆØ�J‰J‹L˜2Ÿ6™6›8Ñ# f§j¡j£lÑ2ˆØˆr3   c                 ó,  — |j                   d   dz   |j                  |j                   d   «      j                  d«      z
  }||j                  «       z
  }| t	        j
                  |«      z   |dd d…f   j                  «       z   j                  d«      }|S )Nr�   r)   .)ri   rŒ  rg  rô   r!  Ú
logsigmoidrO  )r0   rT   rW   r�  ÚdetJs        r2   rc   z+StickBreakingTransform.log_abs_det_jacobianþ  s€   € Ø—‘˜‘˜q‘ 1§:¡:¨a¯g©g°b©kÓ#:×#AÑ#AÀ"Ó#EÑEˆØ�—
‘
“Ñˆà�”Q—\‘\ !“_Ñ$ q¨¨c¨r¨c¨¡{§¡Ó'8Ñ8×=Ñ=¸bÓAˆØˆr3   c                 óR   — t        |«      dk  rt        d«      ‚|d d |d   dz   fz   S ©Nr)   rÎ   r�   rz  rh   s     r2   rj   z$StickBreakingTransform.forward_shape  ó5   € Üˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�S�bˆz˜U 2™Y¨™]Ð,Ñ,Ð,r3   c                 óR   — t        |«      dk  rt        d«      ‚|d d |d   dz
  fz   S r–  rz  rh   s     r2   rm   z$StickBreakingTransform.inverse_shape
  r—  r3   N)re   rp   rq   rr   r   r{  r$   r‡  r%   rs   rO   rS   rZ   rc   rj   rm   rL   r3   r2   r    r    Ø  sB   „ ñð ×$Ñ$€FØ×"Ñ"€HØ€Iò9òòòò-ó
-r3   r    c                   ót   — e Zd ZdZ ej
                  ej                  d«      Zej                  Z	d„ Z
d„ Zd„ Zy)r   zã
    Transform from unconstrained matrices to lower-triangular matrices with
    nonnegative diagonal entries.

    This is useful for parameterizing positive definite matrices in terms of
    their Cholesky factorization.
    r[  c                 ó"   — t        |t        «      S rK   )rŠ   r   rM   s     r2   rO   zLowerCholeskyTransform.__eq__  rŠ  r3   c                 ó„   — |j                  d«      |j                  dd¬«      j                  «       j                  «       z   S ©Nr�   rn  )Údim1Údim2)ÚtrilÚdiagonalrð   Ú
diag_embedr]   s     r2   rS   zLowerCholeskyTransform._call  ó4   € Ø�v‰v�b‹z˜AŸJ™J¨B°R˜JÓ8×<Ñ<Ó>×IÑIÓKÑKÐKr3   c                 ó„   — |j                  d«      |j                  dd¬«      j                  «       j                  «       z   S rœ  )rŸ  r   rô   r¡  r`   s     r2   rZ   zLowerCholeskyTransform._inverse"  r¢  r3   N)re   rp   rq   rr   r   r¢   rŸ   r$   Úlower_choleskyr%   rO   rS   rZ   rL   r3   r2   r   r     s?   „ ñð %ˆ[×$Ñ$ [×%5Ñ%5°qÓ9€FØ×)Ñ)€Hò9òLóLr3   r   c                   ót   — e Zd ZdZ ej
                  ej                  d«      Zej                  Z	d„ Z
d„ Zd„ Zy)r   zN
    Transform from unconstrained matrices to positive-definite matrices.
    r[  c                 ó"   — t        |t        «      S rK   )rŠ   r   rM   s     r2   rO   z PositiveDefiniteTransform.__eq__.  s   € Ü˜%Ô!:Ó;Ð;r3   c                 ó@   —  t        «       |«      }||j                  z  S rK   )r   ÚmTr]   s     r2   rS   zPositiveDefiniteTransform._call1  s   € Ø$Ô"Ó$ QÓ'ˆØ�1—4‘4‰xˆr3   c                 ór   — t         j                  j                  |«      }t        «       j	                  |«      S rK   )r²   ÚlinalgÚcholeskyr   r@   r`   s     r2   rZ   z"PositiveDefiniteTransform._inverse5  s*   € Ü�L‰L×!Ñ! !Ó$ˆÜ%Ó'×+Ñ+¨AÓ.Ð.r3   N)re   rp   rq   rr   r   r¢   rŸ   r$   Úpositive_definiter%   rO   rS   rZ   rL   r3   r2   r   r   &  s=   „ ñð %ˆ[×$Ñ$ [×%5Ñ%5°qÓ9€FØ×,Ñ,€Hò<òó/r3   r   c                   ó  ‡ — e Zd ZU dZee   ed<   	 	 	 ddee   dedee   dz  deddf
ˆ fd	„Z	e
defd
„«       Ze
defd„«       Zdd„Zd„ Zd„ Zd„ Zedefd„«       Zej*                  d„ «       Zej*                  d„ «       Zˆ xZS )r   aá  
    Transform functor that applies a sequence of transforms `tseq`
    component-wise to each submatrix at `dim`, of length `lengths[dim]`,
    in a way compatible with :func:`torch.cat`.

    Example::

       x0 = torch.cat([torch.range(1, 10), torch.range(1, 10)], dim=0)
       x = torch.cat([x0, x0], dim=0)
       t0 = CatTransform([ExpTransform(), identity_transform], dim=0, lengths=[10, 10])
       t = CatTransform([t0, t0], dim=0, lengths=[20, 20])
       y = t(x)
    Ú
transformsNÚtseqrÏ   Úlengthsr&   r'   c                 óô  •— t        d„ |D «       «      st        d«      ‚|r|D �cg c]  }|j                  |«      ‘Œ }}t        ‰| �  |¬«       t        |«      | _        |€dgt        | j                  «      z  }t        |«      | _        t        | j                  «      t        | j                  «      k7  r8t        dt        | j                  «      › dt        | j                  «      › d�«      ‚|| _	        y c c}w )Nc              3   ó<   K  — | ]  }t        |t        «      –— Œ y ­wrK   ©rŠ   r!   ©r§   rl  s     r2   r©   z(CatTransform.__init__.<locals>.<genexpr>R  ó   è ø€ Ò:°”:˜a¤×+Ñ:ùó   ‚ú0All elements of tseq must be Transform instancesrF   r)   z	lengths (z) must match transforms (r�   )
r«   rƒ   rI   r.   r/   rÂ   r®  rÞ   r°  rÏ   )r0   r¯  rÏ   r°  r&   rl  r1   s         €r2   r/   zCatTransform.__init__K  sÛ   ø€ ô Ñ:°TÔ:Ô:Ü Ð!SÓTÐTÙØ6:Ö;°�A—L‘L Õ,Ð;ˆDÐ;Ü‰Ñ JÐÔ/Ü˜t›*ˆŒØˆ?Ø�cœC §¡Ó0Ñ0ˆGÜ˜G“}ˆŒÜˆt�|‰|Ó¤ D§O¡OÓ 4Ò4Ü ØœC §¡Ó-Ð.Ð.GÌÈDÏOÉOÓH\ÐG]Ð]^Ð_óð ð ˆ�ùò <s   ¥C5c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrK   )r:   r´  s     r2   r©   z)CatTransform.event_dim.<locals>.<genexpr>c  ó   è ø€ Ò8 1�1—;•;Ñ8ùrª   )r¡   r®  r;   s    r2   r:   zCatTransform.event_dima  ó   € äÑ8¨¯©Ô8Ó8Ð8r3   c                 ó,   — t        | j                  «      S rK   )rO  r°  r;   s    r2   ÚlengthzCatTransform.lengthe  s   € ä�4—<‘<Ó Ð r3   c                 ó|   — | j                   |k(  r| S t        | j                  | j                  | j                  |«      S rK   )r*   r   r®  rÏ   r°  rH   s     r2   rI   zCatTransform.with_cachei  s2   € Ø×Ñ˜zÒ)ØˆKÜ˜DŸO™O¨T¯X©X°t·|±|ÀZÓPÐPr3   c                 óœ  — |j                  «        | j                   cxk  r|j                  «       k  s,n t        d| j                   › d|j                  «       › d�«      ‚|j                  | j                   «      | j                  k7  rAt        d| j                   › d|j                  | j                   «      › d| j                  › �«      ‚g }d}t	        | j
                  | j                  «      D ]>  \  }}|j                  | j                   ||«      }|j                   ||«      «       ||z   }Œ@ t        j                  || j                   ¬«      S )	Núdim ú out of range for tensor with ú dimensionsúx.size(ú) = ú must equal length r   rf  )rÏ   rƒ   rM  r½  rµ   r®  r°  Únarrowr´   r²   Úcat)r0   rT   ÚyslicesÚstartÚtransr½  Úxslices          r2   rS   zCatTransform._calln  s  € Ø—‘“�˜DŸH™HÔ. q§u¡u£wÔ.Ü Ø�t—x‘x�jÐ >¸q¿u¹u»w¸iÀ{ÐSóð ð �6‰6�$—(‘(Ó˜tŸ{™{Ò*Ü Ø˜$Ÿ(™(˜ 4¨¯©¨t¯x©xÓ(8Ð'9Ð9LÈTÏ[É[ÈMÐZóð ð ˆØˆÜ  §¡°$·,±,Ó?ò 	#‰MˆE�6Ø—X‘X˜dŸh™h¨¨vÓ6ˆFØ�N‰N™5 ›=Ô)Ø˜F‘N‰Eð	#ô �y‰y˜ d§h¡hÔ/Ð/r3   c                 ó®  — |j                  «        | j                   cxk  r|j                  «       k  s,n t        d| j                   › d|j                  «       › d�«      ‚|j                  | j                   «      | j                  k7  rAt        d| j                   › d|j                  | j                   «      › d| j                  › �«      ‚g }d}t	        | j
                  | j                  «      D ]G  \  }}|j                  | j                   ||«      }|j                  |j                  |«      «       ||z   }ŒI t        j                  || j                   ¬«      S )	NrÀ  rÁ  rÂ  úy.size(rÄ  rÅ  r   rf  )rÏ   rƒ   rM  r½  rµ   r®  r°  rÆ  r´   r@   r²   rÇ  )r0   rW   ÚxslicesrÉ  rÊ  r½  Úyslices          r2   rZ   zCatTransform._inverse  s  € Ø—‘“�˜DŸH™HÔ. q§u¡u£wÔ.Ü Ø�t—x‘x�jÐ >¸q¿u¹u»w¸iÀ{ÐSóð ð �6‰6�$—(‘(Ó˜tŸ{™{Ò*Ü Ø˜$Ÿ(™(˜ 4¨¯©¨t¯x©xÓ(8Ð'9Ð9LÈTÏ[É[ÈMÐZóð ð ˆØˆÜ  §¡°$·,±,Ó?ò 	#‰MˆE�6Ø—X‘X˜dŸh™h¨¨vÓ6ˆFØ�N‰N˜5Ÿ9™9 VÓ,Ô-Ø˜F‘N‰Eð	#ô �y‰y˜ d§h¡hÔ/Ð/r3   c                 óf  — |j                  «        | j                   cxk  r|j                  «       k  s,n t        d| j                   › d|j                  «       › d�«      ‚|j                  | j                   «      | j                  k7  rAt        d| j                   › d|j                  | j                   «      › d| j                  › �«      ‚|j                  «        | j                   cxk  r|j                  «       k  s,n t        d| j                   › d|j                  «       › d�«      ‚|j                  | j                   «      | j                  k7  rAt        d| j                   › d|j                  | j                   «      › d| j                  › �«      ‚g }d	}t	        | j
                  | j                  «      D ]£  \  }}|j                  | j                   ||«      }|j                  | j                   ||«      }|j                  ||«      }	|j                  | j                  k  r#t        |	| j                  |j                  z
  «      }	|j                  |	«       ||z   }Œ¥ | j                   }
|
d	k\  r|
|j                  «       z
  }
|
| j                  z   }
|
d	k  rt        j                  ||
¬
«      S t        |«      S )NrÀ  ú out of range for x with rÂ  rÃ  rÄ  rÅ  ú out of range for y with rÍ  r   rf  )rÏ   rƒ   rM  r½  rµ   r®  r°  rÆ  rc   r:   r   r´   r²   rÇ  rO  )r0   rT   rW   Ú
logdetjacsrÉ  rÊ  r½  rË  rÏ  Ú	logdetjacrÏ   s              r2   rc   z!CatTransform.log_abs_det_jacobian�  s=  € Ø—‘“�˜DŸH™HÔ. q§u¡u£wÔ.Ü Ø�t—x‘x�jÐ 9¸!¿%¹%»'¸À+ÐNóð ð �6‰6�$—(‘(Ó˜tŸ{™{Ò*Ü Ø˜$Ÿ(™(˜ 4¨¯©¨t¯x©xÓ(8Ð'9Ð9LÈTÏ[É[ÈMÐZóð ð —‘“�˜DŸH™HÔ. q§u¡u£wÔ.Ü Ø�t—x‘x�jÐ 9¸!¿%¹%»'¸À+ÐNóð ð �6‰6�$—(‘(Ó˜tŸ{™{Ò*Ü Ø˜$Ÿ(™(˜ 4¨¯©¨t¯x©xÓ(8Ð'9Ð9LÈTÏ[É[ÈMÐZóð ð ˆ
ØˆÜ  §¡°$·,±,Ó?ò 	#‰MˆE�6Ø—X‘X˜dŸh™h¨¨vÓ6ˆFØ—X‘X˜dŸh™h¨¨vÓ6ˆFØ×2Ñ2°6¸6ÓBˆIØ�‰ §¡Ò/Ü*¨9°d·n±nÀuÇÁÑ6VÓW�	Ø×Ñ˜iÔ(Ø˜F‘N‰Eð	#ð �h‰hˆØ�!Š8Ø˜Ÿ™›‘-ˆCØ�D—N‘NÑ"ˆØ�Š7Ü—9‘9˜Z¨SÔ1Ð1ä�z“?Ð"r3   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrK   r¦   r´  s     r2   r©   z)CatTransform.bijective.<locals>.<genexpr>·  rº  rª   ©r«   r®  r;   s    r2   rs   zCatTransform.bijectiveµ  r»  r3   c                 ó¦   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  | j
                  «      S c c}w rK   )r   rÇ  r®  r$   rÏ   r°  ©r0   rl  s     r2   r$   zCatTransform.domain¹  s:   € ô �‰Ø#Ÿ™Ö/˜!ˆQ�X‹XÒ/°·±¸4¿<¹<ó
ð 	
ùÚ/ó   žAc                 ó¦   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  | j
                  «      S c c}w rK   )r   rÇ  r®  r%   rÏ   r°  rÙ  s     r2   r%   zCatTransform.codomainÀ  s:   € ô �‰Ø!%§¡Ö1˜AˆQ�Z‹ZÒ1°4·8±8¸T¿\¹\ó
ð 	
ùÚ1rÚ  )r   Nr   ro   )re   rp   rq   rr   rÂ   r!   ru   r   rv   r/   r	   r:   r½  rI   rS   rZ   rc   rw   r•   rs   r   r”   r$   r%   rx   ry   s   @r2   r   r   :  sý   ø… ñð �Y‘Óð
 Ø(,Øñà�yÑ!ðð ðð ˜#‘ Ñ%ð	ð
 ðð 
õð, ð9˜3ò 9ó ð9ð ð!˜ò !ó ð!óQò
0ò"0ò"##ðJ ð9˜4ò 9ó ð9ð ×#Ñ#ñ
ó $ð
ð
 ×#Ñ#ñ
ó $ô
r3   r   c            	       óÎ   ‡ — e Zd ZU dZee   ed<   	 ddee   dededdfˆ fd„Z	dd	„Z
d
„ Zd„ Zd„ Zd„ Zedefd„«       Zej&                  d„ «       Zej&                  d„ «       Zˆ xZS )r   aW  
    Transform functor that applies a sequence of transforms `tseq`
    component-wise to each submatrix at `dim`
    in a way compatible with :func:`torch.stack`.

    Example::

       x = torch.stack([torch.range(1, 10), torch.range(1, 10)], dim=1)
       t = StackTransform([ExpTransform(), identity_transform], dim=1)
       y = t(x)
    r®  r¯  rÏ   r&   r'   Nc                 óØ   •— t        d„ |D «       «      st        d«      ‚|r|D �cg c]  }|j                  |«      ‘Œ }}t        ‰| �  |¬«       t        |«      | _        || _        y c c}w )Nc              3   ó<   K  — | ]  }t        |t        «      –— Œ y ­wrK   r³  r´  s     r2   r©   z*StackTransform.__init__.<locals>.<genexpr>Ú  rµ  r¶  r·  rF   )r«   rƒ   rI   r.   r/   rÂ   r®  rÏ   )r0   r¯  rÏ   r&   rl  r1   s        €r2   r/   zStackTransform.__init__×  se   ø€ ô Ñ:°TÔ:Ô:Ü Ð!SÓTÐTÙØ6:Ö;°�A—L‘L Õ,Ð;ˆDÐ;Ü‰Ñ JÐÔ/Ü˜t›*ˆŒØˆ�ùò <s   ¥A'c                 óf   — | j                   |k(  r| S t        | j                  | j                  |«      S rK   )r*   r   r®  rÏ   rH   s     r2   rI   zStackTransform.with_cacheâ  s,   € Ø×Ñ˜zÒ)ØˆKÜ˜dŸo™o¨t¯x©x¸ÓDÐDr3   c                 ó¤   — t        |j                  | j                  «      «      D �cg c]  }|j                  | j                  |«      ‘Œ  c}S c c}w rK   )ÚrangerM  rÏ   Úselect)r0   rc  Úis      r2   Ú_slicezStackTransform._sliceç  s7   € Ü/4°Q·V±V¸D¿H¹HÓ5EÓ/FÖG¨!�—‘˜Ÿ™ 1Õ%ÒGÐGùÒGs   §#Ac           
      ó‚  — |j                  «        | j                   cxk  r|j                  «       k  s,n t        d| j                   › d|j                  «       › d�«      ‚|j                  | j                   «      t        | j                  «      k7  rJt        d| j                   › d|j                  | j                   «      › dt        | j                  «      › �«      ‚g }t        | j                  |«      | j                  «      D ]  \  }}|j                   ||«      «       Œ t        j                  || j                   ¬«      S )NrÀ  rÁ  rÂ  rÃ  rÄ  ú must equal len(transforms) rf  )
rÏ   rƒ   rM  rÞ   r®  rµ   rä  r´   r²   Ústack)r0   rT   rÈ  rË  rÊ  s        r2   rS   zStackTransform._callê  s  € Ø—‘“�˜DŸH™HÔ. q§u¡u£wÔ.Ü Ø�t—x‘x�jÐ >¸q¿u¹u»w¸iÀ{ÐSóð ð �6‰6�$—(‘(Óœs 4§?¡?Ó3Ò3Ü Ø˜$Ÿ(™(˜ 4¨¯©¨t¯x©xÓ(8Ð'9Ð9UÔVYÐZ^×ZiÑZiÓVjÐUkÐlóð ð ˆÜ  §¡¨Q£°·±ÓAò 	*‰MˆF�EØ�N‰N™5 ›=Õ)ð	*ä�{‰{˜7¨¯©Ô1Ð1r3   c           
      ó”  — |j                  «        | j                   cxk  r|j                  «       k  s,n t        d| j                   › d|j                  «       › d�«      ‚|j                  | j                   «      t        | j                  «      k7  rJt        d| j                   › d|j                  | j                   «      › dt        | j                  «      › �«      ‚g }t        | j                  |«      | j                  «      D ]%  \  }}|j                  |j                  |«      «       Œ' t        j                  || j                   ¬«      S )NrÀ  rÁ  rÂ  rÍ  rÄ  ræ  rf  )rÏ   rƒ   rM  rÞ   r®  rµ   rä  r´   r@   r²   rç  )r0   rW   rÎ  rÏ  rÊ  s        r2   rZ   zStackTransform._inverseø  s  € Ø—‘“�˜DŸH™HÔ. q§u¡u£wÔ.Ü Ø�t—x‘x�jÐ >¸q¿u¹u»w¸iÀ{ÐSóð ð �6‰6�$—(‘(Óœs 4§?¡?Ó3Ò3Ü Ø˜$Ÿ(™(˜ 4¨¯©¨t¯x©xÓ(8Ð'9Ð9UÔVYÐZ^×ZiÑZiÓVjÐUkÐlóð ð ˆÜ  §¡¨Q£°·±ÓAò 	.‰MˆF�EØ�N‰N˜5Ÿ9™9 VÓ,Õ-ð	.ä�{‰{˜7¨¯©Ô1Ð1r3   c           
      ór  — |j                  «        | j                   cxk  r|j                  «       k  s,n t        d| j                   › d|j                  «       › d�«      ‚|j                  | j                   «      t        | j                  «      k7  rJt        d| j                   › d|j                  | j                   «      › dt        | j                  «      › �«      ‚|j                  «        | j                   cxk  r|j                  «       k  s,n t        d| j                   › d|j                  «       › d�«      ‚|j                  | j                   «      t        | j                  «      k7  rJt        d| j                   › d|j                  | j                   «      › dt        | j                  «      › �«      ‚g }| j                  |«      }| j                  |«      }t        ||| j                  «      D ]'  \  }}}|j                  |j                  ||«      «       Œ) t        j                  || j                   ¬	«      S )
NrÀ  rÑ  rÂ  rÃ  rÄ  ræ  rÒ  rÍ  rf  )rÏ   rƒ   rM  rÞ   r®  rä  rµ   r´   rc   r²   rç  )	r0   rT   rW   rÓ  rÈ  rÎ  rË  rÏ  rÊ  s	            r2   rc   z#StackTransform.log_abs_det_jacobian  sÔ  € Ø—‘“�˜DŸH™HÔ. q§u¡u£wÔ.Ü Ø�t—x‘x�jÐ 9¸!¿%¹%»'¸À+ÐNóð ð �6‰6�$—(‘(Óœs 4§?¡?Ó3Ò3Ü Ø˜$Ÿ(™(˜ 4¨¯©¨t¯x©xÓ(8Ð'9Ð9UÔVYÐZ^×ZiÑZiÓVjÐUkÐlóð ð —‘“�˜DŸH™HÔ. q§u¡u£wÔ.Ü Ø�t—x‘x�jÐ 9¸!¿%¹%»'¸À+ÐNóð ð �6‰6�$—(‘(Óœs 4§?¡?Ó3Ò3Ü Ø˜$Ÿ(™(˜ 4¨¯©¨t¯x©xÓ(8Ð'9Ð9UÔVYÐZ^×ZiÑZiÓVjÐUkÐlóð ð ˆ
Ø—+‘+˜a“.ˆØ—+‘+˜a“.ˆÜ%(¨°'¸4¿?¹?Ó%Kò 	JÑ!ˆF�F˜EØ×Ñ˜e×8Ñ8¸ÀÓHÕIð	Jä�{‰{˜:¨4¯8©8Ô4Ð4r3   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrK   r¦   r´  s     r2   r©   z+StackTransform.bijective.<locals>.<genexpr>   rº  rª   r×  r;   s    r2   rs   zStackTransform.bijective  r»  r3   c                 ó�   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  «      S c c}w rK   )r   rç  r®  r$   rÏ   rÙ  s     r2   r$   zStackTransform.domain"  s1   € ô × Ñ °D·O±OÖ!D¨q !§(£(Ò!DÀdÇhÁhÓOÐOùÒ!Dó   žAc                 ó�   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  «      S c c}w rK   )r   rç  r®  r%   rÏ   rÙ  s     r2   r%   zStackTransform.codomain'  s1   € ô × Ñ °d·o±oÖ!F° !§*£*Ò!FÈÏÉÓQÐQùÒ!Frí  rV  ro   )re   rp   rq   rr   rÂ   r!   ru   r   rv   r/   rI   rä  rS   rZ   rc   rw   r•   rs   r   r”   r$   r%   rx   ry   s   @r2   r   r   È  s³   ø… ñ
ð �Y‘Óð JKñ	Ø˜YÑ'ð	Ø.1ð	ØCFð	à	õ	óEò
Hò2ò2ò5ð0 ð9˜4ò 9ó ð9ð ×#Ñ#ñPó $ðPð ×#Ñ#ñRó $ôRr3   r   c                   óœ   ‡ — e Zd ZdZdZej                  ZdZdde	de
ddfˆ fd„Zedej                  dz  fd	„«       Zd
„ Zd„ Zd„ Zdd„Zˆ xZS )r   aA  
    Transform via the cumulative distribution function of a probability distribution.

    Args:
        distribution (Distribution): Distribution whose cumulative distribution function to use for
            the transformation.

    Example::

        # Construct a Gaussian copula from a multivariate normal.
        base_dist = MultivariateNormal(
            loc=torch.zeros(2),
            scale_tril=LKJCholesky(2).sample(),
        )
        transform = CumulativeDistributionTransform(Normal(0, 1))
        copula = TransformedDistribution(base_dist, [transform])
    Tr)   Údistributionr&   r'   Nc                 ó4   •— t         ‰| �  |¬«       || _        y r}   )r.   r/   rð  )r0   rð  r&   r1   s      €r2   r/   z(CumulativeDistributionTransform.__init__D  s   ø€ Ü‰Ñ JÐÔ/Ø(ˆÕr3   c                 ó.   — | j                   j                  S rK   )rð  Úsupportr;   s    r2   r$   z&CumulativeDistributionTransform.domainH  s   € à× Ñ ×(Ñ(Ð(r3   c                 ó8   — | j                   j                  |«      S rK   )rð  Úcdfr]   s     r2   rS   z%CumulativeDistributionTransform._callL  s   € Ø× Ñ ×$Ñ$ QÓ'Ð'r3   c                 ó8   — | j                   j                  |«      S rK   )rð  Úicdfr`   s     r2   rZ   z(CumulativeDistributionTransform._inverseO  s   € Ø× Ñ ×%Ñ% aÓ(Ð(r3   c                 ó8   — | j                   j                  |«      S rK   )rð  Úlog_probrb   s      r2   rc   z4CumulativeDistributionTransform.log_abs_det_jacobianR  s   € Ø× Ñ ×)Ñ)¨!Ó,Ð,r3   c                 óR   — | j                   |k(  r| S t        | j                  |¬«      S r}   )r*   r   rð  rH   s     r2   rI   z*CumulativeDistributionTransform.with_cacheU  s(   € Ø×Ñ˜zÒ)ØˆKÜ.¨t×/@Ñ/@ÈZÔXÐXr3   rn   ro   )re   rp   rq   rr   rs   r   r"  r%   rD   r   rv   r/   rw   rt   r$   rS   rZ   rc   rI   rx   ry   s   @r2   r   r   -  st   ø„ ñð$ €IØ×(Ñ(€HØ€Dñ) \ð )¸sð )È4õ )ð ð)˜×.Ñ.°Ñ5ò )ó ð)ò(ò)ò-÷Yr3   r   )1r¶   r3  r¸   r>   Úcollections.abcr   r²   Útorch.nn.functionalÚnnÚ
functionalr!  r   Útorch.distributionsr   Ú torch.distributions.distributionr   Útorch.distributions.utilsr   r   r	   r
   r   r   r   Útorch.typesr   Ú__all__r!   r=   r   r"   r   r   r   r   r  r   r   r   r   r   r   r   r    r   r   r   r   r   rL   r3   r2   ú<module>r     sg  ðã Û Û Û Ý $ã ß Ð Ý Ý +Ý 9÷õ ÷ .Ý ò€÷0fñ fôRE.˜	ô E.ôP�yô ñD & bÓ)Ð ôK8˜9ô K8ô\I+�yô I+ôX�9ô ô.(R�Yô (RòVNô
/�yô /ô2˜	ô ô0*>�Iô *>ôZ�9ô ô h
�iô h
ôVQ!˜Iô Q!ôh!�yô !ôH5-˜Yô 5-ôpL˜Yô Lô,/ 	ô /ô(K
�9ô K
ô\bR�Yô bRôJ+Y iõ +Yr3   