Ë
    T^(h g ã                   ó`  — d Z ddlZddlmZ ddlmZmZmZ ddlZddl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 ddlmZmZmZ ddlmZ  ej4                  e«      ZdZ G d„ de	j<                  «      Z G d„ de	j<                  «      Z 	 	 	 d^dejB                  de"de#de$de%f
d„Z&	 	 d_dejB                  dee#e%f   de#de%fd„Z' G d„ de	j<                  «      Z( G d„ de	j<                  «      Z) G d„ d e	j<                  «      Z* G d!„ d"e«      Z+ G d#„ d$e	j<                  «      Z, G d%„ d&e	j<                  «      Z- G d'„ d(e+«      Z.d)Z/e G d*„ d+e«      «       Z0e G d,„ d-e«      «       Z1e G d.„ d/e«      «       Z2e G d0„ d1e«      «       Z3e G d2„ d3e«      «       Z4e G d4„ d5e«      «       Z5d6ejl                  jn                  d7ejB                  d8ejB                  fd9„Z8d`d:ejB                  d;eejB                     d8ejB                  fd<„Z9 G d=„ d>e	j<                  «      Z: G d?„ d@e	j<                  «      Z; G dA„ dBe	j<                  «      Z< G dC„ dDe	j<                  «      Z= edEe/«       G dF„ dGe+«      «       Z> G dH„ dIe	j<                  «      Z? edJe/«       G dK„ dLe+«      «       Z@ G dM„ dNe	j<                  «      ZA edOe/«       G dP„ dQe+«      «       ZB edRe/«       G dS„ dTe	j<                  «      «       ZC edUe/«       G dV„ dWe+«      «       ZD G dX„ dYe	j<                  «      ZE edZe/«       G d[„ d\e+«      «       ZFg d]¢ZGy)azPyTorch PatchTST model.é    N)Ú	dataclass)ÚOptionalÚTupleÚUnion)Únné   )ÚACT2CLS)ÚBaseModelOutput)ÚPreTrainedModel)ÚNegativeBinomialOutputÚNormalOutputÚStudentTOutput)ÚModelOutputÚadd_start_docstringsÚloggingé   )ÚPatchTSTConfigr   c                   ó†  ‡ — e Zd ZdZ	 	 	 	 	 ddededededededee   fˆ fd	„Z	d
e
j                  dedefd„Z	 	 	 	 	 dde
j                  dee
j                     deee
j                        dee
j                     dee
j                     dedee
j                  ee
j                     eee
j                        f   fd„Zˆ xZS )ÚPatchTSTAttentionz=Multi-headed attention from 'Attention Is All You Need' paperÚ	embed_dimÚ	num_headsÚdropoutÚ
is_decoderÚbiasÚ	is_causalÚconfigc                 ó
  •— t         ‰| �  «        || _        || _        || _        ||z  | _        || _        | j
                  |z  | j                  k7  rt        d| j                  › d|› d�«      ‚| j
                  dz  | _        || _	        || _
        t        j                  |||¬«      | _        t        j                  |||¬«      | _        t        j                  |||¬«      | _        t        j                  |||¬«      | _        y )Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: ú).g      à¿©r   )ÚsuperÚ__init__r   r   r   Úhead_dimr   Ú
ValueErrorÚscalingr   r   r   ÚLinearÚk_projÚv_projÚq_projÚout_proj)	Úselfr   r   r   r   r   r   r   Ú	__class__s	           €úl/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/patchtst/modeling_patchtst.pyr!   zPatchTSTAttention.__init__)   sä   ø€ ô 	‰ÑÔØ"ˆŒØ"ˆŒØˆŒØ! YÑ.ˆŒØˆŒà�M‰M˜IÑ%¨$¯.©.Ò8ÜØMÈdÏnÉnÐM]Ø$ Y K¨rð3óð ð —}‘} dÑ*ˆŒØ$ˆŒØ"ˆŒä—i‘i 	¨9¸4Ô@ˆŒÜ—i‘i 	¨9¸4Ô@ˆŒÜ—i‘i 	¨9¸4Ô@ˆŒÜŸ	™	 )¨Y¸TÔBˆ�ó    ÚtensorÚseq_lenÚbszc                 óŽ   — |j                  ||| j                  | j                  «      j                  dd«      j	                  «       S )Nr   é   )Úviewr   r"   Ú	transposeÚ
contiguous)r*   r.   r/   r0   s       r,   Ú_shapezPatchTSTAttention._shapeH   s7   € Ø�{‰{˜3 ¨¯©¸¿¹ÓG×QÑQÐRSÐUVÓW×bÑbÓdÐdr-   Úhidden_statesÚkey_value_statesÚpast_key_valueÚattention_maskÚlayer_head_maskÚoutput_attentionsÚreturnc                 ó
  — |du}|j                  «       \  }}	}
| j                  |«      | j                  z  }|r0|�.|d   j                  d   |j                  d   k(  r|d   }|d   }�n
|rE| j	                  | j                  |«      d|«      }| j	                  | j                  |«      d|«      }nÃ|�}| j	                  | j                  |«      d|«      }| j	                  | j                  |«      d|«      }t        j                  |d   |gd¬«      }t        j                  |d   |gd¬«      }nD| j	                  | j                  |«      d|«      }| j	                  | j                  |«      d|«      }| j                  r||f}|| j                  z  d| j                  f} | j	                  ||	|«      j                  |Ž } |j                  |Ž } |j                  |Ž }|j                  d«      }t        j                  ||j                  dd«      «      }|j                  «       || j                  z  |	|fk7  r/t!        d|| j                  z  |	|f› d|j                  «       › �«      ‚|�{|j                  «       |d|	|fk7  r#t!        d	|d|	|f› d|j                  «       › �«      ‚|j                  || j                  |	|«      |z   }|j                  || j                  z  |	|«      }t"        j$                  j'                  |d¬«      }|�›|j                  «       | j                  fk7  r*t!        d
| j                  f› d|j                  «       › �«      ‚|j                  dddd«      |j                  || j                  |	|«      z  }|j                  || j                  z  |	|«      }|r?|j                  || j                  |	|«      }|j                  || j                  z  |	|«      }nd}t"        j$                  j)                  || j(                  | j*                  ¬«      }t        j                  ||«      }|j                  «       || j                  z  |	| j                  fk7  r9t!        d|| j                  z  |	| j                  f› d|j                  «       › �«      ‚|j                  || j                  |	| j                  «      }|j                  dd«      }|j                  ||	| j,                  «      }| j/                  |«      }|||fS )z#Input shape: Batch x Time x ChannelNr   r2   r   éÿÿÿÿ©Údimz$Attention weights should be of size z	, but is z!Attention mask should be of size z/Head mask for a single layer should be of size )ÚpÚtrainingz `attn_output` should be of size )Úsizer(   r$   Úshaper6   r&   r'   ÚtorchÚcatr   r   r"   r3   ÚreshapeÚbmmr4   r#   r   Ú
functionalÚsoftmaxr   rC   r   r)   )r*   r7   r8   r9   r:   r;   r<   Úis_cross_attentionr0   Útgt_lenÚ_Úquery_statesÚ
key_statesÚvalue_statesÚ
proj_shapeÚsrc_lenÚattn_weightsÚattn_weights_reshapedÚ
attn_probsÚattn_outputs                       r,   ÚforwardzPatchTSTAttention.forwardK   s  € ð .°TÐ9Ðà'×,Ñ,Ó.‰ˆˆW�að —{‘{ =Ó1°D·L±LÑ@ˆñ ØÐ*Ø˜qÑ!×'Ñ'¨Ñ*Ð.>×.DÑ.DÀQÑ.GÒGð (¨Ñ*ˆJØ)¨!Ñ,ŠLÙàŸ™ T§[¡[Ð1AÓ%BÀBÈÓLˆJØŸ;™; t§{¡{Ð3CÓ'DÀbÈ#ÓN‰LØÐ'àŸ™ T§[¡[°Ó%?ÀÀSÓIˆJØŸ;™; t§{¡{°=Ó'AÀ2ÀsÓKˆLÜŸ™ N°1Ñ$5°zÐ#BÈÔJˆJÜ Ÿ9™9 n°QÑ&7¸Ð%FÈAÔN‰Lð Ÿ™ T§[¡[°Ó%?ÀÀSÓIˆJØŸ;™; t§{¡{°=Ó'AÀ2ÀsÓKˆLà�?Š?ð )¨,Ð7ˆNà˜DŸN™NÑ*¨B°·±Ð>ˆ
ØC�t—{‘{ <°¸#Ó>×CÑCÀZÐPˆØ'�Z×'Ñ'¨Ð4ˆ
Ø+�|×+Ñ+¨ZÐ8ˆà—/‘/ !Ó$ˆÜ—y‘y ¨z×/CÑ/CÀAÀqÓ/IÓJˆà×ÑÓ 3¨¯©Ñ#7¸À'Ð"JÒJÜØ6¸¸d¿n¹nÑ8LÈgÐW^Ð7_Ð6`ð aØ ×%Ñ%Ó'Ð(ð*óð ð
 Ð%Ø×"Ñ"Ó$¨¨a°¸'Ð(BÒBÜ Ø7¸¸aÀÈ'Ð8RÐ7SÐS\Ð]k×]pÑ]pÓ]rÐ\sÐtóð ð (×,Ñ,¨S°$·.±.À'È7ÓSÐVdÑdˆLØ'×,Ñ,¨S°4·>±>Ñ-AÀ7ÈGÓTˆLä—}‘}×,Ñ,¨\¸rÐ,ÓBˆàÐ&Ø×#Ñ#Ó%¨$¯.©.Ð):Ò:Ü ØEÀtÇ~Á~ÐFWÐEXð YØ'×,Ñ,Ó.Ð/ð1óð ð +×/Ñ/°°2°q¸!Ó<¸|×?PÑ?PÐQTÐVZ×VdÑVdÐfmÐovÓ?wÑwˆLØ'×,Ñ,¨S°4·>±>Ñ-AÀ7ÈGÓTˆLáð
 %1×$5Ñ$5°c¸4¿>¹>È7ÐT[Ó$\Ð!Ø0×5Ñ5°c¸D¿N¹NÑ6JÈGÐU\Ó]‰Là$(Ð!ä—]‘]×*Ñ*¨<¸4¿<¹<ÐRV×R_ÑR_Ð*Ó`ˆ
ä—i‘i 
¨LÓ9ˆà×ÑÓ #¨¯©Ñ"6¸ÀÇÁÐ!OÒOÜØ2°C¸$¿.¹.Ñ4HÈ'ÐSW×S`ÑS`Ð3aÐ2bð cØ×$Ñ$Ó&Ð'ð)óð ð
 "×&Ñ& s¨D¯N©N¸GÀTÇ]Á]ÓSˆØ!×+Ñ+¨A¨qÓ1ˆð "×)Ñ)¨#¨w¸¿¹ÓGˆà—m‘m KÓ0ˆàÐ1°>ÐAÐAr-   )ç        FTFN)NNNNF)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚintÚfloatÚboolr   r   r!   rF   ÚTensorr6   r   rX   Ú__classcell__©r+   s   @r,   r   r   &   sM  ø„ ÙGð Ø ØØØ+/ñCàðCð ðCð ð	Cð
 ðCð ðCð ðCð ˜Ñ(õCð>e˜UŸ\™\ð e°Cð e¸có eð 48Ø8<Ø15Ø26Ø"'ñvBà—|‘|ðvBð # 5§<¡<Ñ0ðvBð !  u§|¡|Ñ!4Ñ5ð	vBð
 ! §¡Ñ.ðvBð " %§,¡,Ñ/ðvBð  ðvBð 
ˆu�|‰|˜X e§l¡lÑ3°X¸eÀEÇLÁLÑ>QÑ5RÐRÑ	S÷vBr-   r   c                   óH   ‡ — e Zd ZdZdefˆ fd„Zdej                  fd„Zˆ xZ	S )ÚPatchTSTBatchNormzP
    Compute batch normalization over the sequence length (time) dimension.
    r   c                 ó‚   •— t         ‰| �  «        t        j                  |j                  |j
                  ¬«      | _        y )N©Úeps)r    r!   r   ÚBatchNorm1dÚd_modelÚnorm_epsÚ	batchnorm©r*   r   r+   s     €r,   r!   zPatchTSTBatchNorm.__init__É   s(   ø€ Ü‰ÑÔÜŸ™¨¯©¸F¿O¹OÔLˆ�r-   Úinputsc                 ól   — |j                  dd«      }| j                  |«      }|j                  dd«      S )a  
        Parameters:
            inputs (`torch.Tensor` of shape `(batch_size, sequence_length, d_model)`):
                input for Batch norm calculation
        Returns:
            `torch.Tensor` of shape `(batch_size, sequence_length, d_model)`
        r   r2   )r4   rl   )r*   rn   Úoutputs      r,   rX   zPatchTSTBatchNorm.forwardÍ   s7   € ð ×!Ñ! ! QÓ'ˆØ—‘ Ó'ˆØ×Ñ  1Ó%Ð%r-   ©
rZ   r[   r\   r]   r   r!   rF   ra   rX   rb   rc   s   @r,   re   re   Ä   s&   ø„ ñðM˜~õ Mð
&˜eŸl™l÷ 
&r-   re   rn   Ú
mask_ratioÚunmasked_channel_indicesÚchannel_consistent_maskingÚ
mask_valuec                 ó°  — |dk  s|dk\  rt        d|› d�«      ‚| j                  \  }}}}| j                  }	t        |d|z
  z  «      }
|r-t	        j
                  |d||	¬«      }|j                  d|d«      }nt	        j
                  ||||	¬«      }t	        j                  ||||	¬«      }d|dd…dd…d|
…f<   t	        j                  |d¬«      }t	        j                  |d¬«      }t	        j                  |d|¬	«      }|j                  d«      j                  ddd|«      }|�d|dd…|dd…dd…f<   | j                  |j                  «       |«      }||d
   fS )aÆ  random_masking: Mask the input considering the control variables.

    Args:
        inputs (`torch.Tensor` of shape `(batch_size, num_channels, sequence_length, num_features)`):
            The input tensor to mask.
        mask_ratio (`float`):
            Masking ratio applied to mask the input data during random pretraining. It is the number between 0 and 1.
        unmasked_channel_indices (list, *optional*):
            Indices of channels that will not be masked.
        channel_consistent_masking (bool, *optional*, defaults to `False`):
            When true, masking will be same across all channels of a timeseries. Otherwise, masking positions will vary
            across channels.
        mask_value (int, *optional*, defaults to 0):
            Define the value of masked patches for pretraining.

    Returns:
        `tuple(torch.Tensor)`: inputs_mask, masked input, same shape as input Tensor and mask tensor of shape [bs x c x
        n]
    r   r   zMask ratio z has to be between 0 and 1.©ÚdeviceNr?   r@   )rA   Úindex©.r   )r#   rE   rx   r^   rF   ÚrandÚrepeatÚonesÚargsortÚgatherÚ	unsqueezeÚmasked_fillr`   )rn   rr   rs   rt   ru   Ú
