§
    ‚Štjø€  ã                   óh  — d dl Z d dlmZmZ d dlmZ d dlZd dlmZ d dl	mc m
Z ddlmZ ddlmZ ddlmZ ddlmZ dd	lmZmZ dd
lmZ ddlmZmZmZmZ ddlm Z  ddl!m"Z" ddl#m$Z$  ej%        e&¦  «        Z'ee G d„ de¦  «        ¦   «         ¦   «         Z(ee G d„ de¦  «        ¦   «         ¦   «         Z) G d„ dej*        ¦  «        Z+ G d„ dej*        ¦  «        Z, ed¦  «         G d„ dej*        ¦  «        ¦   «         Z- G d„ dej*        ¦  «        Z.	 d2dej*        dej/        d ej/        d!ej/        d"ej/        dz  d#e0d$e0e1z  d%ee         fd&„Z2 G d'„ d(ej*        ¦  «        Z3 G d)„ d*ej*        ¦  «        Z4e G d+„ d,e¦  «        ¦   «         Z5e G d-„ d.e5¦  «        ¦   «         Z6 G d/„ d0e5¦  «        Z7g d1¢Z8dS )3é    N)ÚCallableÚSequence)Ú	dataclassé   )Úinitialization)Úuse_kernel_forward_from_hub)ÚFlashAttentionKwargs)ÚBaseModelOutput)ÚALL_ATTENTION_FUNCTIONSÚPreTrainedModel)ÚUnpack)ÚTransformersKwargsÚauto_docstringÚcan_return_tupleÚlogging)Úmerge_with_config_defaults)Úcapture_outputsé   )ÚTimesFmConfigc                   óP   — e Zd ZU dZdZej        dz  ed<   dZej        dz  ed<   dS )ÚTimesFmOutputzÇ
    loc (`torch.Tensor` of shape `(batch_size, )`):
        The mean of the time series inputs.
    scale (`torch.Tensor` of shape `(batch_size,)`):
        The scale of the time series inputs.
    NÚlocÚscale)	Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚtorchÚTensorÚ__annotations__r   © ó    új/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/models/timesfm/modeling_timesfm.pyr   r   ,   sL   € € € € € € ðð ð  $€CˆŒ˜Ñ	Ð#Ð#Ñ#Ø!%€Eˆ5Œ<˜$ÑÐ%Ð%Ñ%Ð%Ð%r"   r   c                   ót   — 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        e
z  dz  ed<   dS )ÚTimesFmOutputForPredictionaµ  
    mean_predictions (`torch.Tensor` of shape `(batch_size, sequence_length)`):
        The mean predictions of the time series.
    full_predictions (`torch.Tensor` of shape `(batch_size, sequence_length)`):
        The full predictions of the time series including the mean and the quantiles.
    loss (`torch.Tensor` of shape `(1,)`, *optional*, returned when `future_values` is provided):
        The loss of the TimesFM model.
    NÚmean_predictionsÚfull_predictionsÚloss)r   r   r   r   r&   r   r   r    r'   r(   Úfloatr!   r"   r#   r%   r%   :   sj   € € € € € € ðð ð -1Ð�e”l TÑ)Ð0Ð0Ñ0Ø,0Ð�e”l TÑ)Ð0Ð0Ñ0Ø(,€Dˆ%Œ,˜Ñ
 Ñ
%Ð,Ð,Ñ,Ð,Ð,r"   r%   c                   ó0   ‡ — e Zd ZdZdefˆ fd„Zdd„Zˆ xZS )Ú
TimesFmMLPzPax MLP in pytorch.Úconfigc                 ó  •— t          ¦   «                              ¦   «          |j        }|j        }t	          j        ||¦  «        | _        t	          j        ||¦  «        | _        t	          j        |d¬¦  «        | _	        d S )Nç�íµ ÷Æ°>)Únormalized_shapeÚeps)
ÚsuperÚ__init__Úhidden_sizeÚintermediate_sizeÚnnÚLinearÚ	gate_projÚ	down_projÚ	LayerNormÚ
layer_norm)Úselfr,   r3   r4   Ú	__class__s       €r#   r2   zTimesFmMLP.__init__N   sl   ø€ Ý‰Œ×ÒÑÔÐØÔ(ˆØ"Ô4Ðåœ ;Ð0AÑBÔBˆŒÝœÐ#4°kÑBÔBˆŒÝœ,¸ÈÐNÑNÔNˆŒˆˆr"   Nc                 óà   — |                       |¦  «        }|                      |¦  «        }t          j        |¦  «        }|                      |¦  «        }|�|d|d d …d d …d f         z
  z  }||z   S )Nç      ð?)r:   r7   ÚFÚrelur8   )r;   ÚxÚpaddingsÚgate_inpÚgateÚoutputss         r#   ÚforwardzTimesFmMLP.forwardW   st   € Ø—?’? 1Ñ%Ô%ˆØ�~Š~˜hÑ'Ô'ˆÝŒv�d‰|Œ|ˆØ—.’. Ñ&Ô&ˆØÐØ  x°°°°1°1°1°d°
Ô';Ñ!;Ñ<ˆGØ˜‰{Ðr"   ©N©r   r   r   r   r   r2   rF   Ú__classcell__©r<   s   @r#   r+   r+   K   se   ø€ € € € € ØÐðO˜}ð Oð Oð Oð Oð Oð Oðð ð ð ð ð ð ð r"   r+   c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )ÚTimesFmResidualBlockzTimesFM residual block.c                 ó>  •— t          ¦   «                              ¦   «          || _        || _        || _        t          j        ||¦  «        | _        t          j        ¦   «         | _	        t          j        ||¦  «        | _
        t          j        ||¦  «        | _        d S rG   )r1   r2   Ú
input_dimsÚhidden_dimsÚoutput_dimsr5   r6   Úinput_layerÚSiLUÚ
activationÚoutput_layerÚresidual_layer)r;   rN   rO   rP   r<   s       €r#   r2   zTimesFmResidualBlock.__init__d   s   ø€ Ý‰Œ×ÒÑÔÐØ$ˆŒØ&ˆÔØ&ˆÔåœ9 Z°Ñ=Ô=ˆÔÝœ'™)œ)ˆŒÝœI k°;Ñ?Ô?ˆÔÝ œi¨
°KÑ@Ô@ˆÔÐÐr"   c                 ó´   — |                       |¦  «        }|                      |¦  «        }|                      |¦  «        }|                      |¦  «        }||z   S rG   )rQ   rS   rT   rU   )r;   rA   ÚhiddenÚoutputÚresiduals        r#   rF   zTimesFmResidualBlock.forwardo   sW   € Ø×!Ò! !Ñ$Ô$ˆØ—’ Ñ(Ô(ˆØ×"Ò" 6Ñ*Ô*ˆØ×&Ò& qÑ)Ô)ˆØ˜Ñ Ð r"   )r   r   r   r   r2   rF   rI   rJ   s   @r#   rL   rL   a   sR   ø€ € € € € Ø!Ð!ð	Að 	Að 	Að 	Að 	Að!ð !ð !ð !ð !ð !ð !r"   rL   ÚRMSNormc                   óT   ‡ — e Zd Zd	deddfˆ fd„Zdej        dej        fd„Zd„ Zˆ xZ	S )
ÚTimesFmRMSNormr.   r0   ÚreturnNc                 ó¬   •— t          ¦   «                              ¦   «          t          j        t	          j        |¦  «        ¦  «        | _        || _        dS )z=
        TimesFmRMSNorm is equivalent to T5LayerNorm
        N)r1   r2   r5   Ú	Parameterr   ÚonesÚweightÚvariance_epsilon)r;   r3   r0   r<   s      €r#   r2   zTimesFmRMSNorm.__init__y   sD   ø€ õ 	‰Œ×ÒÑÔÐÝ”l¥5¤:¨kÑ#:Ô#:Ñ;Ô;ˆŒØ #ˆÔÐÐr"   Úhidden_statesc                 ó  — |j         }|                     t          j        ¦  «        }|                     d¦  «                             dd¬¦  «        }|t          j        || j        z   ¦  «        z  }| j        |                     |¦  «        z  S )Né   éÿÿÿÿT)Úkeepdim)	ÚdtypeÚtor   Úfloat32ÚpowÚmeanÚrsqrtrb   ra   )r;   rc   Úinput_dtypeÚvariances       r#   rF   zTimesFmRMSNorm.forward�   s|   € Ø#Ô)ˆØ%×(Ò(­¬Ñ7Ô7ˆØ ×$Ò$ QÑ'Ô'×,Ò,¨R¸Ð,Ñ>Ô>ˆØ%­¬°H¸tÔ?TÑ4TÑ(UÔ(UÑUˆØŒ{˜]×-Ò-¨kÑ:Ô:Ñ:Ð:r"   c                 óH   — t          | j        j        ¦  «        › d| j        › �S )Nz, eps=)Útuplera   Úshaperb   )r;   s    r#   Ú
