Ë
    óÍ:j±¥  ã                   óŠ  — d dl Z d dlZd dlZd dlZd dlmZ d dlmZ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«      Z0 G d.„ d/e«      Z1 G d0„ d1e«      Z2 G d2„ d3e«      Z3y)4é    N)ÚSequence)ÚOptionalÚUnion)Ú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     €ús/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/torch/distributions/transforms.pyr1   zTransform.__init__b   sB   ø€ Ø%ˆÔØ@DˆŒ	Ø˜Š?ØØ˜1Š_Ø)ˆDÕäÐ8Ó9Ð9Ü‰ÑÕó    c                 óD   — | j                   j                  «       }d |d<   |S )Nr-   )Ú__dict__Úcopy)r2   Ústates     r4   Ú__getstate__zTransform.__getstate__m   s"   € Ø—‘×"Ñ"Ó$ˆØˆˆf‰Øˆr5   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/   ©r2   s    r4   r<   zTransform.event_dimr   s:   € à�;‰;× Ñ  D§M¡M×$;Ñ$;Ò;Ø—;‘;×(Ñ(Ð(ÜÐUÓVÐVr5   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)r2   Úinvs     r4   rB   zTransform.invx   sB   € ð ˆØ�9‰9Ð Ø—)‘)“+ˆCØˆ;Ü# DÓ)ˆCÜŸ™ CÓ(ˆDŒIØˆ
r5   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    r4   ÚsignzTransform.sign†   s
   € ô "Ð!r5   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,   Útyper1   r#   rE   ©r2   r(   s     r4   Ú
with_cachezTransform.with_cacheŽ   sU   € Ø×Ñ˜zÒ)ØˆKÜ�‹:×Ñ¤)×"4Ñ"4Ñ4Ø”4˜“:¨Ô4Ð4Ü!¤T¨$£Z LÐ0NÐ"OÓPÐPr5   c                 ó
   — | |u S ©N© ©r2   Úothers     r4   Ú__eq__zTransform.__eq__•   s   € Ø�uˆ}Ðr5   c                 ó&   — | j                  |«       S rM   )rQ   rO   s     r4   Ú__ne__zTransform.__ne__˜   s   € à—;‘;˜uÓ%Ð%Ð%r5   c                 ó¤   — | j                   dk(  r| j                  |«      S | j                  \  }}||u r|S | j                  |«      }||f| _        |S )z2
        Computes the transform `x => y`.
        r   )r,   Ú_callr.   )r2   ÚxÚx_oldÚy_oldÚys        r4   Ú__call__zTransform.__call__œ   sY   € ð ×Ñ˜qÒ Ø—:‘:˜a“=Ð Ø×'Ñ'‰ˆˆuØ�‰:ØˆLØ�J‰J�q‹MˆØ˜a˜4ˆÔØˆr5   c                 ó¤   — | j                   dk(  r| j                  |«      S | j                  \  }}||u r|S | j                  |«      }||f| _        |S )z1
        Inverts the transform `y => x`.
        r   )r,   Ú_inverser.   )r2   rY   rW   rX   rV   s        r4   Ú	_inv_callzTransform._inv_call©   s[   € ð ×Ñ˜qÒ Ø—=‘= Ó#Ð#Ø×'Ñ'‰ˆˆuØ�‰:ØˆLØ�M‰M˜!ÓˆØ˜a˜4ˆÔØˆr5   c                 ó   — t         ‚)zD
        Abstract method to compute forward transformation.
        rD   ©r2   rV   s     r4   rU   zTransform._call¶   ó
   € ô "Ð!r5   c                 ó   — t         ‚)zD
        Abstract method to compute inverse transformation.
        rD   ©r2   rY   s     r4   r\   zTransform._inverse¼   r`   r5   c                 ó   — t         ‚)zU
        Computes the log det jacobian `log |dy/dx|` given input and output.
        rD   ©r2   rV   rY   s      r4   Úlog_abs_det_jacobianzTransform.log_abs_det_jacobianÂ   r`   r5   c                 ó4   — | j                   j                  dz   S )Nz())r3   Ú__name__r=   s    r4   Ú__repr__zTransform.__repr__È   s   € Ø�~‰~×&Ñ&¨Ñ-Ð-r5   c                 ó   — |S )z{
        Infers the shape of the forward computation, given the input shape.
        Defaults to preserving shape.
        rN   ©r2   Úshapes     r4   Úforward_shapezTransform.forward_shapeË   ó	   € ð
 ˆr5   c                 ó   — |S )z}
        Infers the shapes of the inverse computation, given the output shape.
        Defaults to preserving shape.
        rN   rj   s     r4   Úinverse_shapezTransform.inverse_shapeÒ   rm   r5   ©r   )r)   r#   ©r+   )rg   Ú
__module__Ú__qualname__Ú__doc__Ú	bijectiver   Ú
ConstraintÚ__annotations__Úintr1   r:   Úpropertyr<   rB   rF   rK   rQ   rS   rZ   r]   rU   r\   re   rh   rl   ro   Ú__classcell__©r3   s   @r4   r#   r#   1   sÅ   ø… ñ*ðX €IØ×"Ñ"Ó"Ø×$Ñ$Ó$ñ	 3ð 	¨tõ 	òð
 ðW˜3ò Wó ðWð
 òó ðð ð"�cò "ó ð"óQòò&òòò"ò"ò"ò.òör5   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 ©NrH   )r0   r1   r,   r-   )r2   r}   r3   s     €r4   r1   z_InverseTransform.__init__à   s    ø€ Ü‰Ñ I×$9Ñ$9ÐÔ:Ø(ˆ�	r5   F©Úis_discretec                 óJ   — | j                   €J ‚| j                   j                  S rM   )r-   r'   r=   s    r4   r&   z_InverseTransform.domainä   s"   € à�y‰yÐ$Ð$Ð$Ø�y‰y×!Ñ!Ð!r5   c                 óJ   — | j                   €J ‚| j                   j                  S rM   )r-   r&   r=   s    r4   r'   z_InverseTransform.codomainé   s"   € à�y‰yÐ$Ð$Ð$Ø�y‰y×ÑÐr5   c                 óJ   — | j                   €J ‚| j                   j                  S rM   )r-   ru   r=   s    r4   ru   z_InverseTransform.bijectiveî   s"   € à�y‰yÐ$Ð$Ð$Ø�y‰y×"Ñ"Ð"r5   c                 óJ   — | j                   €J ‚| j                   j                  S rM   )r-   rF   r=   s    r4   rF   z_InverseTransform.signó   s    € à�y‰yÐ$Ð$Ð$Ø�y‰y�~‰~Ðr5   c                 ó   — | j                   S rM   )r-   r=   s    r4   rB   z_InverseTransform.invø   s   € à�y‰yÐr5   c                 óh   — | j                   €J ‚| j                  j                  |«      j                  S rM   )r-   rB   rK   rJ   s     r4   rK   z_InverseTransform.with_cacheü   s-   € Ø�y‰yÐ$Ð$Ð$Ø�x‰x×"Ñ" :Ó.×2Ñ2Ð2r5   c                 ór   — t        |t        «      sy| j                  €J ‚| j                  |j                  k(  S ©NF)Ú
