Ë
    f^(hÒ¢  ã                   ób  — 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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-„ d.e«      Z, G d/„ d0e«      Z-y)1é    N)ÚOptional)Úconstraints)Ú_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ˆ 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Úcodomainc                 ó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__)ÚselfÚ
cache_sizeÚ	__class__s     €ú\/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/torch/distributions/transforms.pyr+   zTransform.__init___   sB   ø€ Ø%ˆÔØ@DˆŒ	Ø˜Š?ØØ˜1Š_Ø)ˆDÕäÐ8Ó9Ð9Ü‰ÑÕó    c                 óD   — | j                   j                  «       }d |d<   |S )Nr'   )Ú__dict__Úcopy)r,   Ústates     r/   Ú__getstate__zTransform.__getstate__j   s"   € Ø—‘×"Ñ"Ó$ˆØˆˆf‰Øˆr0   Úreturnc                 óž   — | 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)   ©r,   s    r/   r8   zTransform.event_dimo   s:   € à�;‰;× Ñ  D§M¡M×$;Ñ$;Ò;Ø—;‘;×(Ñ(Ð(ÜÐUÓVÐVr0   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)r,   Úinvs     r/   r>   zTransform.invu   sB   € ð ˆØ�9‰9Ð Ø—)‘)“+ˆCØˆ;Ü# DÓ)ˆCÜŸ™ CÓ(ˆDŒIØˆ
r0   c                 ó   — t         ‚)z˜
        Returns the sign of the determinant of the Jacobian, if applicable.
        In general this only makes sense for bijective transforms.
        ©ÚNotImplementedErrorr9   s    r/   ÚsignzTransform.signƒ   s
   € ô "Ð!r0   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   rA   ©r,   r-   s     r/   Ú
with_cachezTransform.with_cache‹   sU   € Ø×Ñ˜zÒ)ØˆKÜ�‹:×Ñ¤)×"4Ñ"4Ñ4Ø”4˜“:¨Ô4Ð4Ü!¤T¨$£Z LÐ0NÐ"OÓPÐPr0   c                 ó
   — | |u S ©N© ©r,   Úothers     r/   Ú__eq__zTransform.__eq__’   s   € Ø�uˆ}Ðr0   c                 ó&   — | j                  |«       S rI   )rM   rK   s     r/   Ú__ne__zTransform.__ne__•   s   € à—;‘;˜uÓ%Ð%Ð%r0   c                 ó¤   — | j                   dk(  r| j                  |«      S | j                  \  }}||u r|S | j                  |«      }||f| _        |S )z2
        Computes the transform `x => y`.
        r   )r&   Ú_callr(   )r,   ÚxÚx_oldÚy_oldÚys        r/   Ú__call__zTransform.__call__™   sY   € ð ×Ñ˜qÒ Ø—:‘:˜a“=Ð Ø×'Ñ'‰ˆˆuØ�‰:ØˆLØ�J‰J�q‹MˆØ˜a˜4ˆÔØˆr0   c                 ó¤   — | j                   dk(  r| j                  |«      S | j                  \  }}||u r|S | j                  |«      }||f| _        |S )z1
        Inverts the transform `y => x`.
        r   )r&   Ú_inverser(   )r,   rU   rS   rT   rR   s        r/   Ú	_inv_callzTransform._inv_call¦   s[   € ð ×Ñ˜qÒ Ø—=‘= Ó#Ð#Ø×'Ñ'‰ˆˆuØ�‰:ØˆLØ�M‰M˜!ÓˆØ˜a˜4ˆÔØˆr0   c                 ó   — t         ‚)zD
        Abstract method to compute forward transformation.
        r@   ©r,   rR   s     r/   rQ   zTransform._call³   ó
   € ô "Ð!r0   c                 ó   — t         ‚)zD
        Abstract method to compute inverse transformation.
        r@   ©r,   rU   s     r/   rX   zTransform._inverse¹   r\   r0   c                 ó   — t         ‚)zU
        Computes the log det jacobian `log |dy/dx|` given input and output.
        r@   ©r,   rR   rU   s      r/   Úlog_abs_det_jacobianzTransform.log_abs_det_jacobian¿   r\   r0   c                 ó4   — | j                   j                  dz   S )Nz())r.   Ú__name__r9   s    r/   Ú__repr__zTransform.__repr__Å   s   € Ø�~‰~×&Ñ&¨Ñ-Ð-r0   c                 ó   — |S )z{
        Infers the shape of the forward computation, given the input shape.
        Defaults to preserving shape.
        rJ   ©r,   Úshapes     r/   Úforward_shapezTransform.forward_shapeÈ   ó	   € ð
 ˆr0   c                 ó   — |S )z}
        Infers the shapes of the inverse computation, given the output shape.
        Defaults to preserving shape.
        rJ   rf   s     r/   Úinverse_shapezTransform.inverse_shapeÏ   ri   r0   ©r   )r6   r   ©r%   )rc   Ú
__module__Ú__qualname__Ú__doc__Ú	bijectiver   Ú
ConstraintÚ__annotations__r+   r5   ÚpropertyÚintr8   r>   rB   rG   rM   rO   rV   rY   rQ   rX   ra   rd   rh   rk   Ú__classcell__©r.   s   @r/   r   r   .   s·   ø… ñ*ðX €IØ×"Ñ"Ó"Ø×$Ñ$Ó$õ	òð
 ðW˜3ò Wó ðWð
 òó ðð ð"�cò "ó ð"óQòò&òòò"ò"ò"ò.òör0   r   c                   óú   ‡ — e Zd ZdZde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.
    Ú	transformc                 óH   •— t         ‰| �  |j                  ¬«       || _        y ©NrD   )r*   r+   r&   r'   )r,   ry   r.   s     €r/   r+   z_InverseTransform.__init__Ý   s    ø€ Ü‰Ñ I×$9Ñ$9ÐÔ:Ø(ˆ�	r0   F©Úis_discretec                 óJ   — | j                   €J ‚| j                   j                  S rI   )r'   r#   r9   s    r/   r"   z_InverseTransform.domainá   s"   € à�y‰yÐ$Ð$Ð$Ø�y‰y×!Ñ!Ð!r0   c                 óJ   — | j                   €J ‚| j                   j                  S rI   )r'   r"   r9   s    r/   r#   z_InverseTransform.codomainæ   s"   € à�y‰yÐ$Ð$Ð$Ø�y‰y×ÑÐr0   r6   c                 óJ   — | j                   €J ‚| j                   j                  S rI   )r'   rq   r9   s    r/   rq   z_InverseTransform.bijectiveë   s"   € à�y‰yÐ$Ð$Ð$Ø�y‰y×"Ñ"Ð"r0   c                 óJ   — | j                   €J ‚| j                   j                  S rI   )r'   rB   r9   s    r/   rB   z_InverseTransform.signð   s    € à�y‰yÐ$Ð$Ð$Ø�y‰y�~‰~Ðr0   c                 ó   — | j                   S rI   )r'   r9   s    r/   r>   z_InverseTransform.invõ   s   € à�y‰yÐr0   c                 óh   — | j                   €J ‚| j                  j                  |«      j                  S rI   )r'   r>   rG   rF   s     r/   rG   z_InverseTransform.with_cacheù   s-   € Ø�y‰yÐ$Ð$Ð$Ø�x‰x×"Ñ" :Ó.×2Ñ2Ð2r0   c                 ór   — t        |t        «      sy| j                  €J ‚| j                  |j                  k(  S ©NF)Ú
