§
    ‚Štj¸A ã                   óô  — d Z ddlZddlmZ ddl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 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mZ ddlm Z   ej!        e"¦  «        Z# G d„ dej$        ¦  «        Z% G d„ dej$        ¦  «        Z& G d„ dej$        ¦  «        Z' G d„ dej$        ¦  «        Z( G d„ dej$        ¦  «        Z) G d„ dej$        ¦  «        Z*	 	 d|dej$        dej+        dej+        d ej+        d!ej+        dz  d"e,dz  d#e,d$ee         fd%„Z- G d&„ d'ej$        ¦  «        Z. G d(„ d)ej$        ¦  «        Z/ G d*„ d+ej$        ¦  «        Z0 G d,„ d-ej$        ¦  «        Z1 G d.„ d/ej$        ¦  «        Z2 G d0„ d1ej$        ¦  «        Z3 G d2„ d3ej$        ¦  «        Z4e G d4„ d5e
¦  «        ¦   «         Z5 G d6„ d7ej$        ¦  «        Z6	 	 	 d}d9ej+        d:e,d;e7dz  d<e8d=e9f
d>„Z:	 	 d~d9ej+        d?e7e9z  d;e7dz  d=e9fd@„Z; G dA„ dBej$        ¦  «        Z< G dC„ dDej$        ¦  «        Z= G dE„ dFej$        ¦  «        Z> G dG„ dHej$        ¦  «        Z? G dI„ dJej$        ¦  «        Z@ edK¬L¦  «        e G dM„ dNe¦  «        ¦   «         ¦   «         ZA G dO„ dPe5¦  «        ZB edQ¬L¦  «        e G dR„ dSe¦  «        ¦   «         ¦   «         ZC edT¬L¦  «         G dU„ dVe5¦  «        ¦   «         ZD edW¬L¦  «        e G dX„ dYe¦  «        ¦   «         ¦   «         ZE edZ¬L¦  «         G d[„ d\e5¦  «        ¦   «         ZF ed]¬L¦  «        e G d^„ d_e¦  «        ¦   «         ¦   «         ZG ed`¬L¦  «        e G da„ dbe¦  «        ¦   «         ¦   «         ZH ed`¬L¦  «        e G dc„ dde¦  «        ¦   «         ¦   «         ZIdeejJ        jK        dfej+        dgej+        fdh„ZLddiej+        djej+        dz  dgej+        fdk„ZM G dl„ dme5¦  «        ZN edn¬L¦  «        e G do„ dpe¦  «        ¦   «         ¦   «         ZO G dq„ dre5¦  «        ZP eds¬L¦  «        e G dt„ due¦  «        ¦   «         ¦   «         ZQ G dv„ dwej$        ¦  «        ZR edx¬L¦  «         G dy„ dze5¦  «        ¦   «         ZSg d{¢ZTdS )€zPyTorch PatchTSMixer model.é    N)ÚCallable)Ú	dataclass)ÚPreTrainedModel)ÚModelOutputé   )Úinitialization)ÚFlashAttentionKwargs)ÚALL_ATTENTION_FUNCTIONS)ÚUnpack)ÚNegativeBinomialOutputÚNormalOutputÚStudentTOutput)ÚTransformersKwargsÚauto_docstringÚcan_return_tupleÚloggingé   )ÚPatchTSMixerConfigc                   ó2   ‡ — e Zd ZdZdedefˆ fd„Zd„ Zˆ xZS )ÚPatchTSMixerGatedAttentionz›
    Module that applies gated attention to input data.

    Args:
        in_size (`int`): The input size.
        out_size (`int`): The output size.
    Úin_sizeÚout_sizec                 ó°   •— t          ¦   «                              ¦   «          t          j        ||¦  «        | _        t          j        d¬¦  «        | _        d S )Néÿÿÿÿ©Údim)ÚsuperÚ__init__ÚnnÚLinearÚ
attn_layerÚSoftmaxÚattn_softmax)Úselfr   r   Ú	__class__s      €út/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/patchtsmixer/modeling_patchtsmixer.pyr   z#PatchTSMixerGatedAttention.__init__/   sG   ø€ Ý‰Œ×ÒÑÔÐÝœ) G¨XÑ6Ô6ˆŒÝœJ¨2Ð.Ñ.Ô.ˆÔÐÐó    c                 ó`   — |                       |                      |¦  «        ¦  «        }||z  }|S ©N)r#   r!   )r$   ÚinputsÚattn_weights      r&   Úforwardz"PatchTSMixerGatedAttention.forward4   s0   € Ø×'Ò'¨¯ª¸Ñ(?Ô(?Ñ@Ô@ˆØ˜+Ñ%ˆØˆr'   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úintr   r,   Ú__classcell__©r%   s   @r&   r   r   &   sd   ø€ € € € € ðð ð/ ð /¨sð /ð /ð /ð /ð /ð /ð
ð ð ð ð ð ð r'   r   c                   ó>   ‡ — e Zd ZdZdefˆ fd„Zdej        fd„Zˆ xZ	S )ÚPatchTSMixerBatchNormzP
    Compute batch normalization over the sequence length (time) dimension.
    Úconfigc                 ó’   •— t          ¦   «                              ¦   «          t          j        |j        |j        ¬¦  «        | _        d S )N©Úeps)r   r   r   ÚBatchNorm1dÚd_modelÚnorm_epsÚ	batchnorm©r$   r6   r%   s     €r&   r   zPatchTSMixerBatchNorm.__init__@   s7   ø€ Ý‰Œ×ÒÑÔÐÝœ¨¬¸F¼OÐLÑLÔLˆŒˆˆr'   r*   c                 ó„   — |                      dd¦  «        }|                      |¦  «        }|                      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   é   )Ú	transposer=   )r$   r*   Úoutputs      r&   r,   zPatchTSMixerBatchNorm.forwardD   s@   € ð ×!Ò! ! QÑ'Ô'ˆØ—’ Ñ'Ô'ˆØ×Ò  1Ñ%Ô%Ð%r'   ©
r-   r.   r/   r0   r   r   ÚtorchÚTensorr,   r2   r3   s   @r&   r5   r5   ;   ss   ø€ € € € € ðð ðMÐ1ð Mð Mð Mð Mð Mð Mð
&˜eœlð 
&ð 
&ð 
&ð 
&ð 
&ð 
&ð 
&ð 
&r'   r5   c                   óh   ‡ — e Zd ZdZdefˆ fd„Zededej        fd„¦   «         Z	de
j        fd„Zˆ xZS )ÚPatchTSMixerPositionalEncodingz'
    Class for positional encoding
    r6   c                 óú   •— t          ¦   «                              ¦   «          |j        r|                      |¦  «        | _        d S t          j        t          j        |j	        |j
        ¦  «        ¦  «        | _        d S r)   )r   r   Úuse_positional_encodingÚ_init_peÚposition_encr   Ú	ParameterrD   ÚzerosÚnum_patchesr;   r>   s     €r&   r   z'PatchTSMixerPositionalEncoding.__init__V   sh   ø€ Ý‰Œ×ÒÑÔÐàÔ)ð 	^Ø $§¢¨fÑ 5Ô 5ˆDÔÐÐå "¤­U¬[¸Ô9KÈVÌ^Ñ-\Ô-\Ñ ]Ô ]ˆDÔÐÐr'   Úreturnc                 ó  — | j         dk    r5t          j        t          j        | j        | j        ¦  «        d¬¦  «        }�nD| j         dk    �r!t          j        | 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<   ||                     ¦   «         z
  }||                     ¦   «         d	z  z  }t          j        |d
¬¦  «        }nt#          | j         › d�¦  «        ‚|S )NÚrandomT©Úrequires_gradÚsincosr   r   r@   g     ˆÃ@é
   FzN is not a valid positional encoder. Available types are 'random' and 'sincos'.)Úpositional_encoding_typer   rL   rD   ÚrandnrN   r;   rM   ÚarangeÚ	unsqueezeÚexpÚmathÚlogÚsinÚcosÚmeanÚstdÚ
ValueError)r6   rK   ÚpositionÚdiv_terms       r&   rJ   z'PatchTSMixerPositionalEncoding._init_pe^   s‚  € ð Ô*¨hÒ6Ð6Ýœ<­¬°FÔ4FÈÌÑ(WÔ(WÐgkÐlÑlÔlˆL‰LØÔ,°Ò8Ñ8Ý œ; vÔ'9¸6¼>ÑJÔJˆLÝ”| A vÔ'9Ñ:Ô:×DÒDÀQÑGÔGˆHÝ”y¥¤¨a°´ÀÑ!CÔ!CÍÌÐQXÑHYÔHYÐ\bÔ\jÑHjÐFkÑ!kÑlÔlˆHÝ$)¤I¨h¸Ñ.AÑ$BÔ$BˆL˜˜˜˜A˜D˜q˜D˜Ñ!Ý$)¤I¨h¸Ñ.AÑ$BÔ$BˆL˜˜˜˜A˜D˜q˜D˜Ñ!Ø'¨,×*;Ò*;Ñ*=Ô*=Ñ=ˆLØ'¨<×+;Ò+;Ñ+=Ô+=ÀÑ+BÑCˆLÝœ<¨ÀEÐJÑJÔJˆLˆLåØÔ2ð  Cð  Cð  Cñô ð ð Ðr'   Úpatch_inputc                 ó   — || j         z   }|S r)   )rK   )r$   rd   Úhidden_states      r&   r,   z&PatchTSMixerPositionalEncoding.forwardr   s   € à" TÔ%6Ñ6ˆØÐr'   )r-   r.   r/   r0   r   r   Ústaticmethodr   rL   rJ   rD   rE   r,   r2   r3   s   @r&   rG   rG   Q   s¤   ø€ € € € € ðð ð^Ð1ð ^ð ^ð ^ð ^ð ^ð ^ð ðÐ+ð °´ð ð ð ñ „\ðð& 5¤<ð ð ð ð ð ð ð ð r'   rG   c                   ó>   ‡ — e Zd ZdZdefˆ fd„Zdej        fd„Zˆ xZ	S )ÚPatchTSMixerNormLayerzeNormalization block

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.
    r6   c                 ó  •— t          ¦   «                              ¦   «          |j        | _        d|j                             ¦   «         v rt	          |¦  «        | _        d S t          j        |j        |j	        ¬¦  «        | _        d S )NÚbatchr8   )
r   r   Únorm_mlpÚlowerr5   Únormr   Ú	LayerNormr;   r<   r>   s     €r&   r   zPatchTSMixerNormLayer.__init__€   sl   ø€ Ý‰Œ×ÒÑÔÐàœˆŒà�f”o×+Ò+Ñ-Ô-Ð-Ð-Ý-¨fÑ5Ô5ˆDŒIˆIˆIåœ V¤^¸¼ÐIÑIÔIˆDŒIˆIˆIr'   r*   c                 óT  — d| j                              ¦   «         v rwt          j        ||j        d         |j        d         z  |j        d         |j        d         f¦  «        }|                      |¦  «        }t          j        ||j        ¦  «        }n|                      |¦  «        }|S )a  
        Args:
            inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`):
                Input to the normalization layer.
        Returns:
            `torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`
        rk   r   r   r@   r   )rl   rm   rD   ÚreshapeÚshapern   )r$   r*   Úinputs_reshapeds      r&   r,   zPatchTSMixerNormLayer.forwardŠ   sž   € ð �d”m×)Ò)Ñ+Ô+Ð+Ð+å#œmØà”L ”O f¤l°1¤oÑ5Ø”L ”OØ”L ”Oðñô ˆOð #Ÿiši¨Ñ8Ô8ˆOõ ”] ?°F´LÑAÔAˆFˆFð —Y’Y˜vÑ&Ô&ˆFàˆr'   rC   r3   s   @r&   ri   ri   x   ss   ø€ € € € € ðð ðJÐ1ð Jð Jð Jð Jð Jð Jð˜eœlð ð ð ð ð ð ð ð r'   ri   c                   ó4   ‡ — e Zd Zˆ fd„Zdej        fd„Zˆ xZS )ÚPatchTSMixerMLPc                 ó<  •— t          ¦   «                              ¦   «          ||j        z  }t          j        ||¦  «        | _        t          j        |j        ¦  «        | _        t          j        ||¦  «        | _	        t          j        |j        ¦  «        | _
        d S r)   )r   r   Úexpansion_factorr   r    Úfc1ÚDropoutÚdropoutÚdropout1Úfc2Údropout2)r$   Úin_featuresÚout_featuresr6   Ú
num_hiddenr%   s        €r&   r   zPatchTSMixerMLP.__init__ª   sv   ø€ Ý‰Œ×ÒÑÔÐØ  6Ô#:Ñ:ˆ
Ý”9˜[¨*Ñ5Ô5ˆŒÝœ
 6¤>Ñ2Ô2ˆŒÝ”9˜Z¨Ñ6Ô6ˆŒÝœ
 6¤>Ñ2Ô2ˆŒˆˆr'   r*   c                 óä   — |                       t          j                             |                      |¦  «        ¦  «        ¦  «        }|                      |¦  «        }|                      |¦  «        }|S )zì
        Args:
            inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`):
                Input to the MLP layer.
        Returns:
            `torch.Tensor` of the same shape as `inputs`
        )r{   r   Ú
functionalÚgelurx   r|   r}   )r$   r*   s     r&   r,   zPatchTSMixerMLP.forward²   sX   € ð —’�rœ}×1Ò1°$·(²(¸6Ñ2BÔ2BÑCÔCÑDÔDˆØ—’˜&Ñ!Ô!ˆØ—’˜vÑ&Ô&ˆØˆr'   )r-   r.   r/   r   rD   rE   r,   r2   r3   s   @r&   ru   ru   ©   sU   ø€ € € € € ð3ð 3ð 3ð 3ð 3ð˜eœlð ð ð ð ð ð ð ð r'   ru   c                   ó>   ‡ — e Zd ZdZdefˆ fd„Zdej        fd„Zˆ xZ	S )Ú$PatchTSMixerChannelFeatureMixerBlockzŠThis module mixes the features in the channel dimension.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.
    r6   c                 ó  •— t          ¦   «                              ¦   «          t          |¦  «        | _        |j        | _        t          |j        |j        |¬¦  «        | _        |j        r"t          |j        |j        ¬¦  «        | _	        d S d S ©N©r~   r   r6   ©r   r   )
r   r   ri   rn   Ú
gated_attnru   Únum_input_channelsÚmlpr   Úgating_blockr>   s     €r&   r   z-PatchTSMixerChannelFeatureMixerBlock.__init__È   s–   ø€ Ý‰Œ×ÒÑÔÐå)¨&Ñ1Ô1ˆŒ	Ø Ô+ˆŒÝ"ØÔ1ØÔ2Øð
ñ 
ô 
ˆŒð Ôð 	Ý :ØÔ1¸FÔ<Uð!ñ !ô !ˆDÔÐÐð	ð 	r'   r*   c                 ó   — |}|                       |¦  «        }|                     dddd¦  «        }| j        r|                      |¦  «        }|                      |¦  «        }|                     dddd¦  «        }||z   }|S )zë
        Args:
            inputs (`torch.Tensor` of shape `((batch_size, num_channels, num_patches, d_model))`):
                input to the MLP layer
        Returns:
            `torch.Tensor` of the same shape as `inputs`
        r   r   r@   r   )rn   ÚpermuterŠ   r�   rŒ   )r$   r*   ÚresidualÚouts       r&   r,   z,PatchTSMixerChannelFeatureMixerBlock.forwardØ   s…   € ð ˆØ—’˜6Ñ"Ô"ˆà—’  1 a¨Ñ+Ô+ˆàŒ?ð 	/Ø×&Ò& vÑ.Ô.ˆFà—’˜&Ñ!Ô!ˆà—’  1 a¨Ñ+Ô+ˆà�xÑˆØˆ