extra_reprzTimesFmRMSNorm.extra_reprˆ   s&   € Ý˜œÔ)Ñ*Ô*ÐIÐI°$Ô2GÐIÐIÐIr"   )r.   )
r   r   r   r)   r2   r   r   rF   rs   rI   rJ   s   @r#   r\   r\   w   sŒ   ø€ € € € € ð$ð $¨ð $¸$ð $ð $ð $ð $ð $ð $ð; U¤\ð ;°e´lð ;ð ;ð ;ð ;ðJð Jð Jð Jð Jð Jð Jr"   r\   c                   ó0   ‡ — e Zd ZdZdefˆ fd„Zdd„Zˆ xZS )ÚTimesFmPositionalEmbeddingz6Generates position embedding for a given 1-d sequence.r,   c           
      óÒ  •— t          ¦   «                              ¦   «          |j        }|j        }||c| _        | _        |j        | _        | j        dz  }t          j        t          |¦  «        t          |¦  «        z  ¦  «        t          |dz
  d¦  «        z  }|  
                    d|t          j        t          j        |t          j        ¬¦  «        | z  ¦  «        z  ¦  «         d S )Nre   r   Úinv_timescales©rh   )r1   r2   Úmin_timescaleÚmax_timescaler3   Úembedding_dimsÚmathÚlogr)   ÚmaxÚregister_bufferr   ÚexpÚarangerj   )r;   r,   ry   rz   Únum_timescalesÚlog_timescale_incrementr<   s         €r#   r2   z#TimesFmPositionalEmbedding.__init__�   sá   ø€ Ý‰Œ×ÒÑÔÐØÔ,ˆØÔ,ˆØ1>ÀÐ.ˆÔ˜DÔ.Ø$Ô0ˆÔàÔ,°Ñ1ˆÝ"&¤(­5°Ñ+?Ô+?Å%ÈÑBVÔBVÑ+VÑ"WÔ"WÕZ]Ð^lÐopÑ^pÐrsÑZtÔZtÑ"tÐØ×ÒØØ�EœI¥e¤l°>ÍÌÐ&WÑ&WÔ&WÐ[rÐZrÑ&rÑsÔsÑsñ	
ô 	
ð 	
ð 	
ð 	
r"   Nc                 ó  — |€|€t          d¦  «        ‚|€?t          j        |t          j        | j        j        ¬¦  «                             d¦  «        }n"|j        dk    rt          d|j        › �¦  «        ‚ |j	        g |j        ¢d‘R Ž | j         	                    ddd¦  «        z  }t          j
        t          j        |¦  «        t          j        |¦  «        gd¬	¦  «        }t          j        |ddd| j        dz  f¦  «        }|S )
aÍ  Generates a Tensor of sinusoids with different frequencies.

        Args:
            seq_length: an optional Python int defining the output sequence length.
              if the `position` argument is specified.
            position: [B, seq_length], optional position for each token in the
              sequence, only required when the sequence is packed.

        Returns:
            [B, seqlen, D] if `position` is specified, else [1, seqlen, D]
        Nz.Either position or seq_length must be provided©rh   Údevicer   re   z*position must be 2-dimensional, got shape r   rf   ©Údim)Ú
ValueErrorr   r�   rj   rw   r†   Ú	unsqueezeÚndimrr   ÚviewÚcatÚsinÚcosr?   Úpadr{   )r;   Ú
seq_lengthÚpositionÚscaled_timeÚsignals        r#   rF   z"TimesFmPositionalEmbedding.forward�   s  € ð Ð 
Ð 2ÝÐMÑNÔNÐNàÐå”| Jµe´mÈDÔL_ÔLfÐgÑgÔg×qÒqÐrsÑtÔtˆHˆHØŒ]˜aÒÐÝÐZÈ(Ì.ÐZÐZÑ[Ô[Ð[à#�h”mÐ7 X¤^Ð7°QÐ7Ð7Ð7¸$Ô:M×:RÒ:RÐSTÐVWÐY[Ñ:\Ô:\Ñ\ˆÝ”�EœI kÑ2Ô2µE´I¸kÑ4JÔ4JÐKÐQRÐSÑSÔSˆõ ”�v  1 a¨Ô)<¸qÑ)@ÐAÑBÔBˆØˆr"   ©NNrH   rJ   s   @r#   ru   ru   Œ   s^   ø€ € € € € Ø@Ð@ð
˜}ð 
ð 
ð 
ð 
ð 
ð 
ðð ð ð ð ð ð ð r"   ru   ç        ÚmoduleÚquery_statesÚ
key_statesÚvalue_statesÚattention_maskÚscalingÚdropoutÚkwargsc                 óÀ  — t          j        ||                     dd¦  «        ¦  «        |z  }|�||z   }t          j                             |dt           j        ¬¦  «                             |j        ¦  «        }t          j         	                    ||| j
        ¬¦  «        }t          j        ||¦  «        }	|	                     dd¦  «                             ¦   «         }	|	|fS )Nre   r   rf   )rˆ   rh   )ÚpÚtrainingr   )r   ÚmatmulÚ	transposer5   Ú
functionalÚsoftmaxrj   ri   rh   r�   r¡   Ú
contiguous)
r—   r˜   r™   rš   r›   rœ   r�   rž   Úattn_weightsÚattn_outputs
             r#   Úsimple_eager_attention_forwardr©   º   sÅ   € õ ”< ¨j×.BÒ.BÀ1ÀaÑ.HÔ.HÑIÔIÈGÑS€LØÐ!Ø# nÑ4ˆå”=×(Ò(¨¸2ÅUÄ]Ð(ÑSÔS×VÒVÐWcÔWiÑjÔj€LÝ”=×(Ò(¨¸È6Ì?Ð(Ñ[Ô[€LÝ”,˜|¨\Ñ:Ô:€KØ×'Ò'¨¨1Ñ-Ô-×8Ò8Ñ:Ô:€Kà˜Ð$Ð$r"   c                   ó¼   ‡ — e Zd ZdZdedefˆ fd„Zdej        dej        fd„Z		 dd	ej        d