isinstancer;   r'   rK   s     r/   rM   z_InverseTransform.__eq__ý   s3   € Ü˜%Ô!2Ô3ØØ�y‰yÐ$Ð$Ð$Ø�y‰y˜EŸJ™JÑ&Ð&r0   c                 ó`   — | j                   j                  › dt        | j                  «      › d�S )Nú(ú))r.   rc   Úreprr'   r9   s    r/   rd   z_InverseTransform.__repr__  s)   € Ø—.‘.×)Ñ)Ð*¨!¬D°·±«OÐ+<¸AÐ>Ð>r0   c                 óT   — | j                   €J ‚| j                   j                  |«      S rI   )r'   rY   r[   s     r/   rV   z_InverseTransform.__call__  s'   € Ø�y‰yÐ$Ð$Ð$Ø�y‰y×"Ñ" 1Ó%Ð%r0   c                 óX   — | j                   €J ‚| j                   j                  ||«       S rI   )r'   ra   r`   s      r/   ra   z&_InverseTransform.log_abs_det_jacobian
  s,   € Ø�y‰yÐ$Ð$Ð$Ø—	‘	×.Ñ.¨q°!Ó4Ð4Ð4r0   c                 ó8   — | j                   j                  |«      S rI   )r'   rk   rf   s     r/   rh   z_InverseTransform.forward_shape  ó   € Ø�y‰y×&Ñ& uÓ-Ð-r0   c                 ó8   — | j                   j                  |«      S rI   )r'   rh   rf   s     r/   rk   z_InverseTransform.inverse_shape  rŽ   r0   rm   )rc   rn   ro   rp   r   r+   r   Údependent_propertyr"   r#   rt   Úboolrq   ru   rB   r>   rG   rM   rd   rV   ra   rh   rk   rv   rw   s   @r/   r;   r;   ×   sÊ   ø„ ñð
) )õ )ð $€[×#Ñ#°Ô6ñ"ó 7ð"ð $€[×#Ñ#°Ô6ñ ó 7ð ð ð#˜4ò #ó ð#ð ð�cò ó ðð ð�Yò ó ðó3ò'ò?ò&ò5ò.ö.r0   r;   c                   ó  ‡ — e Zd ZdZddee   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.
    Úpartsc                 ó~   •— |r|D �cg c]  }|j                  |«      ‘Œ }}t        ‰| �	  |¬«       || _        y c c}w r{   )rG   r*   r+   r“   )r,   r“   r-   Úpartr.   s       €r/   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 r…   )r†   r   r“   rK   s     r/   rM   zComposeTransform.__eq__&  s#   € Ü˜%Ô!1Ô2ØØ�z‰z˜UŸ[™[Ñ(Ð(r0   Fr|   c                 ó  — | 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\  sJ ‚||j
                  kD  r#t        j                  |||j
                  z
  «      }|S )Nr   éÿÿÿÿ)	r“   r   Úrealr"   r#   r8   ÚreversedÚmaxÚindependent)r,   r"   r8   r•   s       r/   r"   zComposeTransform.domain+  sØ   € à�zŠzÜ×#Ñ#Ð#Ø—‘˜A‘×%Ñ%ˆà—J‘J˜r‘N×+Ñ+×5Ñ5ˆ	Ü˜TŸZ™ZÓ(ò 	>ˆDØ˜Ÿ™×.Ñ.°·±×1HÑ1HÑHÑHˆIÜ˜I t§{¡{×'<Ñ'<Ó=‰Ið	>ð ˜F×,Ñ,Ò,Ð,Ð,Ø�v×'Ñ'Ò'Ü ×,Ñ,¨V°YÀ×AQÑAQÑ5QÓRˆFØˆr0   c                 óþ  — | 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\  sJ ‚||j
                  kD  r#t        j                  |||j
                  z
  «      }|S )Nr˜   r   )r“   r   r™   r#   r"   r8   r›   rœ   )r,   r#   r8   r•   s       r/   r#   zComposeTransform.codomain:  sÕ   € à�zŠzÜ×#Ñ#Ð#Ø—:‘:˜b‘>×*Ñ*ˆà—J‘J˜q‘M×(Ñ(×2Ñ2ˆ	Ø—J‘Jò 	@ˆDØ˜Ÿ™×0Ñ0°4·;±;×3HÑ3HÑHÑHˆIÜ˜I t§}¡}×'>Ñ'>Ó?‰Ið	@ð ˜H×.Ñ.Ò.Ð.Ð.Ø�x×)Ñ)Ò)Ü"×.Ñ.¨x¸ÀX×EWÑEWÑ9WÓXˆHØˆr0   r6   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrI   ©rq   )Ú.0Úps     r/   ú	<genexpr>z-ComposeTransform.bijective.<locals>.<genexpr>K  s   è ø€ Ò3 1�1—;•;Ñ3ùó   ‚)Úallr“   r9   s    r/   rq   zComposeTransform.bijectiveI  s   € äÑ3¨¯