r'   rC   r3   s   @r&   r…   r…   À   sl   ø€ € € € € ðð ðÐ1ð ð ð ð ð ð ð ˜eœlð ð ð ð ð ð ð ð r'   r…   ç        ÚmoduleÚqueryÚkeyÚvalueÚattention_maskÚscalingrz   Úkwargsc                 ó®  — |€|                      d¦  «        dz  }t          j        ||                     dd¦  «        ¦  «        |z  }|�||z   }t          j                             |d¬¦  «        }t          j                             ||| j        ¬¦  «        }t          j        ||¦  «        }	|	                     dd¦  «         	                    ¦   «         }	|	|fS )Nr   ç      à¿r@   r   r   )ÚpÚtrainingr   )
ÚsizerD   ÚmatmulrA   r   r‚   Úsoftmaxrz   r�   Ú
contiguous)
r“   r”   r•   r–   r—   r˜   rz   r™   Úattn_weightsÚattn_outputs
             r&   Úeager_attention_forwardr¤   ñ   sÈ   € ð €Ø—*’*˜R‘.”. DÑ(ˆõ ”<  s§}¢}°Q¸Ñ':Ô':Ñ;Ô;¸gÑE€LàÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2Ð(Ñ>Ô>€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€Lå”,˜|¨UÑ3Ô3€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r'   c                   óþ   ‡ — e Zd ZdZ	 	 	 	 	 ddededed	ed
edededz  fˆ fd„Z	 	 	 dde	j
        de	j
        dz  de	j
        dz  dedz  dee         dee	j
        e	j
        dz  ee	j
                 dz  f         fd„Zˆ xZS )ÚPatchTSMixerAttentionz=Multi-headed attention from 'Attention Is All You Need' paperr’   FTNÚ	embed_dimÚ	num_headsrz   Ú
is_decoderÚbiasÚ	is_causalr6   c                 ó
  •— t          ¦   «                              ¦   «          || _        || _        || _        ||z  | _        || _        | j        |z  | j        k    rt          d| j        › d|› d�¦  «        ‚| j        dz  | _        || _	        || _
        t          j        |||¬¦  «        | _        t          j        |||¬¦  «        | _        t          j        |||¬¦  «        | _        t          j        |||¬¦  «        | _        d S )Nz;embed_dim must be divisible by num_heads (got `embed_dim`: z and `num_heads`: ú).r›   )rª   )r   r   r§   r¨   rz   Úhead_dimr6   ra   r˜   r©   r«   r   r    Úk_projÚv_projÚq_projÚout_proj)	r$   r§   r¨   rz   r©   rª   r«   r6   r%   s	           €r&   r   zPatchTSMixerAttention.__init__  s  ø€ õ 	‰Œ×ÒÑÔÐØ"ˆŒØ"ˆŒØˆŒØ! YÑ.ˆŒØˆŒàŒM˜IÑ%¨$¬.Ò8Ð8Ýð3ÈdÌnð 3ð 3Ø%.ð3ð 3ð 3ñô ð ð ”} dÑ*ˆŒØ$ˆŒØ"ˆŒå”i 	¨9¸4Ð@Ñ@Ô@ˆŒÝ”i 	¨9¸4Ð@Ñ@Ô@ˆŒÝ”i 	¨9¸4Ð@Ñ@Ô@ˆŒÝœ	 )¨Y¸TÐBÑBÔBˆŒˆˆr'   Úhidden_statesÚkey_value_statesr—   Úoutput_attentionsr™   rO   c                 óú  — |du}|j         dd…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }	|r|n|}
g |
j         dd…         ¢d‘| j        ‘R }|                      |
¦  «                             |¦  «                             dd¦  «        }|                      |
¦  «                             |¦  «                             dd¦  «        }t          j        | j	        j
        t          ¦  «        } || |	|||f| j        sdn| j        | j        |dœ|¤Ž\  }} |j        g |¢d‘R Ž                      ¦   «         }|                      |¦  «        }||dfS )z#Input shape: Batch x Time x ChannelNr   r   r@   r’   )rz   r˜   rµ   )rr   r®   r±   ÚviewrA   r¯   r°   r
   Úget_interfacer6   Ú_attn_implementationr¤   r�   rz   r˜   rq   r¡   r²   )r$   r³   r´   r—   rµ   r™   Úis_cross_attentionÚinput_shapeÚhidden_shapeÚquery_statesÚcurrent_statesÚkv_shapeÚ
key_statesÚvalue_statesÚattention_interfacer£   r¢   s                    r&   r,   zPatchTSMixerAttention.forward0  s½  € ð .°TÐ9Ðð $Ô)¨#¨2¨#Ô.ˆà8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆð —{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆà-?ÐRÐ)Ð)À]ˆØB�^Ô)¨#¨2¨#Ô.ÐB°ÐB°D´MÐBÐBˆØ—[’[ Ñ0Ô0×5Ò5°hÑ?Ô?×IÒIÈ!ÈQÑOÔOˆ
Ø—{’{ >Ñ2Ô2×7Ò7¸ÑAÔA×KÒKÈAÈqÑQÔQˆå(?Ô(MØŒKÔ,Õ.Eñ)
ô )
Ðð %8Ð$7ØØØØØð
%
ð  $œ}Ð>�C�C°$´,Ø”LØ/ð
%
ð 
%
ð ð
%
ð 
%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—m’m KÑ0Ô0ˆà˜L¨$Ð.Ð.r'   )r’   FTFN)NNF)r-   r.   r/   r0   r1   ÚfloatÚboolr   r   rD   rE   r   r	   Útupler,   r2   r3   s   @r&   r¦   r¦     sJ  ø€ € € € € ØGÐGð Ø ØØØ,0ðCð CàðCð ðCð ð	Cð
 ðCð ðCð ðCð # TÑ)ðCð Cð Cð Cð Cð CðD 15Ø.2Ø).ð0/ð 0/à”|ð0/ð  œ,¨Ñ-ð0/ð œ tÑ+ð	0/ð
   $™;ð0/ð Ð-Ô.ð0/ð 
ˆuŒ|˜Uœ\¨DÑ0°%¸¼Ô2EÈÑ2LÐLÔ	Mð0/ð 0/ð 0/ð 0/ð 0/ð 0/ð 0/ð 0/r'   r¦   c                   ó.   ‡ — e Zd ZdZdefˆ fd„Zd„ Zˆ xZS )ÚPatchMixerBlockzxThis module mixes the patch dimension.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.
    r6   c                 ó¸  •— t          ¦   «                              ¦   «          t          |¦  «        | _        |j        | _        |j        | _        t          |j        |j        |¬¦  «        | _        |j        r t          |j        |j        ¬¦  «        | _
        |j        r=t          |j        |j        |j        |¬¦  «        | _        t          |¦  «        | _        d S d S )Nrˆ   r‰   )r§   r¨   rz   r6   )r   r   ri   rn   Ú	self_attnrŠ   ru   rN   rŒ   r   r�   r¦   r;   Úself_attn_headsrz   Úself_attn_layerÚ	norm_attnr>   s     €r&   r   zPatchMixerBlock.__init__k  sß   ø€ Ý‰Œ×ÒÑÔÐå)¨&Ñ1Ô1ˆŒ	àÔ)ˆŒØ Ô+ˆŒå"ØÔ*ØÔ+Øð
ñ 
ô 
ˆŒð Ôð 	tÝ :À6ÔCUÐ`fÔ`rÐ sÑ sÔ sˆDÔàÔð 	;Ý#8Ø œ.Ø Ô0ØœØð	$ñ $ô $ˆDÔ õ 3°6Ñ:Ô:ˆDŒNˆNˆNð	;ð 	;r'   c                 óö  — |}|                       |¦  «        }| j        rY|j        \  }}}}|                     ||z  ||¦  «        }|                      |d¬¦  «        \  }}	}	|                     ||||¦  «        }|                     dd¦  «        }|                      |¦  «        }| j        r|                      |¦  «        }|                     dd¦  «        }| j        r|  	                    ||z   ¦  «        }||z   }
|
S )z’
        Args:
            hidden_state (`torch.Tensor`): Input tensor.

        Returns:
            `torch.Tensor`: Transformed tensor.
        F)rµ   r@   r   )
rn   rÉ   rr   rq   rË   rA   rŒ   rŠ   r�   rÌ   )r$   rf   r�   Ú
batch_sizeÚn_varsrN   r;   Úhidden_state_reshapedÚx_attnÚ_r‘   s              r&   r,   zPatchMixerBlock.forward…  s  € ð  ˆà—y’y Ñ.Ô.ˆàŒ>ð 	NØ7CÔ7IÑ4ˆJ˜ ¨WØ$0×$8Ò$8¸ÀfÑ9LÈkÐ[bÑ$cÔ$cÐ!à×/Ò/Ð0EÐY^Ð/Ñ_Ô_‰LˆF�A�qØ—^’^ J°¸ÀWÑMÔMˆFð $×-Ò-¨a°Ñ3Ô3ˆØ—x’x Ñ-Ô-ˆàŒ?ð 	;Ø×,Ò,¨\Ñ:Ô:ˆLð $×-Ò-¨a°Ñ3Ô3ˆàŒ>ð 	AØŸ>š>¨,¸Ñ*?Ñ@Ô@ˆLà˜XÑ%ˆØˆ
r'   ©r-   r.   r/   r0   r   r   r,   r2   r3   s   @r&   rÇ   rÇ   c  s^   ø€ € € € € ðð ð;Ð1ð ;ð ;ð ;ð ;ð ;ð ;ð4!ð !ð !ð !ð !ð !ð !r'   rÇ   c                   ó>   ‡ — e Zd ZdZdefˆ fd„Zdej        fd„Zˆ xZ	S )ÚFeatureMixerBlockz‚This module mixes the hidden feature dimension.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.

    r6   c                 ó  •— t          ¦   «                              ¦   «          t          |¦  «        | _        |j        | _        t          |j        |j        |¬¦  «        | _        |j        r"t          |j        |j        ¬¦  «        | _	        d S d S r‡   )
r   r   ri   rn   rŠ   ru   r;   rŒ   r   r�   r>   s     €r&   r   zFeatureMixerBlock.__init__²  s�   ø€ Ý‰Œ×ÒÑÔÐå)¨&Ñ1Ô1ˆŒ	à Ô+ˆŒå"ØœØœØð
ñ 
ô 
ˆŒð Ôð 	lÝ :À6Ä>Ð\bÔ\jÐ kÑ kÔ kˆDÔÐÐð	lð 	lr'   Úhiddenc                 ó    — |}|                       |¦  «        }|                      |¦  «        }| j        r|                      |¦  «        }||z   }|S )ú×
        Args:
            hidden (`torch.Tensor` of shape `(batch_size, num_patches, d_model)`):
                Input tensor to the layer.

        Returns:
            `torch.Tensor`: Transformed tensor.
        )rn   rŒ   rŠ   r�   )r$   r×   r�   r‘   s       r&   r,   zFeatureMixerBlock.forwardÂ  sW   € ð ˆØ—’˜6Ñ"Ô"ˆØ—’˜&Ñ!Ô!ˆàŒ?ð 	/Ø×&Ò& vÑ.Ô.ˆFà�xÑˆØˆ