ej        dz  de
e         deej        ej        dz  f         fd„Zˆ xZS )ÚTimesFmAttentionzlImplements the attention used in TimesFM. One key difference is that there is _per_dim_scaling of the query.r,   Ú	layer_idxc                 óä  •— t          ¦   «                              ¦   «          || _        d| _        |j        | _        || _        |j        | _        |j        | _        |j	        | _	        | j        | j	        z  | _
        | j        | j	        z  | _        t          j        t          j        | j	        f¦  «        ¦  «        | _        t          j        | j        | j        | j	        z  ¦  «        | _        t          j        | j        | j        | j	        z  ¦  «        | _        t          j        | j        | j        | j	        z  ¦  «        | _        t          j        | j        | j	        z  | j        ¦  «        | _        d S )NT)r1   r2   r,   Ú	is_causalÚattention_dropoutr¬   Únum_attention_headsÚ	num_headsr3   Úhead_dimÚq_sizeÚkv_sizer5   r_   r   Úemptyrœ   r6   Úq_projÚk_projÚv_projÚo_proj©r;   r,   r¬   r<   s      €r#   r2   zTimesFmAttention.__init__Ó   s  ø€ Ý‰Œ×ÒÑÔÐØˆŒØˆŒØ!'Ô!9ˆÔØ"ˆŒàÔ3ˆŒØ!Ô-ˆÔØœˆŒà”n t¤}Ñ4ˆŒØ”~¨¬Ñ5ˆŒÝ”|¥E¤K°´Ð0@Ñ$AÔ$AÑBÔBˆŒå”i Ô 0°$´.À4Ä=Ñ2PÑQÔQˆŒÝ”i Ô 0°$´.À4Ä=Ñ2PÑQÔQˆŒÝ”i Ô 0°$´.À4Ä=Ñ2PÑQÔQˆŒÝ”i ¤°´Ñ >ÀÔ@PÑQÔQˆŒˆˆr"   Úqueryr]   c                 ó°   — t          j        | j        ¦  «                             dt	          j        | j        ¦  «        z  ¦  «        }||d d d d d …f         z  S )Ng^$3eG÷?)r?   Úsoftplusrœ   Úmulr|   Úsqrtr²   )r;   r»   r   s      r#   Ú_scale_queryzTimesFmAttention._scale_queryç   sO   € Ý”
˜4œ<Ñ(Ô(×,Ò,¨[½4¼9ÀTÄ]Ñ;SÔ;SÑ-SÑTÔTˆØ�u˜T 4¨¨q¨q¨qÐ0Ô1Ñ1Ð1r"   Nrc   r›   rž   c                 óÌ  — |j         d d…         }g |¢d‘| j        ‘R }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }|                      |¦  «                             |¦  «                             dd¦  «        }t          j	        | j
        j        t          ¦  «        }	 |	| ||||f| j        sdn| j        ddœ|¤Ž\  }
} |
j        g |¢d‘R Ž                      ¦   «         }
|                      |
¦  «        }
|
|fS )Nrf   r   re   r–   r>   )r�   rœ   )rr   r²   r¶   rŒ   r£   rÀ   r·   r¸   r   Úget_interfacer,   Ú_attn_implementationr©   r¡   r¯   Úreshaper¦   r¹   )r;   rc   r›   rž   Úinput_shapeÚhidden_shaper˜   r™   rš   Úattention_interfacer¨   r§   s               r#   rF   zTimesFmAttention.forwardë   sˆ  € ð $Ô)¨#¨2¨#Ô.ˆØ8˜Ð8 bÐ8¨$¬-Ð8Ð8ˆà—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆØ×(Ò(¨Ñ6Ô6ˆØ—[’[ Ñ/Ô/×4Ò4°\ÑBÔB×LÒLÈQÐPQÑRÔRˆ
Ø—{’{ =Ñ1Ô1×6Ò6°|ÑDÔD×NÒNÈqÐRSÑTÔTˆå(?Ô(MØŒKÔ,Õ.Lñ)
ô )
Ðð %8Ð$7ØØØØØð	%
ð  $œ}ÐH�C�C°$Ô2HØð	%
ð 	%
ð ð	%
ð 	%
Ñ!ˆ�\ð *�kÔ)Ð;¨;Ð;¸Ð;Ð;Ð;×FÒFÑHÔHˆØ—k’k +Ñ.Ô.ˆØ˜LÐ(Ð(r"   rG   )r   r   r   r   r   Úintr2   r   r   rÀ   r   r	   rq   rF   rI   rJ   s   @r#   r«   r«   Ð   så   ø€ € € € € ØvÐvðR˜}ð R¸ð Rð Rð Rð Rð Rð Rð(2 %¤,ð 2°5´<ð 2ð 2ð 2ð 2ð /3ð)ð )à”|ð)ð œ tÑ+ð)ð Ð-Ô.ð	)ð
 
ˆuŒ|˜Uœ\¨DÑ0Ð0Ô	1ð)ð )ð )ð )ð )ð )ð )ð )r"   r«   c                   ól   ‡ — e Zd ZdZdedefˆ fd„Zdej        dej        dej        dej        fd	„Z	ˆ xZ
S )
ÚTimesFmDecoderLayerzTransformer layer.r,   r¬   c                 óÜ   •— t          ¦   «                              ¦   «          t          ||¬¦  «        | _        t	          |¦  «        | _        t          |j        |j        ¬¦  «        | _	        d S )N)r¬   )r0   )
r1   r2   r«   Ú	self_attnr+   Úmlpr\   r3   Úrms_norm_epsÚinput_layernormrº   s      €r#   r2   zTimesFmDecoderLayer.__init__  s]   ø€ Ý‰Œ×ÒÑÔÐå)¨&¸IÐFÑFÔFˆŒÝ˜fÑ%Ô%ˆŒÝ-¨fÔ.@ÀfÔFYÐZÑZÔZˆÔÐÐr"   rc   r›   rB   r]   c                 ó    — |}|                       |¦  «        }|                      ||¬¦  «        \  }}||z   }|                      ||¬¦  «        }|S )N)rc   r›   )rB   )rÏ   rÌ   rÍ   )r;   rc   r›   rB   rž   rY   Ú_s          r#   rF   zTimesFmDecoderLayer.forward  sh   € ð !ˆØ×,Ò,¨]Ñ;Ô;ˆØŸ>š>Ø'Ø)ð *ñ 
ô 
Ñˆ�qð ! =Ñ0ˆð Ÿš ¸˜ÑBÔBˆàÐr"   )r   r   r   r   r   rÈ   r2   r   r   rF   rI   rJ   s   @r#   rÊ   rÊ     s›   ø€ € € € € ØÐð[˜}ð [¸ð [ð [ð [ð [ð [ð [ðà”|ðð œðð ”,ð	ð 
Œðð ð ð ð ð ð ð r"   rÊ   c                   ót   ‡ — e Zd ZU eed<   dZdgZdZdZdZ	e
edœZ ej        ¦   «         ˆ fd„¦   «         Zˆ xZS )	ÚTimesFmPreTrainedModelr,   ÚtimesfmrÊ   Úpast_values)ÚtimeT)rc   Ú
attentionsc           
      ó4  •— t          ¦   «                              |¦  «         t          |t          ¦  «        rt	          j        |j        ¦  «         d S t          |t          ¦  «        r°|j        dz  }|j	        |j
        }}t          j        t          |¦  «        t          |¦  «        z  ¦  «        t          |dz
  d¦  «        z  }t	          j        |j        |t#          j        t#          j        |t"          j        ¬¦  «        | z  ¦  «        z  ¦  «         d S d S )Nre   r   rx   )r1   Ú_init_weightsÚ