©
Ô3Ó3Ð3r0   c                 óJ   — d}| j                   D ]  }||j                  z  }Œ |S ©Nr%   )r“   rB   )r,   rB   r¢   s      r/   rB   zComposeTransform.signM  s,   € àˆØ—‘ò 	!ˆAØ˜!Ÿ&™&‘=‰Dð	!àˆr0   c                 ó$  — d }| j                   �| j                  «       }|€jt        t        | j                  «      D �cg c]  }|j                  ‘Œ c}«      }t        j                  |«      | _         t        j                  | «      |_         |S c c}w rI   )r'   r   rš   r“   r>   r<   r=   )r,   r>   r¢   s      r/   r>   zComposeTransform.invT  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“   rF   s     r/   rG   zComposeTransform.with_cache_  s&   € Ø×Ñ˜zÒ)ØˆKÜ §
¡
°zÔBÐBr0   c                 ó8   — | j                   D ]
  } ||«      }Œ |S rI   )r“   )r,   rR   r•   s      r/   rV   zComposeTransform.__call__d  s#   € Ø—J‘Jò 	ˆDÙ�Q“‰Að	àˆr0   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"   r8   Úzipr   ra   r#   Ú	functoolsÚreduceÚoperatorÚadd)r,   rR   rU   Úxsr•   Útermsr8   s          r/   ra   z%ComposeTransform.log_abs_det_jacobiani  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Ð4r0   c                 óJ   — | j                   D ]  }|j                  |«      }Œ |S rI   )r“   rh   ©r,   rg   r•   s      r/   rh   zComposeTransform.forward_shape~  s*   € Ø—J‘Jò 	.ˆDØ×&Ñ& uÓ-‰Eð	.àˆr0   c                 ó\   — t        | j                  «      D ]  }|j                  |«      }Œ |S rI   )rš   r“   rk   r·   s      r/   rk   zComposeTransform.inverse_shapeƒ  s/   € Ü˜TŸZ™ZÓ(ò 	.ˆDØ×&Ñ& uÓ-‰Eð	.àˆr0   c                 óÀ   — | j                   j                  dz   }|dj                  | j                  D �cg c]  }|j	                  «       ‘Œ c}«      z  }|dz  }|S c c}w )Nz(
    z,
    z
))r.   rc   Újoinr“   rd   )r,   Ú
fmt_stringr¢   s      r/   rd   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
rl   rm   )rc   rn   ro   rp   Úlistr   r+   rM   r   r�   r"   r#   r   r‘   rq   ru   rB   rt   r>   rG   rV   ra   rh   rk   rd   rv   rw   s   @r/   r   r     sÏ   ø„ ññ˜d 9™oõ ò)ð
 $€[×#Ñ#°Ô6ñó 7ðð $€[×#Ñ#°Ô6ñó 7ðð ð4˜4ò 4ó ð4ð ð�cò ó ðð ð�Yò ó ðóCò
ò
5ò*ò
ö
r0   r   c                   óà   ‡ — e Zd ZdZdˆ 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.
    c                 ó`   •— t         ‰| �  |¬«       |j                  |«      | _        || _        y r{   )r*   r+   rG   Úbase_transformÚreinterpreted_batch_ndims)r,   r¿   rÀ   r-   r.   s       €r/   r+   zIndependentTransform.__init__   s.   ø€ Ü‰Ñ JÐÔ/Ø,×7Ñ7¸
ÓCˆÔØ)BˆÕ&r0   c                 óh   — | j                   |k(  r| S t        | j                  | j                  |¬«      S r{   )r&   r   r¿   rÀ   rF   s     r/   rG   zIndependentTransform.with_cache¥  s5   € Ø×Ñ˜zÒ)ØˆKÜ#Ø×Ñ ×!?Ñ!?ÈJô
ð 	
r0   Fr|   c                 ój   — t        j                  | j                  j                  | j                  «      S rI   )r   rœ   r¿   r"   rÀ   r9   s    r/   r"   zIndependentTransform.domain¬  s,   € ä×&Ñ&Ø×Ñ×&Ñ&¨×(FÑ(Fó
ð 	
r0   c                 ój   — t        j                  | j                  j                  | j                  «      S rI   )r   rœ   r¿   r#   rÀ   r9   s    r/   r#   zIndependentTransform.codomain²  s,   € ä×&Ñ&Ø×Ñ×(Ñ(¨$×*HÑ*Hó
ð 	
r0   r6   c                 ó.   — | j                   j                  S rI   )r¿   rq   r9   s    r/   rq   zIndependentTransform.bijective¸  s   € à×"Ñ"×,Ñ,Ð,r0   c                 ó.   — | j                   j                  S rI   )r¿   rB   r9   s    r/   rB   zIndependentTransform.sign¼  s   € à×"Ñ"×'Ñ'Ð'r0   c                 óˆ   — |j                  «       | j                  j                  k  rt        d«      ‚| j	                  |«      S ©NúToo few dimensions on input)Údimr"   r8   r)   r¿   r[   s     r/   rQ   zIndependentTransform._callÀ  s7   € Ø�5‰5‹7�T—[‘[×*Ñ*Ò*ÜÐ:Ó;Ð;Ø×"Ñ" 1Ó%Ð%r0   c                 óœ   — |j                  «       | j                  j                  k  rt        d«      ‚| j                  j                  |«      S rÇ   )rÉ   r#   r8   r)   r¿   r>   r^   s     r/   rX   zIndependentTransform._inverseÅ  s=   € Ø�5‰5‹7�T—]‘]×,Ñ,Ò,ÜÐ:Ó;Ð;Ø×"Ñ"×&Ñ& qÓ)Ð)r0   c                 ój   — | j                   j                  ||«      }t        || j                  «      }|S rI   )r¿   ra   r   rÀ   )r,   rR   rU   Úresults       r/   ra   z)IndependentTransform.log_abs_det_jacobianÊ  s1   € Ø×$Ñ$×9Ñ9¸!¸QÓ?ˆÜ ¨×(FÑ(FÓGˆØˆr0   c                 óz   — | j                   j                  › dt        | j                  «      › d| j                  › d�S )Nrˆ   z, r‰   )r.   rc   rŠ   r¿   rÀ   r9   s    r/   rd   zIndependentTransform.__repr__Ï  s:   € Ø—.‘.×)Ñ)Ð*¨!¬D°×1DÑ1DÓ,EÐ+FÀbÈ×IgÑIgÐHhÐhiÐjÐjr0   c                 ó8   — | j                   j                  |«      S rI   )r¿   rh   rf   s     r/   rh   z"IndependentTransform.forward_shapeÒ  ó   € Ø×"Ñ"×0Ñ0°Ó7Ð7r0   c                 ó8   — | j                   j                  |«      S rI   )r¿   rk   rf   s     r/   rk   z"IndependentTransform.inverse_shapeÕ  rÏ   r0   rl   rm   )rc   rn   ro   rp   r+   rG   r   r�   r"   r#   rt   r‘   rq   ru   rB   rQ   rX   ra   rd   rh   rk   rv   rw   s   @r/   r   r   ’  sª   ø„ ñõCó

ð $€[×#Ñ#°Ô6ñ
ó 7ð
ð
 $€[×#Ñ#°Ô6ñ
ó 7ð
ð
 ð-˜4ò -ó ð-ð ð(�cò (ó ð(ò&ò
*ò
ò
kò8ö8r0   r   c                   ó–   ‡ — e Zd ZdZdZdˆ 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.)
    Tc                 ó  •— t        j                  |«      | _        t        j                  |«      | _        | j                  j	                  «       | j                  j	                  «       k7  rt        d«      ‚t        ‰| �  |¬«       y )Nz6in_shape, out_shape have different numbers of elementsrD   )r¬   ÚSizeÚin_shapeÚ	out_shapeÚnumelr)   r*   r+   )r,   rÔ   rÕ   r-   r.   s       €r/   r+   zReshapeTransform.__init__é  sa   ø€ ÜŸ
™
 8Ó,ˆŒÜŸ™ IÓ.ˆŒØ�=‰=×ÑÓ  D§N¡N×$8Ñ$8Ó$:Ò:ÜÐUÓVÐVÜ‰Ñ JÐÕ/r0   c                 óp   — t        j                  t         j                  t        | j                  «      «      S rI   )r   rœ   r™   ÚlenrÔ   r9   s    r/   r"   zReshapeTransform.domainð  s$   € ä×&Ñ&¤{×'7Ñ'7¼¸T¿]¹]Ó9KÓLÐLr0   c                 óp   — t        j                  t         j                  t        | j                  «      «      S rI   )r   rœ   r™   rØ   rÕ   r9   s    r/   r#   zReshapeTransform.codomainô  s$   € ä×&Ñ&¤{×'7Ñ'7¼¸T¿^¹^Ó9LÓMÐMr0   c                 óh   — | j                   |k(  r| S t        | j                  | j                  |¬«      S r{   )r&   r   rÔ   rÕ   rF   s     r/   rG   zReshapeTransform.with_cacheø  s,   € Ø×Ñ˜zÒ)ØˆKÜ §¡¨t¯~©~È*ÔUÐUr0   c                 ó¤   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  || j
                  z   «      S rI   )rg   rÉ   rØ   rÔ   ÚreshaperÕ   )r,   rR   Úbatch_shapes      r/   rQ   zReshapeTransform._callý  s?   € Ø—g‘gÐ< §¡£¬#¨d¯m©mÓ*<Ñ <Ð=ˆØ�y‰y˜ t§~¡~Ñ5Ó6Ð6r0   c                 ó¤   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  || j
                  z   «      S rI   )rg   rÉ   rØ   rÕ   rÜ   rÔ   )r,   rU   rÝ   s      r/   rX   zReshapeTransform._inverse  s?   € Ø—g‘gÐ= §¡£¬#¨d¯n©nÓ*=Ñ =Ð>ˆØ�y‰y˜ t§}¡}Ñ4Ó5Ð5r0   c                 óŠ   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  |«      S rI   )rg   rÉ   rØ   rÔ   Ú	new_zeros)r,   rR   rU   rÝ   s       r/   ra   z%ReshapeTransform.log_abs_det_jacobian  s6   € Ø—g‘gÐ< §¡£¬#¨d¯m©mÓ*<Ñ <Ð=ˆØ�{‰{˜;Ó'Ð'r0   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Õ   ©r,   rg   Úcuts      r/   rh   zReshapeTransform.forward_shape	  sŠ   € Üˆu‹:œ˜DŸM™MÓ*Ò*ÜÐ:Ó;Ð;Ü�%‹jœ3˜tŸ}™}Ó-Ñ-ˆØ��ˆ;˜$Ÿ-™-Ò'ÜØ+¨E°#°$¨K¨=¸	À$Ç-Á-ÀÐQóð ð �T�cˆ{˜TŸ^™^Ñ+Ð+r0   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      r/   rk   zReshapeTransform.inverse_shape  s‹   € Üˆu‹:œ˜DŸN™NÓ+Ò+ÜÐ:Ó;Ð;Ü�%‹jœ3˜tŸ~™~Ó.Ñ.ˆØ��ˆ;˜$Ÿ.™.Ò(ÜØ+¨E°#°$¨K¨=¸	À$Ç.Á.ÐAQÐRóð ð �T�cˆ{˜TŸ]™]Ñ*Ð*r0   rl   rm   )rc   rn   ro   rp   rq   r+   r   r�   r"   r#   rG   rQ   rX   ra   rh   rk   rv   rw   s   @r/   r   r   Ù  sk   ø„ ñð €Iõ0ð ×#Ñ#ñMó $ðMð ×#Ñ#ñNó $ðNóVò
7ò6ò(ò,ö+r0   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 rI   )r†   r   rK   s     r/   rM   zExpTransform.__eq__(  ó   € Ü˜%¤Ó.Ð.r0   c                 ó"   — |j                  «       S rI   )Úexpr[   s     r/   rQ   zExpTransform._call+  ó   € Ø�u‰u‹wˆr0   c                 ó"   — |j                  «       S rI   ©Úlogr^   s     r/   rX   zExpTransform._inverse.  rë   r0   c                 ó   — |S rI   rJ   r`   s      r/   ra   z!ExpTransform.log_abs_det_jacobian1  ó   € Øˆr0   N©rc   rn   ro   rp   r   r™   r"   Úpositiver#   rq   rB   rM   rQ   rX   ra   rJ   r0   r/   r   r     s=   „ ñð ×Ñ€FØ×#Ñ#€HØ€IØ€Dò/òòór0   r   c                   óš   ‡ — e Zd ZdZej
                  Zej
                  ZdZdˆ 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}}`.
    Tc                 óJ   •— t         ‰| �  |¬«       t        |«      \  | _        y r{   )r*   r+   r   Úexponent)r,   rõ   r-   r.   s      €r/   r+   zPowerTransform.__init__>  s"   ø€ Ü‰Ñ JÐÔ/Ü(¨Ó2Ñˆ�r0   c                 óR   — | j                   |k(  r| S t        | j                  |¬«      S r{   )r&   r   rõ   rF   s     r/   rG   zPowerTransform.with_cacheB  s&   € Ø×Ñ˜zÒ)ØˆKÜ˜dŸm™m¸
ÔCÐCr0   r6   c                 ó6   — | j                   j                  «       S rI   )rõ   rB   r9   s    r/   rB   zPowerTransform.signG  s   € à�}‰}×!Ñ!Ó#Ð#r0   c                 ó¦   — t        |t        «      sy| j                  j                  |j                  «      j	                  «       j                  «       S r…   )r†   r   rõ   Úeqr¥   ÚitemrK   s     r/   rM   zPowerTransform.__eq__K  s:   € Ü˜%¤Ô0ØØ�}‰}×Ñ §¡Ó/×3Ñ3Ó5×:Ñ:Ó<Ð<r0   c                 ó8   — |j                  | j                  «      S rI   ©Úpowrõ   r[   s     r/   rQ   zPowerTransform._callP  s   € Ø�u‰u�T—]‘]Ó#Ð#r0   c                 ó>   — |j                  d| j                  z  «      S r§   rü   r^   s     r/   rX   zPowerTransform._inverseS  s   € Ø�u‰u�Q˜Ÿ™Ñ&Ó'Ð'r0   c                 ó^   — | j                   |z  |z  j                  «       j                  «       S rI   )rõ   Úabsrî   r`   s      r/   ra   z#PowerTransform.log_abs_det_jacobianV  s(   € Ø—‘ Ñ! AÑ%×*Ñ*Ó,×0Ñ0Ó2Ð2r0   c                 óX   — t        j                  |t        | j                  dd«      «      S ©Nrg   rJ   ©r¬   Úbroadcast_shapesÚgetattrrõ   rf   s     r/   rh   zPowerTransform.forward_shapeY  ó"   € Ü×%Ñ% e¬W°T·]±]ÀGÈRÓ-PÓQÐQr0   c                 óX   — t        j                  |t        | j                  dd«      «      S r  r  rf   s     r/   rk   zPowerTransform.inverse_shape\  r  r0   rl   rm   )rc   rn   ro   rp   r   rò   r"   r#   rq   r+   rG   r   ru   rB   rM   rQ   rX   ra   rh   rk   rv   rw   s   @r/   r   r   5  sk   ø„ ñð ×!Ñ!€FØ×#Ñ#€HØ€Iõ3óDð
 ð$�cò $ó ð$ò=ò