batch_sizeÚnum_channelsÚsequence_lengthÚnum_featuresrx   Úlen_keepÚnoiseÚmaskÚids_shuffleÚids_restoreÚinputs_masks                   r,   Úrandom_maskingrŒ   Ú   sQ  € ð4 �A‚~˜ qšÜ˜; z lÐ2MÐNÓOÐOà>D¿l¹lÑ;€J�˜o¨|Ø�]‰]€Fä�? a¨*¡nÑ5Ó6€Há!Ü—
‘
˜: q¨/À&ÔIˆØ—‘˜Q ¨aÓ0‰ô —
‘
˜: |°_ÈVÔTˆô �:‰:�j ,°ÈÔO€DØ€DŠŠAˆy�ˆyˆÑô —-‘- ¨2Ô.€KÜ—-‘- °Ô4€Kä�<‰<˜ "¨KÔ8€DØ�>‰>˜"Ó×$Ñ$ Q¨¨1¨lÓ;€DØÐ+Ø23ˆŠQÐ(ª!ªQÐ.Ñ/à×$Ñ$ T§Y¡Y£[°*Ó=€KØ˜˜V™Ð$Ð$r-   Únum_forecast_mask_patchesc                 óP  — t        |t        «      r|g}|D �cg c]  }d‘Œ }}| j                  \  }}}}	t        j                  |||| j
                  ¬«      }
g }d}t        |«      }t        ||«      D ]H  \  }}|dk  s||k\  rt        d|› d�«      ‚t        ||z  |z  «      }|j                  |||g«       ||z  }ŒJ t        |d„ ¬«      }||k  r|d   d   ||z
  z   |d   d<   n||kD  r|d	   d   ||z
  z   |d	   d<   d}|D ]  \  }}}||z   }d|
||…d
d
…| d
…f<   |}Œ t        j                  |
j                  d   «      }|
|   }
|
j                  d	«      j                  ddd|	«      }
|�d|
d
d
…|d
d
…d
d
…f<   | j                  |
j                  «       |«      }||
d   fS c c}w )a¡  Forecast masking that masks the last K patches where K is from the num_forecast_mask_patches.
    If num_forecast_mask_patches is a list, samples in the batch will be randomly masked by numbers defined in the list.

    Parameters:
        inputs (`torch.Tensor`):
            Input of shape `(bs, num_channels, num_patch, patch_length)`
        num_forecast_mask_patches (`list`):
            Number of patches to be masked at the end of each batch sample. e.g. 4 or [3, 5].
        unmasked_channel_indices (`list`, *optional*):
            Indices of channels that are not masked.
        mask_value (`int`, *optional*, defaults to 0):
            Values in the masked patches will be filled by `mask_value`.

    Returns:
        `tuple(torch.Tensor)`: inputs_mask, masked input, same shape as inputs Tensor and Mask tensor of shape `(bs,
        num_channels , num_patch)` or `(bs, tsg1, tsg2, num_channels, num_patch)`
    r   rw   r   znum_forecast_mask_patches z6 should be greater than 0 and less than total patches.c                 ó   — | d   S )Nr2   © )Úxs    r,   ú<lambda>z"forecast_masking.<locals>.<lambda>@  s
   € ¨!¨A©$€ r-   )Úkeyr2   r?   Nrz   )Ú
isinstancer^   rE   rF   Úzerosrx   ÚsumÚzipr#   ÚappendÚsortedÚrandpermr€   r|   r�   r`   )rn   r�   rs   ru   rN   Úforecast_mask_ratiosr‚   rƒ   r„   r…   rˆ   Út_listÚtotal_lengthÚtotal_ratioÚpatch_lengthÚratioÚtemp_lenÚbatch1Ú	patch_lenÚbatch2Úpermr‹   s                         r,   Úforecast_maskingr¦     s  € ô0 Ð+¬SÔ1Ø%>Ð$?Ð!Ø'@ÖA !šAÐAÐÐAà>D¿l¹lÑ;€J�˜o¨|Ü�;‰;�z <°ÈÏÉÔW€Dà€FØ€LÜÐ*Ó+€Kä"Ð#<Ð>RÓSò !Ñˆ�eØ˜1Ò °Ò ?ÜØ,¨\¨NÐ:pÐqóð ô �z EÑ)¨KÑ7Ó8ˆØ�‰�| U¨HÐ5Ô6Ø˜Ñ ‰ð!ô �F¡Ô/€Fà�jÒ Ø˜a‘y ‘| z°LÑ'@ÑAˆˆq‰	�!ŠØ	˜
Ò	"Ø˜r™
 1™¨¸
Ñ)BÑCˆˆr‰
�1‰à€FØ"(ò Ñˆ	�1�hØ˜(Ñ"ˆØ./ˆˆV�Fˆ]šA 	˜z™{Ð*Ñ+Ø‰ðô
 �>‰>˜$Ÿ*™* Q™-Ó(€DØ�‰:€Dà�>‰>˜"Ó×$Ñ$ Q¨¨1¨lÓ;€DØÐ+Ø23ˆŠQÐ(ª!ªQÐ.Ñ/à×$Ñ$ T§Y¡Y£[°*Ó=€KØ˜˜V™Ð$Ð$ùòO Bs   ˜	F#c                   óH   ‡ — e Zd ZdZdefˆ fd„Zdej                  fd„Zˆ xZ	S )ÚPatchTSTPatchifyz³
    A class to patchify the time series sequence into different patches

    Returns:
        `torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`
    r   c                 ó  •— t         ‰| �  «        |j                  | _        |j                  | _        |j
                  | _        | j                  | j                  k  r&t        d| j                  › d| j                  › d�«      ‚t        | j                  | j                  «      | j                  z
  | j
                  z  dz   | _        | j                  | j
                  | j                  dz
  z  z   }| j                  |z
  | _	        y )NzSequence length (z+) has to be greater than the patch length (ú)r   )
r    r!   Úcontext_lengthr„   rŸ   Úpatch_strider#   ÚmaxÚnum_patchesÚsequence_start)r*   r   Únew_sequence_lengthr+   s      €r,   r!   zPatchTSTPatchify.__init__`  sò   ø€ Ü‰ÑÔà%×4Ñ4ˆÔØ"×/Ñ/ˆÔØ"×/Ñ/ˆÔà×Ñ 4×#4Ñ#4Ò4ÜØ# D×$8Ñ$8Ð#9Ð9dÐei×evÑevÐdwÐwxÐyóð ô
   × 4Ñ 4°d×6GÑ6GÓHÈ4×K\ÑK\Ñ\Ðae×arÑarÑrÐuvÑvˆÔØ"×/Ñ/°$×2CÑ2CÀt×GWÑGWÐZ[ÑG[Ñ2\Ñ\ÐØ"×2Ñ2Ð5HÑHˆÕr-   Úpast_valuesc                 ó:  — |j                   d   }|| j                  k7  rt        d|› d| j                  › d�«      ‚|dd…| j                  d…dd…f   }|j	                  d| j
                  | j                  ¬«      }|j                  dd«      j                  «       }|S )a!  
        Parameters:
            past_values (`torch.Tensor` of shape `(batch_size, sequence_length, num_channels)`, *required*):
                Input for patchification

        Returns:
            `torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`
        éþÿÿÿzInput sequence length (z%) doesn't match model configuration (r   N)Ú	dimensionrD   Ústepéýÿÿÿ)	rE   r„   r#   r¯   ÚunfoldrŸ   r¬   r4   r5   )r*   r±   r„   rp   s       r,   rX   zPatchTSTPatchify.forwardq  s¨   € ð &×+Ñ+¨BÑ/ˆØ˜d×2Ñ2Ò2ÜØ)¨/Ð):Ð:_Ð`d×`tÑ`tÐ_uÐuwÐxóð ð šQ × 3Ñ 3Ñ 5²qÐ8Ñ9ˆà—‘¨°$×2CÑ2CÈ$×J[ÑJ[�Ó\ˆà×!Ñ! " bÓ)×4Ñ4Ó6ˆØˆr-   rq   rc   s   @r,   r¨   r¨   X  s&   ø„ ñðI˜~õ Ið" 5§<¡<÷ r-   r¨   c                   óH   ‡ — e Zd ZdZdefˆ fd„Zdej                  fd„Zˆ xZ	S )ÚPatchTSTMaskinga�  
    Class to perform random or forecast masking.

    Parameters:
        config (`PatchTSTConfig`): model config
    Returns:
        x_mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`)
            Masked patched input
        mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches)`)
            Bool tensor indicating True on masked points
    r   c                 ó<  •— t         ‰| �  «        |j                  | _        |j                  | _        |j                  | _        |j
                  | _        |j                  | _        |j                  | _        | j                  �t        | j                  «      | _        y y ©N)	r    r!   Úrandom_mask_ratiort   Ú	mask_typer�   rs   ru   r™   rm   s     €r,   r!   zPatchTSTMasking.__init__•  s„   ø€ Ü‰ÑÔØ!'×!9Ñ!9ˆÔØ*0×*KÑ*KˆÔ'Ø×)Ñ)ˆŒØ)/×)IÑ)IˆÔ&Ø(.×(GÑ(GˆÔ%Ø ×+Ñ+ˆŒØ×(Ñ(Ð4Ü,2°4×3PÑ3PÓ,QˆDÕ)ð 5r-   Úpatch_inputc                 ór  — | j                   dk(  r<t        || j                  | j                  | j                  | j
                  ¬«      \  }}nY| j                   dk(  r1t        || j                  | j                  | j
                  ¬«      \  }}nt        d| j                   › d�«      ‚|j                  «       }||fS )aä  
        Parameters:
            patch_input (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`, *required*):
                Patch input

        Return:
            masked_input (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`)
                Masked patched input
            mask (`torch.Tensor` of shape `(batch_size, num_channels, num_patches)`)
                Bool tensor indicating True on masked points

        Úrandom)rn   rr   rs   rt   ru   Úforecast)rn   r�   rs   ru   zInvalid mask type ú.)