isinstancer?   r-   rO   s     r4   rQ   z_InverseTransform.__eq__   s3   € Ü˜%Ô!2Ô3ØØ�y‰yÐ$Ð$Ð$Ø�y‰y˜EŸJ™JÑ&Ð&r5   c                 ó`   — | j                   j                  › dt        | j                  «      › d�S )Nú(ú))r3   rg   Úreprr-   r=   s    r4   rh   z_InverseTransform.__repr__  s)   € Ø—.‘.×)Ñ)Ð*¨!¬D°·±«OÐ+<¸AÐ>Ð>r5   c                 óT   — | j                   €J ‚| j                   j                  |«      S rM   )r-   r]   r_   s     r4   rZ   z_InverseTransform.__call__	  s'   € Ø�y‰yÐ$Ð$Ð$Ø�y‰y×"Ñ" 1Ó%Ð%r5   c                 óX   — | j                   €J ‚| j                   j                  ||«       S rM   )r-   re   rd   s      r4   re   z&_InverseTransform.log_abs_det_jacobian  s,   € Ø�y‰yÐ$Ð$Ð$Ø—	‘	×.Ñ.¨q°!Ó4Ð4Ð4r5   c                 ó8   — | j                   j                  |«      S rM   )r-   ro   rj   s     r4   rl   z_InverseTransform.forward_shape  ó   € Ø�y‰y×&Ñ& uÓ-Ð-r5   c                 ó8   — | j                   j                  |«      S rM   )r-   rl   rj   s     r4   ro   z_InverseTransform.inverse_shape  r’   r5   rq   )rg   rr   rs   rt   r#   r1   r   Údependent_propertyr&   r'   ry   Úboolru   rx   rF   rB   rK   rQ   rh   rZ   re   rl   ro   rz   r{   s   @r4   r?   r?   Ú   sÑ   ø„ ñð
) )ð )°õ )ð $€[×#Ñ#°Ô6ñ"ó 7ð"ð $€[×#Ñ#°Ô6ñ ó 7ð ð ð#˜4ò #ó ð#ð ð�cò ó ðð ð�Yò ó ðó3ò'ò?ò&ò5ò.ö.r5   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   )rK   r0   r1   r—   )r2   r—   r(   Úpartr3   s       €r4   r1   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—   rO   s     r4   rQ   zComposeTransform.__eq__)  s#   € Ü˜%Ô!1Ô2ØØ�z‰z˜UŸ[™[Ñ(Ð(r5   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'   r<   ÚreversedÚmaxÚindependent)r2   r&   r<   r™   s       r4   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Øˆr5   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&   r<   rŸ   r    )r2   r'   r<   r™   s       r4   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Øˆr5   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrM   ©ru   )Ú.0Úps     r4   ú	<genexpr>z-ComposeTransform.bijective.<locals>.<genexpr>N  s   è ø€ Ò3 1�1—;•;Ñ3ùó   ‚)Úallr—   r=   s    r4   ru   zComposeTransform.bijectiveL  s   € äÑ3¨¯
©
Ô3Ó3Ð3r5   c                 óJ   — d}| j                   D ]  }||j                  z  }Œ |S ©Nr+   )r—   rF   )r2   rF   r¦   s      r4   rF   zComposeTransform.signP  s,   € àˆØ—‘ò 	!ˆAØ˜!Ÿ&™&‘=‰Dð	!àˆr5   c                 ó$  — d }| j                   �| j                  «       }|€jt        t        | j                  «      D �cg c]  }|j                  ‘Œ c}«      }t        j                  |«      | _         t        j                  | «      |_         |S c c}w rM   )r-   r   rž   r—   rB   r@   rA   )r2   rB   r¦   s      r4   rB   zComposeTransform.invW  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—   rJ   s     r4   rK   zComposeTransform.with_cacheb  s&   € Ø×Ñ˜zÒ)ØˆKÜ §
¡
°zÔBÐBr5   c                 ó8   — | j                   D ]
  } ||«      }Œ |S rM   )r—   )r2   rV   r™   s      r4   rZ   zComposeTransform.__call__g  s#   € Ø—J‘Jò 	ˆDÙ�Q“‰Að	àˆr5   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	   re   r'   Ú	functoolsÚreduceÚoperatorÚadd)r2   rV   rY   Úxsr™   Útermsr<   s          r4   re   z%ComposeTransform.log_abs_det_jacobianl  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Ð4r5   c                 óJ   — | j                   D ]  }|j                  |«      }Œ |S rM   )r—   rl   ©r2   rk   r™   s      r4   rl   zComposeTransform.forward_shape�  s*   € Ø—J‘Jò 	.ˆDØ×&Ñ& uÓ-‰Eð	.àˆr5   c                 ó\   — t        | j                  «      D ]  }|j                  |«      }Œ |S rM   )rž   r—   ro   r»   s      r4   ro   zComposeTransform.inverse_shape†  s/   € Ü˜TŸZ™ZÓ(ò 	.ˆDØ×&Ñ& uÓ-‰Eð	.àˆr5   c                 óÀ   — | j                   j                  dz   }|dj                  | j                  D �cg c]  }|j	                  «       ‘Œ c}«      z  }|dz  }|S c c}w )Nz(
    z,
    z
))r3   rg   Újoinr—   rh   )r2   Ú
fmt_stringr¦   s      r4   rh   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
rp   rq   )rg   rr   rs   rt   Úlistr#   rx   r1   rQ   r   r”   r&   r'   r   r•   ru   rF   ry   rB   rK   rZ   re   rl   ro   rh   rz   r{   s   @r4   r   r     sÝ   ø„ ññ˜d 9™oð ¸3ð Àtõ ò)ð
 $€[×#Ñ#°Ô6ñó 7ðð $€[×#Ñ#°Ô6ñó 7ðð ð4˜4ò 4ó ð4ð ð�cò ó ðð ð�Yò ó ðóCò
ò
5ò*ò
ö
r5   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   )r0   r1   rK   rÂ   rÃ   )r2   rÂ   rÃ   r(   r3   s       €r4   r1   zIndependentTransform.__init__£  s0   ø€ ô 	‰Ñ JÐÔ/Ø,×7Ñ7¸
ÓCˆÔØ)BˆÕ&r5   c                 óh   — | j                   |k(  r| S t        | j                  | j                  |¬«      S r   )r,   r   rÂ   rÃ   rJ   s     r4   rK   zIndependentTransform.with_cache­  s5   € Ø×Ñ˜zÒ)ØˆKÜ#Ø×Ñ ×!?Ñ!?ÈJô
ð 	
r5   Fr€   c                 ój   — t        j                  | j                  j                  | j                  «      S rM   )r   r    rÂ   r&   rÃ   r=   s    r4   r&   zIndependentTransform.domain´  s,   € ä×&Ñ&Ø×Ñ×&Ñ&¨×(FÑ(Fó
ð 	
r5   c                 ój   — t        j                  | j                  j                  | j                  «      S rM   )r   r    rÂ   r'   rÃ   r=   s    r4   r'   zIndependentTransform.codomainº  s,   € ä×&Ñ&Ø×Ñ×(Ñ(¨$×*HÑ*Hó
ð 	
r5   c                 ó.   — | j                   j                  S rM   )rÂ   ru   r=   s    r4   ru   zIndependentTransform.bijectiveÀ  s   € à×"Ñ"×,Ñ,Ð,r5   c                 ó.   — | j                   j                  S rM   )rÂ   rF   r=   s    r4   rF   zIndependentTransform.signÄ  s   € à×"Ñ"×'Ñ'Ð'r5   c                 óˆ   — |j                  «       | j                  j                  k  rt        d«      ‚| j	                  |«      S ©NúToo few dimensions on input)Údimr&   r<   r/   rÂ   r_   s     r4   rU   zIndependentTransform._callÈ  s7   € Ø�5‰5‹7�T—[‘[×*Ñ*Ò*ÜÐ:Ó;Ð;Ø×"Ñ" 1Ó%Ð%r5   c                 óœ   — |j                  «       | j                  j                  k  rt        d«      ‚| j                  j                  |«      S rË   )rÍ   r'   r<   r/   rÂ   rB   rb   s     r4   r\   zIndependentTransform._inverseÍ  s=   € Ø�5‰5‹7�T—]‘]×,Ñ,Ò,ÜÐ:Ó;Ð;Ø×"Ñ"×&Ñ& qÓ)Ð)r5   c                 ój   — | j                   j                  ||«      }t        || j                  «      }|S rM   )rÂ   re   r	   rÃ   )r2   rV   rY   Úresults       r4   re   z)IndependentTransform.log_abs_det_jacobianÒ  s1   € Ø×$Ñ$×9Ñ9¸!¸QÓ?ˆÜ ¨×(FÑ(FÓGˆØˆr5   c                 óz   — | j                   j                  › dt        | j                  «      › d| j                  › d�S )NrŒ   z, r�   )r3   rg   rŽ   rÂ   rÃ   r=   s    r4   rh   zIndependentTransform.__repr__×  s:   € Ø—.‘.×)Ñ)Ð*¨!¬D°×1DÑ1DÓ,EÐ+FÀbÈ×IgÑIgÐHhÐhiÐjÐjr5   c                 ó8   — | j                   j                  |«      S rM   )rÂ   rl   rj   s     r4   rl   z"IndependentTransform.forward_shapeÚ  ó   € Ø×"Ñ"×0Ñ0°Ó7Ð7r5   c                 ó8   — | j                   j                  |«      S rM   )rÂ   ro   rj   s     r4   ro   z"IndependentTransform.inverse_shapeÝ  rÓ   r5   rp   rq   )rg   rr   rs   rt   r#   rx   r1   rK   r   r”   r&   r'   ry   r•   ru   rF   rU   r\   re   rh   rl   ro   rz   r{   s   @r4   r   r   •  sÙ   ø„ ñð" ñ	Cà!ðCð $'ðCð ð	Cð
 
õCó
ð $€[×#Ñ#°Ô6ñ
ó 7ð
ð
 $€[×#Ñ#°Ô6ñ
ó 7ð
ð
 ð-˜4ò -ó ð-ð ð(�cò (ó ð(ò&ò
*ò
ò
kò8ö8r5   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 elementsrH   )r°   ÚSizerÖ   r×   Únumelr/   r0   r1   )r2   rÖ   r×   r(   r3   s       €r4   r1   zReshapeTransform.__init__ñ  sc   ø€ ô Ÿ
™
 8Ó,ˆŒÜŸ™ IÓ.ˆŒØ�=‰=×ÑÓ  D§N¡N×$8Ñ$8Ó$:Ò:ÜÐUÓVÐVÜ‰Ñ JÐÕ/r5   c                 óp   — t        j                  t         j                  t        | j                  «      «      S rM   )r   r    r�   ÚlenrÖ   r=   s    r4   r&   zReshapeTransform.domainý  s$   € ä×&Ñ&¤{×'7Ñ'7¼¸T¿]¹]Ó9KÓLÐLr5   c                 óp   — t        j                  t         j                  t        | j                  «      «      S rM   )r   r    r�   rÜ   r×   r=   s    r4   r'   zReshapeTransform.codomain  s$   € ä×&Ñ&¤{×'7Ñ'7¼¸T¿^¹^Ó9LÓMÐMr5   c                 óh   — | j                   |k(  r| S t        | j                  | j                  |¬«      S r   )r,   r   rÖ   r×   rJ   s     r4   rK   zReshapeTransform.with_cache  s,   € Ø×Ñ˜zÒ)ØˆKÜ §¡¨t¯~©~È*ÔUÐUr5   c                 ó¤   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  || j
                  z   «      S rM   )rk   rÍ   rÜ   rÖ   Úreshaper×   )r2   rV   Úbatch_shapes      r4   rU   zReshapeTransform._call
  s?   € Ø—g‘gÐ< §¡£¬#¨d¯m©mÓ*<Ñ <Ð=ˆØ�y‰y˜ t§~¡~Ñ5Ó6Ð6r5   c                 ó¤   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  || j
                  z   «      S rM   )rk   rÍ   rÜ   r×   rà   rÖ   )r2   rY   rá   s      r4   r\   zReshapeTransform._inverse  s?   € Ø—g‘gÐ= §¡£¬#¨d¯n©nÓ*=Ñ =Ð>ˆØ�y‰y˜ t§}¡}Ñ4Ó5Ð5r5   c                 óŠ   — |j                   d |j                  «       t        | j                  «      z
   }|j	                  |«      S rM   )rk   rÍ   rÜ   rÖ   Ú	new_zeros)r2   rV   rY   rá   s       r4   re   z%ReshapeTransform.log_abs_det_jacobian  s6   € Ø—g‘gÐ< §¡£¬#¨d¯m©mÓ*<Ñ <Ð=ˆØ�{‰{˜;Ó'Ð'r5   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×   ©r2   rk   Úcuts      r4   rl   zReshapeTransform.forward_shape  sŠ   € Üˆu‹:œ˜DŸM™MÓ*Ò*ÜÐ:Ó;Ð;Ü�%‹jœ3˜tŸ}™}Ó-Ñ-ˆØ��ˆ;˜$Ÿ-™-Ò'ÜØ+¨E°#°$¨K¨=¸	À$Ç-Á-ÀÐQóð ð �T�cˆ{˜TŸ^™^Ñ+Ð+r5   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      r4   ro   zReshapeTransform.inverse_shape   s‹   € Üˆu‹:œ˜DŸN™NÓ+Ò+ÜÐ:Ó;Ð;Ü�%‹jœ3˜tŸ~™~Ó.Ñ.ˆØ��ˆ;˜$Ÿ.™.Ò(ÜØ+¨E°#°$¨K¨=¸	À$Ç.Á.ÐAQÐRóð ð �T�cˆ{˜TŸ]™]Ñ*Ð*r5   rp   rq   )rg   rr   rs   rt   ru   r°   rÙ   rx   r1   r   r”   r&   r'   rK   rU   r\   re   rl   ro   rz   r{   s   @r4   r   r   á  sž   ø„ ñð €Ið ñ	
0à—*‘*ð
0ð —:‘:ð
0ð ð	
0ð
 
õ
0ð ×#Ñ#ñMó $ðMð ×#Ñ#ñNó $ðNóVò
7ò6ò(ò,ö+r5   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 rM   )rŠ   r   rO   s     r4   rQ   zExpTransform.__eq__5  ó   € Ü˜%¤Ó.Ð.r5   c                 ó"   — |j                  «       S rM   )Úexpr_   s     r4   rU   zExpTransform._call8  ó   € Ø�u‰u‹wˆr5   c                 ó"   — |j                  «       S rM   ©Úlogrb   s     r4   r\   zExpTransform._inverse;  rï   r5   c                 ó   — |S rM   rN   rd   s      r4   re   z!ExpTransform.log_abs_det_jacobian>  ó   € Øˆr5   N©rg   rr   rs   rt   r   r�   r&   Úpositiver'   ru   rF   rQ   rU   r\   re   rN   r5   r4   r   r   +  s=   „ ñð ×Ñ€FØ×#Ñ#€HØ€IØ€Dò/òòór5   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   )r0   r1   r
   rø   )r2   rø   r(   r3   s      €r4   r1   zPowerTransform.__init__K  s"   ø€ Ü‰Ñ JÐÔ/Ü(¨Ó2Ñˆ�r5   c                 óR   — | j                   |k(  r| S t        | j                  |¬«      S r   )r,   r   rø   rJ   s     r4   rK   zPowerTransform.with_cacheO  s&   € Ø×Ñ˜zÒ)ØˆKÜ˜dŸm™m¸
ÔCÐCr5   c                 ó6   — | j                   j                  «       S rM   )rø   rF   r=   s    r4   rF   zPowerTransform.signT  s   € à�}‰}×!Ñ!Ó#Ð#r5   c                 ó¦   — t        |t        «      sy| j                  j                  |j                  «      j	                  «       j                  «       S r‰   )rŠ   r   rø   Úeqr©   ÚitemrO   s     r4   rQ   zPowerTransform.__eq__X  s:   € Ü˜%¤Ô0ØØ�}‰}×Ñ §¡Ó/×3Ñ3Ó5×:Ñ:Ó<Ð<r5   c                 ó8   — |j                  | j                  «      S rM   ©Úpowrø   r_   s     r4   rU   zPowerTransform._call]  s   € Ø�u‰u�T—]‘]Ó#Ð#r5   c                 ó>   — |j                  d| j                  z  «      S r«   r   rb   s     r4   r\   zPowerTransform._inverse`  s   € Ø�u‰u�Q˜Ÿ™Ñ&Ó'Ð'r5   c                 ó^   — | j                   |z  |z  j                  «       j                  «       S rM   )rø   Úabsrò   rd   s      r4   re   z#PowerTransform.log_abs_det_jacobianc  s(   € Ø—‘ Ñ! AÑ%×*Ñ*Ó,×0Ñ0Ó2Ð2r5   c                 óX   — t        j                  |t        | j                  dd«      «      S ©Nrk   rN   ©r°   Úbroadcast_shapesÚgetattrrø   rj   s     r4   rl   zPowerTransform.forward_shapef  ó"   € Ü×%Ñ% e¬W°T·]±]ÀGÈRÓ-PÓQÐQr5   c                 óX   — t        j                  |t        | j                  dd«      «      S r  r  rj   s     r4   ro   zPowerTransform.inverse_shapei  r
  r5   rp   rq   )rg   rr   rs   rt   r   rö   r&   r'   ru   r   rx   r1   rK   r   rF   rQ   rU   r\   re   rl   ro   rz   r{   s   @r4   r   r   B  s€   ø„ ñð ×!Ñ!€FØ×#Ñ#€HØ€Iñ3 ð 3°Sð 3Àõ 3óDð
 ð$�cò $ó ð$ò=ò
