Ë
    S^(hE  ã                   ó  — d Z ddlmZmZ ddlZddlmZ ddlmZmZm	Z	m
Z
mZmZmZ  G d„ de«      Z G d„ d	ej                  «      Z G d
„ dej                  «      Z G d„ d«      Z G d„ de«      Z G d„ de«      Z G d„ de«      Zy)z:
Time series distributional output classes and utilities.
é    )ÚCallableÚOptionalN)Únn)ÚAffineTransformÚDistributionÚIndependentÚNegativeBinomialÚNormalÚStudentTÚTransformedDistributionc                   óV   ‡ — e Zd Zddefˆ fd„Zed„ «       Zed„ «       Zed„ «       Zˆ xZ	S )ÚAffineTransformedÚbase_distributionc                 ó”   •— |€dn|| _         |€dn|| _        t        ‰| �  |t	        | j                  | j                   |¬«      g«       y )Ng      ð?ç        ©ÚlocÚscaleÚ	event_dim)r   r   ÚsuperÚ__init__r   )Úselfr   r   r   r   Ú	__class__s        €ú\/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/time_series_utils.pyr   zAffineTransformed.__init__#   sE   ø€ Ø!˜M‘S¨uˆŒ
Ø˜+‘3¨3ˆŒä‰ÑÐ*¬_ÀÇÁÐQU×Q[ÑQ[ÐgpÔ-qÐ,rÕsó    c                 ób   — | j                   j                  | j                  z  | j                  z   S )z7
        Returns the mean of the distribution.
        )Ú	base_distÚmeanr   r   ©r   s    r   r   zAffineTransformed.mean)   s&   € ð
 �~‰~×"Ñ" T§Z¡ZÑ/°$·(±(Ñ:Ð:r   c                 óN   — | j                   j                  | j                  dz  z  S )z;
        Returns the variance of the distribution.
        é   )r   Úvariancer   r   s    r   r"   zAffineTransformed.variance0   s!   € ð
 �~‰~×&Ñ&¨¯©°Q©Ñ6Ð6r   c                 ó6   — | j                   j                  «       S )zE
        Returns the standard deviation of the distribution.
        )r"   Úsqrtr   s    r   ÚstddevzAffineTransformed.stddev7   s   € ð
 �}‰}×!Ñ!Ó#Ð#r   )NNr   )