r½   rŒ   r¼   rs   rt   ru   r¦   r�   r#   r`   )r*   r¾   Úmasked_inputrˆ   s       r,   rX   zPatchTSTMasking.forward   s±   € ð �>‰>˜XÒ%Ü!/Ø"Ø×1Ñ1Ø)-×)FÑ)FØ+/×+JÑ+JØŸ?™?ô"ÑˆL™$ð �^‰^˜zÒ)Ü!1Ø"Ø*.×*HÑ*HØ)-×)FÑ)FØŸ?™?ô	"ÑˆL™$ô Ð1°$·.±.Ð1AÀÐCÓDÐDð �y‰y‹{ˆØ˜TÐ!Ð!r-   rq   rc   s   @r,   r¹   r¹   ˆ  s&   ø„ ñ
ð	R˜~õ 	Rð!" 5§<¡<÷ !"r-   r¹   c                   óT   ‡ — e Zd ZdZdefˆ fd„Zddej                  dee	   fd„Z
ˆ xZS )ÚPatchTSTEncoderLayerz 
    PatchTST encoder layer
    r   c           
      ó  •— t         ‰| �  «        |j                  | _        t        |j                  |j
                  |j                  ¬«      | _        |j                  dkD  rt        j                  |j                  «      nt        j                  «       | _        |j                  dk(  rt        |«      | _        nX|j                  dk(  r1t        j                   |j                  |j"                  ¬«      | _        nt%        |j                  › d�«      ‚| j                  r¿|j                  dkD  rt        j                  |j                  «      nt        j                  «       | _        |j                  dk(  rt        |«      | _        nX|j                  dk(  r1t        j                   |j                  |j"                  ¬«      | _        nt%        |j                  › d�«      ‚t        j*                  t        j,                  |j                  |j.                  |j0                  ¬«      t3        |j4                     «       |j6                  dkD  rt        j                  |j6                  «      nt        j                  «       t        j,                  |j.                  |j                  |j0                  ¬«      «      | _        |j                  dkD  rt        j                  |j                  «      nt        j                  «       | _        |j                  dk(  rt        |«      | _        nX|j                  dk(  r1t        j                   |j                  |j"                  ¬«      | _        nt%        |j                  › d�«      ‚|j>                  | _        y )N)r   r   r   r   rl   Ú	layernormrg   z$ is not a supported norm layer type.r   ) r    r!   Úchannel_attentionr   rj   Únum_attention_headsÚattention_dropoutÚ	self_attnÚpath_dropoutr   ÚDropoutÚIdentityÚdropout_path1Ú	norm_typere   Únorm_sublayer1Ú	LayerNormrk   r#   Údropout_path2Únorm_sublayer2Ú
Sequentialr%   Úffn_dimr   r	   Úactivation_functionÚ
ff_dropoutÚffÚdropout_path3Únorm_sublayer3Úpre_normrm   s     €r,   r!   zPatchTSTEncoderLayer.__init__É  sŽ  ø€ Ü‰ÑÔà!'×!9Ñ!9ˆÔä*Ø—n‘nØ×0Ñ0Ø×,Ñ,ô
ˆŒð AG×@SÑ@SÐVWÒ@WœRŸZ™Z¨×(;Ñ(;Ô<Ô]_×]hÑ]hÓ]jˆÔØ×Ñ˜{Ò*Ü"3°FÓ";ˆDÕØ×Ñ Ò,Ü"$§,¡,¨v¯~©~À6Ç?Á?Ô"SˆDÕä × 0Ñ 0Ð1Ð1UÐVÓWÐWð ×!Ò!ØDJ×DWÑDWÐZ[ÒD[¤§¡¨F×,?Ñ,?Ô!@Ôac×alÑalÓanˆDÔØ×Ñ ;Ò.Ü&7¸Ó&?�Õ#Ø×!Ñ! [Ò0Ü&(§l¡l°6·>±>ÀvÇÁÔ&W�Õ#ä  F×$4Ñ$4Ð#5Ð5YÐ!ZÓ[Ð[ô —-‘-Ü�I‰I�f—n‘n f§n¡n¸6¿;¹;ÔGÜ�F×.Ñ.Ñ/Ó1Ø-3×->Ñ->ÀÒ-BŒB�J‰J�v×(Ñ(Ô)ÌÏÉËÜ�I‰I�f—n‘n f§n¡n¸6¿;¹;ÔGó	
ˆŒð AG×@SÑ@SÐVWÒ@WœRŸZ™Z¨×(;Ñ(;Ô<Ô]_×]hÑ]hÓ]jˆÔØ×Ñ˜{Ò*Ü"3°FÓ";ˆDÕØ×Ñ Ò,Ü"$§,¡,¨v¯~©~À6Ç?Á?Ô"SˆDÕä × 0Ñ 0Ð1Ð1UÐVÓWÐWàŸ™ˆ�r-   Úhidden_stater<   c                 óØ  — |j                   \  }}}}|j                  ||z  ||«      }| j                  r;| j                  | j	                  |«      |¬«      \  }}}	|| j                  |«      z   }n:| j                  ||¬«      \  }}}	| j	                  || j                  |«      z   «      }|j                  ||||«      }| j                  rë|j                  dd«      j                  «       }|j                  ||z  ||«      }| j                  r;| j                  | j                  |«      |¬«      \  }}
}	|| j                  |«      z   }n:| j                  ||¬«      \  }}
}	| j                  || j                  |«      z   «      }|j                  ||||«      }|j                  dd«      j                  «       }|j                  ||z  ||«      }| j                  r3|| j                  | j                  | j                  |«      «      «      z   }n2| j                  || j                  | j                  |«      «      z   «      }|j                  ||||«      }|f}|r|| j                  r|
fn|fz  }|S )a¯  
        Parameters:
            hidden_state (`torch.Tensor` of shape `(batch_size, num_channels, sequence_length, d_model)`, *required*):
                Past values of the time series
            output_attentions (`bool`, *optional*):
                Whether or not to return the output attention of all layers
        Return:
            `torch.Tensor` of shape `(batch_size, num_channels, sequence_length, d_model)`

        )r7   r<   r2   r   )rE   r3   rÜ   rË   rÑ   rÏ   rH   rÈ   r4   r5   rÔ   rÓ   rÚ   rÙ   rÛ   )r*   rÝ   r<   r‚   Únum_input_channelsr„   rj   rW   rT   rN   Úchannel_attn_weightsÚoutputss               r,   rX   zPatchTSTEncoderLayer.forwardú  s®  € ð DP×CUÑCUÑ@ˆ
Ð&¨¸ð $×(Ñ(¨Ð6HÑ)HÈ/Ð[bÓcˆà�=Š=à+/¯>©>Ø"×1Ñ1°,Ó?ÐSdð ,:ó ,Ñ(ˆK˜ qð (¨$×*<Ñ*<¸[Ó*IÑI‰Lð ,0¯>©>Ø*Ð>Oð ,:ó ,Ñ(ˆK˜ qð  ×.Ñ.¨|¸d×>PÑ>PÐQ\Ó>]Ñ/]Ó^ˆLð $×+Ñ+¨JÐ8JÈOÐ]dÓeˆð ×!Ò!à'×1Ñ1°!°QÓ7×BÑBÓDˆLà'×,Ñ,¨Z¸/Ñ-IÐK]Ð_fÓgˆLØ�}Š}à7;·~±~Ø"&×"5Ñ"5°lÓ"CÐWhð 8Fó 8Ñ4�Ð1°1ð  ,¨d×.@Ñ.@ÀÓ.MÑM‘ð 8<·~±~Ø".ÐBSð 8Fó 8Ñ4�Ð1°1ð  $×2Ñ2°<À$×BTÑBTÐU`ÓBaÑ3aÓb�ð (×/Ñ/°
¸OÐM_ÐahÓiˆLà'×1Ñ1°!°QÓ7×BÑBÓDˆLð $×(Ñ(¨Ð6HÑ)HÈ/Ð[bÓcˆØ�=Š=ð (¨$×*<Ñ*<¸T¿W¹WÀT×EXÑEXÐYeÓEfÓ=gÓ*hÑh‰Lð  ×.Ñ.¨|¸d×>PÑ>PÐQU×QXÑQXÐYeÓQfÓ>gÑ/gÓhˆLð $×+Ñ+¨JÐ8JÈOÐ]dÓeˆà�/ˆÙØ¸t×?UÒ?U˜Ð&:Ñ;Ð\hÐ[jÑjˆGàˆr-   r»   )rZ   r[   r\   r]   r   r!   rF   ra   r   r`   rX   rb   rc   s   @r,   rÅ   rÅ   Ä  s3   ø„ ñð/(˜~õ /(ñbQ E§L¡Lð QÀXÈdÁ^÷ Qr-   rÅ   c                   ó*   — e Zd ZeZdZdZdZd„ Zdd„Z	y)ÚPatchTSTPreTrainedModelÚmodelr±   Fc                 ó  — t        |t        «      rˆ| j                  j                  r+t        j
                  j                  |j                  d¬«       | j                  j                  dk(  r-t        j
                  j                  |j                  dd¬«       yyt        |t        j                  «      rJ|j                  j                  j                  «        |j                  j                  j                  d«       yt        |t         «      r^|j"                  j                  j                  j                  «        |j"                  j                  j                  j                  d«       yt        |t        j$                  t        j&                  f«      rm|j                  j                  j                  d| j                  j(                  ¬«       |j                  �%|j                  j                  j                  «        yyy)	z$
        Initialize weights
        g{®Gáz”?)ÚstdrÀ   rY   gš™™™™™¹?)Úmeanræ   ç      ð?N)r”   ÚPatchTSTPositionalEncodingr   Úuse_cls_tokenr   ÚinitÚnormal_Ú	cls_tokenÚpositional_encoding_typeÚposition_encrÒ   r   ÚdataÚzero_ÚweightÚfill_re   rl   r%   ÚConv1dÚinit_std)r*   Úmodules     r,   Ú_init_weightsz%PatchTSTPreTrainedModel._init_weightsT  sW  € ô �fÔ8Ô9à�{‰{×(Ò(Ü—‘—‘ × 0Ñ 0°d�Ô;à�{‰{×3Ñ3°xÒ?Ü—‘—‘ × 3Ñ 3¸#À3�ÕGð @ä˜¤§¡Ô-Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)Ü˜Ô 1Ô2Ø×Ñ×!Ñ!×&Ñ&×,Ñ,Ô.Ø×Ñ×#Ñ#×(Ñ(×.Ñ.¨sÕ3Ü˜¤§¡¬B¯I©IÐ 6Ô7Ø�M‰M×Ñ×&Ñ&¨C°T·[±[×5IÑ5IÐ&ÔJØ�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ð 8r-   c                 ó4   — t        |t        «      r||_        y y r»   )r”   ÚPatchTSTEncoderÚgradient_checkpointing)r*   rö   Úvalues      r,   Ú_set_gradient_checkpointingz3PatchTSTPreTrainedModel._set_gradient_checkpointingj  s   € Ü�fœÔ0Ø,1ˆFÕ)ð 1r-   N)F)
rZ   r[   r\   r   Úconfig_classÚbase_model_prefixÚmain_input_nameÚsupports_gradient_checkpointingr÷   rü   r�   r-   r,   rã   rã   N  s"   „ Ø!€LØÐØ#€OØ&+Ð#ò)ô,2r-   rã   c                   óD   ‡ — e Zd Zdefˆ fd„Zdej                  fd„Zˆ xZS )ÚPatchTSTEmbeddingr   c                 óÊ  •— t         ‰| �  «        |j                  | _        |j                  | _        | j                  r0t	        j
                  |j                  |j                  «      | _        y t	        j                  «       | _        t        |j                  «      D ]E  }| j                  j                  t	        j
                  |j                  |j                  «      «       ŒG y r»   )r    r!   rß   Úshare_embeddingr   r%   rŸ   rj   Úinput_embeddingÚ
ModuleListÚranger˜   )r*   r   rN   r+   s      €r,   r!   zPatchTSTEmbedding.__init__p  s£   ø€ Ü‰ÑÔØ"(×";Ñ";ˆÔØ%×5Ñ5ˆÔà×ÒÜ#%§9¡9¨V×-@Ñ-@À&Ç.Á.Ó#QˆDÕ ä#%§=¡=£?ˆDÔ Ü˜6×4Ñ4Ó5ò \�Ø×$Ñ$×+Ñ+¬B¯I©I°f×6IÑ6IÈ6Ï>É>Ó,ZÕ[ñ\r-   r¾   c                 ó`  — |j                   d   }|| j                  k7  rt        d| j                  › d|› d�«      ‚| j                  r| j	                  |«      }|S t        |«      D �cg c]$  } | j                  |   |dd…|dd…dd…f   «      ‘Œ& }}t        j                  |d¬«      }|S c c}w )a%  
        Parameters:
            patch_input (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`, *required*):
                Patch input for embedding
        return:
            `torch.Tensor` of shape `(batch_size, num_channels, num_patches, d_model)`
        r   z&The defined number of input channels (zQ) in the config has to be the same as the number of channels in the batch input (rª   Nr@   )rE   rß   r#   r  r  r  rF   Ústack)r*   r¾   rß   Ú
embeddingsÚis        r,   rX   zPatchTSTEmbedding.forward|  sÐ   € ð )×.Ñ.¨qÑ1ÐØ ×!8Ñ!8Ò8ÜØ8¸×9PÑ9PÐ8Qð RTØTfÐSgÐghðjóð ð ×ÒØ×-Ñ-¨kÓ:ˆJð Ðô UZÐZlÓTmÖnÈqÐ1˜$×.Ñ.¨qÑ1°+ºaÀÂAÂq¸jÑ2IÕJÐnˆJÐnÜŸ™ Z°QÔ7ˆJØÐùò os   Á')B+©	rZ   r[   r\   r   r!   rF   ra   rX   rb   rc   s   @r,   r  r  o  s!   ø„ ð
\˜~õ 
\ð 5§<¡<÷ r-   r  c                   ó~   ‡ — e Zd ZdZdedefˆ fd„Zedededej                  fd„«       Z
dej                  fd„Zˆ xZS )	ré   z'
    Class for positional encoding
    r   r®   c                 óÄ  •— t         ‰| �  «        |j                  | _        |j                  | _        |j                  r?t	        j
                  t        j                  ddd|j                  «      «      | _	        |dz  }| j                  ||«      | _        |j                  dkD  r%t	        j                  |j                  «      | _        y t	        j                  «       | _        y )Nr   r   )r    r!   rê   rß   r   Ú	ParameterrF   r•   rj   rí   Ú_init_perï   Úpositional_dropoutrÍ   rÎ   ©r*   r   r®   r+   s      €r,   r!   z#PatchTSTPositionalEncoding.__init__˜  s±   ø€ Ü‰ÑÔØ#×1Ñ1ˆÔØ"(×";Ñ";ˆÔØ×ÒäŸ\™\¬%¯+©+°a¸¸A¸v¿~¹~Ó*NÓOˆDŒNØ˜1ÑˆKà ŸM™M¨&°+Ó>ˆÔð 6<×5NÑ5NÐQRÒ5RŒB�J‰J�v×0Ñ0Ó1ð 	ÕÜXZ×XcÑXcÓXeð 	Õr-   r=   c                 ó$  — | j                   dk(  r7t        j                  t        j                  || j
                  «      d¬«      }|S | j                   dk(  �r#t        j                  || j
                  «      }t        j                  d|«      j                  d«      }t        j                  t        j                  d| j
                  d«      t        j                  d«      | j
                  z   z  «      }t        j                  ||z  «      |d d …dd d…f<   t        j                  ||z  «      |d d …dd d…f<   ||j                  «       z
  }||j                  «       d	z  z  }t        j                  |d
¬«      }|S t!        | j                   › d�«      ‚)NrÀ   T©Úrequires_gradÚsincosr   r   r2   g     ˆÃ@é
   FzN is not a valid positional encoder. Available types are 'random' and 'sincos'.)rî   r   r  rF   Úrandnrj   r•   Úaranger€   ÚexpÚmathÚlogÚsinÚcosrç   ræ   r#   )r   r®   rï   ÚpositionÚdiv_terms        r,   r  z#PatchTSTPositionalEncoding._init_pe§  sd  € ð ×*Ñ*¨hÒ6ÜŸ<™<¬¯©°KÀÇÁÓ(PÐ`dÔeˆLð Ðð ×,Ñ,°Ó8Ü Ÿ;™; {°F·N±NÓCˆLÜ—|‘| A {Ó3×=Ñ=¸aÓ@ˆHÜ—y‘y¤§¡¨a°·±ÀÓ!CÌÏÉÐQXÓHYÐ\b×\jÑ\jÑHjÐFkÑ!kÓlˆHÜ$)§I¡I¨h¸Ñ.AÓ$BˆLš˜A˜D˜q˜D˜Ñ!Ü$)§I¡I¨h¸Ñ.AÓ$BˆLš˜A˜D˜q˜D˜Ñ!Ø'¨,×*;Ñ*;Ó*=Ñ=ˆLØ'¨<×+;Ñ+;Ó+=ÀÑ+BÑCˆLÜŸ<™<¨ÀEÔJˆLð
 Ðô Ø×2Ñ2Ð3ð  4Bð  Cóð r-   r¾   c                 óx  — | j                   r�| j                  || j                  dd …d d …f   z   «      }| j                  | j                  d d…d d …f   z   }|j	                  |j
                  d   | j                  dd«      }t        j                  ||fd¬«      }|S | j                  || j                  z   «      }|S )Nr   r   r?   r2   r@   )	rê   r  rï   rí   ÚexpandrE   rß   rF   rG   )r*   r¾   rí   Ú