$ò(ò3òRöRr5   r   c                 óÄ   — t        j                  | j                  «      }t        j                  t        j                  | «      |j
                  d|j                  z
  ¬«      S ©Nç      ð?©ÚminrŸ   )r°   ÚfinfoÚdtypeÚclampÚsigmoidÚtinyÚeps)rV   r  s     r4   Ú_clipped_sigmoidr  m  s<   € Ü�K‰K˜Ÿ™Ó €EÜ�;‰;”u—}‘} QÓ'¨U¯Z©Z¸SÀ5Ç9Á9¹_ÔMÐMr5   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 rM   )rŠ   r   rO   s     r4   rQ   zSigmoidTransform.__eq__|  ó   € Ü˜%Ô!1Ó2Ð2r5   c                 ó   — t        |«      S rM   )r  r_   s     r4   rU   zSigmoidTransform._call  s   € Ü Ó"Ð"r5   c                 óØ   — t        j                  |j                  «      }|j                  |j                  d|j
                  z
  ¬«      }|j                  «       | j                  «       z
  S r  )r°   r  r  r  r  r  rò   Úlog1p)r2   rY   r  s      r4   r\   zSigmoidTransform._inverse‚  sK   € Ü—‘˜AŸG™GÓ$ˆØ�G‰G˜Ÿ