$ò(ò3òRöRr0   r   c                 óÄ   — t        j                  | j                  «      }t        j                  t        j                  | «      |j
                  d|j                  z
  ¬«      S ©Nç      ð?©Úminr›   )r¬   ÚfinfoÚdtypeÚclampÚsigmoidÚtinyÚeps)rR   r  s     r/   Ú_clipped_sigmoidr  `  s<   € Ü�K‰K˜Ÿ™Ó €EÜ�;‰;”u—}‘} QÓ'¨U¯Z©Z¸SÀ5Ç9Á9¹_ÔMÐMr0   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 rI   )r†   r   rK   s     r/   rM   zSigmoidTransform.__eq__o  ó   € Ü˜%Ô!1Ó2Ð2r0   c                 ó   — t        |«      S rI   )r  r[   s     r/   rQ   zSigmoidTransform._callr  s   € Ü Ó"Ð"r0   c                 óØ   — t        j                  |j                  «      }|j                  |j                  d|j
                  z
  ¬«      }|j                  «       | j                  «       z
  S r	  )r¬   r  r  r  r  r  rî   Úlog1p)r,   rU   r  s      r/   rX   zSigmoidTransform._inverseu  sK   € Ü—‘˜AŸG™GÓ$ˆØ�G‰G˜Ÿ
™
¨¨e¯i©i©ˆGÓ8ˆØ�u‰u‹w˜1˜"Ÿ™›Ñ%Ð%r0   c                 ó\   — t        j                  | «       t        j                  |«      z
  S rI   )ÚFr   r`   s      r/   ra   z%SigmoidTransform.log_abs_det_jacobianz  s!   € Ü—