cls_tokensrÝ   s        r,   rX   z"PatchTSTPositionalEncoding.forward»  s¿   € Ø×Òà×1Ñ1°+À×@QÑ@QÐRSÑRTÒVWÐRWÑ@XÑ2XÓYˆKàŸ™¨×):Ñ):¸2¸A¸2ºq¸5Ñ)AÑAˆIà"×)Ñ)¨+×*;Ñ*;¸AÑ*>À×@WÑ@WÐY[Ð]_Ó`ˆJä Ÿ9™9 j°+Ð%>ÀAÔFˆLð Ðð  ×2Ñ2°;À×ARÑARÑ3RÓSˆLØÐr-   )rZ   r[   r\   r]   r   r^   r!   Ústaticmethodr   r  r  rF   ra   rX   rb   rc   s   @r,   ré   ré   “  sX   ø„ ñð
˜~ð 
¸Cõ 
ð ð˜ð °cð ¸b¿l¹lò ó ðð& 5§<¡<÷ r-   ré   c            	       ój   ‡ — e Zd ZdZdedefˆ fd„Z	 	 d
dej                  de	e
   de	e
   defd	„Zˆ xZS )rù   z
    PatchTST Encoder
    r   r®   c                 ó&  •— t         ‰| �  |«       d| _        t        |«      | _        t        ||«      | _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        | j                  «        y c c}w )NF)r    r!   rú   r  Úembedderré   Úpositional_encoderr   r  r  Únum_hidden_layersrÅ   ÚlayersÚ	post_init)r*   r   r®   r  r+   s       €r,   r!   zPatchTSTEncoder.__init__Ð  st   ø€ Ü‰Ñ˜Ô Ø&+ˆÔ#ô *¨&Ó1ˆŒä"<¸VÀ[Ó"QˆÔä—m‘mÌ5ÐQW×QiÑQiÓKjÖ$kÀaÔ%9¸&Õ%AÒ$kÓlˆŒð 	�‰Õùò %ls   ÁBr¾   Úoutput_hidden_statesr<   r=   c                 óJ  — |�|n| j                   j                  }|�|n| j                   j                  }| j                  |«      }| j	                  |«      }|rdnd}|rdnd}| j
                  D ]%  }|r||fz   } |||¬«      }|d   }|sŒ||d   fz   }Œ' t        |||¬«      S )a²  
        Parameters:
            patch_input (`torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`, *required*):
                Past values of the time series
            output_hidden_states (bool, optional): Indicates if hidden states should be outputted.
            output_attentions (bool, optional): Indicates if attentions should be outputted.

        return:
            `BaseModelOutput`
        Nr�   )rÝ   r<   r   r   )Úlast_hidden_stater7   Ú
attentions)r   r<   r,  r'  r(  r*  r
   )	r*   r¾   r,  r<   rÝ   Úencoder_statesÚall_attentionsÚencoder_layerÚlayer_outputss	            r,   rX   zPatchTSTEncoder.forwardÞ  sÐ   € ð  2CÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð
 —m‘m KÓ0ˆà×.Ñ.¨{Ó;ˆá3™¸ˆÙ0™°dˆØ!Ÿ[™[ò 
	FˆMÙ#Ø!/°<°/Ñ!A�á)°|ÐWhÔiˆMð )¨Ñ+ˆLâ Ø!/°=ÀÑ3CÐ2EÑ!E‘ð
	Fô °È^ÐhvÔwÐwr-   ©NN)rZ   r[   r\   r]   r   r^   r!   rF   ra   r   r`   r
   rX   rb   rc   s   @r,   rù   rù   Ë  s_   ø„ ñð˜~ð ¸Cõ ð" 04Ø,0ñ	(xà—\‘\ð(xð ' t™nð(xð $ D™>ð	(xð
 
÷(xr-   rù   aM  
    This model inherits from [`PreTrainedModel`]. Check the superclass documentation for the generic methods the
    library implements for all its model (such as downloading or saving, resizing the input embeddings, pruning heads
    etc.)

    This model is also a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) subclass.
    Use it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage
    and behavior.

    Parameters:
        config ([`PatchTSTConfig`]):
            Model configuration class with all the parameters of the model. Initializing with a config file does not
            load the weights associated with the model, only the configuration. Check out the
            [`~PreTrainedModel.from_pretrained`] method to load the model weights.
c                   ó6  — e Zd ZU dZdZeej                     ed<   dZ	ee
ej                        ed<   dZee
ej                        ed<   dZeej                     ed<   dZeej                     ed<   dZeej                     ed<   dZeej                     ed	<   y)
ÚPatchTSTModelOutputaÉ  
    Base class for model's outputs, with potential hidden states.

    Parameters:
        last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, patch_length)`):
            Sequence of hidden-states at the output of the last layer of the model.
        hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
            one for the output of each layer) of shape `(batch_size, num_channels, height, width)`. Hidden-states of
            the model at the output of each layer plus the optional initial embedding outputs.
        mask: (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches)`, *optional*)
            Bool masked tensor indicating which patches are masked
        loc: (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*)
            Mean of the input data (batch_size, sequence_length, num_channels) over the sequence_length
        scale: (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*)
            Std of the input data (batch_size, sequence_length, num_channels) over the sequence_length
        patch_input (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, patch_length)`):
            Patched input to the Transformer
    Nr.  r7   r/  rˆ   ÚlocÚscaler¾   )rZ   r[   r\   r]   r.  r   rF   ÚFloatTensorÚ__annotations__r7   r   r/  rˆ   r7  r8  r¾   r�   r-   r,   r6  r6    s§   … ñð( 6:Ð�x × 1Ñ 1Ñ2Ó9Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ó9Ø(,€Dˆ(�5×$Ñ$Ñ
%Ó,Ø'+€Cˆ�%×#Ñ#Ñ	$Ó+Ø)-€Eˆ8�E×%Ñ%Ñ&Ó-Ø/3€K�˜%×+Ñ+Ñ,Ô3r-   r6  c                   ó¾   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eeej                        ed<   dZeeej                        ed<   y)ÚPatchTSTForPretrainingOutputaß  
    Output type of [`PatchTSTForPretraining`].

    Parameters:
        loss (*optional*, returned when `labels` is provided, `torch.FloatTensor` of shape `(1,)`):
            MSE loss.
        prediction_outputs (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
            Prediction outputs of the time series modeling heads.
        hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
            shape `(batch_size, sequence_length, hidden_size)`.

            Hidden-states of the model at the output of each layer plus the initial embedding outputs.
        attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
            sequence_length)`.

            Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
            heads.
    NÚlossÚprediction_outputr7   r/  )rZ   r[   r\   r]   r=  r   rF   r9  r:  r>  r7   r   r/  r�   r-   r,   r<  r<  9  sh   … ñð* )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø59Ð�x × 1Ñ 1Ñ2Ó9Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ô9r-   r<  c                   ó¾   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eeej                        ed<   dZeeej                        ed<   y)ÚPatchTSTForRegressionOutputaÇ  
    Output type of [`PatchTSTForRegression`].

    Parameters:
        loss (*optional*, returned when `labels` is provided, `torch.FloatTensor` of shape `(1,)`):
            MSE loss.
        regression_outputs (`torch.FloatTensor` of shape `(batch_size, num_targets)`):
            Regression outputs of the time series modeling heads.
        hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
            shape `(batch_size, sequence_length, hidden_size)`.

            Hidden-states of the model at the output of each layer plus the initial embedding outputs.
        attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
            sequence_length)`.

            Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
            heads.
    Nr=  Úregression_outputsr7   r/  )rZ   r[   r\   r]   r=  r   rF   r9  r:  rA  r7   r   r/  r�   r-   r,   r@  r@  V  sh   … ñð* )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø6:Ð˜ ×!2Ñ!2Ñ3Ó:Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ô9r-   r@  c                   ó  — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eeej                        ed<   dZeeej                        ed<   dZeej                     ed<   dZeej                     ed<   y)	ÚPatchTSTForPredictionOutputaR  
    Output type of [`PatchTSTForPrediction`].

    Parameters:
        loss (*optional*, returned when `labels` is provided, `torch.FloatTensor` of shape `(1,)`):
            MSE loss.
        prediction_outputs (`torch.FloatTensor` of shape `(batch_size, prediction_length, -1)`):
            Prediction outputs of the time series modeling heads.
        hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
            shape `(batch_size, sequence_length, hidden_size)`.

            Hidden-states of the model at the output of each layer plus the initial embedding outputs.
        attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
            sequence_length)`.

            Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
            heads.
        loc: (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*)
            Mean of the input data (batch_size, sequence_length, num_channels) over the sequence_length
        scale: (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*)
            Std of the input data (batch_size, sequence_length, num_channels) over the sequence_length
    Nr=  Úprediction_outputsr7   r/  r7  r8  )rZ   r[   r\   r]   r=  r   rF   r9  r:  rD  r7   r   r/  r7  r8  r�   r-   r,   rC  rC  s  s’   … ñð2 )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø6:Ð˜ ×!2Ñ!2Ñ3Ó:Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ó9Ø'+€Cˆ�%×#Ñ#Ñ	$Ó+Ø)-€Eˆ8�E×%Ñ%Ñ&Ô-r-   rC  c                   ó¾   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eeej                        ed<   dZeeej                        ed<   y)ÚPatchTSTForClassificationOutputaR  
    Output type of [`PatchTSTForClassification`].

    Parameters:
        loss (*optional*, returned when `labels` is provided, `torch.FloatTensor` of shape `(1,)`):
            Total loss as the sum of the masked language modeling loss and the next sequence prediction
            (classification) loss.
        prediction_logits (`torch.FloatTensor` of shape `(batch_size, num_targets)`):
            Prediction scores of the PatchTST modeling head (scores before SoftMax).
        hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
            shape `(batch_size, sequence_length, hidden_size)`.

            Hidden-states of the model at the output of each layer plus the initial embedding outputs.
        attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`):
            Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length,
            sequence_length)`.

            Attentions weights after the attention softmax, used to compute the weighted average in the self-attention
            heads.
    Nr=  Úprediction_logitsr7   r/  )rZ   r[   r\   r]   r=  r   rF   r9  r:  rG  r7   r   r/  r�   r-   r,   rF  rF  –  sh   … ñð, )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø59Ð�x × 1Ñ 1Ñ2Ó9Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ó<Ø59€J�˜˜u×0Ñ0Ñ1Ñ2Ô9r-   rF  c                   ó:   — e Zd ZU dZdZeej                     ed<   y)ÚSamplePatchTSTOutputa!  
    Base class for time series model's predictions outputs that contains the sampled values from the chosen
    distribution.

    Parameters:
        sequences `(batch_size, num_samples, prediction_length, num_targets)`):
                Sampled values from the chosen distribution.
    NÚ	sequences)	rZ   r[   r\   r]   rJ  r   rF   r9  r:  r�   r-   r,   rI  rI  ´  s   … ñð .2€Iˆx˜×)Ñ)Ñ*Ô1r-   rI  ÚinputÚtargetr=   c                 ó&   — | j                  |«       S )zc
    Computes the negative log likelihood loss from input distribution with respect to target.
    )Úlog_prob)rK  rL  s     r,   ÚnllrO  Ã  s   € ð �N‰N˜6Ó"Ð"Ð"r-   Úinput_tensorÚweightsc                 óP  — |�“t        j                  |dk7  | |z  t        j                  | «      «      }t        j                  |r|j	                  |¬«      n|j	                  «       d¬«      }|r|j	                  |¬«      |z  S |j	                  «       |z  S | j                  |¬«      S )aj  
    Computes the weighted average of a given tensor across a given `dim`, masking values associated with weight zero,
    meaning instead of `nan * 0 = nan` you will get `0 * 0 = 0`.

    Args:
        input_tensor (`torch.FloatTensor`):
            Input tensor, of which the average must be computed.
        weights (`torch.FloatTensor`, *optional*):
            Weights tensor, of the same shape as `input_tensor`.
        dim (`int`, *optional*):
            The dim along which to average `input_tensor`.

    Returns:
        `torch.FloatTensor`: The tensor with values averaged along the specified `dim`.
    r   r@   rè   ©Úmin)rF   ÚwhereÚ