r'   rC   r3   s   @r&   rÕ   rÕ   ©  ss   ø€ € € € € ðð ðlÐ1ð lð lð lð lð lð lð ˜eœlð ð ð ð ð ð ð ð r'   rÕ   c                   ó>   ‡ — e Zd ZdZdefˆ fd„Zdej        fd„Zˆ xZ	S )ÚPatchTSMixerLayerz•
    The `PatchTSMixer` layer that does all three kinds of mixing.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.

    r6   c                 óø   •— t          ¦   «                              ¦   «          t          |¬¦  «        | _        t	          |¬¦  «        | _        |j        | _        |j        dk    rt          |¬¦  «        | _        d S d S )N©r6   Úmix_channel)	r   r   rÇ   Úpatch_mixerrÕ   Úfeature_mixerÚmoder…   Úchannel_feature_mixerr>   s     €r&   r   zPatchTSMixerLayer.__init__à  sw   ø€ Ý‰Œ×ÒÑÔÐå*°&Ð9Ñ9Ô9ˆÔÝ.°fÐ=Ñ=Ô=ˆÔà”KˆŒ	àŒ;˜-Ò'Ð'Ý)MÐU[Ð)\Ñ)\Ô)\ˆDÔ&Ð&Ð&ð (Ð'r'   r×   c                 óš   — | j         dk    r|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }|S )rÙ   rÞ   )rá   râ   rß   rà   )r$   r×   s     r&   r,   zPatchTSMixerLayer.forwardë  sO   € ð Œ9˜Ò%Ð%Ø×/Ò/°Ñ7Ô7ˆFà×!Ò! &Ñ)Ô)ˆØ×#Ò# FÑ+Ô+ˆØˆr'   rC   r3   s   @r&   rÛ   rÛ   Ö  ss   ø€ € € € € ðð ð	]Ð1ð 	]ð 	]ð 	]ð 	]ð 	]ð 	]ð˜eœlð ð ð ð ð ð ð ð r'   rÛ   c                   ó6   ‡ — e Zd ZdZdefˆ fd„Zddefd„Zˆ xZS )ÚPatchTSMixerBlockz‹The main computing framework of the `PatchTSMixer` model.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.
    r6   c                 ó¼   •‡— t          ¦   «                              ¦   «          ‰j        }t          j        ˆfd„t          |¦  «        D ¦   «         ¦  «        | _        d S )Nc                 ó0   •— g | ]}t          ‰¬ ¦  «        ‘ŒS )rÝ   )rÛ   )Ú.0rÒ   r6   s     €r&   ú
<listcomp>z.PatchTSMixerBlock.__init__.<locals>.<listcomp>	  s%   ø€ Ð$aÐ$aÐ$aÈ!Õ%6¸fÐ%EÑ%EÔ%EÐ$aÐ$aÐ$ar'   )r   r   Ú
num_layersr   Ú
ModuleListÚrangeÚmixers)r$   r6   rê   r%   s    ` €r&   r   zPatchTSMixerBlock.__init__  sU   øø€ Ý‰Œ×ÒÑÔÐàÔ&ˆ
å”mÐ$aÐ$aÐ$aÐ$aÍuÐU_ÑO`ÔO`Ð$aÑ$aÔ$aÑbÔbˆŒˆˆr'   FÚoutput_hidden_statesc                 óv   — g }|}| j         D ]$} ||¦  «        }|r|                     |¦  «         Œ%|r||fS |dfS )as  
        Args:
            hidden_state (`torch.Tensor`): The input tensor.
            output_hidden_states (`bool`, *optional*, defaults to False.):
                Whether to output the hidden states as well.

        Returns:
            `torch.Tensor`: The embedding. `list`: List of all hidden states if `output_hidden_states` is set to
            `True`.
        N)rí   Úappend)r$   rf   rî   Úall_hidden_statesÚ	embeddingÚmods         r&   r,   zPatchTSMixerBlock.forward  sh   € ð Ðà ˆ	à”;ð 	4ð 	4ˆCØ˜˜I™œˆIØ#ð 4Ø!×(Ò(¨Ñ3Ô3Ð3øàð 	#ØÐ/Ð/Ð/à˜d�?Ð"r'   ©F)	r-   r.   r/   r0   r   r   rÄ   r,   r2   r3   s   @r&   rå   rå   ü  sv   ø€ € € € € ðð ðcÐ1ð cð cð cð cð cð cð#ð #¸$ð #ð #ð #ð #ð #ð #ð #ð #r'   rå   c                   ó0   ‡ — e Zd ZdZddefˆ fd„Zd„ Zˆ xZS )ÚPatchTSMixerForPredictionHeadzqPrediction Head for Forecasting

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.
    Nr6   c                 ó¼  •— t          ¦   «                              ¦   «          |j        | _        | j        �| j                             ¦   «          t	          j        |j        ¦  «        | _        |€-t	          j        |j	        |j
        z  |j        ¦  «        | _        n'|                     |j	        |j
        z  ¦  «        | _        t	          j        d¬¦  «        | _        d S )Néþÿÿÿ©Ú	start_dim)r   r   Úprediction_channel_indicesÚsortr   ry   Úhead_dropoutÚdropout_layerr    rN   r;   Úprediction_lengthÚbase_forecast_blockÚget_parameter_projectionÚFlattenÚflatten)r$   r6   Údistribution_outputr%   s      €r&   r   z&PatchTSMixerForPredictionHead.__init__-  sÁ   ø€ Ý‰Œ×ÒÑÔÐà*0Ô*KˆÔ'àÔ*Ð6ØÔ+×0Ò0Ñ2Ô2Ð2åœZ¨Ô(;Ñ<Ô<ˆÔØÐ&Ý')¤y°&Ô2DÀvÄ~Ñ2UÐX^ÔXpÑ'qÔ'qˆDÔ$Ð$à':×'SÒ'SØÔ" V¤^Ñ3ñ(ô (ˆDÔ$õ ”z¨BÐ/Ñ/Ô/ˆŒˆˆr'   c                 óž  ‡ — ‰                       |¦  «        }‰                      |¦  «        }‰                      |¦  «        }t          |t          ¦  «        rt	          d„ |D ¦   «         ¦  «        }n|                     dd¦  «        }‰ j        �@t          |t          ¦  «        rt	          ˆ fd„|D ¦   «         ¦  «        }n|d‰ j        f         }|S )ar  

        Args:
            hidden_features (`torch.Tensor` of shape `(batch_size, num_patch, d_model)` in `flatten` mode
                or `(batch_size, n_vars, num_patch, d_model)` in `common_channel`/`mix_channel` mode.): Input hidden
                features.

        Returns:
            `torch.Tensor` of shape `(batch_size, prediction_length, nvars)`.

        c              3   óB   K  — | ]}|                      d d¦  «        V — ŒdS )r   rø   N)rA   )rè   Úzs     r&   ú	<genexpr>z8PatchTSMixerForPredictionHead.forward.<locals>.<genexpr>P  s0   è è € ÐCÐC°Q˜QŸ[š[¨¨RÑ0Ô0ÐCÐCÐCÐCÐCÐCr'   r   rø   Nc              3   ó6   •K  — | ]}|d ‰j         f         V — ŒdS ).N)rû   )rè   r  r$   s     €r&   r  z8PatchTSMixerForPredictionHead.forward.<locals>.<genexpr>V  s0   øè è € Ð [Ð [ÈQ  3¨Ô(GÐ#GÔ!HÐ [Ð [Ð [Ð [Ð [Ð [r'   .)r  rþ   r   Ú
isinstancerÅ   rA   rû   ©r$   Úhidden_featuresÚforecasts   `  r&   r,   z%PatchTSMixerForPredictionHead.forward?  sÙ   ø€ ð Ÿ,š, Ñ7Ô7ˆØ×,Ò,¨_Ñ=Ô=ˆØ×+Ò+¨OÑ<Ô<ˆÝ�h¥Ñ&Ô&ð 	2ÝÐCÐC¸(ÐCÑCÔCÑCÔCˆHˆHà×)Ò)¨"¨bÑ1Ô1ˆHàÔ*Ð6Ý˜(¥EÑ*Ô*ð JÝ Ð [Ð [Ð [Ð [ÐRZÐ [Ñ [Ô [Ñ[Ô[��à# C¨Ô)HÐ$HÔI�àˆr'   r)   rÓ   r3   s   @r&   rö   rö   %  sc   ø€ € € € € ðð ð0ð 0Ð1ð 0ð 0ð 0ð 0ð 0ð 0ð$ð ð ð ð ð ð r'   rö   c                   ó0   ‡ — e Zd ZdZddefˆ fd„Zd„ Zˆ xZS )ÚPatchTSMixerLinearHeadz€Linear head for Classification and Regression.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.
    Nr6   c                 ó  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        |j        €|j        }nd}|| _        |€0t          j        |j        |j	        z  |z  |j
        ¦  «        | _        n*|                     |j        |j	        z  |z  ¦  «        | _        |j        €t          j        d¬¦  «        | _        nt          j        d¬¦  «        | _        t          j        |j        ¦  «        | _        d S )Nr   éýÿÿÿrù   rø   )r   r   Úhead_aggregationÚoutput_rangerN   r  r   r    r;   r‹   Únum_targetsÚ