isinstancer«   ÚinitÚones_rœ   ru   r{   rz   ry   r|   r}   r)   r~   Úcopy_rw   r   r€   r�   rj   )r;   r—   r‚   rz   ry   rƒ   r<   s         €r#   rÙ   z$TimesFmPreTrainedModel._init_weights9  s  ø€ å‰Œ×Ò˜fÑ%Ô%Ð%Ý�fÕ.Ñ/Ô/ð 	åŒJ�v”~Ñ&Ô&Ð&Ð&Ð&Ý˜Õ :Ñ;Ô;ð 
	Ø#Ô2°aÑ7ˆNØ+1Ô+?ÀÔAU˜=ˆMÝ&*¤h­u°]Ñ/CÔ/CÅeÈMÑFZÔFZÑ/ZÑ&[Ô&[Õ^aØ Ñ" Añ_ô _ñ 'Ð#õ ŒJØÔ%ØÝ”)�EœL¨½u¼}ÐMÑMÔMÐQhÐPhÑhÑiÔiñjñô ð ð ð ð
	ð 
	r"   )r   r   r   r   r    Úbase_model_prefixÚ_no_split_modulesÚmain_input_nameÚinput_modalitiesÚ_supports_sdparÊ   r«   Ú_can_record_outputsr   Úno_gradrÙ   rI   rJ   s   @r#   rÓ   rÓ   ,  sŒ   ø€ € € € € € àÐÐÑØ!ÐØ.Ð/ÐØ#€OØ ÐØ€Nà,Ø&ðð Ðð
 €U„]�_„_ðð ð ð ñ „_ðð ð ð ð r"   rÓ   c                   ó  ‡ — e Zd Zdefˆ fd„Zdej        dej        deej        eej        ej        f         f         fd„Ze	e
edej        dej        d	ej        d
ee         def
d„¦   «         ¦   «         ¦   «         Ze	 ddej        dz  dedej        dej        dedej        dz  fd„¦   «         Zedej        dej        deej        ej        f         fd„¦   «         Zedej        dej        dej        fd„¦   «         Zˆ xZS )ÚTimesFmModelr,   c                 óÎ  •‡— t          ¦   «                              ‰¦  «         ‰| _        t          d‰j        z  ‰j        ‰j        ¬¦  «        | _        t          j	        ‰j
        ‰j        ¬¦  «        | _        t          j        ˆfd„t          ‰j        ¦  «        D ¦   «         ¦  «        | _        | j        j        rt#          ‰¬¦  «        | _        |                      ¦   «          d S )Nre   ©rN   rP   rO   )Únum_embeddingsÚembedding_dimc                 ó0   •— g | ]}t          ‰|¦  «        ‘ŒS r!   )rÊ   )Ú.0r¬   r,   s     €r#   ú
<listcomp>z)TimesFmModel.__init__.<locals>.<listcomp>Y  s$   ø€ ÐeÐeÐe¸	Õ  ¨Ñ3Ô3ÐeÐeÐer"   )r,   )r1   r2   r,   rL   Úpatch_lengthr3   r4   Úinput_ff_layerr5   Ú	EmbeddingÚ	freq_sizeÚfreq_embÚ
ModuleListÚrangeÚnum_hidden_layersÚlayersÚuse_positional_embeddingru   Úposition_embÚ	post_init©r;   r,   r<   s    `€r#   r2   zTimesFmModel.__init__N  sà   øø€ Ý‰Œ×Ò˜Ñ Ô Ð àˆŒÝ2Ø˜6Ô.Ñ.ØÔ*ØÔ0ð
ñ 
ô 
ˆÔõ
 œ°FÔ4DÐTZÔTfÐgÑgÔgˆŒÝ”mØeÐeÐeÐeÅUÈ6ÔKcÑEdÔEdÐeÑeÔeñ
ô 
ˆŒð Œ;Ô/ð 	JÝ :À&Ð IÑ IÔ IˆDÔð 	�ŠÑÔÐÐÐr"   ÚinputsÚpatched_padsr]   c                 ó”  — |                       ||¦  «        \  }}t          j        || j        j        ¬¦  «        }||dd…ddf         z
  |dd…ddf         z  }t          j        t          j        || j        j        z
  ¦  «        | j        j        k     t          j        | j        j        |j	        |j
        ¬¦  «        |¦  «        }|||ffS )zInput is of shape [B, N, P].©ÚminNr…   )Ú_timesfm_masked_mean_stdr   Úclampr,   Ú	toleranceÚwhereÚabsÚpad_valÚtensorrh   r†   )r;   rû   rü   ÚmuÚsigmarE   s         r#   Ú_forward_transformzTimesFmModel._forward_transforma  sÈ   € ð ×1Ò1°&¸,ÑGÔG‰	ˆˆEÝ”˜E t¤{Ô'<Ð=Ñ=Ô=ˆð ˜B˜q˜q˜q $¨˜}Ô-Ñ-°°q°q°q¸$À°}Ô1EÑEˆÝ”+ÝŒI�f˜tœ{Ô2Ñ2Ñ3Ô3°d´kÔ6KÒKÝŒL˜œÔ,°G´MÈ'Ì.ÐYÑYÔYØñ
ô 
ˆð
 ˜˜U˜Ð#Ð#r"   rÕ   Úpast_values_paddingÚfreqrž   c                 óâ  — |j         d         }|                     |d| j        j        ¦  «        }|                     |d| j        j        ¦  «        }t	          j        t	          j        |dz
  ¦  «        | j        j        k     t	          j        d|j	        |j
        ¬¦  «        |¦  «        }t	          j        t	          j        || j        j        z
  ¦  «        | j        j        k     t	          j        d|j	        |j
        ¬¦  «        |¦  «        }|                      ||¦  «        \  }}|d|z
  z  }t	          j        ||gd¬¦  «        }	|                      |	¦  «        }
t	          j        |d¬¦  «        d         }| j        j        r`|                      |
j         d         ¦  «        }t	          j        |g|
j         d         z  d¬¦  «        }|                      ||¦  «        }|
|z  }
|                      |¦  «        }|
|z  }
|
}|                      ||j         d         |j	        |j
        d¬	¦  «        }| j        d
| j        j        …         D ]} ||f||dœ|¤Ž}Œt1          ||d         |d         ¬¦  «        S )a°  
        past_values (`torch.FloatTensor` of shape `(batch_size, sequence_length)`):
            Past values of the time series that serves as input to the model.
        past_values_padding (`torch.LongTensor` of shape `(batch_size, sequence_length)`):
            The padding indicator of the time series.
        freq (`torch.LongTensor` of shape `(batch_size,)`):
            Frequency indices for the time series data.
        r   rf   r>   r–   r…   r‡   r   T)r›   Úsequence_lengthrh   r†   r®   N)r›   rB   )Úlast_hidden_stater   r   )rr   rŒ   r,   rî   r   r  r  r  r  rh   r†   r  r	  r�   rï   rÿ   r÷   rø   ÚconcatÚ_timesfm_shift_padded_seqrò   Ú_prepare_4d_attention_maskrö   rõ   r   )r;   rÕ   r
  r  rž   ÚbsizeÚpatched_inputsrü   ÚstatsÚconcat_inputsÚmodel_inputÚpatched_paddingÚpos_embÚf_embrc   r›   Úlayers                    r#   rF   zTimesFmModel.forwardq  s˜  € ð& Ô! !Ô$ˆØ$×)Ò)¨%°°T´[Ô5MÑNÔNˆØ*×/Ò/°°r¸4¼;Ô;SÑTÔTˆåœÝŒI�l SÑ(Ñ)Ô)¨D¬KÔ,AÒAÝŒL˜ NÔ$8ÀÔAVÐWÑWÔWØñ