™
¨¨e¯i©i©ˆGÓ8ˆØ�u‰u‹w˜1˜"Ÿ™›Ñ%Ð%r5   c                 ó\   — t        j                  | «       t        j                  |«      z
  S rM   )ÚFr   rd   s      r4   re   z%SigmoidTransform.log_abs_det_jacobian‡  s!   € Ü—
‘
˜A˜2“ˆ¤§¡¨A£Ñ.Ð.r5   N)rg   rr   rs   rt   r   r�   r&   Úunit_intervalr'   ru   rF   rQ   rU   r\   re   rN   r5   r4   r   r   r  s=   „ ñð ×Ñ€FØ×(Ñ(€HØ€IØ€Dò3ò#ò&ó
/r5   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 rM   )rŠ   r   rO   s     r4   rQ   zSoftplusTransform.__eq__–  s   € Ü˜%Ô!2Ó3Ð3r5   c                 ó   — t        |«      S rM   ©r   r_   s     r4   rU   zSoftplusTransform._call™  s   € Ü˜‹{Ðr5   c                 ób   — | j                  «       j                  «       j                  «       |z   S rM   )Úexpm1Únegrò   rb   s     r4   r\   zSoftplusTransform._inverseœ  s'   € Ø��z‰z‹|×ÑÓ!×%Ñ%Ó'¨!Ñ+Ð+r5   c                 ó   — t        | «       S rM   r$  rd   s      r4   re   z&SoftplusTransform.log_abs_det_jacobianŸ  s   € Ü˜!˜“ˆ}Ðr5   Nrõ   rN   r5   r4   r   r   ‹  s=   „ ñð
 ×Ñ€FØ×#Ñ#€HØ€IØ€Dò4òò,ór5   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 rM   )rŠ   r   rO   s     r4   rQ   zTanhTransform.__eq__¿  s   € Ü˜%¤Ó/Ð/r5   c                 ó"   — |j                  «       S rM   )Útanhr_   s     r4   rU   zTanhTransform._callÂ  s   € Ø�v‰v‹xˆr5   c                 ó,   — t        j                  |«      S rM   )r°   Úatanhrb   s     r4   r\   zTanhTransform._inverseÅ  s   € ô �{‰{˜1‹~Ðr5   c                 óV   — dt        j                  d«      |z
  t        d|z  «      z
  z  S )Nç       @g       À)Úmathrò   r   rd   s      r4   re   z"TanhTransform.log_abs_det_jacobianÊ  s*   € ð ”d—h‘h˜s“m aÑ'¬(°4¸!±8Ó*<Ñ<Ñ=Ð=r5   N)rg   rr   rs   rt   r   r�   r&   Úintervalr'   ru   rF   rQ   rU   r\   re   rN   r5   r4   r   r   £  sF   „ ñð, ×Ñ€FØ#ˆ{×#Ñ# D¨#Ó.€HØ€IØ€Dò0òòó