zeros_likeÚclampr–   rç   )rP  rQ  rA   Úweighted_tensorÚsum_weightss        r,   Úweighted_averagerZ  Ë  s›   € ð  ÐÜŸ+™+ g°¡l°LÀ7Ñ4JÌE×L\ÑL\Ð]iÓLjÓkˆÜ—k‘k¹# '§+¡+°# +Ô"6À7Ç;Á;Ã=ÐVYÔZˆÙ03�×#Ñ#¨Ð#Ó,ÐR]Ñ]Ð]¸×9LÑ9LÓ9NÐR]Ñ]Ð]à× Ñ  SÐ Ó)Ð)r-   c            	       ó¬   ‡ — e Zd ZdZdefˆ fd„Zdej                  dej                  deej                  ej                  ej                  f   fd„Z	ˆ xZ
S )ÚPatchTSTStdScalerz½
    Standardize features by calculating the mean and scaling along the first dimension, and then normalizes it by
    subtracting from the mean and dividing by the standard deviation.
    r   c                 óè   •— t         ‰| �  «        t        |d«      r|j                  nd| _        t        |d«      r|j
                  nd| _        t        |d«      r|j                  | _        y d| _        y )NÚscaling_dimr   ÚkeepdimTÚminimum_scalegñhãˆµøä>)r    r!   Úhasattrr^  rA   r_  r`  rm   s     €r,   r!   zPatchTSTStdScaler.__init__ê  s[   ø€ Ü‰ÑÔÜ)0°¸Ô)G�6×%Ò%ÈQˆŒÜ)0°¸Ô)C�v—~’~ÈˆŒÜ5<¸VÀ_Ô5U˜V×1Ñ1ˆÕÐ[_ˆÕr-   rð   Úobserved_indicatorr=   c                 óŒ  — |j                  | j                  | j                  ¬«      }|j                  d«      }||z  j                  | j                  | j                  ¬«      |z  }||z
  |z  dz  j                  | j                  | j                  ¬«      |z  }t	        j
                  || j                  z   «      }||z
  |z  ||fS )áC  
        Parameters:
            data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`):
                input for Batch norm calculation
            observed_indicator (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`):
                Calculating the scale on the observed indicator.
        Returns:
            tuple of `torch.Tensor` of shapes
                (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`,
                `(batch_size, 1, num_input_channels)`)
        ©r_  rè   r2   )r–   rA   r_  Ú	clamp_minrF   Úsqrtr`  )r*   rð   rb  Údenominatorr7  Úvariancer8  s          r,   rX   zPatchTSTStdScaler.forwardð  s¾   € ð )×,Ñ,¨T¯X©X¸t¿|¹|Ð,ÓLˆØ!×+Ñ+¨CÓ0ˆØÐ(Ñ(×-Ñ-¨d¯h©hÀÇÁÐ-ÓMÐP[Ñ[ˆà˜S‘jÐ$6Ñ6¸1Ñ<×AÑAÀ$Ç(Á(ÐTX×T`ÑT`ÐAÓaÐdoÑoˆÜ—
‘
˜8 d×&8Ñ&8Ñ8Ó9ˆØ�s‘
˜eÑ# S¨%Ð/Ð/r-   ©rZ   r[   r\   r]   r   r!   rF   ra   r   rX   rb   rc   s   @r,   r\  r\  ä  sS   ø„ ñð
`˜~õ `ð0Ø—L‘Lð0Ø6;·l±lð0à	ˆu�|‰|˜UŸ\™\¨5¯<©<Ð7Ñ	8÷0r-   r\  c            	       ó¬   ‡ — e Zd ZdZdefˆ fd„Zdej                  dej                  deej                  ej                  ej                  f   fd„Z	ˆ xZ
S )ÚPatchTSTMeanScalerzŠ
    Computes a scaling factor as the weighted average absolute value along the first dimension, and scales the data
    accordingly.
    r   c                 ó&  •— t         ‰| �  «        t        |d«      r|j                  nd| _        t        |d«      r|j
                  nd| _        t        |d«      r|j                  nd| _        t        |d«      r|j                  | _        y d | _        y )Nr^  r   r_  Tr`  ç»½×Ùß|Û=Údefault_scale)r    r!   ra  r^  rA   r_  r`  ro  rm   s     €r,   r!   zPatchTSTMeanScaler.__init__  su   ø€ Ü‰ÑÔÜ)0°¸Ô)G�6×%Ò%ÈQˆŒÜ)0°¸Ô)C�v—~’~ÈˆŒÜ5<¸VÀ_Ô5U˜V×1Ò1Ð[`ˆÔÜ5<¸VÀ_Ô5U˜V×1Ñ1ˆÕÐ[_ˆÕr-   rð   rb  r=   c                 óÊ  — ||z  j                  «       j                  | j                  d¬«      }|j                  | j                  d¬«      }|t        j                  |d¬«      z  }| j
                  €Q|j                  d¬«      }t        j                  |j                  d«      d¬«      }t        j                  ||z  «      }n"| j
                  t        j                  |«      z  }t        j                  |dkD  ||«      }t        j                  || j                  ¬«      }||z  }	| j                  s|j                  | j                  ¬«      }|	t        j                  |«      |fS )rd  Tre  r   rS  r   r@   )Úabsr–   rA   rF   rW  ro  ÚsqueezeÚ	ones_likerU  r`  r_  rV  )
r*   rð   rb  Úts_sumÚnum_observedr8  Ú	batch_sumÚbatch_observationsro  Úscaled_datas
             r,   rX   zPatchTSTMeanScaler.forward  s.  € ð Ð+Ñ+×0Ñ0Ó2×6Ñ6°t·x±xÈÐ6ÓNˆØ)×-Ñ-¨d¯h©hÀÐ-ÓEˆàœŸ™ \°qÔ9Ñ9ˆð ×ÑÐ%ØŸ
™
 q˜
Ó)ˆIÜ!&§¡¨\×-=Ñ-=¸aÓ-@ÀaÔ!HÐÜ!ŸM™M¨)Ð6HÑ*HÓI‰Mà ×.Ñ.´·±ÀÓ1GÑGˆMô —‘˜L¨1Ñ,¨e°]ÓCˆô —‘˜E t×'9Ñ'9Ô:ˆØ˜U‘lˆà�|Š|Ø—M‘M d§h¡h�MÓ/ˆEàœE×,Ñ,¨UÓ3°UÐ:Ð:r-   rj  rc   s   @r,   rl  rl    sS   ø„ ñð
`˜~õ `ð&;Ø—L‘Lð&;Ø6;·l±lð&;à	ˆu�|‰|˜UŸ\™\¨5¯<©<Ð7Ñ	8÷&;r-   rl  c            
       ó¶   ‡ — e Zd ZdZdefˆ fd„Z	 ddej                  deej                     de	ej                  ej                  ej                  f   fd„Z
ˆ xZS )	ÚPatchTSTNOPScalerz|
    Assigns a scaling factor equal to 1 along the first dimension, and therefore applies no scaling to the input data.
    r   c                 óª   •— t         ‰| �  «        t        |d«      r|j                  nd| _        t        |d«      r|j
                  | _        y d| _        y )Nr^  r   r_  T)r    r!   ra  r^  rA   r_  rm   s     €r,   r!   zPatchTSTNOPScaler.__init__D  s@   ø€ Ü‰ÑÔÜ)0°¸Ô)G�6×%Ò%ÈQˆŒÜ)0°¸Ô)C�v—~‘~ˆ�Èˆ�r-   rð   rb  r=   c                 óü   — t        j                  |d¬«      j                  | j                  | j                  ¬«      }t        j
                  |d¬«      j                  | j                  | j                  ¬«      }|||fS )a�  
        Parameters:
            data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`):
                input for Batch norm calculation
        Returns:
            tuple of `torch.Tensor` of shapes
                (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`,
                `(batch_size, 1, num_input_channels)`)
        Fr  )rA   r_  )rF   rs  rç   rA   r_  rV  )r*   rð   rb  r8  r7  s        r,   rX   zPatchTSTNOPScaler.forwardI  si   € ô —‘ °EÔ:×?Ñ?ÀDÇHÁHÐVZ×VbÑVbÐ?ÓcˆÜ×Ñ˜t°5Ô9×>Ñ>À4Ç8Á8ÐUY×UaÑUaÐ>ÓbˆØ�S˜%ÐÐr-   r»   )rZ   r[   r\   r]   r   r!   rF   ra   r   r   rX   rb   rc   s   @r,   rz  rz  ?  s_   ø„ ñðN˜~õ Nð PTñ Ø—L‘Lð Ø6>¸u¿|¹|Ñ6Lð à	ˆu�|‰|˜UŸ\™\¨5¯<©<Ð7Ñ	8÷ r-   rz  c            	       ó¨   ‡ — e Zd Zdefˆ fd„Zdej                  dej                  deej                  ej                  ej                  f   fd„Zˆ xZ	S )ÚPatchTSTScalerr   c                 óÞ   •— t         ‰| �  «        |j                  dk(  s|j                  du rt        |«      | _        y |j                  dk(  rt        |«      | _        y t        |«      | _        y )Nrç   Træ   )r    r!   r$   rl  Úscalerr\  rz  rm   s     €r,   r!   zPatchTSTScaler.__init__[  sU   ø€ Ü‰ÑÔØ�>‰>˜VÒ# v§~¡~¸Ñ'=Ü,¨VÓ4ˆD�KØ�^‰^˜uÒ$Ü+¨FÓ3ˆD�Kä+¨FÓ3ˆD�Kr-   rð   rb  r=   c                 ó8   — | j                  ||«      \  }}}|||fS )a>  
        Parameters:
            data (`torch.Tensor` of shape `(batch_size, sequence_length, num_input_channels)`):
                Input for scaler calculation
            observed_indicator (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`):
                Calculating the scale on the observed indicator.
        Returns:
            tuple of `torch.Tensor` of shapes
                (`(batch_size, sequence_length, num_input_channels)`,`(batch_size, 1, num_input_channels)`,
                `(batch_size, 1, um_input_channels)`)
        )r€  )r*   rð   rb  r7  r8  s        r,   rX   zPatchTSTScaler.forwardd  s)   € ð  Ÿ;™; tÐ-?Ó@Ñˆˆc�5Ø�S˜%ÐÐr-   )
rZ   r[   r\   r   r!   rF   ra   r   rX   rb   rc   s   @r,   r~  r~  Z  sL   ø„ ð4˜~õ 4ð Ø—L‘Lð Ø6;·l±lð à	ˆu�|‰|˜UŸ\™\¨5¯<©<Ð7Ñ	8÷ r-   r~  zOThe bare PatchTST Model outputting raw hidden-states without any specific head.c                   ó¸   ‡ — e Zd Zdefˆ fd„Z	 	 	 	 	 ddej                  deej                     deej                     dee   dee   dee   d	e	e
ef   fd
„Zˆ xZS )ÚPatchTSTModelr   c                 ób  •— t         ‰| �  |«       t        |«      | _        t	        |«      | _        |j                  | _        | j
                  j                  }| j                  rt        |«      | _	        nt        j                  «       | _	        t        ||¬«      | _        | j                  «        y )N)r®   )r    r!   r~  r€  r¨   Ú
patchifierÚdo_mask_inputr®   r¹   Úmaskingr   rÎ   rù   Úencoderr+  r  s      €r,   r!   zPatchTSTModel.__init__{  s�   ø€ Ü‰Ñ˜Ô ä$ VÓ,ˆŒÜ*¨6Ó2ˆŒØ#×1Ñ1ˆÔà—o‘o×1Ñ1ˆà×ÒÜ*¨6Ó2ˆD�LäŸ;™;›=ˆDŒLÜ& v¸;ÔGˆŒð 	�‰Õr-   r±   Úpast_observed_maskÚfuture_valuesr,  r<   Úreturn_dictr=   c           	      óŠ  — |�|n| j                   j                  }|�|n| j                   j                  }|�|n| j                   j                  }|€t	        j
                  |«      }| j                  ||«      \  }}}	| j                  |«      }
| j                  r| j                  |
«      \  }}n| j                  |
«      d}}| j                  |||¬«      }|s>|j                  |j                  |j                  f}||||	|
fz   }t        d„ |D «       «      S t        |j                  |j                  |j                  |||	|
¬«      S )a  
        Parameters:
            past_values (`torch.Tensor` of shape `(bs, sequence_length, num_input_channels)`, *required*):
                Input sequence to the model
            past_observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
                Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
                in `[0, 1]`:

                - 1 for values that are **observed**,
                - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
            future_values (`torch.BoolTensor` of shape `(batch_size, prediction_length, num_input_channels)`, *optional*):
                Future target values associated with the `past_values`
            output_hidden_states (`bool`, *optional*):
                Whether or not to return the hidden states of all layers
            output_attentions (`bool`, *optional*):
                Whether or not to return the output attention of all layers
            return_dict (`bool`, *optional*):
                Whether or not to return a `ModelOutput` instead of a plain tuple.

        Returns:
            `PatchTSTModelOutput` or tuple of `torch.Tensor` (if `return_dict`=False or `config.return_dict`=False)

        Examples:

        ```python
        >>> from huggingface_hub import hf_hub_download
        >>> import torch
        >>> from transformers import PatchTSTModel

        >>> file = hf_hub_download(
        ...     repo_id="hf-internal-testing/etth1-hourly-batch", filename="train-batch.pt", repo_type="dataset"
        ... )
        >>> batch = torch.load(file)

        >>> model = PatchTSTModel.from_pretrained("namctin/patchtst_etth1_pretrain")

        >>> # during training, one provides both past and future values
        >>> outputs = model(
        ...     past_values=batch["past_values"],
        ...     future_values=batch["future_values"],
        ... )

        >>> last_hidden_state = outputs.last_hidden_state
        ```N)r¾   r,  r<   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wr»   r�   )Ú.0Úvs     r,   ú	<genexpr>z(PatchTSTModel.forward.<locals>.<genexpr>Ý  s   è ø€ Ò=˜q¨q©}œÑ=ùs   ‚Š)r.  r7   r/  rˆ   r7  r8  r¾   )r   Úuse_return_dictr<   r,  rF   rs  r€  r…  r†  r‡  rˆ  r.  r7   r/  Útupler6  )r*   r±   r‰  rŠ  r,  r<   r‹  Úscaled_past_valuesr7  r8  Úpatched_valuesÚmasked_valuesrˆ   Úencoder_outputrá   s                  r,   rX   zPatchTSTModel.forward�  sY  € ðl &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ1BÐ1NÑ-ÐTX×T_ÑT_×TqÑTqÐà$8Ð$DÑ È$Ï+É+×JjÑJjð 	ð Ð%Ü!&§¡°Ó!=Ðð *.¯©°[ÐBTÓ)UÑ&Ð˜C ð Ÿ™Ð);Ó<ˆØ×ÒØ"&§,¡,¨~Ó">ÑˆM™4à"&§,¡,¨~Ó">À˜4ˆMàŸ™Ø%Ð<PÐduð &ó 