‘
˜A˜2“ˆ¤§¡¨A£Ñ.Ð.r0   N)rc   rn   ro   rp   r   r™   r"   Úunit_intervalr#   rq   rB   rM   rQ   rX   ra   rJ   r0   r/   r   r   e  s=   „ ñð ×Ñ€FØ×(Ñ(€HØ€IØ€Dò3ò#ò&ó
/r0   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 rI   )r†   r   rK   s     r/   rM   zSoftplusTransform.__eq__‰  s   € Ü˜%Ô!2Ó3Ð3r0   c                 ó   — t        |«      S rI   ©r   r[   s     r/   rQ   zSoftplusTransform._callŒ  s   € Ü˜‹{Ðr0   c                 ób   — | j                  «       j                  «       j                  «       |z   S rI   )Úexpm1Únegrî   r^   s     r/   rX   zSoftplusTransform._inverse�  s'   € Ø��z‰z‹|×ÑÓ!×%Ñ%Ó'¨!Ñ+Ð+r0   c                 ó   — t        | «       S rI   r   r`   s      r/   ra   z&SoftplusTransform.log_abs_det_jacobian’  s   € Ü˜!˜“ˆ}Ðr0   Nrñ   rJ   r0   r/   r   r   ~  s=   „ ñð
 ×Ñ€FØ×#Ñ#€HØ€IØ€Dò4òò,ór0   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 rI   )r†   r   rK   s     r/   rM   zTanhTransform.__eq__²  s   € Ü˜%¤Ó/Ð/r0   c                 ó"   — |j                  «       S rI   )Útanhr[   s     r/   rQ   zTanhTransform._callµ  s   € Ø�v‰v‹xˆr0   c                 ó,   — t        j                  |«      S rI   )r¬   Úatanhr^   s     r/   rX   zTanhTransform._inverse¸  s   € ô �{‰{˜1‹~Ðr0   c                 óV   — dt        j                  d«      |z
  t        d|z  «      z
  z  S )Nç       @g       À)Úmathrî   r   r`   s      r/   ra   z"TanhTransform.log_abs_det_jacobian½  s*   € ð ”d—h‘h˜s“m aÑ'¬(°4¸!±8Ó*<Ñ<Ñ=Ð=r0   N)rc   rn   ro   rp   r   r™   r"   Úintervalr#   rq   rB   rM   rQ   rX   ra   rJ   r0   r/   r   r   –  sF   „ ñð, ×Ñ€FØ#ˆ{×#Ñ# D¨#Ó.€HØ€IØ€Dò0òòó
>r0   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 rI   )r†   r   rK   s     r/   rM   zAbsTransform.__eq__É  rè   r0   c                 ó"   — |j                  «       S rI   )r   r[   s     r/   rQ   zAbsTransform._callÌ  rë   r0   c                 ó   — |S rI   rJ   r^   s     r/   rX   zAbsTransform._inverseÏ  rð   r0   N)rc   rn   ro   rp   r   r™   r"   rò   r#   rM   rQ   rX   rJ   r0   r/   r   r   Ã  s*   „ Ù5à×Ñ€FØ×#Ñ#€Hò/òór0   r   c                   óä   ‡ — e Zd ZdZdZdˆ 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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.
    Tc                 óP   •— t         ‰| �  |¬«       || _        || _        || _        y r{   )r*   r+   ÚlocÚscaleÚ
_event_dim)r,   r5  r6  r8   r-   r.   s        €r/   r+   zAffineTransform.__init__á  s(   ø€ Ü‰Ñ JÐÔ/ØˆŒØˆŒ
Ø#ˆ�r0   r6   c                 ó   — | j                   S rI   )r7  r9   s    r/   r8   zAffineTransform.event_dimç  s   € à�‰Ðr0   Fr|   c                 óœ   — | j                   dk(  rt        j                  S t        j                  t        j                  | j                   «      S ©Nr   ©r8   r   r™   rœ   r9   s    r/   r"   zAffineTransform.domainë  ó7   € à�>‰>˜QÒÜ×#Ñ#Ð#Ü×&Ñ&¤{×'7Ñ'7¸¿¹ÓHÐHr0   c                 óœ   — | j                   dk(  rt        j                  S t        j                  t        j                  | j                   «      S r:  r;  r9   s    r/   r#   zAffineTransform.codomainñ  r<  r0   c                 ó~   — | j                   |k(  r| S t        | j                  | j                  | j                  |¬«      S r{   )r&   r   r5  r6  r8   rF   s     r/   rG   zAffineTransform.with_cache÷  s7   € Ø×Ñ˜zÒ)ØˆKÜØ�H‰H�d—j‘j $§.¡.¸Zô
ð 	
r0   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   r5  r   r¥   rú   r6  rK   s     r/   rM   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Ô;Øàr0   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†   r6  r   ÚfloatrB   r9   s    r/   rB   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�‰Ó Ð r0   c                 ó:   — | j                   | j                  |z  z   S rI   ©r5  r6  r[   s     r/   rQ   zAffineTransform._call  s   € Ø�x‰x˜$Ÿ*™* q™.Ñ(Ð(r0   c                 ó:   — || j                   z
  | j                  z  S rI   rC  r^   s     r/   rX   zAffineTransform._inverse  s   € Ø�D—H‘H‘ §
¡
Ñ*Ð*r0   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˜   )rg   r6  r†   r   r¬   Ú	full_liker-  rî   r   r8   ÚsizeÚviewÚsumÚexpand)r,   rR   rU   rg   r6  rÌ   Úresult_sizes          r/   ra   z$AffineTransform.log_abs_det_jacobian  s²   € Ø—‘ˆØ—
‘
ˆÜ�eœWÔ%Ü—_‘_ Q¬¯©´°U³Ó(<Ó=‰Fä—Y‘Y˜uÓ%×)Ñ)Ó+ˆFØ�>Š>Ø Ÿ+™+›-Ð(9¨4¯>©>¨/Ð:¸UÑBˆKØ—[‘[ Ó-×1Ñ1°"Ó5ˆFØÐ+˜TŸ^™^˜OÐ,ˆEØ�}‰}˜UÓ#Ð#r0   c           	      ó„   — t        j                  |t        | j                  dd«      t        | j                  dd«      «      S r  ©r¬   r  r  r5  r6  rf   s     r/   rh   zAffineTransform.forward_shape+  ó7   € Ü×%Ñ%Ø”7˜4Ÿ8™8 W¨bÓ1´7¸4¿:¹:ÀwÐPRÓ3Só
ð 	
r0   c           	      ó„   — t        j                  |t        | j                  dd«      t        | j                  dd«      «      S r  rM  rf   s     r/   rk   zAffineTransform.inverse_shape0  rN  r0   ©r   r   rm   )rc   rn   ro   rp   rq   r+   rt   ru   r8   r   r�   r"   r#   rG   rM   rB   rQ   rX   ra   rh   rk   rv   rw   s   @r/   r   r   Ó  s³   ø„ ñ	ð €Iõ$ð ð˜3ò ó ðð $€[×#Ñ#°Ô6ñIó 7ðIð
 $€[×#Ñ#°Ô6ñIó 7ðIó

òð( ð!�cò !ó ð!ò
)ò+ò$ò
ö

r0   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 uncontrained 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Úeyerg   rV  r
   )r,   rR   r  ÚrÚzÚz1m_cumprod_sqrtrU   s          r/   rQ   zCorrCholeskyTransform._callK  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Ô@Ñ@ˆØˆr0   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   rW  rS  rU  )r¬   Úcumsumr
   r   rY  r  r#  )r,   rU   Úy_cumsumÚy_cumsum_shiftedÚy_vecÚy_cumsum_vecÚtrR   s           r/   rX   zCorrCholeskyTransform._inverseZ  s�   € ð ”u—|‘| A¨¡E¨rÔ2Ñ2ˆÜ˜x¨¨S¨b¨S¨Ñ1°A°q°6ÀÔCÐÜ" 1¨2Ô.ˆÜ)Ð*:ÀÔDˆØ�\×'Ñ'Ó)Ñ)ˆà�W‰W‹Y˜Ÿ™›Ÿ™›Ñ(¨AÑ-ˆØˆr0   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˜   r`  éþÿÿÿrS  ç      à?r,  )ra  r   rî   rI  r   r-  )r,   rR   rU   ÚintermediatesÚ
y1m_cumsumÚy1m_cumsum_trilÚstick_breaking_logdetÚtanh_logdets           r/   ra   z*CorrCholeskyTransform.log_abs_det_jacobianf  sˆ   € ð ˜!˜a™%Ÿ™¨B˜Ó/Ñ/ˆ
ô -¨Z¸bÔAˆØ # ×&;Ñ&;Ó&=×&AÑ&AÀ"Ó&EÑ EÐØ˜A¤¨¨a©Ó 0Ñ0´4·8±8¸C³=Ñ@×EÑEÈ"ÐEÓMÑMˆØ$ {Ñ2Ð2r0   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      Ð?rU  ri  z-Input is not a flattend lower-diagonal number)rØ   r)   Úround)r,   rg   ÚNÚDs       r/   rh   z#CorrCholeskyTransform.forward_shapet  st   € äˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�"‰IˆÜ�4˜!˜a™%‘< CÑ'¨#Ñ-Ó.ˆØ��A‘‰;˜!Ñ˜qÒ ÜÐLÓMÐMØ�S�bˆz˜Q ˜FÑ"Ð"r0   c                 ó’   — t        |«      dk  rt        d«      ‚|d   |d   k7  rt        d«      ‚|d   }||dz
  z  dz  }|d d |fz   S )NrU  rÈ   rh  r˜   zInput is not squarer%   ©rØ   r)   )r,   rg   rr  rq  s       r/   rk   z#CorrCholeskyTransform.inverse_shape~  sc   € äˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�‰9˜˜b™	Ò!ÜÐ2Ó3Ð3Ø�"‰IˆØ��Q‘‰K˜1ÑˆØ�S�bˆz˜Q˜DÑ Ð r0   rI   )rc   rn   ro   rp   r   Úreal_vectorr"   Úcorr_choleskyr#   rq   rQ   rX   ra   rh   rk   rJ   r0   r/   r   r   6  s=   „ ñð  ×$Ñ$€FØ×(Ñ(€HØ€Iòò