projectionr  r  r  ry   rý   rz   )r$   r6   r  Ú
mul_factorr%   s       €r&   r   zPatchTSMixerLinearHead.__init__e  sú   ø€ Ý‰Œ×ÒÑÔÐà &Ô 7ˆÔØ"Ô/ˆÔàÔ"Ð*ØÔ+ˆJˆJàˆJØ#6ˆÔ ØÐ&Ý œiØ” Ô!:Ñ:¸ZÑGØÔ"ñô ˆDŒOˆOð
 2×JÒJØ” Ô!:Ñ:¸ZÑGñô ˆDŒOð Ô"Ð*Ýœ:°Ð3Ñ3Ô3ˆDŒLˆLåœ:°Ð3Ñ3Ô3ˆDŒLå”z &Ô"5Ñ6Ô6ˆŒˆˆr'   c                 ó  — |                      dd¦  «        }| j        dk    r	|d         }nH| j        dk    r|                     d¬¦  «        j        }n!| j        dk    r|                     d¬¦  «        }| j        r|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }| j        €E| j	        �>t          j        |¦  «        | j	        d	         | j	        d
         z
  z  | j	        d
         z   }|S )ai  
        Args:
            hidden_features (`torch.Tensor` of shape `(batch_size x num_patch x d_model)` in `flatten` mode
                or `(batch_size x n_vars x num_patch x d_model)` in `common_channel`/`mix_channel` mode.): Input hidden
                features.

        Returns:
            `torch.Tensor` of shape `(batch_size x num_targets)`.
        r   rø   Úuse_last).r   Úmax_poolr   Úavg_poolNr   r   )rA   r  ÚmaxÚvaluesr_   r  rz   r  r  r  rD   Úsigmoid)r$   r  s     r&   r,   zPatchTSMixerLinearHead.forward�  s  € ð *×3Ò3°B¸Ñ;Ô;ˆØÔ  JÒ.Ð.à-¨gÔ6ˆOˆOØÔ" jÒ0Ð0à-×1Ò1°bÐ1Ñ9Ô9Ô@ˆOˆOØÔ" jÒ0Ð0à-×2Ò2°rÐ2Ñ:Ô:ˆOàŒ<ð 	<Ø"Ÿlšl¨?Ñ;Ô;ˆOØŸ,š, Ñ7Ô7ˆØŸ/š/¨/Ñ:Ô:ˆàÔ$Ð,°4Ô3DÐ3På”˜oÑ.Ô.°$Ô2CÀAÔ2FÈÔIZÐ[\ÔI]Ñ2]Ñ^ÐaeÔarÐstÔauÑuð ð Ðr'   r)   rÓ   r3   s   @r&   r  r  ]  sc   ø€ € € € € ðð ð7ð 7Ð1ð 7ð 7ð 7ð 7ð 7ð 7ð8 ð  ð  ð  ð  ð  ð  r'   r  c                   ód   ‡ — e Zd ZU eed<   dZdZdZdZ e	j
        ¦   «         ˆ fd„¦   «         Zˆ xZS )ÚPatchTSMixerPreTrainedModelr6   ÚmodelÚpast_values)ÚtimeFc                 óz  •— t          ¦   «                              |¦  «         t          |t          ¦  «        r0| j        j        dk    rt          j        |j        dd¬¦  «         dS dS t          |t          ¦  «        r>t          j
        |j        j        ¦  «         t          j        |j        j        ¦  «         dS dS )zInitialize weightsrQ   r’   gš™™™™™¹?)r_   r`   N)r   Ú_init_weightsr
  rG   r6   rV   ÚinitÚnormal_rK   r5   Úzeros_r=   rª   Úones_Úweight)r$   r“   r%   s     €r&   r$  z)PatchTSMixerPreTrainedModel._init_weights­  s¸   ø€ õ 	‰Œ×Ò˜fÑ%Ô%Ð%Ý�fÕ<Ñ=Ô=ð 	0àŒ{Ô3°xÒ?Ð?Ý”˜VÔ0°sÀÐDÑDÔDÐDÐDÐDð @Ð?å˜Õ 5Ñ6Ô6ð 	0ÝŒK˜Ô(Ô-Ñ.Ô.Ð.ÝŒJ�vÔ'Ô.Ñ/Ô/Ð/Ð/Ð/ð	0ð 	0r'   )r-   r.   r/   r   Ú__annotations__Úbase_model_prefixÚmain_input_nameÚinput_modalitiesÚsupports_gradient_checkpointingrD   Úno_gradr$  r2   r3   s   @r&   r  r  ¤  sq   ø€ € € € € € ð ÐÐÑØÐØ#€OØ ÐØ&+Ð#à€U„]�_„_ð	0ð 	0ð 	0ð 	0ñ „_ð	0ð 	0ð 	0ð 	0ð 	0r'   r  c                   ó.   ‡ — e Zd ZdZdefˆ fd„Zd„ Zˆ xZS )ÚPatchTSMixerPretrainHeadzcPretraining head.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.
    r6   c                 óÌ   •— t          ¦   «                              ¦   «          t          j        |j        ¦  «        | _        t          j        |j        |j        ¦  «        | _	        d S r)   )
r   r   r   ry   rý   rþ   r    r;   Úpatch_lengthÚbase_pt_blockr>   s     €r&   r   z!PatchTSMixerPretrainHead.__init__Â  sM   ø€ Ý‰Œ×ÒÑÔÐåœZ¨Ô(;Ñ<Ô<ˆÔÝœY v¤~°vÔ7JÑKÔKˆÔÐÐr'   c                 óZ   — |                       |¦  «        }|                      |¦  «        }|S )a  
        Args:
            hidden_features (`torch.Tensor` of shape `(batch_size x num_patch x d_model)` in `flatten` mode
                or `(batch_size x n_vars x num_patch x d_model)` in `common_channel`/`mix_channel` mode.): Input hidden
                features.

        Returns:
            `torch.Tensor` of shape `(batch_size x n_vars x num_patch x patch_length)`.
        )rþ   r4  r  s      r&   r,   z PatchTSMixerPretrainHead.forwardÈ  s/   € ð ×,Ò,¨_Ñ=Ô=ˆØ×%Ò% oÑ6Ô6ˆØˆr'   rÓ   r3   s   @r&   r1  r1  º  se   ø€ € € € € ðð ðLÐ1ð Lð Lð Lð Lð Lð Lðð ð ð ð ð ð r'   r1  Fr*   Ú
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  ¦  «        }
|r0t	          j        |d||	¬¦  «        }|                     d|d¦  «        }nt	          j        ||||	¬¦  «        }t	          j        ||||	¬¦  «        }d|dd…dd…d|
…f<   t	          j        |d¬¦  «        }t	          j        |d¬¦  «        }t	          j	        |d|¬	¦  «        }| 
                    d¦  «                             ddd|¦  «        }|�d|dd…|dd…dd…f<   |                      |                     ¦   «         |¦  «        }||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   )r   Úindex©.r   )ra   rr   r<  r1   rD   ÚrandÚrepeatÚonesÚargsortÚgatherrY   Úmasked_fillrÄ   )r*   r6  r7  r8  r9  rÎ   Únum_channelsÚsequence_lengthÚnum_featuresr<  Úlen_keepÚnoiseÚmaskÚids_shuffleÚids_restoreÚinputs_masks                   r&   Úrandom_maskingrN  Ù  s›  € ð4 �A‚~€~˜ qš˜ÝÐN zÐNÐNÐNÑOÔOÐOà>D¼lÑ;€J�˜o¨|ØŒ]€Få�? a¨*¡nÑ5Ñ6Ô6€Hà!ð UÝ”
˜: q¨/À&ÐIÑIÔIˆØ—’˜Q ¨aÑ0Ô0ˆˆõ ”
˜: |°_ÈVÐTÑTÔTˆõ Œ:�j ,°ÈÐOÑOÔO€DØ€DˆˆˆˆAˆAˆAˆy�ˆyˆÑõ ”- ¨2Ð.Ñ.Ô.€KÝ”- °Ð4Ñ4Ô4€KåŒ<˜ "¨KÐ8Ñ8Ô8€DØ�>Š>˜"ÑÔ×$Ò$ Q¨¨1¨lÑ;Ô;€DØÐ+Ø23ˆˆQˆQˆQÐ(¨!¨!¨!¨Q¨Q¨QÐ.Ñ/à×$Ò$ T§Y¢Y¡[¤[°*Ñ=Ô=€KØ˜˜VœÐ$Ð$r'   Únum_forecast_mask_patchesc                 ó®  — t          |t          ¦  «        r|g}d„ |D ¦   «         }| j        \  }}}}t          j        |||| j        ¬¦  «        }	g }
d}t          |¦  «        }t          ||¦  «        D ]V\  }}|dk    s||k    rt          d|› d�¦  «        ‚t          ||z  |z  ¦  «        }|
 	                    |||g¦  «         ||z  }ŒWt          |
d„ ¬¦  «        }
||k     r|
d         d         ||z
  z   |
d         d<   n#||k    r|
d	         d         ||z
  z   |
d	         d<   d}|
D ]\  }}}||z   }d
|	||…dd…| d…f<   |}Œt          j        |	j        d         ¦  «        }|	|         }	|	                     d	¦  «                             d
d
d
|¦  «        }	|�d|	dd…|dd…dd…f<   |                      |	                     ¦   «         |¦  «        }||	d         fS )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)`
    c                 ó   — g | ]}d ‘ŒS )r   © )rè   rÒ   s     r&   ré   z$forecast_masking.<locals>.<listcomp>.  s   € ÐAÐAÐA !˜AÐAÐAÐAr'   r;  r   znum_forecast_mask_patches z6 should be greater than 0 and less than total patches.c                 ó   — | d         S ©Nr@   rR  )Úxs    r&   ú<lambda>z"forecast_masking.<locals>.<lambda>@  s
   € ¨!¨A¬$€ r'   )r•   r@   r   r   Nr>  )r
  r1   rr   rD   rM   r<  ÚsumÚzipra   rð   ÚsortedÚrandpermrY   r@  rD  rÄ   )r*   rO  r7  r9  Úforecast_mask_ratiosrÎ   rE  rF  rG  rJ  Út_listÚtotal_lengthÚtotal_ratior3  ÚratioÚtemp_lenÚbatch1Ú	patch_lenrÒ   Úbatch2ÚpermrM  s                         r&   Úforecast_maskingre    sU  € õ0 Ð+­SÑ1Ô1ð @Ø%>Ð$?Ð!ØAÐAÐ'@ÐAÑAÔAÐà>D¼lÑ;€J�˜o¨|ÝŒ;�z <°ÈÌÐWÑWÔW€Dà€FØ€LÝÐ*Ñ+Ô+€Kå"Ð#<Ð>RÑSÔSð !ð !Ñˆ�eØ˜1ÒÐ °Ò ?Ð ?ÝØq¨\ÐqÐqÐqñô ð õ �z EÑ)¨KÑ7Ñ8Ô8ˆØ�Š�| U¨HÐ5Ñ6Ô6Ð6Ø˜Ñ ˆˆå�F  Ð/Ñ/Ô/€Fà�jÒ Ð Ø˜a”y ”| z°LÑ'@ÑAˆˆqŒ	�!‰ˆØ	˜
Ò	"Ð	"Ø˜rœ
 1œ¨¸
Ñ)BÑCˆˆrŒ
�1‰à€FØ"(ð ð Ñˆ	�1�hØ˜(Ñ"ˆØ./ˆˆV�Fˆ]˜A˜A˜A 	˜z˜{˜{Ð*Ñ+ØˆˆåŒ>˜$œ* Qœ-Ñ(Ô(€DØ�Œ:€Dà�>Š>˜"ÑÔ×$Ò$ Q¨¨1¨lÑ;Ô;€DØÐ+Ø23ˆˆQˆQˆQÐ(¨!¨!¨!¨Q¨Q¨QÐ.Ñ/à×$Ò$ T§Y¢Y¡[¤[°*Ñ=Ô=€KØ˜˜VœÐ$Ð$r'   c                   ó>   ‡ — e Zd ZdZdefˆ fd„Zdej        fd„Zˆ xZ	S )ÚPatchTSMixerPatchifyz³
    A class to patchify the time series sequence into different patches

    Returns:
        `torch.Tensor` of shape `(batch_size, num_channels, num_patches, patch_length)`
    r6   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
  | _	        d S )NzSequence length (z+) has to be greater than the patch length (ú)r   )
r   r   Úcontext_lengthrF  r3  Úpatch_stridera   r  rN   Úsequence_start)r$   r6   Únew_sequence_lengthr%   s      €r&   r   zPatchTSMixerPatchify.__init__a  sá   ø€ Ý‰Œ×ÒÑÔÐà%Ô4ˆÔØ"Ô/ˆÔØ"Ô/ˆÔàÔ 4Ô#4Ò4Ð4ÝØy DÔ$8ÐyÐyÐeiÔevÐyÐyÐyñô ð õ
   Ô 4°dÔ6GÑHÔHÈ4ÔK\Ñ\ÐaeÔarÑrÐuvÑvˆÔØ"Ô/°$Ô2CÀtÔGWÐZ[ÑG[Ñ2\Ñ\ÐØ"Ô2Ð5HÑHˆÔÐÐr'   r!  c                 ó,  — |j         d         }|| j        k    rt          d|› d| j        › d�¦  «        ‚|dd…| j        d…dd…f         }|                     d| j        | j        ¬¦  «        }|                     dd¦  «                             ¦   «         }|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)`
        rø   zInput sequence length (z%) doesn't match model configuration (r­   N)Ú	dimensionrž   Ústepr  )	rr   rF  ra   rl  Úunfoldr3  rk  rA   r¡   )r$   r!  rF  rB   s       r&   r,   zPatchTSMixerPatchify.forwardr  s²   € ð &Ô+¨BÔ/ˆØ˜dÔ2Ò2Ð2ÝØx¨/ÐxÐxÐ`dÔ`tÐxÐxÐxñô ð ð ˜Q˜Q˜Q Ô 3Ð 5Ð 5°q°q°qÐ8Ô9ˆà—’¨°$Ô2CÈ$ÔJ[�Ñ\Ô\ˆà×!Ò! " bÑ)Ô)×4Ò4Ñ6Ô6ˆØˆr'   rC   r3   s   @r&   rg  rg  Y  ss   ø€ € € € € ðð ðIÐ1ð Ið Ið Ið Ið Ið Ið" 5¤<ð ð ð ð ð ð ð ð r'   rg  c                   ó>   ‡ — e Zd ZdZdefˆ fd„Zdej        fd„Zˆ xZ	S )ÚPatchTSMixerMaskinga”  
    Class to perform random or forecast masking.

    Parameters:
        config (`PatchTSMixerConfig`): 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
    r6   c                 ó  •— t          ¦   «                              ¦   «          |j        | _        |j        | _        |j        | _        |j        | _        |j        | _        |j        | _        | j        �t          | j        ¦  «        | _        d S d S r)   )	r   r   Úrandom_mask_ratior8  Ú	mask_typerO  r7  r9  rY  r>   s     €r&   r   zPatchTSMixerMasking.__init__—  sƒ   ø€ Ý‰Œ×ÒÑÔÐØ!'Ô!9ˆÔØ*0Ô*KˆÔ'ØÔ)ˆŒØ)/Ô)IˆÔ&Ø(.Ô(GˆÔ%Ø Ô+ˆŒØÔ(Ð4Ý,2°4Ô3PÑ,QÔ,QˆDÔ)Ð)Ð)ð 5Ð4r'   rd   c                 ó2  — | j         dk    r,t          || j        | j        | j        | j        ¬¦  «        \  }}nI| j         dk    r&t          || j        | j        | j        ¬¦  «        \  }}nt          d| j         › d�¦  «        ‚| 	                    ¦   «         }||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

        rQ   )r*   r6  r7  r8  r9  r  )r*   rO  r7  r9  zInvalid mask type ú.)
rv  rN  ru  r7  r8  r9  re  rO  ra   rÄ   )r$   rd   Úmasked_inputrJ  s       r&   r,   zPatchTSMixerMasking.forward¢  s¼   € ð Œ>˜XÒ%Ð%Ý!/Ø"ØÔ1Ø)-Ô)FØ+/Ô+JØœ?ð"ñ "ô "ÑˆL˜$˜$ð Œ^˜zÒ)Ð)Ý!1Ø"Ø*.Ô*HØ)-Ô)FØœ?ð	"ñ "ô "ÑˆL˜$˜$õ ÐC°$´.ÐCÐCÐCÑDÔDÐDð �yŠy‰{Œ{ˆØ˜TÐ!Ð!r'   rC   r3   s   @r&   rs  rs  Š  ss   ø€ € € € € ð
ð 
ð	RÐ1ð 	Rð 	Rð 	Rð 	Rð 	Rð 	Rð!" 5¤<ð !"ð !"ð !"ð !"ð !"ð !"ð !"ð !"r'   rs  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 )ÚPatchTSMixerStdScalerz½
    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.
    r6   c                 óü   •— t          ¦   «                              ¦   «          t          |d¦  «        r|j        nd| _        t          |d¦  «        r|j        nd| _        t          |d¦  «        r|j        nd| _        d S )NÚscaling_dimr   ÚkeepdimTÚminimum_scalegñhãˆµøä>)r   r   Úhasattrr}  r   r~  r  r>   s     €r&   r   zPatchTSMixerStdScaler.__init__Í  sy   ø€ Ý‰Œ×ÒÑÔÐÝ)0°¸Ñ)GÔ)GÐN�6Ô%Ð%ÈQˆŒÝ)0°¸Ñ)CÔ)CÐM�v”~�~ÈˆŒÝ5<¸VÀ_Ñ5UÔ5UÐ_˜VÔ1Ð1Ð[_ˆÔÐÐr'   ÚdataÚobserved_indicatorrO   c                 ód  — |                      | j        | j        ¬¦  «        }|                     d¦  «        }||z                        | j        | j        ¬¦  «        |z  }||z
  |z  dz                        | 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@   )rW  r   r~  Ú	clamp_minrD   Úsqrtr  )r$   r�  r‚  ÚdenominatorÚlocÚvarianceÚscales          r&   r,   zPatchTSMixerStdScaler.forwardÓ  s»   € ð )×,Ò,¨T¬X¸t¼|Ð,ÑLÔLˆØ!×+Ò+¨CÑ0Ô0ˆØÐ(Ñ(×-Ò-¨d¬hÀÄÐ-ÑMÔMÐP[Ñ[ˆà˜S‘jÐ$6Ñ6¸1Ñ<×AÒAÀ$Ä(ÐTXÔT`ÐAÑaÔaÐdoÑoˆÝ”
˜8 dÔ&8Ñ8Ñ9Ô9ˆØ�s‘
˜eÑ# S¨%Ð/Ð/r'   ©r-   r.   r/   r0   r   r   rD   rE   rÅ   r,   r2   r3   s   @r&   r{  r{  Ç  s˜   ø€ € € € € ðð ð
`Ð1ð `ð `ð `ð `ð `ð `ð0Ø”Lð0Ø6;´lð0à	ˆuŒ|˜Uœ\¨5¬<Ð7Ô	8ð0ð 0ð 0ð 0ð 0ð 0ð 0ð 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 )ÚPatchTSMixerMeanScalerzŠ
    Computes a scaling factor as the weighted average absolute value along the first dimension, and scales the data
    accordingly.
    r6   c                 ó8  •— t          ¦   «                              ¦   «          t          |d¦  «        r|j        nd| _        t          |d¦  «        r|j        nd| _        t          |d¦  «        r|j        nd| _        t          |d¦  «        r|j        nd | _        d S )Nr}  r   r~  Tr  ç»½×Ùß|Û=Údefault_scale)r   r   r€  r}  r   r~  r  r’  r>   s     €r&   r   zPatchTSMixerMeanScaler.__init__ñ  s™   ø€ Ý‰Œ×ÒÑÔÐÝ)0°¸Ñ)GÔ)GÐN�6Ô%Ð%ÈQˆŒÝ)0°¸Ñ)CÔ)CÐM�v”~�~ÈˆŒÝ5<¸VÀ_Ñ5UÔ5UÐ`˜VÔ1Ð1Ð[`ˆÔÝ5<¸VÀ_Ñ5UÔ5UÐ_˜VÔ1Ð1Ð[_ˆÔÐÐr'   r�  r‚  rO   c                 ó¨  — ||z                        ¦   «                              | j        d¬¦  «        }|                     | j        d¬¦  «        }|t          j        |d¬¦  «        z  }| j        €W|                     d¬¦  «        }t          j        |                     d¦  «        d¬¦  «        }t          j        ||z  ¦  «        }n| j        t          j        |¦  «        z  }t          j        |dk    ||¦  «        }t          j        || j	        ¬¦  «        }||z  }	| j
        s|                     | j        ¬¦  «        }|	t          j        |¦  «        |fS )r„  Tr…  r   ©ÚminNr   r   )ÚabsrW  r   rD   Úclampr’  ÚsqueezeÚ	ones_likeÚwherer  r~  Ú