ˆñ Ø%×7Ñ7¸×9UÑ9UÐWe×WpÑWpÐqˆGØ  s¨E°>Ð BÑBˆGÜÑ= GÔ=Ó=Ð=ä"Ø,×>Ñ>Ø(×6Ñ6Ø%×0Ñ0ØØØØ&ô
ð 	
r-   ©NNNNN)rZ   r[   r\   r   r!   rF   ra   r   r`   r   r   r6  rX   rb   rc   s   @r,   rƒ  rƒ  v  sž   ø„ ð
˜~õ ð* 6:Ø04Ø/3Ø,0Ø&*ñZ
à—\‘\ðZ
ð % U§\¡\Ñ2ðZ
ð   §¡Ñ-ð	Z
ð
 ' t™nðZ
ð $ D™>ðZ
ð ˜d‘^ðZ
ð 
ˆuÐ)Ð)Ñ	*÷Z
r-   rƒ  c                   ó`   ‡ — e Zd ZdZdefˆ fd„Zdej                  dej                  fd„Zˆ xZ	S )ÚPatchTSTMaskPretrainHeadz-
    Pretraining head for mask modelling
    r   c                 ó0  •— t         ‰| �  «        |j                  dkD  rt        j                  |j                  «      nt        j
                  «       | _        t        j                  |j                  |j                  «      | _
        |j                  | _        y ©Nr   )r    r!   Úhead_dropoutr   rÍ   rÎ   r   r%   rj   rŸ   Úlinearrê   rm   s     €r,   r!   z!PatchTSTMaskPretrainHead.__init__ï  sh   ø€ Ü‰ÑÔØ:@×:MÑ:MÐPQÒ:Q”r—z‘z &×"5Ñ"5Ô6ÔWY×WbÑWbÓWdˆŒÜ—i‘i §¡°×0CÑ0CÓDˆŒØ#×1Ñ1ˆÕr-   Ú	embeddingr=   c                 ó€   — | j                  | j                  |«      «      }| j                  r|dd…dd…dd…dd…f   }|S )aÛ  
        Parameters:
            embedding (`torch.Tensor` of shape `(bs, num_channels, num_patches, d_model)` or
                    `(bs, num_channels, num_patches+1, d_model)` if `cls_token` is set to True, *required*):
                Embedding from the model
        Returns:
            `torch.Tensor` of shape `(bs, num_channels, num_patches, d_model)` or
                            `(bs, num_channels, num_patches+1, d_model)` if `cls_token` is set to True

        Nr   )r�  r   rê   )r*   rž  s     r,   rX   z PatchTSTMaskPretrainHead.forwardõ  s>   € ð —K‘K §¡¨YÓ 7Ó8ˆ	Ø×ÒØ!¢!¢Q¨©ªA +Ñ.ˆIØÐr-   rq   rc   s   @r,   r™  r™  ê  s/   ø„ ñð2˜~õ 2ð §¡ð °%·,±,÷ r-   r™  z The PatchTST for pretrain model.c                   ó˜   ‡ — e Zd Zdefˆ fd„Z	 	 	 	 d
dej                  deej                     dee   dee   dee   de	e
ef   fd	„Zˆ xZS )ÚPatchTSTForPretrainingr   c                 ó”   •— t         ‰| �  |«       d|_        t        |¬«      | _        t        |«      | _        | j                  «        y )NT)r   )r    r!   r†  rƒ  rä   r™  Úheadr+  rm   s     €r,   r!   zPatchTSTForPretraining.__init__  s<   ø€ Ü‰Ñ˜Ô à#ˆÔÜ"¨&Ô1ˆŒ
Ü,¨VÓ4ˆŒ	ð 	�‰Õr-   r±   r‰  r,  r<   r‹  r=   c                 óü  — |�|n| j                   j                  }| j                  ||||d¬«      }| j                  |j                  «      }t        j                  d¬«      } |||j                  «      }	|	j                  d¬«      |j                  z  j                  «       |j                  j                  «       dz   z  }
|j                  }|s|f|dd	 z   }|
�|
f|z   }|S |}|S t        |
|||j                  ¬
«      S )aª	  
        Parameters:
            past_values (`torch.Tensor` of shape `(bs, sequence_length, num_input_channels)`, *required*):
                Input sequence to the model
            past_observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
                Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
                in `[0, 1]`:

                - 1 for values that are **observed**,
                - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
            output_hidden_states (`bool`, *optional*):
                Whether or not to return the hidden states of all layers
            output_attentions (`bool`, *optional*):
                Whether or not to return the output attention of all layers
            return_dict (`bool`, *optional*): Whether or not to return a `ModelOutput` instead of a plain tuple.

        Returns:
            `PatchTSTForPretrainingOutput` or tuple of `torch.Tensor` (if `return_dict`=False or
            `config.return_dict`=False)

        Examples:

        ```python
        >>> from huggingface_hub import hf_hub_download
        >>> import torch
        >>> from transformers import PatchTSTConfig, PatchTSTForPretraining

        >>> file = hf_hub_download(
        ...     repo_id="hf-internal-testing/etth1-hourly-batch", filename="train-batch.pt", repo_type="dataset"
        ... )
        >>> batch = torch.load(file)

        >>> # Config for random mask pretraining
        >>> config = PatchTSTConfig(
        ...     num_input_channels=7,
        ...     context_length=512,
        ...     patch_length=12,
        ...     stride=12,
        ...     mask_type='random',
        ...     random_mask_ratio=0.4,
        ...     use_cls_token=True,
        ... )
        >>> # Config for forecast mask pretraining
        >>> config = PatchTSTConfig(
        ...     num_input_channels=7,
        ...     context_length=512,
        ...     patch_length=12,
        ...     stride=12,
        ...     mask_type='forecast',
        ...     num_forecast_mask_patches=5,
        ...     use_cls_token=True,
        ... )
        >>> model = PatchTSTForPretraining(config)

        >>> # during training, one provides both past and future values
        >>> outputs = model(past_values=batch["past_values"])

        >>> loss = outputs.loss
        >>> loss.backward()
        ```T©r±   r‰  r,  r<   r‹  Únone©Ú	reductionr?   r@   rn  r   éüÿÿÿ)r=  r>  r7   r/  )r   r‘  rä   r£  r.  r   ÚMSELossr¾   rç   rˆ   r–   r7   r<  r/  )r*   r±   r‰  r,  r<   r‹  Úmodel_outputÚx_hatr=  Úloss_valÚmasked_lossr0  rá   s                r,   rX   zPatchTSTForPretraining.forward  s  € ðJ &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆð —z‘zØ#Ø1Ø!5Ø/Øð "ó 
ˆð —	‘	˜,×8Ñ8Ó9ˆô �z‰z FÔ+ˆÙ˜˜|×7Ñ7Ó8ˆØ—}‘}¨�}Ó,¨|×/@Ñ/@Ñ@×EÑEÓGÈ<×K\ÑK\×K`ÑK`ÓKbÐejÑKjÑkˆà%×3Ñ3ˆÙØ�h ¨a°Ð!3Ñ3ˆGØ2=Ð2I�{�n wÑ.ˆGØˆNð PWˆGØˆNÜ+Ø°À^Ð`l×`wÑ`wô
ð 	
r-   )NNNN)rZ   r[   r\   r   r!   rF   ra   r   r`   r   r   r<  rX   rb   rc   s   @r,   r¡  r¡    s‡   ø„ ð
˜~õ ð 6:Ø/3Ø,0Ø&*ña
à—\‘\ða
ð % U§\¡\Ñ2ða
ð ' t™nð	a
ð
 $ D™>ða
ð ˜d‘^ða
ð 
ˆuÐ2Ð2Ñ	3÷a
r-   r¡  c                   óD   ‡ — e Zd Zdefˆ fd„Zdej                  fd„Zˆ xZS )ÚPatchTSTClassificationHeadr   c                 ó¢  •— t         ‰| �  «        |j                  | _        |j                  | _        t	        j
                  d¬«      | _        |j                  dkD  rt	        j                  |j                  «      nt	        j                  «       | _
        t	        j                  |j                  |j                  z  |j                  «      | _        y ©Nr   ©Ú	start_dimr   )r    r!   rê   Úpooling_typer   ÚFlattenÚflattenrœ  rÍ   rÎ   r   r%   rß   rj   Únum_targetsr�  rm   s     €r,   r!   z#PatchTSTClassificationHead.__init__z  s‘   ø€ Ü‰ÑÔØ#×1Ñ1ˆÔØ"×/Ñ/ˆÔÜ—z‘z¨AÔ.ˆŒØ:@×:MÑ:MÐPQÒ:Q”r—z‘z &×"5Ñ"5Ô6ÔWY×WbÑWbÓWdˆŒÜ—i‘i × 9Ñ 9¸F¿N¹NÑ JÈF×L^ÑL^Ó_ˆ�r-   rž  c                 ón  — | j                   r|dd…dd…ddd…f   }ng| j                  dk(  r|j                  d¬«      }nE| j                  dk(  r|j                  d¬«      j                  }nt        d| j                  › d�«      ‚| j                  |«      }| j                  | j                  |«      «      }|S )	a[  
        Parameters:
            embedding (`torch.Tensor` of shape `(bs, num_channels, num_patches, d_model)` or
                     `(bs, num_channels, num_patches+1, d_model)` if `cls_token` is set to True, *required*):
                Embedding from the model
        Returns:
            `torch.Tensor` of shape `(bs, num_targets)`

        Nr   rç   r2   r@   r­   úpooling operator ú is not implemented yet)	rê   rµ  rç   r­   Úvaluesr#   r·  r�  r   ©r*   rž  Úpooled_embeddingrp   s       r,   rX   z"PatchTSTClassificationHead.forward‚  s®   € ð ×Òà(ªªA¨q²!¨Ñ4ÑØ×Ñ &Ò(à(Ÿ~™~°!˜~Ó4ÑØ×Ñ %Ò'à(Ÿ}™}°˜}Ó3×:Ñ:ÑäÐ0°×1BÑ1BÐ0CÐCZÐ[Ó\Ð\àŸ<™<Ð(8Ó9Ðà—‘˜TŸ\™\Ð*:Ó;Ó<ˆØˆr-   r  rc   s   @r,   r°  r°  y  s!   ø„ ð`˜~õ `ð §¡÷ r-   r°  z&The PatchTST for classification model.c                   ó¤   ‡ — e Zd Zdefˆ fd„Z	 	 	 	 	 ddej                  deej                     dee   dee   dee   dee   d	e	e
ef   fd
„Zˆ xZS )ÚPatchTSTForClassificationr   c                 óÔ   •— t         ‰| �  |«       |j                  rt        j	                  d«       d|_        t        |«      | _        t        |«      | _        | j                  «        y )Nú+Setting `do_mask_input` parameter to False.F)
r    r!   r†  ÚloggerÚwarningrƒ  rä   r°  r£  r+  rm   s     €r,   r!   z"PatchTSTForClassification.__init__£  sT   ø€ Ü‰Ñ˜Ô ð ×ÒÜ�N‰NÐHÔIØ#(ˆFÔ ä" 6Ó*ˆŒ
Ü.¨vÓ6ˆŒ	ð 	�‰Õr-   r±   Útarget_valuesr‰  r,  r<   r‹  r=   c                 óR  — |�|n| j                   j                  }| j                  ||||d¬«      }| j                  |j                  «      }d}	|�t        j                  «       }
 |
||«      }	|s|f|dd z   }|	�|	f|z   }|S |}|S t        |	||j                  |j                  ¬«      S )aº  
        Parameters:
            past_values (`torch.Tensor` of shape `(bs, sequence_length, num_input_channels)`, *required*):
                Input sequence to the model
            target_values (`torch.Tensor`, *optional*):
                Labels associates with the `past_values`
            past_observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
                Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
                in `[0, 1]`:

                - 1 for values that are **observed**,
                - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
            output_hidden_states (`bool`, *optional*):
                Whether or not to return the hidden states of all layers
            output_attentions (`bool`, *optional*):
                Whether or not to return the output attention of all layers
            return_dict (`bool`, *optional*):
                Whether or not to return a `ModelOutput` instead of a plain tuple.

        Returns:
            `PatchTSTForClassificationOutput` or tuple of `torch.Tensor` (if `return_dict`=False or
            `config.return_dict`=False)

        Examples:

        ```python
        >>> from transformers import PatchTSTConfig, PatchTSTForClassification

        >>> # classification task with two input channel2 and 3 classes
        >>> config = PatchTSTConfig(
        ...     num_input_channels=2,
        ...     num_targets=3,
        ...     context_length=512,
        ...     patch_length=12,
        ...     stride=12,
        ...     use_cls_token=True,
        ... )
        >>> model = PatchTSTForClassification(config=config)

        >>> # during inference, one only provides past values
        >>> past_values = torch.randn(20, 512, 2)
        >>> outputs = model(past_values=past_values)
        >>> labels = outputs.prediction_logits
        ```NTr¥  r   r¶   )r=  rG  r7   r/  )
