Ë
    Gêñi8  ã                   ó  — d Z ddl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.
é    )ÚCallableN)Ú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/pod-logistic/pod-ai/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Úprojr/   )r   r-   r.   r/   Ú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)r6   r/   )r   r9   r6   Ú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   rG   r   s     €r   r   zLambdaLayer.__init__O   s   ø€ Ü‰ÑÔØ ˆ�r   c                 ó(   —  | j                   |g|¢­Ž S r;   )rG   )r   r9   Úargss      r   r=   zLambdaLayer.forwardS   s   € Øˆt�}‰}˜QÐ& Ò&Ð&r   )r%   r&   r'   r   r=   r)   r*   s   @r   rE   rE   N   s   ø„ ô!ö'r   rE   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
j                  dz  d
e
j                  dz  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.   r8   r0   Nc                 ó|   — || _         | j                  D �ci c]  }||| j                  |   z  “Œ c}| _        y c c}w r;   )r8   r.   )r   r8   Ú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é   ©r8   rL   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   )rS   r   r   )r   rR   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.
        rP   r2   )r8   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.
        )ÚlenrX   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   r2   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.   r/   )r,   r.   rE   r/   )r   r-   s     r   Úget_parameter_projectionz+DistributionOutput.get_parameter_projection‰   s'   € ô #Ø#Ø—]‘]Ü" 4§?¡?Ó3ô
ð 	
r   rI   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   rI   s     r   r/   zDistributionOutput.domain_map“   s   € ô "Ó#Ð#r   r9   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      @ç       @)rB   r#   Úsquare)r9   s    r   Ú
squarepluszDistributionOutput.squareplus›   s*   € ð ”E—J‘JœuŸ|™|¨A›°Ñ4Ó5Ñ5¸Ñ<Ð<r   )rP   ©NN)r%   r&   r'   ÚtypeÚ__annotations__r>   r?   r@   r   rS   rB   rC   r   rV   r(   rA   rX   r   Úfloatr\   r   ÚModuler^   r/   Ústaticmethodrd   r2   r   r   rK   rK   W   s  … ØÓØÓØ�3˜�8‰nÓñK˜Cð K¨ó KòHð $(Ø%)ñ	
\ð �\‰\˜DÑ ð
\ð �|‰|˜dÑ"ð	
\ð
 
ó
\ð ð4˜Uò 4ó ð4ð ð%˜3ò %ó ð%ð ð %ò ó ðð
°Cð 
¸B¿I¹Ió 
ð$ §¡ó $ð ð=�e—l‘lð = u§|¡|ò =ó ñ=r   rK   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.
    rP   )Údfr   r   r.   rL   rm   r   r   c                 ó  — | j                  |«      j                  t        j                  |j                  «      j
                  «      }d| j                  |«      z   }|j                  d«      |j                  d«      |j                  d«      fS )Nrb   éÿÿÿÿ©rd   Ú	clamp_minrB   ÚfinfoÚdtypeÚepsÚsqueeze)Úclsrm   r   r   s       r   r/   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?   r@   r>   rg   r
   rL   rf   ÚclassmethodrB   rC   r/   r2   r   r   rl   rl   ¤   se   … ñð '(°¸AÑ>€Hˆd�3˜�8‰nÓ>Ø'Ð˜Ó'àðB˜EŸL™Lð B¨u¯|©|ð BÀEÇLÁLò Bó ñBr   rl   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.
    rP   )r   r   r.   rL   r   r   c                 óÔ   — | j                  |«      j                  t        j                  |j                  «      j
                  «      }|j                  d«      |j                  d«      fS ©Nro   rp   )rv   r   r   s      r   r/   zNormalOutput.domain_map»   sJ   € à—‘˜uÓ%×/Ñ/´·±¸E¿K¹KÓ0H×0LÑ0LÓMˆØ�{‰{˜2‹ §¡¨bÓ 1Ð1Ð1r   N)r%   r&   r'   rw   r.   r?   r@   r>   rg   r	   rL   rf   rx   rB   rC   r/   r2   r   r   rz   rz   ³   sS   … ñð ()°1Ñ5€Hˆd�3˜�8‰nÓ5Ø%Ð˜Ó%àð2˜UŸ\™\ð 2°%·,±,ò 2ó ñ2r   rz   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j                  dz  dej                  dz  d	efd„Zy)ÚNegativeBinomialOutputz6
    Negative Binomial distribution output class.
    rP   ©Útotal_countÚlogitsr.   rL   r€   r�   c                 óh   — | j                  |«      }|j                  d«      |j                  d«      fS r|   )rd   ru   )rv   r€   r�   s      r   r/   z!NegativeBinomialOutput.domain_mapÉ   s/   € à—n‘n [Ó1ˆØ×"Ñ" 2Ó&¨¯©°rÓ(:Ð:Ð:r   r0   c                 óŠ   — |\  }}| j                   dk(  r| j                  ||¬«      S t        | j                  ||¬«      d«      S )NrP   r   rQ   )r   rR   r€   r�   s       r   rS   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;   )ÚlogrS   )r   rR   r   r   r€   r�   s         r   rV   z#NegativeBinomialOutput.distributionØ   s:   € ð )Ñˆ�VàÐà�e—i‘i“kÑ!ˆFà×&Ñ&¨°VÐ'<Ó=Ð=r   re   )r%   r&   r'   rw   r.   r?   r@   r>   rg   r   rL   rf   rx   rB   rC   r/   r   rS   rV   r2   r   r   r~   r~   Á   s˜   … ñð 01¸AÑ>€Hˆd�3˜�8‰nÓ>Ø/Ð˜Ó/àð; U§\¡\ð ;¸5¿<¹<ò ;ó ð;ðc°ó cð Y]ñ	>Ø$Ÿ|™|¨dÑ2ð	>ØBGÇ,Á,ÐQUÑBUð	>à	ô	>r   r~   )rw   Úcollections.abcr   rB   r   Útorch.distributionsr   r   r   r   r	   r
   r   r   ri   r,   rE   rK   rl   rz   r~   r2   r   r   ú<module>rˆ      s‡   ðñõ %ã Ý ÷÷ ñ ô$Ð/ô $ô:2˜"Ÿ)™)ô 2ô'�"—)‘)ô '÷J=ñ J=ôZBÐ'ô Bô2Ð%ô 2ô >Ð/õ  >r   