zeros_like)
r$   r�  r‚  Úts_sumÚnum_observedrŒ  Ú	batch_sumÚbatch_observationsr’  Úscaled_datas
             r&   r,   zPatchTSMixerMeanScaler.forwardø  sE  € ð Ð+Ñ+×0Ò0Ñ2Ô2×6Ò6°t´xÈÐ6ÑNÔNˆØ)×-Ò-¨d¬hÀÐ-ÑEÔEˆà�œ \°qÐ9Ñ9Ô9Ñ9ˆð ÔÐ%ØŸ
š
 q˜
Ñ)Ô)ˆIÝ!&¤¨\×-=Ò-=¸aÑ-@Ô-@ÀaÐ!HÑ!HÔ!HÐÝ!œM¨)Ð6HÑ*HÑIÔIˆMˆMà Ô.µ´ÀÑ1GÔ1GÑGˆMõ ”˜L¨1Ò,¨e°]ÑCÔCˆõ ”˜E tÔ'9Ð:Ñ:Ô:ˆØ˜U‘lˆàŒ|ð 	0Ø—M’M d¤h�MÑ/Ô/ˆEà�EÔ,¨UÑ3Ô3°UÐ:Ð:r'   r�  r3   s   @r&   r�  r�  ë  s˜   ø€ € € € € ðð ð
`Ð1ð `ð `ð `ð `ð `ð `ð&;Ø”Lð&;Ø6;´lð&;à	ˆuŒ|˜Uœ\¨5¬<Ð7Ô	8ð&;ð &;ð &;ð &;ð &;ð &;ð &;ð &;r'   r�  c            
       óŠ   ‡ — e Zd ZdZdefˆ fd„Z	 d	dej        dej        dz  deej        ej        ej        f         fd„Z	ˆ xZ
S )
ÚPatchTSMixerNOPScalerz|
    Assigns a scaling factor equal to 1 along the first dimension, and therefore applies no scaling to the input data.
    r6   c                 óÀ   •— t          ¦   «                              ¦   «          t          |d¦  «        r|j        nd| _        t          |d¦  «        r|j        nd| _        d S )Nr}  r   r~  T)r   r   r€  r}  r   r~  r>   s     €r&   r   zPatchTSMixerNOPScaler.__init__'  sW   ø€ Ý‰Œ×ÒÑÔÐÝ)0°¸Ñ)GÔ)GÐN�6Ô%Ð%ÈQˆŒÝ)0°¸Ñ)CÔ)CÐM�v”~�~ÈˆŒˆˆr'   Nr�  r‚  rO   c                 óà   — t          j        |d¬¦  «                             | j        | j        ¬¦  «        }t          j        |d¬¦  «                             | 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)`)
        FrR   )r   r~  )rD   r™  r_   r   r~  r›  )r$   r�  r‚  rŒ  rŠ  s        r&   r,   zPatchTSMixerNOPScaler.forward,  sl   € õ ” °EÐ:Ñ:Ô:×?Ò?ÀDÄHÐVZÔVbÐ?ÑcÔcˆÝÔ˜t°5Ð9Ñ9Ô9×>Ò>À4Ä8ÐUYÔUaÐ>ÑbÔbˆØ�S˜%ÐÐr'   r)   r�  r3   s   @r&   r¢  r¢  "  s©   ø€ € € € € ðð ðNÐ1ð Nð Nð Nð Nð Nð Nð MQð ð  Ø”Lð Ø6;´lÀTÑ6Ið à	ˆuŒ|˜Uœ\¨5¬<Ð7Ô	8ð ð  ð  ð  ð  ð  ð  ð  r'   r¢  zS
    Base class for `PatchTSMixerEncoderOutput`, with potential hidden states.
    )Úcustom_introc                   ó\   — e Zd ZU dZdZej        dz  ed<   dZe	ej                 dz  ed<   dS )ÚPatchTSMixerEncoderOutputa-  
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, d_model)`):
        Hidden-state at the output of the last layer of the model.
    hidden_states (`tuple(torch.FloatTensor)`, *optional*):
        Hidden-states of the model at the output of each layer.
    NÚlast_hidden_stater³   )
r-   r.   r/   r0   r¨  rD   ÚFloatTensorr*  r³   rÅ   rR  r'   r&   r§  r§  =  sT   € € € € € € ðð ð 37Ð�uÔ(¨4Ñ/Ð6Ð6Ñ6Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ð9Ð9r'   r§  c            	       óp   ‡ — e Zd ZdZdefˆ fd„Zee	 d	dej	        de
dz  defd„¦   «         ¦   «         Zˆ xZS )
ÚPatchTSMixerEncoderz°
    Encoder for PatchTSMixer which inputs patched time-series and outputs patched embeddings.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.
    r6   c                 ó,  •— t          ¦   «                              |¦  «         t          j        |j        |j        ¦  «        | _        |j        rt          |¬¦  «        | _	        nd | _	        t          |¬¦  «        | _        |                      ¦   «          d S )NrÝ   )r   r   r   r    r3  r;   ÚpatcherrI   rG   Úpositional_encoderrå   Úmlp_mixer_encoderÚ	post_initr>   s     €r&   r   zPatchTSMixerEncoder.__init__X  s‡   ø€ Ý‰Œ×Ò˜Ñ Ô Ð å”y Ô!4°f´nÑEÔEˆŒØÔ)ð 	+Ý&DÈFÐ&SÑ&SÔ&SˆDÔ#Ð#à&*ˆDÔ#Ý!2¸&Ð!AÑ!AÔ!AˆÔð 	�ŠÑÔÐÐÐr'   Nr!  rî   rO   c                 óÚ   — |�|n| j         j        }|                      |¦  «        }| j        �|                      |¦  «        }|                      ||¬¦  «        \  }}t          ||¬¦  «        S )aÑ  
        past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
            Context values of the time series. For a pretraining task, this denotes the input time series to
            predict the masked portion. For a forecasting task, this denotes the history/past time series values.
            Similarly, for classification or regression tasks, it denotes the appropriate context values of the
            time series.

            For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series,
            it is greater than 1.

        Returns:
            `torch.FloatTensor` of shape `(batch_size, n_vars, num_patches, d_model)`
        N©rî   )r¨  r³   )r6   rî   r­  r®  r¯  r§  )r$   r!  rî   r™   Úpatchesr¨  r³   s          r&   r,   zPatchTSMixerEncoder.forwarde  s�   € ð. %9Ð$DÐ Ð È$Ì+ÔJjð 	ð
 —,’,˜{Ñ+Ô+ˆð Ô"Ð.Ø×-Ò-¨gÑ6Ô6ˆGà+/×+AÒ+AÀ'Ð`tÐ+AÑ+uÔ+uÑ(Ð˜=å(Ð;LÐ\iÐjÑjÔjÐjr'   r)   )r-   r.   r/   r0   r   r   r   r   rD   rE   rÄ   r§  r,   r2   r3   s   @r&   r«  r«  O  s²   ø€ € € € € ðð ðÐ1ð ð ð ð ð ð ð Øð -1ð!kð !kà”\ð!kð # T™kð!kð
 
#ð!kð !kð !kñ „^ñ Ôð!kð !kð !kð !kð !kr'   r«  zG
    Base class for model's outputs, with potential hidden states.
    c                   óÔ   — e Zd ZU dZdZej        dz  ed<   dZe	ej                 dz  ed<   dZ
ej        dz  ed<   dZej        dz  ed<   dZej        dz  ed<   dZej        dz  ed<   dS )	ÚPatchTSMixerModelOutputa  
    last_hidden_state (`torch.FloatTensor`  of shape `(batch_size, num_channels, num_patches, d_model)`):
        Hidden-state at the output of the last layer of the model.
    hidden_states (`tuple(torch.FloatTensor)`, *optional*):
        Hidden-states of the model at the output of each layer.
    patch_input (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches, patch_length)`):
        Patched input data to the model.
    mask (`torch.FloatTensor` of shape `(batch_size, num_channels, num_patches)`, *optional*):
        Bool Tensor indicating True in masked patches and False otherwise.
    loc (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*):
        Gives the mean of the context window per channel. Used for revin denorm outside the model, if revin
        enabled.
    scale (`torch.FloatTensor` of shape `(batch_size, 1, num_channels)`, *optional*):
        Gives the std dev of the context window per channel. Used for revin denorm outside the model, if revin
        enabled.
    Nr¨  r³   rd   rJ  rŠ  rŒ  )r-   r.   r/   r0   r¨  rD   r©  r*  r³   rÅ   rd   rJ  rŠ  rŒ  rR  r'   r&   rµ  rµ  ‹  s´   € € € € € € ðð ð" 37Ð�uÔ(¨4Ñ/Ð6Ð6Ñ6Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ø,0€K�Ô" TÑ)Ð0Ð0Ñ0Ø%)€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø$(€CˆÔ	˜TÑ	!Ð(Ð(Ñ(Ø&*€Eˆ5Ô˜tÑ#Ð*Ð*Ñ*Ð*Ð*r'   rµ  z=
    The PatchTSMixer Model for time-series forecasting.
    c                   óˆ   ‡ — e Zd Zddedefˆ fd„Zee	 	 ddej	        dej	        dz  dedz  d	e
fd
„¦   «         ¦   «         Zˆ xZS )ÚPatchTSMixerModelFr6   Ú
mask_inputc                 ó¼  •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          |¦  «        | _        |du rt          |¦  «        | _        nd| _        |j        dk    rt          |¦  «        | _
        n=|j        dk    s	|j        du rt          |¦  «        | _
        nt          |¦  «        | _
        |                      ¦   «          dS )z•
        mask_input (bool, *optional*, defaults to `False`):
            Whether to mask the input using the [`PatchTSMixerMasking`] module.
        TNr_   r`   )r   r   r«  Úencoderrg  Úpatchingrs  Úmaskingr˜   r�  Úscalerr{  r¢  r°  )r$   r6   r¸  r%   s      €r&   r   zPatchTSMixerModel.__init__±  sÍ   ø€ õ
 	‰Œ×Ò˜Ñ Ô Ð å*¨6Ñ2Ô2ˆŒÝ,¨VÑ4Ô4ˆŒà˜ÐÐÝ.¨vÑ6Ô6ˆDŒLˆLàˆDŒLàŒ>˜VÒ#Ð#Ý0°Ñ8Ô8ˆDŒKˆKØŒ^˜uÒ$Ð$¨¬¸$Ð(>Ð(>Ý/°Ñ7Ô7ˆDŒKˆKå/°Ñ7Ô7ˆDŒKð 	�ŠÑÔÐÐÐr'   Nr!  Úobserved_maskrî   rO   c                 óš  — |�|n| j         j        }d}|€t          j        |¦  «        }|                      ||¦  «        \  }}}|                      |¦  «        }	|	}
| j        �|                      |	¦  «        \  }
}|                      |
|¬¦  «        }t          |t          ¦  «        r	t          |Ž }t          |j        |j        |	|||¬¦  «        S )aë  
        past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
            Context values of the time series. For a pretraining task, this denotes the input time series to predict
            the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
            for classification or regression tasks, it denotes the appropriate context values of the time series.

            For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
            greater than 1.
        observed_mask (`torch.FloatTensor` 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).
        Nr²  )r¨  r³   rd   rJ  rŠ  rŒ  )r6   rî   rD   r™  r½  r»  r¼  rº  r
  rÅ   r§  rµ  r¨  r³   )r$   r!  r¾  rî   r™   rJ  Úscaled_past_valuesrŠ  rŒ  Ú	patched_xÚ	enc_inputÚencoder_outputs               r&   r,   zPatchTSMixerModel.forwardÊ  sõ   € ð0 %9Ð$DÐ Ð È$Ì+ÔJjð 	ð ˆØÐ Ý!œO¨KÑ8Ô8ˆMØ)-¯ª°[À-Ñ)PÔ)PÑ&Ð˜C à—M’MÐ"4Ñ5Ô5ˆ	àˆ	ØŒ<Ð#Ø"Ÿlšl¨9Ñ5Ô5‰OˆI�tð ŸšØØ!5ð &ñ 