r   r‘  rä   r£  r.  r   ÚCrossEntropyLossrF  r7   r/  )r*   r±   rÅ  r‰  r,  r<   r‹  r«  Úy_hatr­  r=  rá   s               r,   rX   z!PatchTSTForClassification.forward±  s×   € ðl &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—z‘zØ#Ø1Ø!5Ø/Øð "ó 
ˆð —	‘	˜,×8Ñ8Ó9ˆàˆØÐ$Ü×&Ñ&Ó(ˆDÙ˜E =Ó1ˆHáØ�h ¨a°Ð!3Ñ3ˆGØ/7Ð/C�x�k GÑ+ˆGØˆNð JQˆGØˆNÜ.ØØ#Ø&×4Ñ4Ø#×.Ñ.ô	
ð 	
r-   r—  )rZ   r[   r\   r   r!   rF   ra   r   r`   r   r’  rF  rX   rb   rc   s   @r,   rÀ  rÀ  ž  s™   ø„ ð
˜~õ ð" 15Ø-1Ø/3Ø,0Ø&*ñO
à—\‘\ðO
ð   §¡Ñ-ðO
ð % T™Nð	O
ð
 ' t™nðO
ð $ D™>ðO
ð ˜d‘^ðO
ð 
ˆuÐ5Ð5Ñ	6÷O
r-   rÀ  z"The PatchTST for regression Model.c                   óF   ‡ — e Zd Zddefˆ fd„Zdej                  fd„Zˆ xZS )ÚPatchTSTPredictionHeadr   c                 ó  •— t         ‰| �  «        |j                  | _        |j                  | _        |j                  | _        |j
                  | _        | j
                  s| j                  r|j                  }n|j                  |z  }| j                  �sVt        j                  «       | _	        t        j                  «       | _
        t        j                  «       | _        t        | j                  «      D ]ò  }| j                  j                  t        j                  d¬«      «       |€:| j                  j                  t        j                  ||j                   «      «       n*| j                  j                  |j#                  |«      «       | j                  j                  |j$                  dkD  rt        j&                  |j$                  «      nt        j(                  «       «       Œô y t        j                  d¬«      | _        |€&t        j                  ||j                   «      | _        n|j#                  |«      | _        |j$                  dkD  rt        j&                  |j$                  «      nt        j(                  «       | _        y )Nr2   r³  r   )r    r!   Úshare_projectionrß   rê   rµ  rj   r   r  ÚprojectionsÚdropoutsÚflattensr  r˜   r¶  r%   Úprediction_lengthÚget_parameter_projectionrœ  rÍ   rÎ   r·  Ú
projectionr   )r*   r   r®   Údistribution_outputr"   r  r+   s         €r,   r!   zPatchTSTPredictionHead.__init__  sÓ  ø€ Ü‰ÑÔà &× 7Ñ 7ˆÔØ"(×";Ñ";ˆÔØ#×1Ñ1ˆÔØ"×/Ñ/ˆÔØ×Ò × 2Ò 2Ø—~‘~‰Hà—~‘~¨Ñ3ˆHà×$Ó$ä!Ÿ}™}›ˆDÔÜŸM™M›OˆDŒMÜŸM™M›OˆDŒMÜ˜4×2Ñ2Ó3ò t�Ø—‘×$Ñ$¤R§Z¡Z¸!Ô%<Ô=Ø&Ð.à×$Ñ$×+Ñ+¬B¯I©I°hÀ×@XÑ@XÓ,YÕZð ×$Ñ$×+Ñ+Ð,?×,XÑ,XÐYaÓ,bÔcØ—‘×$Ñ$È×H[ÑH[Ð^_ÒH_¤R§Z¡Z°×0CÑ0CÔ%DÔeg×epÑepÓerÕsñtô Ÿ:™:°Ô2ˆDŒLØ"Ð*ä"$§)¡)¨H°f×6NÑ6NÓ"O�•ð #6×"NÑ"NÈxÓ"X�”Ø>D×>QÑ>QÐTUÒ>Uœ2Ÿ:™: f×&9Ñ&9Ô:Ô[]×[fÑ[fÓ[hˆD�Lr-   rž  c                 óä  — | j                   r|dd…dd…ddd…f   }nP| j                  dk(  r|j                  d¬«      }n.| j                  dk(  r|j                  d¬«      j                  }n|}| j
                  sŽg }t        | j                  «      D ]\  } | j                  |   |dd…|dd…f   «      } | j                  |   |«      } | j                  |   |«      }|j                  |«       Œ^ t        j                  |d¬«      }n3| j                  |«      }| j                  |«      }| j!                  |«      }t#        |t$        «      rt%        d„ |D «       «      }|S |j'                  dd«      }|S )	aj  
        Parameters:
            embedding (`torch.Tensor` of shape `(bs, num_channels, num_patches, d_model)` or
                     `(bs, num_channels, num_patches+1, d_model)` if `cls_token` is set to True, *required*):
                Embedding from the model
        Returns:
            `torch.Tensor` of shape `(bs, forecast_len, num_channels)`

        Nr   rç   r2   r@   r­   r   c              3   ó@   K  — | ]  }|j                  d d«      –— Œ y­w)r2   r   N)r4   )rŽ  Úzs     r,   r�  z1PatchTSTPredictionHead.forward.<locals>.<genexpr>[  s   è ø€ Ò=°˜1Ÿ;™; q¨!×,Ñ=ùs   ‚)rê   rµ  rç   r­   r¼  rÌ  r  rß   rÏ  rÎ  rÍ  r˜   rF   r	  r·  r   rÒ  r”   r’  r4   )r*   rž  r¾  rp   r  s        r,   rX   zPatchTSTPredictionHead.forward-  st  € ð ×Òà(ªªA¨q²!¨Ñ4Ñà× Ñ  FÒ*à#,§>¡>°a >Ó#8Ñ Ø×"Ñ" eÒ+à#,§=¡=°Q =Ó#7×#>Ñ#>Ñ ð $-Ð à×$Ò$ØˆFÜ˜4×2Ñ2Ó3ò 0�à#3 4§=¡=°Ñ#3Ð4DÂQÈÊ1ÀWÑ4MÓ#NÐ Ø#3 4§=¡=°Ñ#3Ð4DÓ#EÐ ð $7 4×#3Ñ#3°AÑ#6Ð7GÓ#HÐ Ø—‘Ð.Õ/ð0ô —[‘[ ¨QÔ/‰Fð  $Ÿ|™|Ð,<Ó=ÐØ#Ÿ|™|Ð,<Ó=Ðð —_‘_Ð%5Ó6ˆFä�fœeÔ$äÑ=°fÔ=Ó=ˆFð ˆð ×%Ñ% a¨Ó+ˆFØˆr-   r»   r  rc   s   @r,   rÊ  rÊ    s"   ø„ ñ
#i˜~õ #iðJ1 §¡÷ 1r-   rÊ  z"The PatchTST for prediction model.c                   óþ   ‡ — e Zd Zdefˆ fd„Z	 	 	 	 	 ddej                  deej                     deej                     dee   dee   dee   d	e	e
ef   fd
„Z	 ddej                  deej                     d	efd„Zˆ xZS )ÚPatchTSTForPredictionr   c                 óŠ  •— t         ‰| �  |«       |j                  rt        j	                  d«       d|_        t        |«      | _        |j                  dk(  rd | _        n™|j                  dk(  rt        |j                  ¬«      | _        nn|j                  dk(  rt        |j                  ¬«      | _        nC|j                  dk(  rt        |j                  ¬«      | _        nt        d|j                  › �«      ‚t        || j                  j                  j                   | j                  ¬	«      | _        | j%                  «        y )
NrÂ  FÚmseÚ	student_tr@   ÚnormalÚnegative_binomialúUnknown distribution output )rÓ  )r    r!   r†  rÃ  rÄ  rƒ  rä   r=  rÓ  r   rÐ  r   r   r#   rÊ  r…  r®   r£  r+  rm   s     €r,   r!   zPatchTSTForPrediction.__init__f  s  ø€ Ü‰Ñ˜Ô ð ×ÒÜ�N‰NÐHÔIØ#(ˆFÔ ä" 6Ó*ˆŒ
à�;‰;˜%ÒØ'+ˆDÕ$à×)Ñ)¨[Ò8Ü+9¸f×>VÑ>VÔ+W�Õ(Ø×+Ñ+¨xÒ7Ü+7¸F×<TÑ<TÔ+U�Õ(Ø×+Ñ+Ð/BÒBÜ+AÀf×F^ÑF^Ô+_�Õ(ä Ð#?À×@ZÑ@ZÐ?[Ð!\Ó]Ð]ä*Ø�D—J‘J×)Ñ)×5Ñ5È4×KcÑKcô
ˆŒ	ð
 	�‰Õr-   r±   r‰  rŠ  r,  r<   r‹  r=   c                 óŒ  — |�|n| j                   j                  }| j                  ||||d¬«      }| j                  |j                  «      }d}	| j
                  r|}
n||j                  z  |j                  z   }
|�u| j
                  rJ| j
                  j                  ||j                  |j                  ¬«      }t        ||«      }	t        |	«      }	nt        j                  d¬«      } ||
|«      }	|j                  }|j                  }|s|
f|dd z   }|	�|	f|z   }|S |}|S t        |	|
|j                  |j                  ||¬	«      S )
aV	  
        Parameters:
            past_values (`torch.Tensor` of shape `(bs, sequence_length, num_input_channels)`, *required*):
                Input sequence to the model
            past_observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
                Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
                in `[0, 1]`:

                - 1 for values that are **observed**,
                - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
            future_values (`torch.Tensor` of shape `(bs, forecast_len, num_input_channels)`, *optional*):
                Future target values associated with the `past_values`
            output_hidden_states (`bool`, *optional*):
                Whether or not to return the hidden states of all layers
            output_attentions (`bool`, *optional*):
                Whether or not to return the output attention of all layers
            return_dict (`bool`, *optional*):
                Whether or not to return a `ModelOutput` instead of a plain tuple.

        Returns:
            `PatchTSTForPredictionOutput` or tuple of `torch.Tensor` (if `return_dict`=False or
            `config.return_dict`=False)

        Examples:

        ```python
        >>> from huggingface_hub import hf_hub_download
        >>> import torch
        >>> from transformers import PatchTSTConfig, PatchTSTForPrediction

        >>> file = hf_hub_download(
        ...     repo_id="hf-internal-testing/etth1-hourly-batch", filename="train-batch.pt", repo_type="dataset"
        ... )
        >>> batch = torch.load(file)

        >>> # Prediction task with 7 input channels and prediction length is 96
        >>> model = PatchTSTForPrediction.from_pretrained("namctin/patchtst_etth1_forecast")

        >>> # during training, one provides both past and future values
        >>> outputs = model(
        ...     past_values=batch["past_values"],
        ...     future_values=batch["future_values"],
        ... )

        >>> loss = outputs.loss
        >>> loss.backward()

        >>> # during inference, one only provides past values, the model outputs future values
        >>> outputs = model(past_values=batch["past_values"])
        >>> prediction_outputs = outputs.prediction_outputs
        ```NTr¥  ©r7  r8  rç   r§  r   r?   )r=  rD  r7   r/  r7  r8  )r   r‘  rä   r£  r.  rÓ  r8  r7  ÚdistributionrO  rZ  r   rª  rC  r7   r/  )r*   r±   r‰  rŠ  r,  r<   r‹  r«  rÈ  r­  Ú	y_hat_outrá  r=  r7  r8  rá   s                   r,   rX   zPatchTSTForPrediction.forwardƒ  sn  € ðz &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆð —z‘zØ#Ø1Ø!5Ø/Øð "ó 
ˆð —	‘	˜,×8Ñ8Ó9ˆàˆà×#Ò#Ø‰Ià × 2Ñ 2Ñ2°\×5EÑ5EÑEˆIàÐ$Ø×'Ò'Ø#×7Ñ7×DÑDØ˜|×/Ñ/°|×7IÑ7Ið  Eó  �ô ˜|¨]Ó;�ä+¨HÓ5‘ä—z‘z¨FÔ3�Ù 	¨=Ó9�à×ÑˆØ×"Ñ"ˆáØ �l \°!°BÐ%7Ñ7ˆGØ/7Ð/C�x�k GÑ+ˆGØˆNð JQˆGØˆNÜ*ØØ(Ø&×4Ñ4Ø#×.Ñ.ØØô
ð 	
r-   c                 óª  — | j                   j                  } | |d|d¬«      }| j                  rz| j                  j                  |j                  |j
                  |j                  ¬«      }t        |«      D �cg c]  }|j                  «       ‘Œ }}t        j                  |d¬«      }n|j                  j                  d«      }t        |¬«      S c c}w )a   
        Generate sequences of sample predictions from a model with a probability distribution head.

        Parameters:
            past_values (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`):
                Past values of the time series that serves as context in order to predict the future.
            past_observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
                Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
                in `[0, 1]`:

                - 1 for values that are **observed**,
                - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).

        Return:
            [`SamplePatchTSTOutput`] where the outputs `sequences` tensor will have shape `(batch_size, number of
            samples, prediction_length, 1)` or `(batch_size, number of samples, prediction_length, num_input_channels)`
            for multivariate predictions.
        NF)r±   rŠ  r‰  r,  rà  r   r@   ©rJ  )r   Únum_parallel_samplesrÓ  rá  rD  r7  r8  r  ÚsamplerF   r	  r€   rI  ©r*   r±   r‰  rå  rá   rá  rN   Úsampless           r,   ÚgeneratezPatchTSTForPrediction.generateð  sÃ   € ð0  $Ÿ{™{×?Ñ?Ðñ Ø#ØØ1Ø!&ô	
ˆð ×#Ò#à×3Ñ3×@Ñ@Ø×*Ñ*°·±À7Ç=Á=ð Aó ˆLô 7<Ð<PÓ6QÖR°�|×*Ñ*Õ,ÐRˆGÐRä—k‘k '¨qÔ1‰Gà×0Ñ0×:Ñ:¸1Ó=ˆGä#¨gÔ6Ð6ùò Ss   Á8Cr—  r»   )rZ   r[   r\   r   r!   rF   ra   r   r`   r   r   rC  rX   rI  ré  rb   rc   s   @r,   rØ  rØ  a  sÓ   ø„ ð
˜~õ ð@ 6:Ø04Ø/3Ø,0Ø&*ñk
à—\‘\ðk
ð % U§\¡\Ñ2ðk
ð   §¡Ñ-ð	k
ð
 ' t™nðk
ð $ D™>ðk
ð ˜d‘^ðk
ð 
ˆuÐ1Ð1Ñ	2ók
ð` 6:ñ-7à—\‘\ð-7ð % U§\¡\Ñ2ð-7ð 
÷	-7r-   rØ  c                   óJ   ‡ — e Zd ZdZddefˆ fd„Zdej                  fd„Zˆ xZ	S )ÚPatchTSTRegressionHeadz
    Regression head
    r   c                 ó  •— t         ‰| �  «        |j                  | _        |j                  | _        |j
                  | _        || _        |j                  |j                  z  }t        j                  d¬«      | _        |j                  dkD  rt        j                  |j                  «      nt        j                  «       | _        |€&t        j                   ||j"                  «      | _        y |j'                  |«      | _        y r²  )r    r!   Úoutput_rangeÚy_rangerê   rµ  rÓ  rß   rj   r   r¶  r·  rœ  rÍ   rÎ   r   r%   r¸  rÒ  rÑ  )r*   r   rÓ  r"   r+   s       €r,   r!   zPatchTSTRegressionHead.__init__%  sÃ   ø€ Ü‰ÑÔØ×*Ñ*ˆŒØ#×1Ñ1ˆÔØ"×/Ñ/ˆÔØ#6ˆÔ à×,Ñ,¨v¯~©~Ñ=ˆä—z‘z¨AÔ.ˆŒØ:@×:MÑ:MÐPQÒ:Q”r—z‘z &×"5Ñ"5Ô6ÔWY×WbÑWbÓWdˆŒàÐ&Ü Ÿi™i¨°&×2DÑ2DÓEˆD�Oà1×JÑJÈ8ÓTˆD�Or-   rž  c                 ó2  — | j                   r|dd…dd…ddd…f   }ng| j                  dk(  r|j                  d¬«      }nE| j                  dk(  r|j                  d¬«      j                  }nt        d| j                  › d�«      ‚| j                  | j                  |«      «      }| j                  |«      }| j                  du | j                  duz  rEt        j                  |«      | j                  d	   | j                  d   z
  z  | j                  d   z   }|S )
aY  
        Parameters:
            embedding (`torch.Tensor` of shape `(bs, num_channels, num_patches, d_model)` or
                    `(bs, num_channels, num_patches+1, d_model)` if `cls_token` is set to True, *required*):
                Embedding from the model
        Returns:
            `torch.Tensor` of shape `(bs, output_dim)`

        Nr   rç   r2   r@   r­   rº  r»  r   )rê   rµ  rç   r­   r¼  r#   r   r·  rÒ  rÓ  rî  rF   Úsigmoidr½  s       r,   rX   zPatchTSTRegressionHead.forward6  s  € ð ×Òà(ªªA¨q²!¨Ñ4ÑØ×Ñ &Ò(à(Ÿ~™~°!˜~Ó4ÑØ×Ñ %Ò'à(Ÿ}™}°˜}Ó3×:Ñ:ÑäÐ0°×1BÑ1BÐ0CÐCZÐ[Ó\Ð\ð  Ÿ<™<¨¯©Ð5EÓ(FÓGÐð —‘Ð!1Ó2ˆà×$Ñ$¨Ð,°·±ÀTÐ1IÒJÜ—]‘] 6Ó*¨d¯l©l¸1©oÀÇÁÈQÁÑ.OÑPÐSW×S_ÑS_Ð`aÑSbÑbˆFØˆr-   r»   rq   rc   s   @r,   rë  rë     s&   ø„ ññU˜~õ Uð" §¡÷ r-   rë  z"The PatchTST for regression model.c                   óþ   ‡ — e Zd Zdefˆ fd„Z	 	 	 	 	 ddej                  deej                     deej                     dee   dee   dee   d	e	e
ef   fd
„Z	 ddej                  deej                     d	efd„Zˆ xZS )ÚPatchTSTForRegressionr   c                 óJ  •— t         ‰| �  |«       |j                  rt        j	                  d«       d|_        t        |«      | _        |j                  dk(  rd | _        n™|j                  dk(  rt        |j                  ¬«      | _        nn|j                  dk(  rt        |j                  ¬«      | _        nC|j                  dk(  rt        |j                  ¬«      | _        nt        d|j                  › �«      ‚t        || j                  «      | _        | j!                  «        y )	NrÂ  FrÚ  rÛ  r@   rÜ  rÝ  rÞ  )r    r!   r†  rÃ  rÄ  rƒ  rä   r=  rÓ  r   r¸  r   r   r#   rë  r£  r+  rm   s     €r,   r!   zPatchTSTForRegression.__init__\  sî   ø€ Ü‰Ñ˜Ô ð ×ÒÜ�N‰NÐHÔIØ#(ˆFÔ ä" 6Ó*ˆŒ
Ø�;‰;˜%ÒØ'+ˆDÕ$à×)Ñ)¨[Ò8Ü+9¸f×>PÑ>PÔ+Q�Õ(Ø×+Ñ+¨xÒ7Ü+7¸F×<NÑ<NÔ+O�Õ(Ø×+Ñ+Ð/BÒBÜ+AÀf×FXÑFXÔ+Y�Õ(ä Ð#?À×@ZÑ@ZÐ?[Ð!\Ó]Ð]ä*¨6°4×3KÑ3KÓLˆŒ	ð 	�‰Õr-   r±   rÅ  r‰  r,  r<   r‹  r=   c           	      óX  — |�|n| j                   j                  }| j                  ||||d¬«      }| j                  |j                  «      }d}	|�›| j
                  rp| j
                  j                  |«      }
t        |D �cg c](  }|j                  d| j                   j                  «      ‘Œ* c}«      }t        |
|«      }	t        |	«      }	nt        j                  d¬«      }	 |	||«      }	|s|f|dd z   }|	�|	f|z   }|S |}|S t        |	||j                  |j                   ¬	«      S c c}w )
a'  
        Parameters:
            past_values (`torch.Tensor` of shape `(bs, sequence_length, num_input_channels)`, *required*):
                Input sequence to the model
            target_values (`torch.Tensor` of shape `(bs, num_input_channels)`):
                Target values associates with the `past_values`
            past_observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
                Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
                in `[0, 1]`:

                - 1 for values that are **observed**,
                - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).
            output_hidden_states (`bool`, *optional*):
                Whether or not to return the hidden states of all layers
            output_attentions (`bool`, *optional*):
                Whether or not to return the output attention of all layers
            return_dict (`bool`, *optional*):
                Whether or not to return a `ModelOutput` instead of a plain tuple.

        Returns:
            `PatchTSTForRegressionOutput` or tuple of `torch.Tensor` (if `return_dict`=False or
            `config.return_dict`=False)

        Examples:

        ```python
        >>> from transformers import PatchTSTConfig, PatchTSTForRegression

        >>> # Regression task with 6 input channels and regress 2 targets
        >>> model = PatchTSTForRegression.from_pretrained("namctin/patchtst_etth1_regression")

        >>> # during inference, one only provides past values, the model outputs future values
        >>> past_values = torch.randn(20, 512, 6)
        >>> outputs = model(past_values=past_values)
        >>> regression_outputs = outputs.regression_outputs
        ```NTr¥  r?   rç   r§  r   r¶   )r=  rA  r7   r/  )r   r‘  rä   r£  r.  rÓ  rá  r’  r3   r¸  rO  rZ  r   rª  r@  r7   r/  )r*   r±   rÅ  r‰  r,  r<   r‹  r«  rÈ  r=  rá  Úitemrá   s                r,   rX   zPatchTSTForRegression.forwardv  s=  € ð\ &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—z‘zØ#Ø1Ø!5Ø/Øð "ó 