ô 
ˆõ
 ”{ÝŒI�n t¤{Ô':Ñ:Ñ;Ô;¸d¼kÔ>SÒSÝŒL˜ LÔ$6¸|Ô?RÐSÑSÔSØñ
ô 
ˆð
 !%× 7Ò 7¸ÈÑ UÔ UÑˆ˜ð (¨3°Ñ+=Ñ>ˆÝœ	 >°<Ð"@ÀbÐIÑIÔIˆØ×)Ò)¨-Ñ8Ô8ˆõ  œ) L°bÐ9Ñ9Ô9¸!Ô<ˆØŒ;Ô/ð 	#Ø×'Ò'¨Ô(9¸!Ô(<Ñ=Ô=ˆGÝ”l G 9¨{Ô/@ÀÔ/CÑ#CÈÐKÑKÔKˆGØ×4Ò4°_ÀgÑNÔNˆGØ˜7Ñ"ˆKà—’˜dÑ#Ô#ˆØ�uÑˆð $ˆØ×8Ò8Ø*Ø)Ô/°Ô2ØÔ%Ø Ô'Øð 9ñ 
ô 
ˆð ”[Ð!@ 4¤;Ô#@Ð!@ÔAð 	ð 	ˆEØ!˜EØðà-Ø(ðð ð ð	ð ˆMˆMõ Ø+Ø�a”Ø˜”(ð
ñ 
ô 
ð 	
r"   Tr›   Nr  rh   r†   r®   c                 ó”  — |j         rt          j        |¦  «        j        nt          j        |¦  «        j        }| �(|                      | j        d         ddd¦  «        } | |z  } |rbt          j        t          j        ||f||¬¦  «        |z  d¬¦  «        }|                     dd||¦  «        }| �t          j	        | |¦  «        } n|} | S )aí  
        Creates 4D attention mask and combines causal and padding masks if needed.

        Args:
            attention_mask: Optional tensor of shape (batch_size, seq_length) containing padding mask
            sequence_length: Length of the sequence
            dtype: Data type of the mask
            device: Device of the mask
            is_causal: Whether to apply causal masking

        Returns:
            4D attention mask of shape (batch_size, 1, seq_length, seq_length)
        Nr   r   rf   r…   )Údiagonal)
Úis_floating_pointr   Úfinforÿ   ÚiinforŒ   rr   Útriur`   Úminimum)r›   r  rh   r†   r®   Ú	min_valueÚcausal_masks          r#   r  z'TimesFmModel._prepare_4d_attention_mask¼  sé   € ð, /4Ô.EÐa•E”K Ñ&Ô&Ô*Ð*Í5Ì;ÐW\ÑK]ÔK]ÔKaˆ	ð Ð%à+×0Ò0°Ô1EÀaÔ1HÈ!ÈQÐPRÑSÔSˆNØ+¨iÑ7ˆNð ð 	-Ýœ*Ý”
˜O¨_Ð=ÀUÐSYÐZÑZÔZÐ]fÑfØðñ ô ˆKð &×*Ò*¨1¨a°À/ÑRÔRˆKð Ð)Ý!&¤¨~¸{Ñ!KÔ!K��à!,�àÐr"   Úpaddingc                 óD  — dt           j        fd„}t          j        d|z
  d¬¦  «        } ||¦  «        }t          j        | j        d         ¦  «        }| ||dd…f         }|||dd…f         }d|z
  }t          j        |d¬¦  «        }	t          j        |	d¬	¦  «        }	t          j        ||z  d¬¦  «        }
|
|	z  }||                     d
¦  «        z
  |z  }t          j        |dz  d¬¦  «        |	z  }t          j        |d¬	¦  «        }t          j        |¦  «        }||fS )aÃ  Calculates mean and standard deviation of `inputs` across axis 1.

        It excludes values where `padding` is 1.

        Args:
            inputs: A PyTorch tensor of shape [b, n, p].
            padding: A PyTorch tensor of shape [b, n, p] with values 0 or 1.

        Returns:
            A tuple containing the mean and standard deviation.
            We return the statistics of the first patch with more than three non-padded values.
        Úarrc                 ó.  — t          j        | dk                         t           j        ¦  «        d¬¦  «        }| dk                         t           j        ¦  «                             d¬¦  «        }t          j        |dk    | j        d         dz
  |¦  «        S )Nr   r   r‡   r   )r   Úargmaxri   Úint32Úsumr  rr   )r&  ÚindicesÚrow_sums      r#   Ú_get_patch_indexz?TimesFmModel._timesfm_masked_mean_std.<locals>._get_patch_indexú  ss   € Ý”l C¨1¢H§=¢=µ´Ñ#=Ô#=À1ÐEÑEÔEˆGØ˜a’x—m’m¥E¤KÑ0Ô0×4Ò4¸Ð4Ñ;Ô;ˆGÝ”;˜w¨!š|¨S¬Y°q¬\¸AÑ-=¸wÑGÔGÐGr"   r   re   r‡   r   Nr>   rþ   rf   r–   )r   r   r*  r�   rr   r  rŠ   r¿   )rû   r$  r-  Úpad_sumÚpatch_indicesÚbidxsr&  r�   ÚmaskÚnum_valid_elementsÚ
masked_sumÚmasked_meanÚmasked_centered_arrÚ
masked_varÚ
masked_stds                  r#   r   z%TimesFmModel._timesfm_masked_mean_stdê  sX  € ð 	H¥%¤,ð 	Hð 	Hð 	Hð 	Hõ
 ”)˜A ™K¨QÐ/Ñ/Ô/ˆØ(Ð(¨Ñ1Ô1ˆÝ”˜Vœ\¨!œ_Ñ-Ô-ˆà�U˜M¨1¨1¨1Ð,Ô-ˆØ�e˜]¨A¨A¨AÐ-Ô.ˆð �3‰wˆõ #œY t°Ð3Ñ3Ô3ÐÝ"œ[Ð);ÀÐEÑEÔEÐõ ”Y˜s T™z¨qÐ1Ñ1Ô1ˆ
Ø Ð#5Ñ5ˆð  # [×%:Ò%:¸2Ñ%>Ô%>Ñ>À$ÑFÐÝ”YÐ2°AÑ5¸1Ð=Ñ=Ô=Ð@RÑRˆ
Ý”[ °Ð5Ñ5Ô5ˆ
Ý”Z 
Ñ+Ô+ˆ
à˜JÐ&Ð&r"   r1  Úseqc                 óž  — |j         \  }}}| dk    }|                     t          j        ¦  «                             d¬¦  «        }d||                     d¬¦  «         <   t          j        ||j        ¬¦  «                             ddd¦  «         	                    |d|¦  «        }||dd…ddf         z
  |z  }| 
                    d|¦  «        }	|	S )zéShifts rows of seq based on the first 0 in each row of the mask.

        Args:
            mask: mask tensor of shape [B, N]
            seq: seq tensor of shape [B, N, P]

        Returns:
            The shifted sequence.
        r   r   r‡   rf   )r†   N)rr   ri   r   r)  r(  Úanyr�   r†   rŒ   ÚexpandÚgather)
r1  r8  Ú
batch_sizeÚnum_seqÚfeature_dimÚnew_maskr+  Ú	idx_rangeÚshifted_idxÚshifted_seqs
             r#   r  z&TimesFmModel._timesfm_shift_padded_seq  sØ   € ð ,/¬9Ñ(ˆ
�G˜[à%)¨Q¢Yˆð —+’+�eœkÑ*Ô*×1Ò1°aÐ1Ñ8Ô8ˆð )+ˆ�—’ !�Ñ$Ô$Ð$Ñ%õ ”L °´Ð<Ñ<Ô<×AÒAÀ!ÀRÈÑKÔK×RÒRÐS]Ð_aÐcnÑoÔoˆ	ð ! 7¨1¨1¨1¨d°D¨=Ô#9Ñ9¸WÑDˆð —j’j  KÑ0Ô0ˆàÐr"   )T)r   r   r   r   r2   r   r   rq   r	  r   r   r   Ú
LongTensorr   r   r   rF   ÚstaticmethodrÈ   rh   r†   Úboolr  r   r  rI   rJ   s   @r#   ræ   ræ   L  s  ø€ € € € € ð˜}ð ð ð ð ð ð ð&$Ø”lð$Ø27´,ð$à	ˆuŒ|˜U 5¤<°´Ð#=Ô>Ð>Ô	?ð$ð $ð $ð $ð   ØØðF
à”\ðF
ð #Ô-ðF
ð Œlð	F
ð
 Ð+Ô,ðF
ð 
ðF
ð F
ð F
ñ „^ñ „_ñ  ÔðF
ðP ð ð+ð +Øœ tÑ+ð+àð+ð Œ{ð+ð ”ð	+ð
 ð+ð 
Œ˜Ñ	ð+ð +ð +ñ „\ð+ðZ ð,'¨¬ð ,'ÀÄð ,'ÐQVÐW\ÔWcÐejÔeqÐWqÔQrð ,'ð ,'ð ,'ñ „\ð,'ð\ ð¨¬ð ¸5¼<ð ÈEÌLð ð ð ñ „\ðð ð ð ð r"   ræ   c                   ó  ‡ — e Zd ZdZdefˆ fd„Z	 ddeej                 dee	         dz  de	dz  de
ej        d	f         fd
„Zdej        de
ej        ej        f         dej        fd„Zdej        dej        dej        fd„Zee	 	 	 	 	 	 ddeej                 deej        e	z           dz  de	dz  dej        dz  de	dz  dededee         defd„¦   «         ¦   «         Zedej        de	deej                 fd„¦   «         Zˆ xZS )ÚTimesFmModelForPredictionz/TimesFM model for quantile and mean prediction.r,   c                 óT  •— t          ¦   «                              |¦  «         || _        |j        | _        |j        | _        t          |¦  «        | _        t          |j
        |j        dt          |j        ¦  «        z   z  |j        ¬¦  «        | _        |                      ¦   «          d S )Nr   rè   )r1   r2   r,   Úcontext_lengthÚcontext_lenÚhorizon_lengthÚhorizon_lenræ   ÚdecoderrL   r3   ÚlenÚ	quantilesr4   Úhorizon_ff_layerrù   rú   s     €r#   r2   z"TimesFmModelForPrediction.__init__=  sŸ   ø€ Ý‰Œ×Ò˜Ñ Ô Ð àˆŒØ!Ô0ˆÔØ!Ô0ˆÔå# FÑ+Ô+ˆŒõ !5ØÔ)ØÔ-°µS¸Ô9IÑ5JÔ5JÑ1JÑKØÔ0ð!
ñ !
ô !
ˆÔð 	�ŠÑÔÐÐÐr"   Nrû   r  rK  r]   .c                 ó   — |€| j         }g g }}|D �]}|j        d         }t          j        || j        z   |j        |j        ¬¦  «        }||k     rt||z
  }	t          j        t          j        |	|j        |j        ¬¦  «        |gd¬¦  «        }t          j        t          j        |	|j        |j        ¬¦  «        |gd¬¦  «        }n$||k    r|| d…         }||| j        z    d…         }| 	                    |¦  «         | 	                    |¦  «         �Œt          j
        |d¬¦  «        t          j
        |d¬¦  «        f}
|�M|
t          j        |dt          |¦  «        …         t          j        ¬¦  «                             dd¦  «        fz   }
|
S )aá  Pad/truncate input time series to `context_len` and build a padding mask.

        Args:
            inputs: A list of 1d Tensors. Each Tensor is the context time series of a single forecast task.
            freq: Optional list of frequencies (returned as a tensor when provided).
            context_len: Optional context length override (defaults to `self.context_len`).

        Returns:
            Tuple of (padded_inputs, padding_mask) and optionally a freq tensor.
        Nr   r…   r‡   rx   rf   r   )rK  rr   r   ÚzerosrM  rh   r†   r�   r`   ÚappendÚstackr  rO  r)  rÄ   )r;   rû   r  rK  Úinput_tsÚinput_paddingÚtsÚ	input_lenr$  Únum_front_padÚresults              r#   Ú_preprocessz%TimesFmModelForPrediction._preprocessP  s¥  € ð ÐØÔ*ˆKà"$ b�-ˆàð 	*ñ 	*ˆBØœ œˆIÝ”k )¨dÔ.>Ñ">ÀbÄhÐWYÔW`ÐaÑaÔaˆGØ˜;Ò&Ð&Ø +¨iÑ 7�Ý”Y¥¤¨MÀÄÐRTÔR[Ð \Ñ \Ô \Ð^`ÐaÐghÐiÑiÔi�Ýœ)¥U¤Z°ÀRÄXÐV]ÔVdÐ%eÑ%eÔ%eÐgnÐ$oÐuvÐwÑwÔw��Ø˜[Ò(Ð(Ø˜˜˜˜Ô&�Ø! K°$Ô2BÑ$BÐ"CÐ"EÐ"EÔF�à�OŠO˜BÑÔÐØ× Ò  Ñ)Ô)Ð)Ñ)å”+˜h¨AÐ.Ñ.Ô.µ´¸MÈqÐ0QÑ0QÔ0QÐRˆØÐØ�uœ|¨D°µ3°v±;´;°Ô,?ÅuÄ{ÐSÑSÔS×[Ò[Ð\^Ð`aÑbÔbÐdÑdˆFØˆr"   Úmodel_outputr  c                 ó  — |                       |¦  «        }|j        \  }}}|                     ||| j        j        t          | j        j        ¦  «        dz   ¦  «        }|\  }}||dd…dddf         z  |dd…dddf         z   S )z*Postprocess output of stacked transformer.r   N)rQ  rr   rŒ   r,   rL  rO  rP  )	r;   r]  r  Ú	output_tsÚbÚnrÑ   r  r  s	            r#   Ú_postprocess_outputz-TimesFmModelForPrediction._postprocess_outputu  s’   € ð ×)Ò)¨,Ñ7Ô7ˆ	ð ”/‰ˆˆ1ˆaØ—N’N 1 a¨¬Ô)CÅSÈÌÔI^ÑE_ÔE_ÐbcÑEcÑdÔdˆ	à‰	ˆˆEØ˜5    D¨$°Ð!4Ô5Ñ5¸¸1¸1¸1¸dÀDÈ$Ð;NÔ8OÑOÐOr"   ÚpredictionsÚtargetsc                 ó4  — g }t          | j        j        ¦  «        D ]W\  }}||d|f         z
  }t          j        |dz
  |z  ||z  ¦  «        }|                     |                     ¦   «         ¦  «         ŒXt          j        |¦  «                             ¦   «         S )N.r   )Ú	enumerater,   rP  r   r~   rT  rl   rU  )r;   rc  rd  ÚlossesÚiÚqÚerrorsr(   s           r#   Ú_quantile_lossz(TimesFmModelForPrediction._quantile_loss„  s‘   € ØˆÝ˜dœkÔ3Ñ4Ô4ð 	'ð 	'‰DˆAˆqØ˜{¨3°¨6Ô2Ñ2ˆFÝ”9˜a !™e vÑ-¨q°6©zÑ:Ô:ˆDØ�MŠM˜$Ÿ)š)™+œ+Ñ&Ô&Ð&Ð&ÝŒ{˜6Ñ"Ô"×'Ò'Ñ)Ô)Ð)r"   FrÕ   Úwindow_sizeÚfuture_valuesÚforecast_context_lenÚreturn_forecast_on_contextÚtruncate_negativerž   c                 ó	  ‡#— |€| j         Š#n|Š#|d         j        }	ˆ#fd„|D ¦   «         }
t          j        t          j        d„ |
D ¦   «         ¦  «        ¦  «        }|�ig }g }t          |
¦  «        D ]O\  }}|                     |                      ||¦  «        ¦  «         |�|                     ||         gdz  ¦  «         ŒP|}
|�|}|€-t           	                    d¦  «         dgt          |
¦  «        z  }|                      |
|¦  «        \  }}}|                     |	¦  «        }|                     |	¦  «        }|                     |	¦  «        }|}|j        d         }g }|j        d         |j        d         | j        z   k    r3t          d|j        d         › d	|j        d         › d
| j        › �¦  «        ‚| j        j        }| j        |z   dz
  |z  }t%          |¦  «        D �]9}|dd…d|j        d         …f         }|dd…‰# d…f         }|dd…‰# d…f         } | j        d|||dœ|¤Ž}|                      |j        |j        |j        f¦  «        }|rv|dk    rp|dd…dd…d| j        j        …dd…f         }|                     |                     d¦  «        d|                     d¦  «        ¦  «        }|                     |¦  «         |dd…dd|…df         }|dd…dd|…dd…f         }|                     |¦  «         t          j        ||gd¬¦  «        }�Œ;|r;t          j        |d¬¦  «        dd…d|| j        j        z
  | j        z   …dd…f         }n*t          j        |d¬¦  «        dd…d| j        …dd…f         }|dd…dd…df         }|�6|ddd…df         |ddd…df         z   }|ddd…df         |ddd…df         z   }|rX|dk    }t          j        ||                     d¦  «        |¦  «        }t          j        ||                     d¦  «        |¦  «        }d} |�?t?          j         ||¦  «        }!|  !                    |dd…dd…dd…f         |¦  «        }"|!|"z   } tE          |j        |j#        |j$        ||| ¬¦  «        S )aŸ  
        past_values (`torch.FloatTensor` of shape `(batch_size, sequence_length)`):
            Past values of the time series that serves as input to the model.
        freq (`torch.LongTensor` of shape `(batch_size,)`):
            Frequency indices for the time series data.
        window_size (`int`, *optional*):
            Window size of trend + residual decomposition. If None then we do not do decomposition.
        future_values (`torch.Tensor`, *optional*):
            Optional future time series values to be used for loss computation.
        forecast_context_len (`int`, *optional*):
            Optional max context length.
        return_forecast_on_context (`bool`, *optional*):
            True to return the forecast on the context when available, i.e. after the first input patch.
        truncate_negative (`bool`, *optional*):
            Truncate to only non-negative values if any of the contexts have non-negative values,
            otherwise do nothing.

        Example:

        ```python
        >>> from transformers import TimesFmModelForPrediction

        >>> model = TimesFmModelForPrediction.from_pretrained("google/timesfm-2.0-500m-pytorch")

        >>> forecast_input = [torch.linspace(0, 20, 100).sin(), torch.linspace(0, 20, 200).sin(), torch.linspace(0, 20, 400).sin()]
        >>> frequency_input = torch.tensor([0, 1, 2], dtype=torch.long)

        >>> # Generate
        >>> with torch.no_grad():
        >>>     outputs = model(past_values=forecast_input, freq=frequency_input, return_dict=True)
        >>>     point_forecast_conv = outputs.mean_predictions
        >>>     quantile_forecast_conv = outputs.full_predictions
        ```
        Nr   c                 ó&   •— g | ]}|‰ d …         ‘ŒS rG   r!   )rì   rX  Úfcontext_lens     €r#   rí   z5TimesFmModelForPrediction.forward.<locals>.<listcomp>Â  s$   ø€ Ð;Ð;Ð;¨�"�l�]�^�^Ô$Ð;Ð;Ð;r"   c                 ó6   — g | ]}t          j        |¦  «        ‘ŒS r!   )r   rÿ   )rì   rX  s     r#   rí   z5TimesFmModelForPrediction.forward.<locals>.<listcomp>Ã  s    € Ð(HÐ(HÐ(H¸2­¬°2©¬Ð(HÐ(HÐ(Hr"   re   z6No frequency provided via `freq`. Default to high (0).r   z=Length of paddings must match length of input + horizon_len: z != z + )rÕ   r
  r  rf   r   )Úaxis.r–   )r  r×   rc   r&   r'   r(   r!   )%rK  r†   r   rÿ   rU  rf  ÚextendÚ_timesfm_moving_averageÚloggerÚinforO  r\  ri   rr   rM  r‰   r,   rL  rô   rN  rb  r  r   r   rî   rÄ   ÚsizerT  Úconcatenater  Ú	clamp_minr?   Úmse_lossrk  r%   r×   rc   )$r;   rÕ   r  rl  rm  rn  ro  rp  rž   r†   rû   Úinp_minÚ