ó3ò#ó!r0   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 rI   )r†   r   rK   s     r/   rM   zSoftmaxTransform.__eq__–  r  r0   c                 ó|   — |}||j                  dd«      d   z
  j                  «       }||j                  dd«      z  S )Nr˜   Tr   )r›   rê   rI  )r,   rR   ÚlogprobsÚprobss       r/   rQ   zSoftmaxTransform._call™  s@   € ØˆØ˜HŸL™L¨¨TÓ2°1Ñ5Ñ5×:Ñ:Ó<ˆØ�u—y‘y  TÓ*Ñ*Ð*r0   c                 ó&   — |}|j                  «       S rI   rí   )r,   rU   r{  s      r/   rX   zSoftmaxTransform._inversež  s   € ØˆØ�y‰y‹{Ðr0   c                 ó8   — t        |«      dk  rt        d«      ‚|S ©Nr%   rÈ   rt  rf   s     r/   rh   zSoftmaxTransform.forward_shape¢  ó   € Üˆu‹:˜Š>ÜÐ:Ó;Ð;Øˆr0   c                 ó8   — t        |«      dk  rt        d«      ‚|S r~  rt  rf   s     r/   rk   zSoftmaxTransform.inverse_shape§  r  r0   N)rc   rn   ro   rp   r   ru  r"   Úsimplexr#   rM   rQ   rX   rh   rk   rJ   r0   r/   r   r   ‰  s8   „ ñð ×$Ñ$€FØ×"Ñ"€Hò3ò+ò
òó
r0   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 rI   )r†   r   rK   s     r/   rM   zStickBreakingTransform.__eq__¿  ó   € Ü˜%Ô!7Ó8Ð8r0   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   rW  )rg   Únew_onesra  r  rî   rZ  r
   )r,   rR   Úoffsetr]  Ú	z_cumprodrU   s         r/   rQ   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ˆØˆr0   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  )	rg   r†  ra  r¬   r  r  r  r  rî   )r,   rU   Úy_cropr‡  ÚsfrR   s         r/   rX   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ˆØˆr0   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%   .)rg   r†  ra  rî   r  Ú
logsigmoidrI  )r,   rR   rU   r‡  ÚdetJs        r/   ra   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ˆØˆr0   c                 óR   — t        |«      dk  rt        d«      ‚|d d |d   dz   fz   S ©Nr%   rÈ   r˜   rt  rf   s     r/   rh   z$StickBreakingTransform.forward_shapeÚ  ó5   € Üˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�S�bˆz˜U 2™Y¨™]Ð,Ñ,Ð,r0   c                 óR   — t        |«      dk  rt        d«      ‚|d d |d   dz
  fz   S r�  rt  rf   s     r/   rk   z$StickBreakingTransform.inverse_shapeß  r‘  r0   N)rc   rn   ro   rp   r   ru  r"   r�  r#   rq   rM   rQ   rX   ra   rh   rk   rJ   r0   r/   r   r   ­  sB   „ ñð ×$Ñ$€FØ×"Ñ"€HØ€Iò9òòòò-ó
-r0   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.
    rU  c                 ó"   — t        |t        «      S rI   )r†   r   rK   s     r/   rM   zLowerCholeskyTransform.__eq__ñ  r„  r0   c                 ó„   — |j                  d«      |j                  dd¬«      j                  «       j                  «       z   S ©Nr˜   rh  )Údim1Údim2)ÚtrilÚdiagonalrê   Ú
diag_embedr[   s     r/   rQ   zLowerCholeskyTransform._callô  ó4   € Ø�v‰v�b‹z˜AŸJ™J¨B°R˜JÓ8×<Ñ<Ó>×IÑIÓKÑKÐKr0   c                 ó„   — |j                  d«      |j                  dd¬«      j                  «       j                  «       z   S r–  )r™  rš  rî   r›  r^   s     r/   rX   zLowerCholeskyTransform._inverse÷  rœ  r0   N)rc   rn   ro   rp   r   rœ   r™   r"   Úlower_choleskyr#   rM   rQ   rX   rJ   r0   r/   r   r   å  s?   „ ñð %ˆ[×$Ñ$ [×%5Ñ%5°qÓ9€FØ×)Ñ)€Hò9òLóLr0   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.
    rU  c                 ó"   — t        |t        «      S rI   )r†   r   rK   s     r/   rM   z PositiveDefiniteTransform.__eq__  s   € Ü˜%Ô!:Ó;Ð;r0   c                 ó@   —  t        «       |«      }||j                  z  S rI   )r   ÚmTr[   s     r/   rQ   zPositiveDefiniteTransform._call  s   € Ø$Ô"Ó$ QÓ'ˆØ�1—4‘4‰xˆr0   c                 ór   — t         j                  j                  |«      }t        «       j	                  |«      S rI   )r¬   ÚlinalgÚcholeskyr   r>   r^   s     r/   rX   z"PositiveDefiniteTransform._inverse
  s*   € Ü�L‰L×!Ñ! !Ó$ˆÜ%Ó'×+Ñ+¨AÓ.Ð.r0   N)rc   rn   ro   rp   r   rœ   r™   r"   Úpositive_definiter#   rM   rQ   rX   rJ   r0   r/   r   r   û  s=   „ ñð %ˆ[×$Ñ$ [×%5Ñ%5°qÓ9€FØ×,Ñ,€Hò<òó/r0   r   c                   óÚ   ‡ — e Zd ZU dZee   ed<   dˆ 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)
    Ú