ô 
ˆõ
 �n¥eÑ,Ô,ð 	HÝ6¸ÐGˆNå&Ø,Ô>Ø(Ô6Ø!ØØØð
ñ 
ô 
ð 	
r'   rô   ©NN)r-   r.   r/   r   rÄ   r   r   r   rD   rE   rµ  r,   r2   r3   s   @r&   r·  r·  «  s¼   ø€ € € € € ðð Ð1ð ¸tð ð ð ð ð ð ð2 Øð .2Ø,0ð	5
ð 5
à”\ð5
ð ”| dÑ*ð5
ð # T™kð	5
ð 
!ð5
ð 5
ð 5
ñ „^ñ Ôð5
ð 5
ð 5
ð 5
ð 5
r'   r·  z>
    Output type of [`PatchTSMixerForPreTrainingOutput`].
    c                   ó˜   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
eej                 dz  ed<   dS )Ú PatchTSMixerForPreTrainingOutputa@  
    loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`):
        Total loss
    prediction_outputs (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, patch_length)`):
        Prediction output from the pretrain head.
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`):
        Backbone embeddings before passing through the head.
    hidden_states (`tuple(torch.FloatTensor)`, *optional*):
        Hidden-states of the model at the output of each layer.
    NÚlossÚprediction_outputsr¨  r³   ©r-   r.   r/   r0   rÇ  rD   r©  r*  rÈ  r¨  r³   rÅ   rR  r'   r&   rÆ  rÆ    ó…   € € € € € € ð	ð 	ð &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø37Ð˜Ô)¨DÑ0Ð7Ð7Ñ7Ø26Ð�uÔ(¨4Ñ/Ð6Ð6Ñ6Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ð9Ð9r'   rÆ  z.
    `PatchTSMixer` for mask pretraining.
    c                   óˆ   ‡ — e Zd Zdefˆ fd„Zee	 	 	 ddej        dej        dz  de	dz  de	d	e
f
d
„¦   «         ¦   «         Zˆ xZS )ÚPatchTSMixerForPretrainingr6   c                 óà   •— t          ¦   «                              |¦  «         t          |d¬¦  «        | _        t	          |¬¦  «        | _        |j        | _        |                      ¦   «          d S )NT)r¸  rÝ   )r   r   r·  r   r1  ÚheadÚmasked_lossr°  r>   s     €r&   r   z#PatchTSMixerForPretraining.__init__"  sd   ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ý& v¸$Ð?Ñ?Ô?ˆŒ
Ý,°FÐ;Ñ;Ô;ˆŒ	Ø!Ô-ˆÔð 	�ŠÑÔÐÐÐr'   NTr!  r¾  rî   Úreturn_lossrO   c                 óp  — |�|n| j         j        }| j        du r!t          j                             d¬¦  «        }n t          j                             d¬¦  «        }|                      |||¬¦  «        }t          |t          ¦  «        r	t          |Ž }|  
                    |j        ¦  «        }|du r |||j        ¦  «        }	nd}	| j        du rO|	�M|	                     d¬¦  «        |j        z                       ¦   «         |j                             ¦   «         d	z   z  }	t!          |	||j        |j        ¬
¦  «        S )aT  
        past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
            Context values of the time series. For a pretraining task, this denotes the input time series to predict
            the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
            for classification or regression tasks, it denotes the appropriate context values of the time series.

            For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
            greater than 1.
        observed_mask (`torch.FloatTensor` 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_loss (`bool`,  *optional*):
            Whether to return the loss in the `forward` call.
        NTÚnone©Ú	reductionr_   ©r¾  rî   r   r   r‘  ©rÇ  rÈ  r¨  r³   )r6   rî   rÏ  rD   r   ÚMSELossr   r
  rÅ   rµ  rÎ  r¨  rd   r_   rJ  rW  rÆ  r³   )
r$   r!  r¾  rî   rÐ  r™   rÇ  Úmodel_outputÚx_hatÚloss_vals
             r&   r,   z"PatchTSMixerForPretraining.forward+  sY  € ð6 %9Ð$DÐ Ð È$Ì+ÔJjð 	ð Ô˜tÐ#Ð#Ý”8×#Ò#¨fÐ#Ñ5Ô5ˆDˆDå”8×#Ò#¨fÐ#Ñ5Ô5ˆDð —z’zØØ'Ø!5ð "ñ 
ô 
ˆõ
 �l¥EÑ*Ô*ð 	BÝ2°LÐAˆLà—	’	˜,Ô8Ñ9Ô9ˆà˜$ÐÐØ�t˜E <Ô#;Ñ<Ô<ˆHˆHàˆHð Ô˜tÐ#Ð#¨Ð(<Ø Ÿš¨"˜Ñ-Ô-°Ô0AÑA×FÒFÑHÔHÈLÔL]×LaÒLaÑLcÔLcÐfkÑLkÑlˆHå/ØØ$Ø*Ô<Ø&Ô4ð	
ñ 
ô 
ð 	
r'   ©NNT)r-   r.   r/   r   r   r   r   rD   rE   rÄ   rÆ  r,   r2   r3   s   @r&   rÌ  rÌ    s½   ø€ € € € € ðÐ1ð ð ð ð ð ð ð Øð .2Ø,0Ø ð:
ð :
à”\ð:
ð ”| dÑ*ð:
ð # T™kð	:
ð
 ð:
ð 
*ð:
ð :
ð :
ñ „^ñ Ôð:
ð :
ð :
ð :
ð :
r'   rÌ  z=
    Output type of [`PatchTSMixerForPredictionOutput`].
    c                   óÔ   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
eej                 dz  ed<   dZej        dz  ed<   dZej        dz  ed<   dS )	ÚPatchTSMixerForPredictionOutputaD  
    loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`):
        Total loss.
    prediction_outputs (`torch.FloatTensor` of shape `(batch_size, prediction_length, num_input_channels)`):
        Prediction output from the forecast head.
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`):
        Backbone embeddings before passing through the head.
    hidden_states (`tuple(torch.FloatTensor)`, *optional*):
        Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
    loc (`torch.FloatTensor`, *optional* of shape `(batch_size, 1, num_input_channels)`):
        Input mean
    scale (`torch.FloatTensor`, *optional* of shape `(batch_size, 1, num_input_channels)`):
        Input std dev
    NrÇ  rÈ  r¨  r³   rŠ  rŒ  )r-   r.   r/   r0   rÇ  rD   r©  r*  rÈ  r¨  r³   rÅ   rŠ  rŒ  rR  r'   r&   rÝ  rÝ  j  sµ   € € € € € € ðð ð &*€Dˆ%Ô
˜dÑ
"Ð)Ð)Ñ)Ø37Ð˜Ô)¨DÑ0Ð7Ð7Ñ7Ø26Ð�uÔ(¨4Ñ/Ð6Ð6Ñ6Ø59€M�5˜Ô*Ô+¨dÑ2Ð9Ð9Ñ9Ø$(€CˆÔ	˜TÑ	!Ð(Ð(Ñ(Ø&*€Eˆ5Ô˜tÑ#Ð*Ð*Ñ*Ð*Ð*r'   rÝ  zƒ
    Base class for time series model's predictions outputs that contains the sampled values from the chosen
    distribution.
    c                   ó2   — e Zd ZU dZdZej        dz  ed<   dS )Ú"SamplePatchTSMixerPredictionOutputú¨
    sequences (`torch.FloatTensor` of shape `(batch_size, num_samples, prediction_length, number_channels)`):
        Sampled values from the chosen distribution.
    NÚ	sequences©r-   r.   r/   r0   rá  rD   r©  r*  rR  r'   r&   rß  rß  ˆ  ó6   € € € € € € ðð ð
 +/€IˆuÔ  4Ñ'Ð.Ð.Ñ.Ð.Ð.r'   rß  c                   ó2   — e Zd ZU dZdZej        dz  ed<   dS )Ú"SamplePatchTSMixerRegressionOutputrà  Nrá  râ  rR  r'   r&   rå  rå  ˜  rã  r'   rå  ÚinputÚtargetrO   c                 ó.   — |                       |¦  «         S )zc
    Computes the negative log likelihood loss from input distribution with respect to target.
    )Úlog_prob)ræ  rç  s     r&   Únllrê  ©  s   € ð �NŠN˜6Ñ"Ô"Ð"Ð"r'   Úinput_tensorÚweightsc                 ón  — |�žt          j        |dk    | |z  t          j        | ¦  «        ¦  «        }t          j        |r|                     |¬¦  «        n|                     ¦   «         d¬¦  «        }|r|                     |¬¦  «        n|                     ¦   «         |z  S |                      |¬¦  «        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`.
    Nr   r   r†  r”  )rD   rš  r›  r—  rW  r_   )rë  rì  r   Úweighted_tensorÚsum_weightss        r&   Úweighted_averagerð  ±  s±   € ð  ÐÝœ+ g°¢l°LÀ7Ñ4JÍEÔL\Ð]iÑLjÔLjÑkÔkˆÝ”k¸#Ð"P '§+¢+°# +Ñ"6Ô"6Ð"6À7Ç;Â;Á=Ä=ÐVYÐZÑZÔZˆØ03ÐN�×#Ò#¨Ð#Ñ,Ô,Ð,¸×9LÒ9LÑ9NÔ9NÐR]Ñ]Ð]à× Ò  SÐ Ñ)Ô)Ð)r'   c                   óþ   ‡ — e Zd ZdZdefˆ fd„Zee	 	 	 	 ddej	        dej	        dz  dej	        dz  d	e
dz  d
e
defd„¦   «         ¦   «         Z ej        ¦   «         	 ddej	        dej	        dz  defd„¦   «         Zˆ xZS )ÚPatchTSMixerForPredictionz 
    `PatchTSMixer` for forecasting application.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.

    Returns:
        `None`.
    r6   c                 ó   •— t          ¦   «                              |¦  «         |j        | _        |j        | _        |j        | _        |j        dk    rd | _        na|j        }t          t          t          dœ}| 
                    |j        ¦  «        }|� ||¬¦  «        | _        nt          d|j        › �¦  «        ‚t          |¦  «        | _        t          || j        ¬¦  «        | _        |                      ¦   «          d S )NÚmse©Ú	student_tÚnormalÚnegative_binomialr   úUnknown distribution output ©r6   r  )r   r   rÇ  rû   Únum_parallel_samplesr  rÿ   r   r   r   Úgetra   r·  r   rö   rÎ  r°  )r$   r6   r   Údistribution_output_mapÚoutput_classr%   s        €r&   r   z"PatchTSMixerForPrediction.__init__Õ  s  ø€ Ý‰Œ×Ò˜Ñ Ô Ð Ø”KˆŒ	Ø*0Ô*KˆÔ'Ø$*Ô$?ˆÔ!àŒ;˜%ÒÐØ'+ˆDÔ$Ð$àÔ*ˆCå+Ý&Ý%;ð'ð 'Ð#ð
 3×6Ò6°vÔ7QÑRÔRˆLØÐ'Ø+7¨<¸CÐ+@Ñ+@Ô+@�Ô(Ð(å Ð!\ÀÔ@ZÐ!\Ð!\Ñ]Ô]Ð]å& vÑ.Ô.ˆŒ