>r5   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 rM   )rŠ   r   rO   s     r4   rQ   zAbsTransform.__eq__Ö  rì   r5   c                 ó"   — |j                  «       S rM   )r  r_   s     r4   rU   zAbsTransform._callÙ  rï   r5   c                 ó   — |S rM   rN   rb   s     r4   r\   zAbsTransform._inverseÜ  rô   r5   N)rg   rr   rs   rt   r   r�   r&   rö   r'   rQ   rU   r\   rN   r5   r4   r   r   Ð  s*   „ Ù5à×Ñ€FØ×#Ñ#€Hò/òór5   r   c                   ó  ‡ — e Zd ZdZdZ	 	 ddeeef   deeef   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ef   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   )r0   r1   r8  r9  Ú
_event_dim)r2   r8  r9  r<   r(   r3   s        €r4   r1   zAffineTransform.__init__î  s*   ø€ ô 	‰Ñ JÐÔ/ØˆŒØˆŒ
Ø#ˆ�r5   c                 ó   — | j                   S rM   )r;  r=   s    r4   r<   zAffineTransform.event_dimú  s   € à�‰Ðr5   Fr€   c                 óœ   — | j                   dk(  rt        j                  S t        j                  t        j                  | j                   «      S ©Nr   ©r<   r   r�   r    r=   s    r4   r&   zAffineTransform.domainþ  ó7   € à�>‰>˜QÒÜ×#Ñ#Ð#Ü×&Ñ&¤{×'7Ñ'7¸¿¹ÓHÐHr5   c                 óœ   — | j                   dk(  rt        j                  S t        j                  t        j                  | j                   «      S r>  r?  r=   s    r4   r'   zAffineTransform.codomain  r@  r5   c                 ó~   — | j                   |k(  r| S t        | j                  | j                  | j                  |¬«      S r   )r,   r   r8  r9  r<   rJ   s     r4   rK   zAffineTransform.with_cache
  s7   € Ø×Ñ˜zÒ)ØˆKÜØ�H‰H�d—j‘j $§.¡.¸Zô