new_inputsÚ	new_freqsrh  rX  rV  rW  Úinp_freqÚ	final_outrK  Úfull_outputsÚoutput_patch_lenÚnum_decode_patchesÚ
step_indexÚcurrent_paddingÚdecoder_outputÚfprop_outputsÚnew_full_tsÚnew_tsÚmean_outputsr  r(   r}  Úquantile_lossrs  s$                                      @r#   rF   z!TimesFmModelForPrediction.forwardŒ  s¨  ø€ ð^  Ð'ØÔ+ˆLˆLà/ˆLà˜Q”Ô&ˆà;Ð;Ð;Ð;¨{Ð;Ñ;Ô;ˆÝ”)�EœKÐ(HÐ(HÀÐ(HÑ(HÔ(HÑIÔIÑJÔJˆàÐ"ØˆJØˆIÝ" 6Ñ*Ô*ð 4ð 4‘��2Ø×!Ò! $×">Ò">¸rÀ;Ñ"OÔ"OÑPÔPÐPØÐ#Ø×$Ò$ d¨1¤g Y°¡]Ñ3Ô3Ð3øØˆFØÐØ �àˆ<Ý�KŠKÐPÑQÔQÐQØ�3�˜V™œÑ$ˆDà,0×,<Ò,<¸VÀTÑ,JÔ,JÑ)ˆ�- Ø—;’;˜vÑ&Ô&ˆØ%×(Ò(¨Ñ0Ô0ˆØ—;’;˜vÑ&Ô&ˆàˆ	Ø”o aÔ(ˆØˆàÔ˜qÔ! Y¤_°QÔ%7¸$Ô:JÑ%JÒJÐJÝðZØ!Ô'¨Ô*ðZð ZØ09´ÀÔ0BðZð ZØGKÔGWðZð Zñô ð ð  œ;Ô5Ðà"Ô.Ð1AÑAÀAÑEÐJZÑZÐÝÐ 2Ñ3Ô3ð 	Hñ 	HˆJØ+¨A¨A¨A¨q°9´?À1Ô3EÐ/EÐ,EÔFˆOØ     \ M N NÐ!2Ô3ˆHØ+¨A¨A¨A°¨}¨~¨~Ð,=Ô>ˆMØ,8¨D¬Lð -Ø$Ø$1Øð-ð -ð ð	-ð -ˆNð !×4Ò4ØÔ0ØÔ# ^Ô%9Ð:ñô ˆMð
 *ð 1¨j¸Aªo¨oØ+¨A¨A¨A¨s°¨sÐ4N°d´kÔ6NÐ4NÐPQÐPQÐPQÐ,QÔR�Ø)×1Ò1°+×2BÒ2BÀ1Ñ2EÔ2EÀrÈ;×K[ÒK[Ð\]ÑK^ÔK^Ñ_Ô_�Ø×#Ò# KÑ0Ô0Ð0à" 1 1 1 bÐ*;Ð+;Ð*;¸QÐ#>Ô?ˆFØ'¨¨¨¨2Ð/@Ð0@Ð/@À!À!À!Ð(CÔDˆKØ×Ò Ñ,Ô,Ð,ÝÔ)¨9°fÐ*=ÀBÐGÑGÔGˆI‰Ià%ð 	_Ý Ô,¨\ÀÐBÑBÔBØ��ÐP�k D¤KÔ$<Ñ<¸tÔ?OÑOÐPÐRSÐRSÐRSÐSôˆLˆLõ !Ô,¨\ÀÐBÑBÔBÀ1À1À1ÀaÈ$ÔJZÐFZÐ\]Ð\]Ð\]ÐC]Ô^ˆLà# A A A q q q¨! GÔ,ˆØÐ"Ø'¨¨¨1¨¨c¨	Ô2°\À!À$ÀQÀ$ÈÀ)Ô5LÑLˆLØ'¨¨¨1¨¨c¨	Ô2°\À!À$ÀQÀ$ÈÀ)Ô5LÑLˆLØð 	Yà˜q’LˆEÝ œ; u¨l×.DÒ.DÀSÑ.IÔ.IÈ<ÑXÔXˆLÝ œ; u¨l×.DÒ.DÀSÑ.IÔ.IÈ<ÑXÔXˆLàˆØÐ$Ý”z ,°Ñ>Ô>ˆHØ ×/Ò/°¸Q¸Q¸QÀÀÀÀ1À2À2¸XÔ0FÈÑVÔVˆMØ˜mÑ+ˆDå)Ø,Ô>Ø%Ô0Ø(Ô6Ø)Ø)Øð
ñ 
ô 
ð 	
r"   r&  c                 ó2  — t          j        | |dz
  dfdd¦  «        }t          j        || j        | j        ¬¦  «        |z  }t          j        |                     ddd¦  «        |                     ddd¦  «        ¦  «                             ¦   «         }|| |z
  gS )zCCalculates the moving average using PyTorch's convolution function.r   r   Úconstantr…   rf   )	r?   r�   r   r`   rh   r†   Úconv1drŒ   Úsqueeze)r&  rl  Ú