Ý1ØØ $Ô 8ð
ñ 
ô 
ˆŒ	ð 	�ŠÑÔÐÐÐr'   NTr!  r¾  Úfuture_valuesrî   rÐ  rO   c                 ó‚  — | j         dk    rt          j        d¬¦  «        }n"| j         dk    rt          }nt	          d¦  «        ‚|�|n| j        j        }|                      |||¬¦  «        }t          |t          ¦  «        r	t          |Ž }|                      |j        ¦  «        }	d}
| j        �Ã| j        rp| j                             |	|j        d| j        f         |j        d| j        f         ¬	¦  «        }|�,|d
u r( |||d| j        f         ¦  «        }
t%          |
¦  «        }
nÀ|	|j        d| j        f         z  |j        d| j        f         z   }	|�|d
u r ||	|d| j        f         ¦  «        }
nt| j        rI| j                             |	|j        |j        ¬	¦  «        }|�|d
u r |||¦  «        }
t%          |
¦  «        }
n$|	|j        z  |j        z   }	|�|d
u r ||	|¦  «        }
| j        �)|j        d| j        f         }|j        d| j        f         }n|j        }|j        }t'          |
|	|j        |j        ||¬¦  «        S )aÖ  
        past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
            Context values of the time series. For a pretraining task, this denotes the input time series to predict
            the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
            for classification or regression tasks, it denotes the appropriate context values of the time series.

            For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
            greater than 1.
        observed_mask (`torch.FloatTensor` 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.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting,:
            `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*):
            Target values of the time series, that serve as labels for the model. The `future_values` is what the
            Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT
            required for a pretraining task.

            For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want
            to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter,
            pass the target data with all channels, as channel Filtering for both prediction and target will be
            manually applied before the loss computation.
        return_loss (`bool`,  *optional*):
            Whether to return the loss in the `forward` call.
        rô  r_   rÓ  rê  ú2Invalid loss function: Allowed values: mse and nllNrÕ  .©rŠ  rŒ  T)rÇ  rÈ  r¨  r³   rŠ  rŒ  )rÇ  r   r×  rê  ra   r6   rî   r   r
  rÅ   rµ  rÎ  r¨  rû   r  ÚdistributionrŠ  rŒ  rð  rÝ  r³   )r$   r!  r¾  rÿ  rî   rÐ  r™   rÇ  rØ  Úy_hatrÚ  r  rŠ  rŒ  s                 r&   r,   z!PatchTSMixerForPrediction.forwardó  sÜ  € ðJ Œ9˜ÒÐÝ”:¨Ð/Ñ/Ô/ˆDˆDØŒY˜%ÒÐÝˆDˆDåÐQÑRÔRÐRð %9Ð$DÐ Ð È$Ì+ÔJjð 	ð
 —z’zØØ'Ø!5ð "ñ 
ô 
ˆõ
 �l¥EÑ*Ô*ð 	BÝ2°LÐAˆLð —	’	˜,Ô8Ñ9Ô9ˆàˆØÔ*Ð6ØÔ'ð `Ø#Ô7×DÒDØØ$Ô(¨¨dÔ.MÐ)MÔNØ&Ô,¨S°$Ô2QÐ-QÔRð  Eñ  ô  �ð
 !Ð,°ÀÐ1DÐ1DØ#˜tØ$Ø% c¨4Ô+JÐ&JÔKñ ô  �Hõ
  0°Ñ9Ô9�Høð ˜LÔ.¨s°DÔ4SÐ/SÔTÑTØ"Ô& s¨DÔ,KÐ'KÔLñMð ð !Ð,°ÀÐ1DÐ1DØ#˜t E¨=¸¸dÔ>]Ð9]Ô+^Ñ_Ô_�HøàÔ'ð 
:Ø#Ô7×DÒDØ˜|Ô/°|Ô7Ið  Eñ  ô  �ð !Ð,°ÀÐ1DÐ1DØ#˜t L°-Ñ@Ô@�HÝ/°Ñ9Ô9�Høà Ô 2Ñ2°\Ô5EÑE�Ø Ð,°ÀÐ1DÐ1DØ#˜t E¨=Ñ9Ô9�HàÔ*Ð6ØÔ" 3¨Ô(GÐ#GÔHˆCØ Ô& s¨DÔ,KÐ'KÔLˆEˆEàÔ"ˆCØ Ô&ˆEå.ØØ$Ø*Ô<Ø&Ô4ØØð
ñ 
ô 
ð 	
r'   c                 ó
  ‡— | j         } | |d|d¬¦  «        }| j                             |j        |j        |j        ¬¦  «        Šˆfd„t          |¦  «        D ¦   «         }t          j        |d¬¦  «        }t          |¬¦  «        S )	aÀ  
        Generate sequences of sample predictions from a model with a probability distribution head.

        Args:
            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.

            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:
            [`SamplePatchTSMixerPredictionOutput`] where the outputs `sequences` tensor will have shape `(batch_size,
            number of samples, prediction_length, num_input_channels)`.
        NF)r!  rÿ  r¾  rî   r  c                 ó8   •— g | ]}‰                      ¦   «         ‘ŒS rR  ©Úsample©rè   rÒ   r  s     €r&   ré   z6PatchTSMixerForPrediction.generate.<locals>.<listcomp>Œ  s%   ø€ ÐNÐNÐN¨Q�<×&Ò&Ñ(Ô(ÐNÐNÐNr'   r   r   ©rá  )
rû  r  r  rÈ  rŠ  rŒ  rì   rD   Ústackrß  )r$   r!  r¾  rû  ÚoutputsÚsamplesr  s         @r&   Úgeneratez"PatchTSMixerForPrediction.generateb  s«   ø€ ð2  $Ô8Ðð �$Ø#ØØ'Ø!&ð	
ñ 
ô 
ˆð Ô/×<Ò<ØÔ&¨G¬K¸w¼}ð =ñ 
ô 
ˆð
 OÐNÐNÐNµ%Ð8LÑ2MÔ2MÐNÑNÔNˆõ ”+˜g¨1Ð-Ñ-Ô-ˆÝ1¸GÐDÑDÔDÐDr'   )NNNTr)   )r-   r.   r/   r0   r   r   r   r   rD   rE   rÄ   rÝ  r,   r/  rß  r  r2   r3   s   @r&   rò  rò  É  sJ  ø€ € € € € ð	ð 	ðÐ1ð ð ð ð ð ð ð< Øð .2Ø-1Ø,0Ø ðk
ð k
à”\ðk
ð ”| dÑ*ðk
ð ”| dÑ*ð	k
ð
 # T™kðk
ð ðk
ð 
)ðk
ð k
ð k
ñ „^ñ Ôðk
ðZ €U„]�_„_ð .2ð-Eð -Eà”\ð-Eð ”| dÑ*ð-Eð 
,ð	-Eð -Eð -Eñ „_ð-Eð -Eð -Eð -Eð -Er'   rò  zK
    Output type of [`PatchTSMixerForTimeSeriesClassificationOutput`].
    c                   ó˜   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
eej                 dz  ed<   dS )Ú-PatchTSMixerForTimeSeriesClassificationOutputaP  
    loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`):
        Total loss.
    prediction_outputs (`torch.FloatTensor` of shape `(batch_size, num_labels)`):
        Prediction output from the classification head.
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`):
        Backbone embeddings before passing through the head.
    hidden_states (`tuple(torch.FloatTensor)`, *optional*):
        Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
    NrÇ  rÈ  r¨  r³   rÉ  rR  r'   r&   r  r  “  rÊ  r'   r  c                   óŒ   ‡ — e Zd ZdZdefˆ fd„Zee	 	 	 ddej	        dej	        dz  de
dz  d	e
d
ef
d„¦   «         ¦   «         Zˆ xZS )Ú'PatchTSMixerForTimeSeriesClassificationz£
    `PatchTSMixer` for classification application.

    Args:
        config (`PatchTSMixerConfig`):
            Configuration.

    Returns:
        `None`.
    r6   c                 ó&  •— t          ¦   «                              |¦  «         t          |¦  «        | _        t	          |¬¦  «        | _        |j        dv r!t          |j        |j	        ¬¦  «        | _
        nd | _
        |                      ¦   «          d S )NrÝ   ©r`   r_   T©r;   rN   )r   r   r·  r   r  rÎ  r˜   ÚInjectScalerStatistics4Dr;   rN   Úinject_scaler°  r>   s     €r&   r   z0PatchTSMixerForTimeSeriesClassification.__init__·  s‘   ø€ Ý‰Œ×Ò˜Ñ Ô Ð å& vÑ.Ô.ˆŒ
Ý*Øð
ñ 
ô 
ˆŒ	ð Œ>Ð2Ð2Ð2Ý 8ÀÄÐ]cÔ]oÐ pÑ pÔ pˆDÔÐà $ˆDÔð 	�ŠÑÔÐÐÐr'   NTr!  Útarget_valuesrî   rÐ  rO   c                 óÆ  — t           j                             ¦   «         }|�|n| j        j        }|                      ||¬¦  «        }t          |t          ¦  «        r	t          |Ž }| j	        �,|  	                    |j
        |j        |j        ¬¦  «        |_
        |                      |j
        ¦  «        }|�|du r |||¦  «        }	nd}	t          |	||j
        |j        ¬¦  «        S )að  
        past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
            Context values of the time series. For a pretraining task, this denotes the input time series to predict
            the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
            for classification or regression tasks, it denotes the appropriate context values of the time series.

            For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
            greater than 1.
        target_values (`torch.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting,
            `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*):
            Target
            values of the time series, that serve as labels for the model. The `target_values` is what the
            Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT
            required for a pretraining task.

            For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want
            to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter,
            pass the target data with all channels, as channel Filtering for both prediction and target will be
            manually applied before the loss computation.

            For a classification task, it has a shape of `(batch_size,)`.

            For a regression task, it has a shape of `(batch_size, num_targets)`.
        return_loss (`bool`, *optional*):
            Whether to return the loss in the `forward` call.
        Nr²  r  TrÖ  )rD   r   ÚCrossEntropyLossr6   rî   r   r
  rÅ   rµ  r  r¨  rŠ  rŒ  rÎ  r  r³   )
r$   r!  r  rî   rÐ  r™   rÇ  rØ  r  rÚ  s
             r&   r,   z/PatchTSMixerForTimeSeriesClassification.forwardÆ  s  € õJ Œx×(Ò(Ñ*Ô*ˆð %9Ð$DÐ Ð È$Ì+ÔJjð 	ð —z’zØØ!5ð "ñ 
ô 
ˆõ �l¥EÑ*Ô*ð 	BÝ2°LÐAˆLàÔÐ(Ø-1×->Ò->ØÔ.Ø Ô$Ø"Ô(ð .?ñ .ô .ˆLÔ*ð —	’	˜,Ô8Ñ9Ô9ˆàÐ$¨¸Ð)<Ð)<Ø�t˜E =Ñ1Ô1ˆHˆHàˆHå<ØØ$Ø*Ô<Ø&Ô4ð	
ñ 
ô 
ð 	
r'   rÛ  )r-   r.   r/   r0   r   r   r   r   rD   rE   rÄ   r  r,   r2   r3   s   @r&   r  r  «  sÕ   ø€ € € € € ð	ð 	ðÐ1ð ð ð ð ð ð ð Øð .2Ø,0Ø ðC
ð C
à”\ðC
ð ”| dÑ*ðC
ð # T™kð	C
ð
 ðC
ð 
7ðC
ð C
ð C
ñ „^ñ ÔðC
ð C
ð C
ð C
ð C
r'   r  z=
    Output type of [`PatchTSMixerForRegressionOutput`].
    c                   ó˜   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dZ	ej        dz  ed<   dZ
eej                 dz  ed<   dS )ÚPatchTSMixerForRegressionOutputaM  
    loss (*optional*, returned when `y` is provided, `torch.FloatTensor` of shape `()`):
        Total loss.
    regression_outputs (`torch.FloatTensor` of shape `(batch_size, num_targets)`):
        Prediction output from the regression head.
    last_hidden_state (`torch.FloatTensor` of shape `(batch_size, num_input_channels, num_patches, d_model)`):
        Backbone embeddings before passing through the head.
    hidden_states (`tuple(torch.FloatTensor)`, *optional*):
        Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
    NrÇ  Úregression_outputsr¨  r³   )r-   r.   r/   r0   rÇ  rD   r©  r*  r  r¨  r³   rÅ   rR  r'   r&   r  r    rÊ  r'   r  c                   ó`   ‡ — e Zd Zd
dededefˆ fd„Zdej        dej        dej        fd	„Zˆ xZS )r  r@   r;   rN   Ú	expansionc                 óD  •— t          ¦   «                              ¦   «          t          j        |dz   ||z  ¦  «        | _        t          j        ||z  |¦  «        | _        t          j        dd|z  ¦  «        | _        t          j        d|z  d¦  «        | _        || _        d S rT  )	r   r   r   r    Úinverse_trans_expansionÚinverse_trans_compressionÚmap_scale_expansionÚmap_scale_compressionrN   )r$   r;   rN   r  r%   s       €r&   r   z!InjectScalerStatistics4D.__init__'  s�   ø€ Ý‰Œ×ÒÑÔÐå')¤y°¸1±¸iÈ'Ñ>QÑ'RÔ'RˆÔ$Ý)+¬°9¸wÑ3FÈÑ)PÔ)PˆÔ&Ý#%¤9¨Q°°I±Ñ#>Ô#>ˆÔ Ý%'¤Y¨q°9©}¸aÑ%@Ô%@ˆÔ"Ø&ˆÔÐÐr'   r*   rŠ  rŒ  c                 ó.  — |                      dd¦  «        }|                     d¦  «        }|                     dd| j        d¦  «        }|                      dd¦  «        }|                     d¦  «        }|                     dd| j        d¦  «        }t	          j        ||gd¬¦  «        }|                      |¦  «        }|                      |¦  «        }t	          j        ||gd¬¦  «        }|                      |¦  «        }|  	                    |¦  «        }|S )a‰  
        Args:
            inputs (`torch.Tensor` of shape `(batch_size, num_input_channels, num_patch, d_model)`)
            loc (`torch.Tensor` of shape `(batch_size, 1, num_input_channels)`)
            scale (`torch.Tensor` of shape `(batch_size, 1, num_input_channels)`)
        Returns:
            `torch.Tensor` of shape `(batch_size, num_input_channels, num_patch, d_model)`
        r   rø   r   r   )
rA   rY   r@  rN   rD   Úcatr#  r$  r!  r"  )r$   r*   rŠ  rŒ  r_   ÚstdevÚconcat_statss          r&   r,   z InjectScalerStatistics4D.forward0  s  € ð �}Š}˜R Ñ$Ô$ˆØ�~Š~˜bÑ!Ô!ˆØ�{Š{˜1˜a Ô!1°1Ñ5Ô5ˆà—’  BÑ'Ô'ˆØ—’ Ñ#Ô#ˆØ—’˜Q  4Ô#3°QÑ7Ô7ˆå”y $¨ °BÐ7Ñ7Ô7ˆà×/Ò/°Ñ=Ô=ˆØ×1Ò1°,Ñ?Ô?ˆå”˜F LÐ1°rÐ:Ñ:Ô:ˆØ×-Ò-¨fÑ5Ô5ˆØ×/Ò/°Ñ7Ô7ˆàˆr'   )r@   )	r-   r.   r/   r1   r   rD   rE   r,   r2   r3   s   @r&   r  r  &  s†   ø€ € € € € ð'ð ' ð '°#ð 'À#ð 'ð 'ð 'ð 'ð 'ð 'ð˜eœlð °´ð ÀeÄlð ð ð ð ð ð ð ð r'   r  z4
    `PatchTSMixer` for regression application.
    c                   óÌ   ‡ — e Zd Zdefˆ fd„Zee	 	 	 ddej        dej        dz  de	dz  de	d	e