ˆð —	‘	˜,×8Ñ8Ó9ˆàˆØÐ$Ø×'Ò'Ø#×7Ñ7×DÑDÀUÓK�äÐRWÖXÈ$˜tŸy™y¨¨T¯[©[×-DÑ-DÕEÒXÓY�Ü˜<¨Ó7�ä'¨Ó-‘ä—z‘z¨FÔ3�Ù˜E =Ó1�áà�h ¨a°Ð!3Ñ3ˆGØ+/Ð+;�t�g Ñ'ˆGØˆNð BIˆGØˆNÜ*ØØ$Ø&×4Ñ4Ø#×.Ñ.ô	
ð 	
ùò Ys   Â -D'c                 óv  — | j                   j                  } | |d|d¬«      }| j                  j                  |j                  «      }t        |«      D �cg c]  }|j                  «       ‘Œ }}t        j                  |d¬«      j                  d|| j                   j                  «      }t        |¬«      S c c}w )a¢  
        Generate sequences of sample predictions from a model with a probability distribution head.

        Parameters:
            past_values (`torch.FloatTensor` of shape `(batch_size, sequence_length, num_input_channels)`):
                Past values of the time series that serves as context in order to predict the future.
            past_observed_mask (`torch.BoolTensor` of shape `(batch_size, sequence_length, num_input_channels)`, *optional*):
                Boolean mask to indicate which `past_values` were observed and which were missing. Mask values selected
                in `[0, 1]`:

                - 1 for values that are **observed**,
                - 0 for values that are **missing** (i.e. NaNs that were replaced by zeros).

        Return:
            [`SamplePatchTSTOutput`] where the outputs `sequences` tensor will have shape `(batch_size, number of
            samples, num_targets)`.
        NF)r±   rÅ  r‰  r,  r   r@   r?   rä  )r   rå  rÓ  rá  rA  r  ræ  rF   r	  r3   r¸  rI  rç  s           r,   ré  zPatchTSTForRegression.generateÉ  s§   € ð.  $Ÿ{™{×?Ñ?Ðñ Ø#ØØ1Ø!&ô	
ˆð ×/Ñ/×<Ñ<¸W×=WÑ=WÓXˆä27Ð8LÓ2MÖN¨Q�<×&Ñ&Õ(ÐNˆÐNä—+‘+˜g¨1Ô-×2Ñ2°2Ð7KÈTÏ[É[×MdÑMdÓeˆÜ#¨gÔ6Ð6ùò Os   ÁB6r—  r»   )rZ   r[   r\   r   r!   rF   ra   r   r`   r   r’  r@  rX   rI  ré  rb   rc   s   @r,   rò  rò  W  sÒ   ø„ ð
˜~õ ð: 15Ø59Ø/3Ø,0Ø&*ñQ
à—\‘\ðQ
ð   §¡Ñ-ðQ
ð % U§\¡\Ñ2ð	Q
ð
 ' t™nðQ
ð $ D™>ðQ
ð ˜d‘^ðQ
ð 
ˆuÐ1Ð1Ñ	2óQ
ðl 6:ñ'7à—\‘\ð'7ð % U§\¡\Ñ2ð'7ð 
÷	'7r-   rò  )rƒ  rã   rØ  r¡  rò  rÀ  )NFr   r›  r4  )Hr]   r  Údataclassesr   Útypingr   r   r   rF   r   Úactivationsr	   Úmodeling_outputsr
   Úmodeling_utilsr   Útime_series_utilsr   r   r   Úutilsr   r   r   Úconfiguration_patchtstr   Ú
get_loggerrZ   rÃ  Ú_CONFIG_FOR_DOCÚModuler   re   ra   r_   Úlistr`   r^   rŒ   r¦   r¨   r¹   rÅ   rã   r  ré   rù   ÚPATCHTST_START_DOCSTRINGr6  r<  r@  rC  rF  rI  ÚdistributionsÚDistributionrO  rZ  r\  rl  rz  r~  rƒ  r™  r¡  r°  rÀ  rÊ  rØ  rë  rò  Ú__all__r�   r-   r,   ú<module>r     sô  ðñ ã Ý !ß )Ñ )ã Ý å "Ý /Ý -ß UÑ Uß ?Ñ ?Ý 2ð 
ˆ×	Ñ	˜HÓ	%€à"€ô[B˜Ÿ	™	ô [Bô|&˜Ÿ	™	ô &ð2 &*Ø',Øñ7%Ø�L‰Lð7%àð7%ð #ð7%ð !%ð	7%ð
 ó7%ðz &*Øñ	A%Ø�L‰LðA%à$ T¨3 YÑ/ðA%ð #ðA%ð ó	A%ôH-�r—y‘yô -ô`9"�b—i‘iô 9"ôxG˜2Ÿ9™9ô GôT2˜oô 2ôB!˜Ÿ	™	ô !ôH5 §¡ô 5ôp;xÐ-ô ;xð|Ð ð" ô4˜+ó 4ó ð4ð< ô: ;ó :ó ð:ð8 ô: +ó :ó ð:ð8 ô. +ó .ó ð.ðD ô: kó :ó ð:ð: ô
2˜;ó 
2ó ð
2ð#ˆu×"Ñ"×/Ñ/ð #¸¿¹ð #È%Ï,É,ó #ñ* 5§<¡<ð *¸(À5Ç<Á<Ñ:Pð *Ðfk×frÑfró *ô2 0˜Ÿ	™	ô  0ôH3;˜Ÿ™ô 3;ôn ˜Ÿ	™	ô  ô6 �R—Y‘Yô  ñ8 ØUØóôm
Ð+ó m
ó	ðm
ô`˜rŸy™yô ñ8 Ø&Øóôl
Ð4ó l
ó	ðl
ô^" §¡ô "ñJ Ø,Øóô^
Ð 7ó ^
ó	ð^
ñB Ø(ØóôW˜RŸY™Yó Wó	ðWñt Ø(Øóôx7Ð3ó x7ó	ðx7ôv4˜RŸY™Yô 4ñn Ø(ØóôU7Ð3ó U7ó	ðU7òp�r-   