arr_paddedÚkernelÚsmoothed_arrs        r#   rw  z1TimesFmModelForPrediction._timesfm_moving_average  sŽ   € õ ”U˜3 ¨q¡°!Ð 4°jÀ!ÑDÔDˆ
å”˜K¨s¬yÀÄÐLÑLÔLÈ{ÑZˆå”x 
§¢°°1°bÑ 9Ô 9¸6¿;º;ÀqÈ!ÈRÑ;PÔ;PÑQÔQ×YÒYÑ[Ô[ˆØ˜c LÑ0Ð1Ð1r"   r•   )NNNNFF)r   r   r   r   r   r2   r   r   r   rÈ   rq   r\  rb  rk  r   r   rF  r   r   r%   rF   rE  Úlistrw  rI   rJ   s   @r#   rH  rH  :  s9  ø€ € € € € Ø9Ð9ð˜}ð ð ð ð ð ð ð( lpð#ð #Ø˜uœ|Ô,ð#Ø4<¸S´MÀDÑ4Hð#Ø^aÐdhÑ^hð#à	ˆuŒ|˜SÐ Ô	!ð#ð #ð #ð #ðJPØ!œLðPØ16°u´|ÀUÄ\Ð7QÔ1RðPà	ŒðPð Pð Pð Pð*¨%¬,ð *ÀÄð *ÐRWÔR^ð *ð *ð *ð *ð Øð 59Ø"&Ø-1Ø+/Ø+0Ø"'ðN
ð N
à˜eœlÔ+ðN
ð �u”| cÑ)Ô*¨TÑ1ðN
ð ˜4‘Zð	N
ð
 ”| dÑ*ðN
ð " D™jðN
ð %)ðN
ð  ðN
ð Ð+Ô,ðN
ð 
$ðN
ð N
ð N
ñ „^ñ ÔðN
ð` ð2 U¤\ð 2Àð 2ÈÈUÌ\ÔHZð 2ð 2ð 2ñ „\ð2ð 2ð 2ð 2ð 2r"   rH  )rH  rÓ   ræ   )r–   )9r|   Úcollections.abcr   r   Údataclassesr   r   Útorch.nnr5   Útorch.nn.functionalr¤   r?   Ú r   rÛ   Úintegrationsr   Úmodeling_flash_attention_utilsr	   Úmodeling_outputsr
   Úmodeling_utilsr   r   Úprocessing_utilsr   Úutilsr   r   r   r   Úutils.genericr   Úutils.output_capturingr   Úconfiguration_timesfmr   Ú
get_loggerr   rx  r   r%   ÚModuler+   rL   r\   ru   r   r)   rÈ   r©   r«   rÊ   rÓ   ræ   rH  Ú__all__r!   r"   r#   ú<module>r§     s:  ðð* €€€Ø .Ð .Ð .Ð .Ð .Ð .Ð .Ð .Ø !Ð !Ð !Ð !Ð !Ð !à €€€Ø Ð Ð Ð Ð Ð Ø Ð Ð Ð Ð Ð Ð Ð Ð à &Ð &Ð &Ð &Ð &Ð &Ø 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø BÐ BÐ BÐ BÐ BÐ BØ /Ð /Ð /Ð /Ð /Ð /Ø FÐ FÐ FÐ FÐ FÐ FÐ FÐ FØ &Ð &Ð &Ð &Ð &Ð &Ø RÐ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RÐ RØ 7Ð 7Ð 7Ð 7Ð 7Ð 7Ø 5Ð 5Ð 5Ð 5Ð 5Ð 5Ø 0Ð 0Ð 0Ð 0Ð 0Ð 0ð 
ˆÔ	˜HÑ	%Ô	%€ð Ø
ð	&ð 	&ð 	&ð 	&ð 	&�Oñ 	&ô 	&ñ „ñ „ð	&ð Ø
ð-ð -ð -ð -ð - ñ -ô -ñ „ñ „ð-ðð ð ð ð �”ñ ô ð ð,!ð !ð !ð !ð !˜2œ9ñ !ô !ð !ð, Ð˜YÑ'Ô'ðJð Jð Jð Jð J�R”Yñ Jô Jñ (Ô'ðJð(+ð +ð +ð +ð + ¤ñ +ô +ð +ðj ð%ð %ØŒIð%à”,ð%ð ”ð%ð ”,ð	%ð
 ”L 4Ñ'ð%ð ð%ð �S‰[ð%ð Ð'Ô(ð%ð %ð %ð %ð,9)ð 9)ð 9)ð 9)ð 9)�r”yñ 9)ô 9)ð 9)ðxð ð ð ð ˜"œ)ñ ô ð ð@ ðð ð ð ð ˜_ñ ô ñ „ðð> ðjð jð jð jð jÐ)ñ jô jñ „ðjðZm2ð m2ð m2ð m2ð m2Ð 6ñ m2ô m2ð m2ð` RÐ
QÐ
Q€€€r"   