ð 	
r5   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   r8  r   r©   rþ   r9  rO   s     r4   rQ   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Ô;Øàr5   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Š   r9  r   ÚfloatrF   r=   s    r4   rF   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�‰Ó Ð r5   c                 ó:   — | j                   | j                  |z  z   S rM   ©r8  r9  r_   s     r4   rU   zAffineTransform._call+  s   € Ø�x‰x˜$Ÿ*™* q™.Ñ(Ð(r5   c                 ó:   — || j                   z
  | j                  z  S rM   rG  rb   s     r4   r\   zAffineTransform._inverse.  s   € Ø�D—H‘H‘ §
¡
Ñ*Ð*r5   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œ   )rk   r9  rŠ   r   r°   Ú	full_liker1  rò   r  r<   ÚsizeÚviewÚsumÚexpand)r2   rV   rY   rk   r9  rÐ   Úresult_sizes          r4   re   z$AffineTransform.log_abs_det_jacobian1  s²   € Ø—‘ˆØ—
‘
ˆÜ�eœWÔ%Ü—_‘_ Q¬¯©´°U³Ó(<Ó=‰Fä—Y‘Y˜uÓ%×)Ñ)Ó+ˆFØ�>Š>Ø Ÿ+™+›-Ð(9¨4¯>©>¨/Ð:¸UÑBˆKØ—[‘[ Ó-×1Ñ1°"Ó5ˆFØÐ+˜TŸ^™^˜OÐ,ˆEØ�}‰}˜UÓ#Ð#r5   c           	      ó„   — t        j                  |t        | j                  dd«      t        | j                  dd«      «      S r  ©r°   r  r	  r8  r9  rj   s     r4   rl   zAffineTransform.forward_shape>  ó7   € Ü×%Ñ%Ø”7˜4Ÿ8™8 W¨bÓ1´7¸4¿:¹:ÀwÐPRÓ3Só
ð 	
r5   c           	      ó„   — t        j                  |t        | j                  dd«      t        | j                  dd«      «      S r  rQ  rj   s     r4   ro   zAffineTransform.inverse_shapeC  rR  r5   ©r   r   rq   )rg   rr   rs   rt   ru   r   r   rE  rx   r1   ry   r<   r   r”   r&   r'   rK   rQ   rF   rU   r\   re   rl   ro   rz   r{   s   @r4   r   r   à  s  ø„ ñ	ð €Ið Øñ
$à�6˜5�=Ñ!ð
$ð �V˜U�]Ñ#ð
$ð ð	
$ð
 ð
$ð 
õ
$ð ð˜3ò ó ðð $€[×#Ñ#°Ô6ñIó 7ðIð
 $€[×#Ñ#°Ô6ñIó 7ðIó

òð( ð!�e˜F C˜KÑ(ò !ó ð!ò
)ò+ò$ò
ö

r5   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Úeyerk   rZ  r   )r2   rV   r  ÚrÚzÚz1m_cumprod_sqrtrY   s          r4   rU   zCorrCholeskyTransform._call^  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Ô@Ñ@ˆØˆr5   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[  rW  rY  )r°   Úcumsumr   r   r]  r  r'  )r2   rY   Úy_cumsumÚy_cumsum_shiftedÚy_vecÚy_cumsum_vecÚtrV   s           r4   r\   zCorrCholeskyTransform._inversem  s�   € ð ”u—|‘| A¨¡E¨rÔ2Ñ2ˆÜ˜x¨¨S¨b¨S¨Ñ1°A°q°6ÀÔCÐÜ" 1¨2Ô.ˆÜ)Ð*:ÀÔDˆØ�\×'Ñ'Ó)Ñ)ˆà�W‰W‹Y˜Ÿ™›Ÿ™›Ñ(¨AÑ-ˆØˆr5   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œ   rd  éþÿÿÿrW  ç      à?r0  )re  r   rò   rM  r   r1  )r2   rV   rY   ÚintermediatesÚ
y1m_cumsumÚy1m_cumsum_trilÚstick_breaking_logdetÚtanh_logdets           r4   re   z*CorrCholeskyTransform.log_abs_det_jacobiany  sˆ   € ð ˜!˜a™%Ÿ™¨B˜Ó/Ñ/ˆ
ô -¨Z¸bÔAˆØ # ×&;Ñ&;Ó&=×&AÑ&AÀ"Ó&EÑ EÐØ˜A¤¨¨a©Ó 0Ñ0´4·8±8¸C³=Ñ@×EÑEÈ"ÐEÓMÑMˆØ$ {Ñ2Ð2r5   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      Ð?rY  rm  z-Input is not a flattend lower-diagonal number)rÜ   r/   Úround)r2   rk   ÚNÚDs       r4   rl   z#CorrCholeskyTransform.forward_shape‡  st   € äˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�"‰IˆÜ�4˜!˜a™%‘< CÑ'¨#Ñ-Ó.ˆØ��A‘‰;˜!Ñ˜qÒ ÜÐLÓMÐMØ�S�bˆz˜Q ˜FÑ"Ð"r5   c                 ó’   — t        |«      dk  rt        d«      ‚|d   |d   k7  rt        d«      ‚|d   }||dz
  z  dz  }|d d |fz   S )NrY  rÌ   rl  rœ   zInput is not squarer+   ©rÜ   r/   )r2   rk   rv  ru  s       r4   ro   z#CorrCholeskyTransform.inverse_shape‘  sc   € äˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�‰9˜˜b™	Ò!ÜÐ2Ó3Ð3Ø�"‰IˆØ��Q‘‰K˜1ÑˆØ�S�bˆz˜Q˜DÑ Ð r5   rM   )rg   rr   rs   rt   r   Úreal_vectorr&   Úcorr_choleskyr'   ru   rU   r\   re   rl   ro   rN   r5   r4   r   r   I  s=   „ ñð  ×$Ñ$€FØ×(Ñ(€HØ€Iòò
ó3ò#ó!r5   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 rM   )rŠ   r    rO   s     r4   rQ   zSoftmaxTransform.__eq__©  r  r5   c                 ó|   — |}||j                  dd«      d   z
  j                  «       }||j                  dd«      z  S )Nrœ   Tr   )rŸ   rî   rM  )r2   rV   ÚlogprobsÚprobss       r4   rU   zSoftmaxTransform._call¬  s@   € ØˆØ˜HŸL™L¨¨TÓ2°1Ñ5Ñ5×:Ñ:Ó<ˆØ�u—y‘y  TÓ*Ñ*Ð*r5   c                 ó&   — |}|j                  «       S rM   rñ   )r2   rY   r  s      r4   r\   zSoftmaxTransform._inverse±  s   € ØˆØ�y‰y‹{Ðr5   c                 ó8   — t        |«      dk  rt        d«      ‚|S ©Nr+   rÌ   rx  rj   s     r4   rl   zSoftmaxTransform.forward_shapeµ  ó   € Üˆu‹:˜Š>ÜÐ:Ó;Ð;Øˆr5   c                 ó8   — t        |«      dk  rt        d«      ‚|S r‚  rx  rj   s     r4   ro   zSoftmaxTransform.inverse_shapeº  rƒ  r5   N)rg   rr   rs   rt   r   ry  r&   Úsimplexr'   rQ   rU   r\   rl   ro   rN   r5   r4   r    r    œ  s8   „ ñð ×$Ñ$€FØ×"Ñ"€Hò3ò+ò
òó
r5   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 rM   )rŠ   r"   rO   s     r4   rQ   zStickBreakingTransform.__eq__Ò  ó   € Ü˜%Ô!7Ó8Ð8r5   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[  )rk   Únew_onesre  r  rò   r^  r   )r2   rV   Úoffsetra  Ú	z_cumprodrY   s         r4   rU   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ˆØˆr5   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  )	rk   rŠ  re  r°   r  r  r  r  rò   )r2   rY   Úy_cropr‹  ÚsfrV   s         r4   r\   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ˆØˆr5   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+   .)rk   rŠ  re  rò   r  Ú
logsigmoidrM  )r2   rV   rY   r‹  ÚdetJs        r4   re   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ˆØˆr5   c                 óR   — t        |«      dk  rt        d«      ‚|d d |d   dz   fz   S ©Nr+   rÌ   rœ   rx  rj   s     r4   rl   z$StickBreakingTransform.forward_shapeí  ó5   € Üˆu‹:˜Š>ÜÐ:Ó;Ð;Ø�S�bˆz˜U 2™Y¨™]Ð,Ñ,Ð,r5   c                 óR   — t        |«      dk  rt        d«      ‚|d d |d   dz
  fz   S r”  rx  rj   s     r4   ro   z$StickBreakingTransform.inverse_shapeò  r•  r5   N)rg   rr   rs   rt   r   ry  r&   r…  r'   ru   rQ   rU   r\   re   rl   ro   rN   r5   r4   r"   r"   À  sB   „ ñð ×$Ñ$€FØ×"Ñ"€HØ€Iò9òòòò-ó