Ú__name__Ú
__module__Ú__qualname__r   r   Úpropertyr   r"   r%   Ú__classcell__©r   s   @r   r   r   "   sM   ø„ ñt¨,õ tð ñ;ó ð;ð ñ7ó ð7ð ñ$ó ô$r   r   c            	       óœ   ‡ — e Zd Zdedeeef   dedeej                     f   ddfˆ fd„Z
dej                  deej                     fd	„Zˆ xZS )
ÚParameterProjectionÚin_featuresÚargs_dimÚ
domain_map.ÚreturnNc           	      óÞ   •— t        ‰| �  di |¤Ž || _        t        j                  |j                  «       D �cg c]  }t        j                  ||«      ‘Œ c}«      | _        || _        y c c}w )N© )	r   r   r/   r   Ú
ModuleListÚvaluesÚLinearÚprojr0   )r   r.   r/   r0   ÚkwargsÚdimr   s         €r   r   zParameterProjection.__init__@   sV   ø€ ô 	‰ÑÑ"˜6Ò"Ø ˆŒÜ—M‘MÈ(Ï/É/ÓJ[Ö"\À3¤2§9¡9¨[¸#Õ#>Ò"\Ó]ˆŒ	Ø$ˆ�ùò #]s   ¹A*Úxc                 óh   — | j                   D �cg c]
  } ||«      ‘Œ }} | j                  |Ž S c c}w ©N)r7   r0   )r   r:   r7   Úparams_unboundeds       r   ÚforwardzParameterProjection.forwardH   s5   € Ø04·	±	Ö:¨™D �GÐ:ÐÐ:àˆt�‰Ð 0Ð1Ð1ùò ;s   �/)r&   r'   r(   ÚintÚdictÚstrr   ÚtupleÚtorchÚTensorr   r>   r*   r+   s   @r   r-   r-   ?   sg   ø„ ð%Øð%Ø*.¨s°C¨x©.ð%ØFNÈsÐTYÐZ_×ZfÑZfÑTgÐOgÑFhð%à	õ%ð2˜Ÿ™ð 2¨%°·±Ñ*=÷ 2r   r-   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚLambdaLayerc                 ó0   •— t         ‰| �  «        || _        y r<   )r   r   Úfunction)r   rH   r   s     €r   r   zLambdaLayer.__init__O   s   ø€ Ü‰ÑÔØ ˆ�r   c                 ó(   —  | j                   |g|¢­Ž S r<   )rH   )r   r:   Úargss      r   r>   zLambdaLayer.forwardS   s   € Øˆt�}‰}˜QÐ& Ò&Ð&r   )r&   r'   r(   r   r>   r*   r+   s   @r   rF   rF   N   s   ø„ ô!ö'r   rF   c                   ód  — e Zd ZU eed<   eed<   eeef   ed<   ddeddfd„Zd„ Z		 	 dd	e
ej                     d
e
ej                     defd„Zedefd„«       Zedefd„«       Zedefd„«       Zdedej,                  fd„Zdej                  fd„Zedej                  dej                  fd„«       Zy)ÚDistributionOutputÚdistribution_classr.   r/   r9   r1   Nc                 ó|   — || _         | j                  D �ci c]  }||| j                  |   z  “Œ c}| _        y c c}w r<   )r9   r/   )r   r9   Úks      r   r   zDistributionOutput.__init__\   s5   € ØˆŒØ<@¿M¹MÖJ°q˜˜C $§-¡-°Ñ"2Ñ2Ñ2ÒJˆ�ùÒJs   –9c                 óp   — | j                   dk(  r | j                  |Ž S t         | j                  |Ž d«      S )Né   ©r9   rM   r   )r   Ú
distr_argss     r   Ú_base_distributionz%DistributionOutput._base_distribution`   s;   € Ø�8‰8�qŠ=Ø*�4×*Ñ*¨JÐ7Ð7äÐ6˜t×6Ñ6¸
ÐCÀQÓGÐGr   r   r   c                 ób   — | j                  |«      }|€|€|S t        |||| j                  ¬«      S )Nr   )rT   r   r   )r   rS   r   r   Údistrs        r   ÚdistributionzDistributionOutput.distributionf   s7   € ð ×'Ñ'¨
Ó3ˆØˆ;˜5˜=ØˆLä$ U°¸5ÈDÏNÉNÔ[Ð[r   c                 ó>   — | j                   dk(  rdS | j                   fS )zo
        Shape of each individual event contemplated by the distributions that this object constructs.
        rQ   r3   )r9   r   s    r   Úevent_shapezDistributionOutput.event_shaper   s   € ð
 —X‘X ’]ˆrÐ3¨¯©¨Ð3r   c                 ó,   — t        | j                  «      S )z�
        Number of event dimensions, i.e., length of the `event_shape` tuple, of the distributions that this object
        constructs.
        )ÚlenrY   r   s    r   r   zDistributionOutput.event_dimy   s   € ô �4×#Ñ#Ó$Ð$r   c                  ó   — y)zÇ
        A float that will have a valid numeric value when computing the log-loss of the corresponding distribution. By
        default 0.0. This value will be used when padding data series.
        r   r3   r   s    r   Úvalue_in_supportz#DistributionOutput.value_in_support�   s   € ð r   c                 óX   — t        || j                  t        | j                  «      ¬«      S )z~
        Return the parameter projection layer that maps the input to the appropriate parameters of the distribution.
        )r.   r/   r0   )r-   r/   rF   r0   )r   r.   s     r   Úget_parameter_projectionz+DistributionOutput.get_parameter_projection‰   s'   € ô #Ø#Ø—]‘]Ü" 4§?¡?Ó3ô
ð 	
r   rJ   c                 ó   — t        «       ‚)a  
        Converts arguments to the right shape and domain. The domain depends on the type of distribution, while the
        correct shape is obtained by reshaping the trailing axis in such a way that the returned tensors define a
        distribution of the right event_shape.
        )ÚNotImplementedError)r   rJ   s     r   r0   zDistributionOutput.domain_map“   s   € ô "Ó#Ð#r   r:   c                 ód   — | t        j                  t        j                  | «      dz   «      z   dz  S )z²
        Helper to map inputs to the positive orthant by applying the square-plus operation. Reference:
        https://twitter.com/jon_barron/status/1387167648669048833
        g      @ç       @)rC   r$   Úsquare)r:   s    r   Ú
squarepluszDistributionOutput.squareplus›   s*   € ð ”E—J‘JœuŸ|™|¨A›°Ñ4Ó5Ñ5¸Ñ<Ð<r   )rQ   ©NN)r&   r'   r(   ÚtypeÚ__annotations__r?   r@   rA   r   rT   r   rC   rD   r   rW   r)   rB   rY   r   Úfloatr]   r   ÚModuler_   r0   Ústaticmethodre   r3   r   r   rL   rL   W   s  … ØÓØÓØ�3˜�8‰nÓñK˜Cð K¨ó KòHð '+Ø(,ñ	
\ð �e—l‘lÑ#ð
\ð ˜Ÿ™Ñ%ð	
\ð
 
ó
\ð ð4˜Uò 4ó ð4ð ð%˜3ò %ó ð%ð ð %ò ó ðð
°Cð 
¸B¿I¹Ió 
ð$ §¡ó $ð ð=�e—l‘lð = u§|¡|ò =ó ñ=r   rL   c                   óš   — e Zd ZU dZddddœZeeef   ed<   e	Z
eed<   edej                  dej                  dej                  fd	„«       Zy
)ÚStudentTOutputz.
    Student-T distribution output class.
    rQ   )Údfr   r   r/   rM   rn   r   r   c                 ó  — | j                  |«      j                  t        j                  |j                  «      j
                  «      }d| j                  |«      z   }|j                  d«      |j                  d«      |j                  d«      fS )Nrc   éÿÿÿÿ©re   Ú	clamp_minrC   ÚfinfoÚdtypeÚepsÚsqueeze)Úclsrn   r   r   s       r   r0   zStudentTOutput.domain_map¬   sg   € à—‘˜uÓ%×/Ñ/´·±¸E¿K¹KÓ0H×0LÑ0LÓMˆØ�3—>‘> "Ó%Ñ%ˆØ�z‰z˜"‹~˜sŸ{™{¨2›°·±¸bÓ0AÐAÐAr   N)r&   r'   r(   Ú__doc__r/   r@   rA   r?   rh   r   rM   rg   ÚclassmethodrC   rD   r0   r3   r   r   rm   rm   ¤   se   … ñð '(°¸AÑ>€Hˆd�3˜�8‰nÓ>Ø'Ð˜Ó'àðB˜EŸL™Lð B¨u¯|©|ð BÀEÇLÁLò Bó ñBr   rm   c                   ó€   — e Zd ZU dZdddœZeeef   ed<   e	Z
eed<   edej                  dej                  fd„«       Zy	)
ÚNormalOutputz+
    Normal distribution output class.
    rQ   )r   r   r/   rM   r   r   c                 óÔ   — | j                  |«      j                  t        j                  |j                  «      j
                  «      }|j                  d«      |j                  d«      fS ©Nrp   rq   )rw   r   r   s      r   r0   zNormalOutput.domain_map»   sJ   € à—‘˜uÓ%×/Ñ/´·±¸E¿K¹KÓ0H×0LÑ0LÓMˆØ�{‰{˜2‹ §¡¨bÓ 1Ð1Ð1r   N)r&   r'   r(   rx   r/   r@   rA   r?   rh   r
   rM   rg   ry   rC   rD   r0   r3   r   r   r{   r{   ³   sS   … ñð ()°1Ñ5€Hˆd�3˜�8‰nÓ5Ø%Ð˜Ó%àð2˜UŸ\™\ð 2°%·,±,ò 2ó ñ2r   r{   c                   óØ   — e Zd ZU dZdddœZeeef   ed<   e	Z
eed<   edej                  dej                  fd„«       Zd	efd
„Z	 ddeej                     deej                     d	efd„Zy)ÚNegativeBinomialOutputz6
    Negative Binomial distribution output class.
    rQ   ©Útotal_countÚlogitsr/   rM   r�   r‚   c                 óh   — | j                  |«      }|j                  d«      |j                  d«      fS r}   )re   rv   )rw   r�   r‚   s      r   r0   z!NegativeBinomialOutput.domain_mapÉ   s/   € à—n‘n [Ó1ˆØ×"Ñ" 2Ó&¨¯©°rÓ(:Ð:Ð:r   r1   c                 óŠ   — |\  }}| j                   dk(  r| j                  ||¬«      S t        | j                  ||¬«      d«      S )NrQ   r€   rR   )r   rS   r�   r‚   s       r   rT   z)NegativeBinomialOutput._base_distributionÎ   sL   € Ø(Ñˆ�VØ�8‰8�qŠ=Ø×*Ñ*°{È6Ð*ÓRÐRä˜t×6Ñ6À;ÐW]Ð6Ó^Ð`aÓbÐbr   Nr   r   c                 ó\   — |\  }}|�||j                  «       z  }| j                  ||f«      S r<   )ÚlogrT   )r   rS   r   r   r�   r‚   s         r   rW   z#NegativeBinomialOutput.distributionØ   s:   € ð )Ñˆ�VàÐà�e—i‘i“kÑ!ˆFà×&Ñ&¨°VÐ'<Ó=Ð=r   rf   )r&   r'   r(   rx   r/   r@   rA   r?   rh   r	   rM   rg   ry   rC   rD   r0   r   rT   r   rW   r3   r   r   r   r   Á   s—   … ñð 01¸AÑ>€Hˆd�3˜�8‰nÓ>Ø/Ð˜Ó/àð; U§\¡\ð ;¸5¿<¹<ò ;ó ð;ðc°ó cð _cñ	>Ø'¨¯©Ñ5ð	>ØEMÈeÏlÉlÑE[ð	>à	ô	>r   r   )rx   Útypingr   r   rC   r   Útorch.distributionsr   r   r   r	   r
   r   r   r   rj   r-   rF   rL   rm   r{   r   r3   r   r   ú<module>r‰      s‡   ðñ÷ &ã Ý ÷÷ ñ ô$Ð/ô $ô:2˜"Ÿ)™)ô 2ô'�"—)‘)ô '÷J=ñ J=ôZBÐ'ô Bô2Ð%ô 2ô >Ð/õ  >r   