transformsc                 óv  •— t        d„ |D «       «      sJ ‚|r|D �cg c]  }|j                  |«      ‘Œ }}t        ‰| �  |¬«       t	        |«      | _        |€dgt        | j
                  «      z  }t	        |«      | _        t        | j                  «      t        | j
                  «      k(  sJ ‚|| _        y c c}w )Nc              3   ó<   K  — | ]  }t        |t        «      –— Œ y ­wrI   ©r†   r   ©r¡   rf  s     r/   r£   z(CatTransform.__init__.<locals>.<genexpr>!  ó   è ø€ Ò:°”:˜a¤×+Ñ:ùó   ‚rD   r%   )	r¥   rG   r*   r+   r¼   r¨  rØ   ÚlengthsrÉ   )r,   ÚtseqrÉ   r¯  r-   rf  r.   s         €r/   r+   zCatTransform.__init__   s¢   ø€ ÜÑ:°TÔ:Ô:Ð:Ð:ÙØ6:Ö;°�A—L‘L Õ,Ð;ˆDÐ;Ü‰Ñ JÐÔ/Ü˜t›*ˆŒØˆ?Ø�cœC §¡Ó0Ñ0ˆGÜ˜G“}ˆŒÜ�4—<‘<Ó ¤C¨¯©Ó$8Ò8Ð8Ð8Øˆ�ùò <s   œB6r6   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrI   )r8   r¬  s     r/   r£   z)CatTransform.event_dim.<locals>.<genexpr>.  ó   è ø€ Ò8 1�1—;•;Ñ8ùr¤   )r›   r¨  r9   s    r/   r8   zCatTransform.event_dim,  ó   € äÑ8¨¯©Ô8Ó8Ð8r0   c                 ó,   — t        | j                  «      S rI   )rI  r¯  r9   s    r/   ÚlengthzCatTransform.length0  s   € ä�4—<‘<Ó Ð r0   c                 ó|   — | j                   |k(  r| S t        | j                  | j                  | j                  |«      S rI   )r&   r   r¨  rÉ   r¯  rF   s     r/   rG   zCatTransform.with_cache4  s2   € Ø×Ñ˜zÒ)ØˆKÜ˜DŸO™O¨T¯X©X°t·|±|ÀZÓPÐPr0   c                 óÐ  — |j                  «        | j                   cxk  r|j                  «       k  sJ ‚ J ‚|j                  | j                   «      | j                  k(  sJ ‚g }d}t        | j                  | j
                  «      D ]>  \  }}|j                  | j                   ||«      }|j                   ||«      «       ||z   }Œ@ t        j                  || j                   ¬«      S ©Nr   r`  )
rÉ   rG  r¶  r¯   r¨  r¯  Únarrowr®   r¬   Úcat)r,   rR   ÚyslicesÚstartÚtransr¶  Úxslices          r/   rQ   zCatTransform._call9  s¾   € Ø—‘“ˆx˜4Ÿ8™8Ô- a§e¡e£gÒ-Ð-Ñ-Ð-Ð-Ø�v‰v�d—h‘hÓ 4§;¡;Ò.Ð.Ð.ØˆØˆÜ  §¡°$·,±,Ó?ò 	#‰MˆE�6Ø—X‘X˜dŸh™h¨¨vÓ6ˆFØ�N‰N™5 ›=Ô)Ø˜F‘N‰Eð	#ô �y‰y˜ d§h¡hÔ/Ð/r0   c                 óâ  — |j                  «        | j                   cxk  r|j                  «       k  sJ ‚ J ‚|j                  | j                   «      | j                  k(  sJ ‚g }d}t        | j                  | j
                  «      D ]G  \  }}|j                  | j                   ||«      }|j                  |j                  |«      «       ||z   }ŒI t        j                  || j                   ¬«      S r¹  )rÉ   rG  r¶  r¯   r¨  r¯  rº  r®   r>   r¬   r»  )r,   rU   Úxslicesr½  r¾  r¶  Úyslices          r/   rX   zCatTransform._inverseD  sÃ   € Ø—‘“ˆx˜4Ÿ8™8Ô- a§e¡e£gÒ-Ð-Ñ-Ð-Ð-Ø�v‰v�d—h‘hÓ 4§;¡;Ò.Ð.Ð.ØˆØˆÜ  §¡°$·,±,Ó?ò 	#‰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Ô/Ð/r0   c                 óÎ  — |j                  «        | j                   cxk  r|j                  «       k  sJ ‚ J ‚|j                  | j                   «      | j                  k(  sJ ‚|j                  «        | j                   cxk  r|j                  «       k  sJ ‚ J ‚|j                  | j                   «      | j                  k(  s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 r¹  )rÉ   rG  r¶  r¯   r¨  r¯  rº  ra   r8   r   r®   r¬   r»  rI  )r,   rR   rU   Ú
logdetjacsr½  r¾  r¶  r¿  rÂ  Ú	logdetjacrÉ   s              r/   ra   z!CatTransform.log_abs_det_jacobianO  s‘  € Ø—‘“ˆx˜4Ÿ8™8Ô- a§e¡e£gÒ-Ð-Ñ-Ð-Ð-Ø�v‰v�d—h‘hÓ 4§;¡;Ò.Ð.Ð.Ø—‘“ˆx˜4Ÿ8™8Ô- a§e¡e£gÒ-Ð-Ñ-Ð-Ð-Ø�v‰v�d—h‘hÓ 4§;¡;Ò.Ð.Ð.Øˆ
ØˆÜ  §¡°$·,±,Ó?ò 	#‰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“?Ð"r0   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrI   r    r¬  s     r/   r£   z)CatTransform.bijective.<locals>.<genexpr>j  r³  r¤   ©r¥   r¨  r9   s    r/   rq   zCatTransform.bijectiveh  r´  r0   c                 ó¦   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  | j
                  «      S c c}w rI   )r   r»  r¨  r"   rÉ   r¯  ©r,   rf  s     r/   r"   zCatTransform.domainl  s8   € ä�‰Ø#Ÿ™Ö/˜!ˆQ�X‹XÒ/°·±¸4¿<¹<ó
ð 	
ùÚ/ó   žAc                 ó¦   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  | j
                  «      S c c}w rI   )r   r»  r¨  r#   rÉ   r¯  rÊ  s     r/   r#   zCatTransform.codomainr  s8   € ä�‰Ø!%§¡Ö1˜AˆQ�Z‹ZÒ1°4·8±8¸T¿\¹\ó
ð 	
ùÚ1rË  )r   Nr   rm   )rc   rn   ro   rp   r¼   r   rs   r+   r   ru   r8   r¶  rG   rQ   rX   ra   rt   r‘   rq   r   r�   r"   r#   rv   rw   s   @r/   r   r     s³   ø… ñð �Y‘Óõ
ð ð9˜3ò 9ó ð9ð ð!˜ò !ó ð!óQò
	0ò	0ò#ð2 ð9˜4ò 9ó ð9ð ×#Ñ#ñ
ó $ð
ð
 ×#Ñ#ñ
ó $ô
r0   r   c                   ó´   ‡ — e Zd ZU dZee   ed<   dˆ 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¨  c                 óÆ   •— t        d„ |D «       «      sJ ‚|r|D �cg c]  }|j                  |«      ‘Œ }}t        ‰| �  |¬«       t	        |«      | _        || _        y c c}w )Nc              3   ó<   K  — | ]  }t        |t        «      –— Œ y ­wrI   r«  r¬  s     r/   r£   z*StackTransform.__init__.<locals>.<genexpr>‰  r­  r®  rD   )r¥   rG   r*   r+   r¼   r¨  rÉ   )r,   r°  rÉ   r-   rf  r.   s        €r/   r+   zStackTransform.__init__ˆ  s]   ø€ ÜÑ:°TÔ:Ô:Ð:Ð:ÙØ6:Ö;°�A—L‘L Õ,Ð;ˆDÐ;Ü‰Ñ JÐÔ/Ü˜t›*ˆŒØˆ�ùò <s   œAc                 óf   — | j                   |k(  r| S t        | j                  | j                  |«      S rI   )r&   r   r¨  rÉ   rF   s     r/   rG   zStackTransform.with_cache�  s,   € Ø×Ñ˜zÒ)ØˆKÜ˜dŸo™o¨t¯x©x¸ÓDÐDr0   c                 ó¤   — t        |j                  | j                  «      «      D �cg c]  }|j                  | j                  |«      ‘Œ  c}S c c}w rI   )ÚrangerG  rÉ   Úselect)r,   r]  Úis      r/   Ú_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J ‚ J ‚|j                  | j                   «      t        | j                  «      k(  sJ ‚g }t	        | j                  |«      | j                  «      D ]  \  }}|j                   ||«      «       Œ t        j                  || j                   ¬«      S ©Nr`  )	rÉ   rG  rØ   r¨  r¯   rÕ  r®   r¬   Ústack)r,   rR   r¼  r¿  r¾  s        r/   rQ   zStackTransform._call˜  s¡   € Ø—‘“ˆx˜4Ÿ8™8Ô- a§e¡e£gÒ-Ð-Ñ-Ð-Ð-Ø�v‰v�d—h‘hÓ¤3 t§¡Ó#7Ò7Ð7Ð7ØˆÜ  §¡¨Q£°·±ÓAò 	*‰MˆF�EØ�N‰N™5 ›=Õ)ð	*ä�{‰{˜7¨¯©Ô1Ð1r0   c                 ó¶  — |j                  «        | j                   cxk  r|j                  «       k  sJ ‚ J ‚|j                  | j                   «      t        | j                  «      k(  sJ ‚g }t	        | j                  |«      | j                  «      D ]%  \  }}|j                  |j                  |«      «       Œ' t        j                  || j                   ¬«      S r×  )
rÉ   rG  rØ   r¨  r¯   rÕ  r®   r>   r¬   rØ  )r,   rU   rÁ  rÂ  r¾  s        r/   rX   zStackTransform._inverse   s¦   € Ø—‘“ˆx˜4Ÿ8™8Ô- a§e¡e£gÒ-Ð-Ñ-Ð-Ð-Ø�v‰v�d—h‘hÓ¤3 t§¡Ó#7Ò7Ð7Ð7ØˆÜ  §¡¨Q£°·±ÓAò 	.‰MˆF�EØ�N‰N˜5Ÿ9™9 VÓ,Õ-ð	.ä�{‰{˜7¨¯©Ô1Ð1r0   c                 ó¶  — |j                  «        | j                   cxk  r|j                  «       k  sJ ‚ J ‚|j                  | j                   «      t        | j                  «      k(  sJ ‚|j                  «        | j                   cxk  r|j                  «       k  sJ ‚ J ‚|j                  | j                   «      t        | j                  «      k(  sJ ‚g }| j	                  |«      }| j	                  |«      }t        ||| j                  «      D ]'  \  }}}|j                  |j                  ||«      «       Œ) t        j                  || j                   ¬«      S r×  )
rÉ   rG  rØ   r¨  rÕ  r¯   r®   ra   r¬   rØ  )	r,   rR   rU   rÄ  r¼  rÁ  r¿  rÂ  r¾  s	            r/   ra   z#StackTransform.log_abs_det_jacobian¨  s  € Ø—‘“ˆx˜4Ÿ8™8Ô- a§e¡e£gÒ-Ð-Ñ-Ð-Ð-Ø�v‰v�d—h‘hÓ¤3 t§¡Ó#7Ò7Ð7Ð7Ø—‘“ˆx˜4Ÿ8™8Ô- a§e¡e£gÒ-Ð-Ñ-Ð-Ð-Ø�v‰v�d—h‘hÓ¤3 t§¡Ó#7Ò7Ð7Ð7Øˆ
Ø—+‘+˜a“.ˆØ—+‘+˜a“.ˆÜ%(¨°'¸4¿?¹?Ó%Kò 	JÑ!ˆF�F˜EØ×Ñ˜e×8Ñ8¸ÀÓHÕIð	Jä�{‰{˜:¨4¯8©8Ô4Ð4r0   r6   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrI   r    r¬  s     r/   r£   z+StackTransform.bijective.<locals>.<genexpr>¶  r³  r¤   rÈ  r9   s    r/   rq   zStackTransform.bijective´  r´  r0   c                 ó�   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  «      S c c}w rI   )r   rØ  r¨  r"   rÉ   rÊ  s     r/   r"   zStackTransform.domain¸  s/   € ä× Ñ °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 rI   )r   rØ  r¨  r#   rÉ   rÊ  s     r/   r#   zStackTransform.codomain¼  s/   € ä× Ñ °d·o±oÖ!F° !§*£*Ò!FÈÏÉÓQÐQùÒ!FrÞ  rP  rm   )rc   rn   ro   rp   r¼   r   rs   r+   rG   rÕ  rQ   rX   ra   rt   r‘   rq   r   r�   r"   r#   rv   rw   s   @r/   r   r   y  s‡   ø… ñ
ð �Y‘ÓõóEò
Hò2ò2ò
5ð ð9˜4ò 9ó ð9ð ×#Ñ#ñPó $ðPð ×#Ñ#ñRó $ôRr0   r   c                   óˆ   ‡ — e Zd ZdZdZej                  ZdZdˆ fd„	Z	e
dej                  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%   c                 ó4   •— t         ‰| �  |¬«       || _        y r{   )r*   r+   Údistribution)r,   râ  r-   r.   s      €r/   r+   z(CumulativeDistributionTransform.__init__Ø  s   ø€ Ü‰Ñ JÐÔ/Ø(ˆÕr0   r6   c                 ó.   — | j                   j                  S rI   )râ  Úsupportr9   s    r/   r"   z&CumulativeDistributionTransform.domainÜ  s   € à× Ñ ×(Ñ(Ð(r0   c                 ó8   — | j                   j                  |«      S rI   )râ  Úcdfr[   s     r/   rQ   z%CumulativeDistributionTransform._callà  s   € Ø× Ñ ×$Ñ$ QÓ'Ð'r0   c                 ó8   — | j                   j                  |«      S rI   )râ  Úicdfr^   s     r/   rX   z(CumulativeDistributionTransform._inverseã  s   € Ø× Ñ ×%Ñ% aÓ(Ð(r0   c                 ó8   — | j                   j                  |«      S rI   )râ  Úlog_probr`   s      r/   ra   z4CumulativeDistributionTransform.log_abs_det_jacobianæ  s   € Ø× Ñ ×)Ñ)¨!Ó,Ð,r0   c                 óR   — | j                   |k(  r| S t        | j                  |¬«      S r{   )r&   r   râ  rF   s     r/   rG   z*CumulativeDistributionTransform.with_cacheé  s(   € Ø×Ñ˜zÒ)ØˆKÜ.¨t×/@Ñ/@ÈZÔXÐXr0   rl   rm   )rc   rn   ro   rp   rq   r   r  r#   rB   r+   rt   rr   r"   rQ   rX   ra   rG   rv   rw   s   @r/   r   r   Á  sZ   ø„ ñð$ €IØ×(Ñ(€HØ€Dõ)ð ð)˜×.Ñ.ò )ó ð)ò(ò)ò-÷Yr0   r   ).r°   r-  r²   r<   Útypingr   r¬   Útorch.nn.functionalÚnnÚ
functionalr  Útorch.distributionsr   Ú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   rJ   r0   r/   ú<module>rô     s_  ðã Û Û Û Ý ã ß Ð Ý +÷õ ÷ .Ý ò€÷0fñ fôR;.˜	ô ;.ô|w�yô wñt & bÓ)Ð ôD8˜9ô D8ôNB+�yô B+ôJ�9ô ô.(R�Yô (RòVNô
/�yô /ô2˜	ô ô0*>�Iô *>ôZ�9ô ô `
�iô `
ôFP!˜Iô P!ôf!�yô !ôH5-˜Yô 5-ôpL˜Yô Lô,/ 	ô /ô(g
�9ô g
ôTER�Yô ERôP+Y iõ +Yr0   