-r5   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.
    rY  c                 ó"   — t        |t        «      S rM   )rŠ   r   rO   s     r4   rQ   zLowerCholeskyTransform.__eq__  rˆ  r5   c                 ó„   — |j                  d«      |j                  dd¬«      j                  «       j                  «       z   S ©Nrœ   rl  )Údim1Údim2)ÚtrilÚdiagonalrî   Ú
diag_embedr_   s     r4   rU   zLowerCholeskyTransform._call  ó4   € Ø�v‰v�b‹z˜AŸJ™J¨B°R˜JÓ8×<Ñ<Ó>×IÑIÓKÑKÐKr5   c                 ó„   — |j                  d«      |j                  dd¬«      j                  «       j                  «       z   S rš  )r�  rž  rò   rŸ  rb   s     r4   r\   zLowerCholeskyTransform._inverse
  r   r5   N)rg   rr   rs   rt   r   r    r�   r&   Úlower_choleskyr'   rQ   rU   r\   rN   r5   r4   r   r   ø  s?   „ ñð %ˆ[×$Ñ$ [×%5Ñ%5°qÓ9€FØ×)Ñ)€Hò9òLóLr5   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.
    rY  c                 ó"   — t        |t        «      S rM   )rŠ   r   rO   s     r4   rQ   z PositiveDefiniteTransform.__eq__  s   € Ü˜%Ô!:Ó;Ð;r5   c                 ó@   —  t        «       |«      }||j                  z  S rM   )r   ÚmTr_   s     r4   rU   zPositiveDefiniteTransform._call  s   € Ø$Ô"Ó$ QÓ'ˆØ�1—4‘4‰xˆr5   c                 ór   — t         j                  j                  |«      }t        «       j	                  |«      S rM   )r°   ÚlinalgÚcholeskyr   rB   rb   s     r4   r\   z"PositiveDefiniteTransform._inverse  s*   € Ü�L‰L×!Ñ! !Ó$ˆÜ%Ó'×+Ñ+¨AÓ.Ð.r5   N)rg   rr   rs   rt   r   r    r�   r&   Úpositive_definiter'   rQ   rU   r\   rN   r5   r4   r   r     s=   „ ñð %ˆ[×$Ñ$ [×%5Ñ%5°qÓ9€FØ×,Ñ,€Hò<òó/r5   r   c                   ó  ‡ — e Zd ZU dZee   ed<   	 	 	 ddee   dede	ee      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                 ó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rM   ©rŠ   r#   ©r¥   rj  s     r4   r§   z(CatTransform.__init__.<locals>.<genexpr>:  ó   è ø€ Ò:°”:˜a¤×+Ñ:ùó   ‚rH   r+   )	r©   rK   r0   r1   rÀ   r¬  rÜ   r®  rÍ   )r2   r­  rÍ   r®  r(   rj  r3   s         €r4   r1   zCatTransform.__init__3  s¤   ø€ ô Ñ:°TÔ:Ô:Ð:Ð:ÙØ6:Ö;°�A—L‘L Õ,Ð;ˆDÐ;Ü‰Ñ JÐÔ/Ü˜t›*ˆŒØˆ?Ø�cœC §¡Ó0Ñ0ˆGÜ˜G“}ˆŒÜ�4—<‘<Ó ¤C¨¯©Ó$8Ò8Ð8Ð8Øˆ�ùò <s   œB6c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrM   )r<   r²  s     r4   r§   z)CatTransform.event_dim.<locals>.<genexpr>G  ó   è ø€ Ò8 1�1—;•;Ñ8ùr¨   )rŸ   r¬  r=   s    r4   r<   zCatTransform.event_dimE  ó   € äÑ8¨¯©Ô8Ó8Ð8r5   c                 ó,   — t        | j                  «      S rM   )rM  r®  r=   s    r4   ÚlengthzCatTransform.lengthI  s   € ä�4—<‘<Ó Ð r5   c                 ó|   — | j                   |k(  r| S t        | j                  | j                  | j                  |«      S rM   )r,   r   r¬  rÍ   r®  rJ   s     r4   rK   zCatTransform.with_cacheM  s2   € Ø×Ñ˜zÒ)ØˆKÜ˜DŸO™O¨T¯X©X°t·|±|ÀZÓPÐPr5   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   rd  )
rÍ   rK  rº  r³   r¬  r®  Únarrowr²   r°   Úcat)r2   rV   ÚyslicesÚstartÚtransrº  Úxslices          r4   rU   zCatTransform._callR  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Ô/Ð/r5   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Í   rK  rº  r³   r¬  r®  r¾  r²   rB   r°   r¿  )r2   rY   ÚxslicesrÁ  rÂ  rº  Úyslices          r4   r\   zCatTransform._inverse]  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Ô/Ð/r5   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Í   rK  rº  r³   r¬  r®  r¾  re   r<   r	   r²   r°   r¿  rM  )r2   rV   rY   Ú
logdetjacsrÁ  rÂ  rº  rÃ  rÆ  Ú	logdetjacrÍ   s              r4   re   z!CatTransform.log_abs_det_jacobianh  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“?Ð"r5   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrM   r¤   r²  s     r4   r§   z)CatTransform.bijective.<locals>.<genexpr>ƒ  r·  r¨   ©r©   r¬  r=   s    r4   ru   zCatTransform.bijective�  r¸  r5   c                 ó¦   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  | j
                  «      S c c}w rM   )r   r¿  r¬  r&   rÍ   r®  ©r2   rj  s     r4   r&   zCatTransform.domain…  s8   € ä�‰Ø#Ÿ™Ö/˜!ˆQ�X‹XÒ/°·±¸4¿<¹<ó
ð 	
ùÚ/ó   žAc                 ó¦   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  | j
                  «      S c c}w rM   )r   r¿  r¬  r'   rÍ   r®  rÎ  s     r4   r'   zCatTransform.codomain‹  s8   € ä�‰Ø!%§¡Ö1˜AˆQ�Z‹ZÒ1°4·8±8¸T¿\¹\ó