f
d
„¦   «         ¦   «         Z ej        ¦   «         dej        d	efd„¦   «         Zˆ xZS )ÚPatchTSMixerForRegressionr6   c                 ó^  •— t          ¦   «                              |¦  «         t          |¦  «        | _        |j        | _        |j        | _        |j        | _        |j        dk    rd | _        n_t          t          t          dœ}| 
                    |j        ¦  «        }|� ||j        ¬¦  «        | _        nt          d|j        › �¦  «        ‚|j        dv r!t          |j        |j        ¬¦  «        | _        nd | _        t%          || j        ¬¦  «        | _        |                      ¦   «          d S )Nrô  rõ  r   rù  r  r  rú  )r   r   r·  r   rÇ  r  rû  r   r   r   rü  r  ra   r˜   r  r;   rN   r  r  rÎ  r°  )r$   r6   rý  rþ  r%   s       €r&   r   z"PatchTSMixerForRegression.__init__T  s7  ø€ Ý‰Œ×Ò˜Ñ Ô Ð å& vÑ.Ô.ˆŒ
à”KˆŒ	Ø#)Ô#=ˆÔ à$*Ô$?ˆÔ!àŒ;˜%ÒÐØ'+ˆDÔ$Ð$õ ,Ý&Ý%;ð'ð 'Ð#ð
 3×6Ò6°vÔ7QÑRÔRˆLØÐ'Ø+7¨<¸FÔ<NÐ+OÑ+OÔ+O�Ô(Ð(å Ð!\ÀÔ@ZÐ!\Ð!\Ñ]Ô]Ð]àŒ>Ð2Ð2Ð2Ý 8ÀÄÐ]cÔ]oÐ pÑ pÔ pˆDÔÐà $ˆDÔå*ØØ $Ô 8ð
ñ 
ô 
ˆŒ	ð 	�ŠÑÔÐÐÐr'   NTr!  r  rî   rÐ  rO   c                 ó&  ‡ — ‰ j         dk    rt          j        d¬¦  «        }n"‰ j         dk    rt          }nt	          d¦  «        ‚|�|n‰ j        j        }‰                      ||¬¦  «        }t          |t          ¦  «        r	t          |Ž }‰ j        �,‰                      |j        |j        |j        ¬¦  «        |_        ‰                      |j        ¦  «        }|�›|d	u r—‰ j        rƒ‰ j        d
k    r't#          j        |dk     ¦  «        rt'          d¦  «        ‚‰ j                             |¦  «        }	t          ˆ fd„|D ¦   «         ¦  «        } ||	|¦  «        }
t+          |
¦  «        }
n |||¦  «        }
nd}
t-          |
||j        |j        ¬¦  «        S )aä  
        past_values (`torch.FloatTensor` of shape `(batch_size, seq_length, num_input_channels)`):
            Context values of the time series. For a pretraining task, this denotes the input time series to predict
            the masked portion. For a forecasting task, this denotes the history/past time series values. Similarly,
            for classification or regression tasks, it denotes the appropriate context values of the time series.

            For univariate time series, `num_input_channels` dimension should be 1. For multivariate time series, it is
            greater than 1.
        target_values (`torch.FloatTensor` of shape `(batch_size, target_len, num_input_channels)` for forecasting,
            `(batch_size, num_targets)` for regression, or `(batch_size,)` for classification, *optional*):
            Target values of the time series, that serve as labels for the model. The `target_values` is what the
            Transformer needs during training to learn to output, given the `past_values`. Note that, this is NOT
            required for a pretraining task.

            For a forecasting task, the shape is be `(batch_size, target_len, num_input_channels)`. Even if we want
            to forecast only specific channels by setting the indices in `prediction_channel_indices` parameter,
            pass the target data with all channels, as channel Filtering for both prediction and target will be
            manually applied before the loss computation.

            For a classification task, it has a shape of `(batch_size,)`.

            For a regression task, it has a shape of `(batch_size, num_targets)`.
        return_loss (`bool`, *optional*):
            Whether to return the loss in the `forward` call.
        rô  r_   rÓ  rê  r  Nr²  r  Trø  r   zDtarget_values cannot be negative for negative_binomial distribution.c              3   óX   •K  — | ]$}|                      d ‰j        j        ¦  «        V — Œ%dS )r   N)r·   r6   r  )rè   Úitemr$   s     €r&   r  z4PatchTSMixerForRegression.forward.<locals>.<genexpr>½  s6   øè è € ÐWÐWÈ˜dŸiši¨¨D¬KÔ,CÑDÔDÐWÐWÐWÐWÐWÐWr'   )rÇ  r  r¨  r³   )rÇ  r   r×  rê  ra   r6   rî   r   r
  rÅ   rµ  r  r¨  rŠ  rŒ  rÎ  r  rD   ÚanyÚ	Exceptionr  rð  r  r³   )r$   r!  r  rî   rÐ  r™   rÇ  rØ  r  r  rÚ  s   `          r&   r,   z!PatchTSMixerForRegression.forwardy  sÚ  ø€ ðH Œ9˜ÒÐÝ”:¨Ð/Ñ/Ô/ˆDˆDØŒY˜%ÒÐÝˆDˆDåÐQÑRÔRÐRð %9Ð$DÐ Ð È$Ì+ÔJjð 	ð —z’zØØ!5ð "ñ 
ô 
ˆõ �l¥EÑ*Ô*ð 	BÝ2°LÐAˆLàÔÐ(Ø-1×->Ò->ØÔ.Ø Ô$Ø"Ô(ð .?ñ .ô .ˆLÔ*ð —	’	˜,Ô8Ñ9Ô9ˆàÐ$¨¸Ð)<Ð)<ØÔ'ð 
6ØÔ+Ð/BÒBÐBÅuÄyÐQ^ÐabÒQbÑGcÔGcÐBÝ#Ð$jÑkÔkÐkØ#Ô7×DÒDÀUÑKÔK�åÐWÐWÐWÐWÐQVÐWÑWÔWÑWÔW�Ø˜4 ¨mÑ<Ô<�å+¨HÑ5Ô5��à˜4  }Ñ5Ô5��àˆHå.ØØ$Ø*Ô<Ø&Ô4ð	
ñ 
ô 
ð 	
r'   c                 ó,  ‡— | j         } | |dd¬¦  «        }| j                             |j        ¦  «        Šˆfd„t	          |¦  «        D ¦   «         }t          j        |d¬¦  «                             d|| j        j	        ¦  «        }t          |¬¦  «        S )	a
  
        Generate sequences of sample predictions from a model with a probability distribution head.

        Args:
            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 target values.

        Return:
            [`SamplePatchTSMixerRegressionOutput`] where the outputs `sequences` tensor will have shape `(batch_size,
            number of samples, num_targets)`.
        NF)r!  r  rî   c                 ó8   •— g | ]}‰                      ¦   «         ‘ŒS rR  r  r	  s     €r&   ré   z6PatchTSMixerForRegression.generate.<locals>.<listcomp>ë  s2   ø€ ð 
ð 
ð 
Ø&'ˆL×ÒÑ!Ô!ð
ð 
ð 
r'   r   r   r   r
  )rû  r  r  r  rì   rD   r  r·   r6   r  rå  )r$   r!  rû  r  r  r  s        @r&   r  z"PatchTSMixerForRegression.generateÍ  sº   ø€ ð"  $Ô8Ðð �$Ø#ØØ!&ð
ñ 
ô 
ˆð Ô/×<Ò<¸WÔ=WÑXÔXˆð
ð 
ð 
ð 
Ý+0Ð1EÑ+FÔ+Fð
ñ 
ô 
ˆõ
 ”+˜g¨1Ð-Ñ-Ô-×2Ò2°2Ð7KÈTÌ[ÔMdÑeÔeˆÝ1¸GÐDÑDÔDÐDr'   rÛ  )r-   r.   r/   r   r   r   r   rD   rE   rÄ   r  r,   r/  rå  r  r2   r3   s   @r&   r*  r*  N  s  ø€ € € € € ð#Ð1ð #ð #ð #ð #ð #ð #ðJ Øð .2Ø,0Ø ðP
ð P
à”\ðP
ð ”| dÑ*ðP
ð # T™kð	P
ð
 ðP
ð 
)ðP
ð P
ð P
ñ „^ñ ÔðP
ðd €U„]�_„_ð#Eà”\ð#Eð 
,ð#Eð #Eð #Eñ „_ð#Eð #Eð #Eð #Eð #Er'   r*  )r  r·  rÌ  rò  r  r*  )Nr’   )NFr   )Nr   rÄ  )Ur0   r[   Úcollections.abcr   Údataclassesr   rD   Útorch.nnr   Útransformers.modeling_utilsr   Útransformers.utilsr   Ú r   r%  Úmodeling_flash_attention_utilsr	   Úmodeling_utilsr
   Úprocessing_utilsr   Útime_series_utilsr   r   r   Úutilsr   r   r   r   Úconfiguration_patchtsmixerr   Ú
get_loggerr-   ÚloggerÚModuler   r5   rG   ri   ru   r…   rE   rÃ   r¤   r¦   rÇ   rÕ   rÛ   rå   rö   r  r  r1  ÚlistrÄ   r1   rN  re  rg  rs  r{  r�  r¢  r§  r«  rµ  r·  rÆ  rÌ  rÝ  rß  rå  ÚdistributionsÚDistributionrê  rð  rò  r  r  r  r  r*  Ú__all__rR  r'   r&   ú<module>rF     sA  ðð "Ð !à €€€Ø $Ð $Ð $Ð $Ð $Ð $Ø !Ð !Ð !Ð !Ð !Ð !à €€€Ø Ð Ð Ð Ð Ð à 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø *Ð *Ð *Ð *Ð *Ð *à &Ð &Ð &Ð &Ð &Ð &Ø BÐ BÐ BÐ BÐ BÐ BØ 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø &Ð &Ð &Ð &Ð &Ð &Ø UÐ UÐ UÐ UÐ UÐ UÐ UÐ UÐ UÐ UØ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RØ :Ð :Ð :Ð :Ð :Ð :ð 
ˆÔ	˜HÑ	%Ô	%€ðð ð ð ð  ¤ñ ô ð ð*&ð &ð &ð &ð &˜BœIñ &ô &ð &ð,$ð $ð $ð $ð $ R¤Yñ $ô $ð $ðN.ð .ð .ð .ð .˜BœIñ .ô .ð .ðbð ð ð ð �b”iñ ô ð ð.-ð -ð -ð -ð -¨2¬9ñ -ô -ð -ðn !Øð%ð %ØŒIð%àŒ<ð%ð 
Œð%ð Œ<ð	%ð
 ”L 4Ñ'ð%ð �T‰\ð%ð ð%ð Ð'Ô(ð%ð %ð %ð %ð:R/ð R/ð R/ð R/ð R/˜BœIñ R/ô R/ð R/ðjCð Cð Cð Cð C�b”iñ Cô Cð CðL*ð *ð *ð *ð *˜œ	ñ *ô *ð *ðZ#ð #ð #ð #ð #˜œ	ñ #ô #ð #ðL&#ð &#ð &#ð &#ð &#˜œ	ñ &#ô &#ð &#ðR5ð 5ð 5ð 5ð 5 B¤Iñ 5ô 5ð 5ðpDð Dð Dð Dð D˜RœYñ Dô Dð DðN ð0ð 0ð 0ð 0ð 0 /ñ 0ô 0ñ „ð0ð*ð ð ð ð ˜rœyñ ô ð ðD -1Ø',Øð7%ð 7%ØŒLð7%àð7%ð # T™kð7%ð !%ð	7%ð
 ð7%ð 7%ð 7%ð 7%ð| -1Øð	A%ð A%ØŒLðA%à# c™zðA%ð # T™kðA%ð ð	A%ð A%ð A%ð A%ðJ-ð -ð -ð -ð -˜2œ9ñ -ô -ð -ðb9"ð 9"ð 9"ð 9"ð 9"˜"œ)ñ 9"ô 9"ð 9"ðz 0ð  0ð  0ð  0ð  0˜BœIñ  0ô  0ð  0ðH3;ð 3;ð 3;ð 3;ð 3;˜RœYñ 3;ô 3;ð 3;ðn ð  ð  ð  ð  ˜BœIñ  ô  ð  ð6 €ððñ ô ð
 ð	:ð 	:ð 	:ð 	:ð 	: ñ 	:ô 	:ñ „ñô ð	:ð9kð 9kð 9kð 9kð 9kÐ5ñ 9kô 9kð 9kðx €ððñ ô ð
 ð+ð +ð +ð +ð +˜kñ +ô +ñ „ñô ð+ð4 €ððñ ô ð
Q
ð Q
ð Q
ð Q
ð Q
Ð3ñ Q
ô Q
ñô ð
Q
ðh €ððñ ô ð
 ð:ð :ð :ð :ð : {ñ :ô :ñ „ñô ð:ð$ €ððñ ô ð
F
ð F
ð F
ð F
ð F
Ð!<ñ F
ô F
ñô ð
F
ðR €ððñ ô ð
 ð+ð +ð +ð +ð + kñ +ô +ñ „ñô ð+ð0 €ððñ ô ð ð/ð /ð /ð /ð /¨ñ /ô /ñ „ñô ð/ð €ððñ ô ð ð/ð /ð /ð /ð /¨ñ /ô /ñ „ñô ð/ð#ˆuÔ"Ô/ð #¸¼ð #È%Ì,ð #ð #ð #ð #ð*ð * 5¤<ð *¸%¼,ÈÑ:Mð *ÐchÔcoð *ð *ð *ð *ð0GEð GEð GEð GEð GEÐ ;ñ GEô GEð GEðT €ððñ ô ð
 ð:ð :ð :ð :ð :°Kñ :ô :ñ „ñô ð:ð$`
ð `
ð `
ð `
ð `
Ð.Iñ `
ô `
ð `
ðF €ððñ ô ð
 ð:ð :ð :ð :ð : kñ :ô :ñ „ñô ð:ð$%ð %ð %ð %ð %˜rœyñ %ô %ð %ðP €ððñ ô ð
^Eð ^Eð ^Eð ^Eð ^EÐ ;ñ ^Eô ^Eñô ð
^EðBð ð €€€r'   