ð 	
ùÚ1rÏ  )r   Nr   rq   )rg   rr   rs   rt   rÀ   r#   rw   r   rx   r   r1   r   r<   rº  rK   rU   r\   re   ry   r•   ru   r   r”   r&   r'   rz   r{   s   @r4   r   r   "  sü   ø… ñð �Y‘Óð
 Ø+/Øñà�yÑ!ðð ðð ˜( 3™-Ñ(ð	ð
 ðð 
õð$ ð9˜3ò 9ó ð9ð ð!˜ò !ó ð!óQò
	0ò	0ò#ð2 ð9˜4ò 9ó ð9ð ×#Ñ#ñ
ó $ð
ð
 ×#Ñ#ñ
ó $ô
r5   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J ‚|r|D �cg c]  }|j                  |«      ‘Œ }}t        ‰| �  |¬«       t	        |«      | _        || _        y c c}w )Nc              3   ó<   K  — | ]  }t        |t        «      –— Œ y ­wrM   r±  r²  s     r4   r§   z*StackTransform.__init__.<locals>.<genexpr>¤  r³  r´  rH   )r©   rK   r0   r1   rÀ   r¬  rÍ   )r2   r­  rÍ   r(   rj  r3   s        €r4   r1   zStackTransform.__init__¡  s_   ø€ ô Ñ:°TÔ:Ô:Ð:Ð:ÙØ6:Ö;°�A—L‘L Õ,Ð;ˆDÐ;Ü‰Ñ JÐÔ/Ü˜t›*ˆŒØˆ�ùò <s   œAc                 óf   — | j                   |k(  r| S t        | j                  | j                  |«      S rM   )r,   r!   r¬  rÍ   rJ   s     r4   rK   zStackTransform.with_cache«  s,   € Ø×Ñ˜zÒ)ØˆKÜ˜dŸo™o¨t¯x©x¸ÓDÐDr5   c                 ó¤   — t        |j                  | j                  «      «      D �cg c]  }|j                  | j                  |«      ‘Œ  c}S c c}w rM   )ÚrangerK  rÍ   Úselect)r2   ra  Úis      r4   Ú_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 ©Nrd  )	rÍ   rK  rÜ   r¬  r³   rÙ  r²   r°   Ústack)r2   rV   rÀ  rÃ  rÂ  s        r4   rU   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Ð1r5   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Í   rK  rÜ   r¬  r³   rÙ  r²   rB   r°   rÜ  )r2   rY   rÅ  rÆ  rÂ  s        r4   r\   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Ð1r5   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Í   rK  rÜ   r¬  rÙ  r³   r²   re   r°   rÜ  )	r2   rV   rY   rÈ  rÀ  rÅ  rÃ  rÆ  rÂ  s	            r4   re   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Ð4r5   c                 ó:   — t        d„ | j                  D «       «      S )Nc              3   ó4   K  — | ]  }|j                   –— Œ y ­wrM   r¤   r²  s     r4   r§   z+StackTransform.bijective.<locals>.<genexpr>Ñ  r·  r¨   rÌ  r=   s    r4   ru   zStackTransform.bijectiveÏ  r¸  r5   c                 ó�   — t        j                  | j                  D �cg c]  }|j                  ‘Œ c}| j                  «      S c c}w rM   )r   rÜ  r¬  r&   rÍ   rÎ  s     r4   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 rM   )r   rÜ  r¬  r'   rÍ   rÎ  s     r4   r'   zStackTransform.codomain×  s/   € ä× Ñ °d·o±oÖ!F° !§*£*Ò!FÈÏÉÓQÐQùÒ!Frâ  rT  rq   )rg   rr   rs   rt   rÀ   r#   rw   r   rx   r1   rK   rÙ  rU   r\   re   ry   r•   ru   r   r”   r&   r'   rz   r{   s   @r4   r!   r!   ’  s³   ø… ñ
ð �Y‘Óð JKñØ˜YÑ'ðØ.1ðØCFðà	õóEò
Hò2ò2ò
5ð ð9˜4ò 9ó ð9ð ×#Ñ#ñPó $ðPð ×#Ñ#ñRó $ôRr5   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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+   Údistributionr(   r)   Nc                 ó4   •— t         ‰| �  |¬«       || _        y r   )r0   r1   rå  )r2   rå  r(   r3   s      €r4   r1   z(CumulativeDistributionTransform.__init__ó  s   ø€ Ü‰Ñ JÐÔ/Ø(ˆÕr5   c                 ó.   — | j                   j                  S rM   )rå  Úsupportr=   s    r4   r&   z&CumulativeDistributionTransform.domain÷  s   € à× Ñ ×(Ñ(Ð(r5   c                 ó8   — | j                   j                  |«      S rM   )rå  Úcdfr_   s     r4   rU   z%CumulativeDistributionTransform._callû  s   € Ø× Ñ ×$Ñ$ QÓ'Ð'r5   c                 ó8   — | j                   j                  |«      S rM   )rå  Úicdfrb   s     r4   r\   z(CumulativeDistributionTransform._inverseþ  s   € Ø× Ñ ×%Ñ% aÓ(Ð(r5   c                 ó8   — | j                   j                  |«      S rM   )rå  Úlog_probrd   s      r4   re   z4CumulativeDistributionTransform.log_abs_det_jacobian  s   € Ø× Ñ ×)Ñ)¨!Ó,Ð,r5   c                 óR   — | j                   |k(  r| S t        | j                  |¬«      S r   )r,   r   rå  rJ   s     r4   rK   z*CumulativeDistributionTransform.with_cache  s(   € Ø×Ñ˜zÒ)ØˆKÜ.¨t×/@Ñ/@ÈZÔXÐXr5   rp   rq   )rg   rr   rs   rt   ru   r   r   r'   rF   r   rx   r1   ry   r   rv   r&   rU   r\   re   rK   rz   r{   s   @r4   r   r   Ü  st   ø„ ñð$ €IØ×(Ñ(€HØ€Dñ) \ð )¸sð )È4õ )ð ð)˜ ×!7Ñ!7Ñ8ò )ó ð)ò(ò)ò-÷Yr5   r   )4r´   r1  r¶   r@   Úcollections.abcr   Útypingr   r   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   rN   r5   r4   ú<module>rú     sh  ðã Û Û Û Ý $ß "ã ß Ð Ý Ý +Ý 9÷õ ÷ .Ý ò€÷0fñ fôR;.˜	ô ;.ô|w�yô wñt & bÓ)Ð ôI8˜9ô I8ôXG+�yô G+ôT�9ô ô.(R�Yô (RòVNô
/�yô /ô2˜	ô ô0*>�Iô *>ôZ�9ô ô f
�iô f
ôRP!˜Iô P!ôf!�yô !ôH5-˜Yô 5-ôpL˜Yô Lô,/ 	ô /ô(m
�9ô m
ô`GR�Yô GRôT+Y iõ +Yr5   