Ë
    S^(hªÁ  ã                   óJ  — d Z ddlZddlmZ ddlmZ ddlmZmZm	Z	 ddl
Z
ddlZ
ddl
mZ ddlmZmZmZ dd	lmZ  e«       rdd
lmZ ddlmZ ddlmZmZmZmZmZmZmZmZm Z  ddl!m"Z" ddl#m$Z$ ddlm%Z%m&Z&m'Z'm(Z(m)Z) ddl*m+Z+  e(jX                  e-«      Z.dZ/dZ0d„ Z1d„ Z2d„ Z3 G d„ dejh                  «      Z5 G d„ dejh                  «      Z6 G d„ dejh                  «      Z7 G d„ dejh                  «      Z8 G d„ d ejh                  «      Z9 G d!„ d"ejh                  «      Z: G d#„ d$ejh                  «      Z; G d%„ d&ejh                  «      Z< G d'„ d(ejh                  «      Z= G d)„ d*ejh                  «      Z> G d+„ d,ejh                  «      Z? G d-„ d.ejh                  «      Z@ G d/„ d0ejh                  «      ZA G d1„ d2ejh                  «      ZB G d3„ d4e"«      ZCe G d5„ d6e«      «       ZDd7ZEd8ZF e&d9eE«       G d:„ d;eC«      «       ZG e&d<eE«       G d=„ d>eC«      «       ZH e&d?eE«       G d@„ dAeC«      «       ZI e&dBeE«       G dC„ dDeC«      «       ZJ e&dEeE«       G dF„ dGeC«      «       ZK e&dHeE«       G dI„ dJeC«      «       ZL e&dKeE«       G dL„ dMeC«      «       ZM e&dNeE«       G dO„ dPeC«      «       ZNg dQ¢ZOy)RzPyTorch FNet model.é    N)Ú	dataclass)Úpartial)ÚOptionalÚTupleÚUnion)Únn)ÚBCEWithLogitsLossÚCrossEntropyLossÚMSELossé   )Úis_scipy_available)Úlinalg)ÚACT2FN)	ÚBaseModelOutputÚBaseModelOutputWithPoolingÚMaskedLMOutputÚModelOutputÚMultipleChoiceModelOutputÚNextSentencePredictorOutputÚQuestionAnsweringModelOutputÚSequenceClassifierOutputÚTokenClassifierOutput)ÚPreTrainedModel)Úapply_chunking_to_forward)Úadd_code_sample_docstringsÚadd_start_docstringsÚ%add_start_docstrings_to_model_forwardÚloggingÚreplace_return_docstringsé   )Ú
FNetConfigzgoogle/fnet-baser!   c                 ó¤   — | j                   d   }|d|…d|…f   }| j                  t        j                  «      } t        j                  d| ||«      S )z4Applies 2D matrix multiplication to 3D input arrays.r    Nzbij,jk,ni->bnk)ÚshapeÚtypeÚtorchÚ	complex64Úeinsum)ÚxÚmatrix_dim_oneÚmatrix_dim_twoÚ
seq_lengths       úd/var/www/skyplay_api_hub/venv/lib/python3.12/site-packages/transformers/models/fnet/modeling_fnet.pyÚ_two_dim_matmulr-   @   sN   € à—‘˜‘€JØ# K Z K°°*°Ð$<Ñ=€NØ	�‰Œu�‰Ó€AÜ�<‰<Ð(¨!¨^¸^ÓLÐLó    c                 ó   — t        | ||«      S ©N)r-   )r(   r)   r*   s      r,   Útwo_dim_matmulr1   I   s   € Ü˜1˜n¨nÓ=Ð=r.   c                 ó˜   — | }t        t        | j                  «      dd «      D ]#  }t        j                  j	                  ||¬«      }Œ% |S )zÑ
    Applies n-dimensional Fast Fourier Transform (FFT) to input array.

    Args:
        x: Input n-dimensional array.

    Returns:
        n-dimensional Fourier transform of input n-dimensional array.
    r    N)Úaxis)ÚreversedÚrangeÚndimr%   Úfft)r(   Úoutr3   s      r,   Úfftnr9   N   sG   € ð €CÜœ˜qŸv™v› q rÐ*Ó+ò ,ˆÜ�i‰i�m‰m˜C dˆmÓ+‰ð,à€Jr.   c                   ó*   ‡ — e Zd ZdZˆ fd„Zdd„Zˆ xZS )ÚFNetEmbeddingszGConstruct the embeddings from word, position and token_type embeddings.c                 óx  •— t         ‰| �  «        t        j                  |j                  |j
                  |j                  ¬«      | _        t        j                  |j                  |j
                  «      | _	        t        j                  |j                  |j
                  «      | _        t        j                  |j
                  |j                  ¬«      | _        t        j                  |j
                  |j
                  «      | _        t        j                   |j"                  «      | _        | j'                  dt)        j*                  |j                  «      j-                  d«      d¬«       | j'                  dt)        j.                  | j0                  j3                  «       t(        j4                  ¬«      d¬«       y )	N)Úpadding_idx©ÚepsÚposition_ids)r    éÿÿÿÿF)Ú
persistentÚtoken_type_ids©Údtype)ÚsuperÚ__init__r   Ú	EmbeddingÚ
vocab_sizeÚhidden_sizeÚpad_token_idÚword_embeddingsÚmax_position_embeddingsÚposition_embeddingsÚtype_vocab_sizeÚtoken_type_embeddingsÚ	LayerNormÚlayer_norm_epsÚLinearÚ
projectionÚDropoutÚhidden_dropout_probÚdropoutÚregister_bufferr%   ÚarangeÚexpandÚzerosr@   ÚsizeÚlong©ÚselfÚconfigÚ	__class__s     €r,   rG   zFNetEmbeddings.__init__a   s<  ø€ Ü‰ÑÔÜ!Ÿ|™|¨F×,=Ñ,=¸v×?QÑ?QÐ_e×_rÑ_rÔsˆÔÜ#%§<¡<°×0NÑ0NÐPV×PbÑPbÓ#cˆÔ Ü%'§\¡\°&×2HÑ2HÈ&×J\ÑJ\Ó%]ˆÔ"ô Ÿ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒäŸ)™) F×$6Ñ$6¸×8JÑ8JÓKˆŒÜ—z‘z &×"<Ñ"<Ó=ˆŒð 	×ÑØœEŸL™L¨×)GÑ)GÓH×OÑOÐPWÓXÐejð 	ô 	
ð 	×ÑØœeŸk™k¨$×*;Ñ*;×*@Ñ*@Ó*BÌ%Ï*É*ÔUÐbgð 	õ 	
r.   c                 óX  — |�|j                  «       }n|j                  «       d d }|d   }|€| j                  d d …d |…f   }|€st        | d«      r-| j                  d d …d |…f   }|j	                  |d   |«      }|}n:t        j                  |t
        j                  | j                  j                  ¬«      }|€| j                  |«      }| j                  |«      }	||	z   }
| j                  |«      }|
|z  }
| j                  |
«      }
| j                  |
«      }
| j                  |
«      }
|
S )NrA   r    rC   r   ©rE   Údevice)r\   r@   ÚhasattrrC   rZ   r%   r[   r]   rd   rL   rP   rN   rQ   rT   rW   )r_   Ú	input_idsrC   r@   Úinputs_embedsÚinput_shaper+   Úbuffered_token_type_idsÚ buffered_token_type_ids_expandedrP   Ú
embeddingsrN   s               r,   ÚforwardzFNetEmbeddings.forwardw   s=  € ØÐ Ø#Ÿ.™.Ó*‰Kà'×,Ñ,Ó.¨s°Ð3ˆKà  ‘^ˆ
àÐØ×,Ñ,ªQ°°°¨^Ñ<ˆLð
 Ð!Ü�tÐ-Ô.Ø*.×*=Ñ*=ºaÀÀ*À¸nÑ*MÐ'Ø3J×3QÑ3QÐR]Ð^_ÑR`ÐblÓ3mÐ0Ø!A‘ä!&§¡¨[ÄÇ
Á
ÐSW×SdÑSd×SkÑSkÔ!l�àÐ Ø ×0Ñ0°Ó;ˆMØ $× :Ñ :¸>Ó JÐà"Ð%:Ñ:ˆ
à"×6Ñ6°|ÓDÐØÐ)Ñ)ˆ
Ø—^‘^ JÓ/ˆ
Ø—_‘_ ZÓ0ˆ
Ø—\‘\ *Ó-ˆ
ØÐr.   )NNNN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__rG   rl   Ú__classcell__©ra   s   @r,   r;   r;   ^   s   ø„ ÙQô
÷,!r.   r;   c                   ó*   ‡ — e Zd Zˆ fd„Zd„ Zd„ Zˆ xZS )ÚFNetBasicFourierTransformc                 óD   •— t         ‰| �  «        | j                  |«       y r0   )rF   rG   Ú_init_fourier_transformr^   s     €r,   rG   z"FNetBasicFourierTransform.__init__œ   s   ø€ Ü‰ÑÔØ×$Ñ$ VÕ,r.   c                 óœ  — |j                   s+t        t        j                  j                  d¬«      | _        y |j                  dk  rût        «       rÐ| j                  dt        j                  t        j                  |j                  «      t        j                  ¬«      «       | j                  dt        j                  t        j                  |j                  «      t        j                  ¬«      «       t        t        | j                   | j"                  ¬«      | _        y t%        j&                  d«       t        | _        y t        | _        y )	N)r    é   ©Údimé   Údft_mat_hiddenrD   Údft_mat_seq)r)   r*   zpSciPy is needed for DFT matrix calculation and is not found. Using TPU optimized fast fourier transform instead.)Úuse_tpu_fourier_optimizationsr   r%   r7   r9   Úfourier_transformrM   r   rX   Útensorr   ÚdftrJ   r&   Útpu_short_seq_lengthr1   r}   r|   r   Úwarning)r_   r`   s     r,   rv   z1FNetBasicFourierTransform._init_fourier_transform    sè   € Ø×3Ò3Ü%,¬U¯Y©Y¯^©^ÀÔ%HˆDÕ"Ø×+Ñ+¨tÒ3Ü!Ô#Ø×$Ñ$Ø$¤e§l¡l´6·:±:¸f×>PÑ>PÓ3QÔY^×YhÑYhÔ&iôð ×$Ñ$Ø!¤5§<¡<´·
±
¸6×;VÑ;VÓ0WÔ_d×_nÑ_nÔ#oôô *1Ü"°4×3CÑ3CÐTX×TgÑTgô*�Õ&ô —‘ð*ôô *.�Õ&ä%)ˆDÕ"r.   c                 ó>   — | j                  |«      j                  }|fS r0   )r   Úreal)r_   Úhidden_statesÚoutputss      r,   rl   z!FNetBasicFourierTransform.forward·   s"   € ð ×(Ñ(¨Ó7×<Ñ<ˆØˆzÐr.   )rm   rn   ro   rG   rv   rl   rq   rr   s   @r,   rt   rt   ›   s   ø„ ô-ò*ö.r.   rt   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚFNetBasicOutputc                 ó‚   •— t         ‰| �  «        t        j                  |j                  |j
                  ¬«      | _        y ©Nr>   )rF   rG   r   rQ   rJ   rR   r^   s     €r,   rG   zFNetBasicOutput.__init__Â   s,   ø€ Ü‰ÑÔÜŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆ�r.   c                 ó.   — | j                  ||z   «      }|S r0   )rQ   ©r_   r†   Úinput_tensors      r,   rl   zFNetBasicOutput.forwardÆ   s   € ØŸ™ |°mÑ'CÓDˆØÐr.   ©rm   rn   ro   rG   rl   rq   rr   s   @r,   r‰   r‰   Á   s   ø„ ôUör.   r‰   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚFNetFourierTransformc                 ób   •— t         ‰| �  «        t        |«      | _        t	        |«      | _        y r0   )rF   rG   rt   r_   r‰   Úoutputr^   s     €r,   rG   zFNetFourierTransform.__init__Ì   s&   ø€ Ü‰ÑÔÜ-¨fÓ5ˆŒ	Ü% fÓ-ˆ�r.   c                 óX   — | j                  |«      }| j                  |d   |«      }|f}|S ©Nr   )r_   r“   )r_   r†   Úself_outputsÚfourier_outputr‡   s        r,   rl   zFNetFourierTransform.forwardÑ   s1   € Ø—y‘y Ó/ˆØŸ™ \°!¡_°mÓDˆØ!Ð#ˆØˆr.   r�   rr   s   @r,   r‘   r‘   Ë   s   ø„ ô.ö
r.   r‘   c                   óV   ‡ — e Zd Zˆ fd„Zdej
                  dej
                  fd„Zˆ xZS )ÚFNetIntermediatec                 ó  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        |j                  t        «      rt        |j                     | _        y |j                  | _        y r0   )rF   rG   r   rS   rJ   Úintermediate_sizeÚdenseÚ
isinstanceÚ
hidden_actÚstrr   Úintermediate_act_fnr^   s     €r,   rG   zFNetIntermediate.__init__Ú   s]   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3KÑ3KÓLˆŒ
Ü�f×'Ñ'¬Ô-Ü'-¨f×.?Ñ.?Ñ'@ˆDÕ$à'-×'8Ñ'8ˆDÕ$r.   r†   Úreturnc                 óJ   — | j                  |«      }| j                  |«      }|S r0   )rœ   r    ©r_   r†   s     r,   rl   zFNetIntermediate.forwardâ   s&   € ØŸ
™
 =Ó1ˆØ×0Ñ0°Ó?ˆØÐr.   ©rm   rn   ro   rG   r%   ÚTensorrl   rq   rr   s   @r,   r™   r™   Ù   s#   ø„ ô9ð U§\¡\ð °e·l±l÷ r.   r™   c                   ón   ‡ — e Zd Zˆ fd„Zdej
                  dej
                  dej
                  fd„Zˆ xZS )Ú
FNetOutputc                 ó(  •— t         ‰| �  «        t        j                  |j                  |j
                  «      | _        t        j                  |j
                  |j                  ¬«      | _        t        j                  |j                  «      | _        y r‹   )rF   rG   r   rS   r›   rJ   rœ   rQ   rR   rU   rV   rW   r^   s     €r,   rG   zFNetOutput.__init__ê   s`   ø€ Ü‰ÑÔÜ—Y‘Y˜v×7Ñ7¸×9KÑ9KÓLˆŒ
ÜŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆŒÜ—z‘z &×"<Ñ"<Ó=ˆ�r.   r†   rŽ   r¡   c                 ór   — | j                  |«      }| j                  |«      }| j                  ||z   «      }|S r0   )rœ   rW   rQ   r�   s      r,   rl   zFNetOutput.forwardð   s7   € ØŸ
™
 =Ó1ˆØŸ™ ]Ó3ˆØŸ™ }°|Ñ'CÓDˆØÐr.   r¤   rr   s   @r,   r§   r§   é   s1   ø„ ô>ð U§\¡\ð ÀÇÁð ÐRW×R^ÑR^÷ r.   r§   c                   ó*   ‡ — e Zd Zˆ fd„Zd„ Zd„ Zˆ xZS )Ú	FNetLayerc                 ó²   •— t         ‰| �  «        |j                  | _        d| _        t	        |«      | _        t        |«      | _        t        |«      | _	        y ©Nr    )
rF   rG   Úchunk_size_feed_forwardÚseq_len_dimr‘   Úfourierr™   Úintermediater§   r“   r^   s     €r,   rG   zFNetLayer.__init__ø   sI   ø€ Ü‰ÑÔØ'-×'EÑ'EˆÔ$ØˆÔÜ+¨FÓ3ˆŒÜ,¨VÓ4ˆÔÜ  Ó(ˆ�r.   c                 ó�   — | j                  |«      }|d   }t        | j                  | j                  | j                  |«      }|f}|S r•   )r°   r   Úfeed_forward_chunkr®   r¯   )r_   r†   Úself_fourier_outputsr—   Úlayer_outputr‡   s         r,   rl   zFNetLayer.forward   sO   € Ø#Ÿ|™|¨MÓ:ÐØ-¨aÑ0ˆä0Ø×#Ñ# T×%AÑ%AÀ4×CSÑCSÐUcó
ˆð  �/ˆàˆr.   c                 óL   — | j                  |«      }| j                  ||«      }|S r0   )r±   r“   )r_   r—   Úintermediate_outputrµ   s       r,   r³   zFNetLayer.feed_forward_chunk  s*   € Ø"×/Ñ/°Ó?ÐØ—{‘{Ð#6¸ÓGˆØÐr.   )rm   rn   ro   rG   rl   r³   rq   rr   s   @r,   r«   r«   ÷   s   ø„ ô)ò
ör.   r«   c                   ó&   ‡ — e Zd Zˆ fd„Zdd„Zˆ xZS )ÚFNetEncoderc                 óÐ   •— t         ‰| �  «        || _        t        j                  t        |j                  «      D �cg c]  }t        |«      ‘Œ c}«      | _        d| _	        y c c}w )NF)
rF   rG   r`   r   Ú
ModuleListr5   Únum_hidden_layersr«   ÚlayerÚgradient_checkpointing)r_   r`   Ú_ra   s      €r,   rG   zFNetEncoder.__init__  sN   ø€ Ü‰ÑÔØˆŒÜ—]‘]¼uÀV×E]ÑE]Ó?^Ö#_¸!¤I¨fÕ$5Ò#_Ó`ˆŒ
Ø&+ˆÕ#ùò $`s   ½A#c                 ó2  — |rdnd }t        | j                  «      D ]O  \  }}|r||fz   }| j                  r)| j                  r| j	                  |j
                  |«      }n ||«      }|d   }ŒQ |r||fz   }|st        d„ ||fD «       «      S t        ||¬«      S )N© r   c              3   ó&   K  — | ]	  }|€Œ|–— Œ y ­wr0   rÁ   )Ú.0Úvs     r,   ú	<genexpr>z&FNetEncoder.forward.<locals>.<genexpr>+  s   è ø€ ÒX˜qÈ!É-œÑXùs   ‚Š)Úlast_hidden_stater†   )Ú	enumerater½   r¾   ÚtrainingÚ_gradient_checkpointing_funcÚ__call__Útupler   )r_   r†   Úoutput_hidden_statesÚreturn_dictÚall_hidden_statesÚiÚlayer_moduleÚlayer_outputss           r,   rl   zFNetEncoder.forward  s°   € Ù"6™B¸DÐä(¨¯©Ó4ò 		-‰OˆAˆ|Ù#Ø$5¸Ð8HÑ$HÐ!à×*Ò*¨t¯}ª}Ø $× AÑ AÀ,×BWÑBWÐYfÓ g‘á ,¨]Ó ;�à)¨!Ñ,‰Mð		-ñ  Ø 1°]Ð4DÑ DÐáÜÑX ]Ð4EÐ$FÔXÓXÐXä°ÐN_Ô`Ð`r.   )FTr�   rr   s   @r,   r¹   r¹     s   ø„ ô,÷ar.   r¹   c                   óV   ‡ — e Zd Zˆ fd„Zdej
                  dej
                  fd„Zˆ xZS )Ú
FNetPoolerc                 ó²   •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        j                  «       | _        y r0   )rF   rG   r   rS   rJ   rœ   ÚTanhÚ
activationr^   s     €r,   rG   zFNetPooler.__init__2  s9   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
ÜŸ'™'›)ˆ�r.   r†   r¡   c                 ó\   — |d d …df   }| j                  |«      }| j                  |«      }|S r•   )rœ   rÖ   )r_   r†   Úfirst_token_tensorÚpooled_outputs       r,   rl   zFNetPooler.forward7  s6   € ð +ª1¨a¨4Ñ0ÐØŸ
™
Ð#5Ó6ˆØŸ™¨Ó6ˆØÐr.   r¤   rr   s   @r,   rÓ   rÓ   1  s#   ø„ ô$ð
 U§\¡\ð °e·l±l÷ r.   rÓ   c                   óV   ‡ — e Zd Zˆ fd„Zdej
                  dej
                  fd„Zˆ xZS )ÚFNetPredictionHeadTransformc                 óh  •— t         ‰| �  «        t        j                  |j                  |j                  «      | _        t        |j                  t        «      rt        |j                     | _
        n|j                  | _
        t        j                  |j                  |j                  ¬«      | _        y r‹   )rF   rG   r   rS   rJ   rœ   r�   rž   rŸ   r   Útransform_act_fnrQ   rR   r^   s     €r,   rG   z$FNetPredictionHeadTransform.__init__B  s{   ø€ Ü‰ÑÔÜ—Y‘Y˜v×1Ñ1°6×3EÑ3EÓFˆŒ
Ü�f×'Ñ'¬Ô-Ü$*¨6×+<Ñ+<Ñ$=ˆDÕ!à$*×$5Ñ$5ˆDÔ!ÜŸ™ f×&8Ñ&8¸f×>SÑ>SÔTˆ�r.   r†   r¡   c                 ól   — | j                  |«      }| j                  |«      }| j                  |«      }|S r0   )rœ   rÝ   rQ   r£   s     r,   rl   z#FNetPredictionHeadTransform.forwardK  s4   € ØŸ
™
 =Ó1ˆØ×-Ñ-¨mÓ<ˆØŸ™ }Ó5ˆØÐr.   r¤   rr   s   @r,   rÛ   rÛ   A  s$   ø„ ôUð U§\¡\ð °e·l±l÷ r.   rÛ   c                   ó,   ‡ — e Zd Zˆ fd„Zd„ Zdd„Zˆ xZS )ÚFNetLMPredictionHeadc                 óD  •— t         ‰| �  «        t        |«      | _        t	        j
                  |j                  |j                  «      | _        t	        j                  t        j                  |j                  «      «      | _        | j                  | j                  _        y r0   )rF   rG   rÛ   Ú	transformr   rS   rJ   rI   ÚdecoderÚ	Parameterr%   r[   Úbiasr^   s     €r,   rG   zFNetLMPredictionHead.__init__S  si   ø€ Ü‰ÑÔÜ4°VÓ<ˆŒô —y‘y ×!3Ñ!3°V×5FÑ5FÓGˆŒä—L‘L¤§¡¨V×->Ñ->Ó!?Ó@ˆŒ	Ø ŸI™Iˆ�‰Õr.   c                 óJ   — | j                  |«      }| j                  |«      }|S r0   )râ   rã   r£   s     r,   rl   zFNetLMPredictionHead.forward^  s$   € ØŸ™ }Ó5ˆØŸ™ ]Ó3ˆØÐr.   c                 óÌ   — | j                   j                  j                  j                  dk(  r| j                  | j                   _        y | j                   j                  | _        y )NÚmeta)rã   rå   rd   r$   ©r_   s    r,   Ú_tie_weightsz!FNetLMPredictionHead._tie_weightsc  sC   € à�<‰<×Ñ×#Ñ#×(Ñ(¨FÒ2Ø $§	¡	ˆD�L‰LÕð Ÿ™×)Ñ)ˆD�Ir.   )r¡   N)rm   rn   ro   rG   rl   rê   rq   rr   s   @r,   rà   rà   R  s   ø„ ô	&ò÷
*r.   rà   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚFNetOnlyMLMHeadc                 óB   •— t         ‰| �  «        t        |«      | _        y r0   )rF   rG   rà   Úpredictionsr^   s     €r,   rG   zFNetOnlyMLMHead.__init__m  s   ø€ Ü‰ÑÔÜ/°Ó7ˆÕr.   c                 ó(   — | j                  |«      }|S r0   )rî   )r_   Úsequence_outputÚprediction_scoress      r,   rl   zFNetOnlyMLMHead.forwardq  s   € Ø ×,Ñ,¨_Ó=ÐØ Ð r.   r�   rr   s   @r,   rì   rì   l  s   ø„ ô8ö!r.   rì   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚFNetOnlyNSPHeadc                 ól   •— t         ‰| �  «        t        j                  |j                  d«      | _        y ©Nrx   )rF   rG   r   rS   rJ   Úseq_relationshipr^   s     €r,   rG   zFNetOnlyNSPHead.__init__x  s'   ø€ Ü‰ÑÔÜ "§	¡	¨&×*<Ñ*<¸aÓ @ˆÕr.   c                 ó(   — | j                  |«      }|S r0   )rö   )r_   rÙ   Úseq_relationship_scores      r,   rl   zFNetOnlyNSPHead.forward|  s   € Ø!%×!6Ñ!6°}Ó!EÐØ%Ð%r.   r�   rr   s   @r,   ró   ró   w  s   ø„ ôAö&r.   ró   c                   ó$   ‡ — e Zd Zˆ fd„Zd„ Zˆ xZS )ÚFNetPreTrainingHeadsc                 óŒ   •— t         ‰| �  «        t        |«      | _        t	        j
                  |j                  d«      | _        y rõ   )rF   rG   rà   rî   r   rS   rJ   rö   r^   s     €r,   rG   zFNetPreTrainingHeads.__init__ƒ  s4   ø€ Ü‰ÑÔÜ/°Ó7ˆÔÜ "§	¡	¨&×*<Ñ*<¸aÓ @ˆÕr.   c                 óN   — | j                  |«      }| j                  |«      }||fS r0   )rî   rö   )r_   rð   rÙ   rñ   rø   s        r,   rl   zFNetPreTrainingHeads.forwardˆ  s0   € Ø ×,Ñ,¨_Ó=ÐØ!%×!6Ñ!6°}Ó!EÐØ Ð"8Ð8Ð8r.   r�   rr   s   @r,   rú   rú   ‚  s   ø„ ôAö
9r.   rú   c                   ó"   — e Zd ZdZeZdZdZd„ Zy)ÚFNetPreTrainedModelz†
    An abstract class to handle weights initialization and a simple interface for downloading and loading pretrained
    models.
    ÚfnetTc                 ó  — t        |t        j                  «      rm|j                  j                  j                  d| j                  j                  ¬«       |j                  �%|j                  j                  j                  «        yyt        |t        j                  «      rz|j                  j                  j                  d| j                  j                  ¬«       |j                  �2|j                  j                  |j                     j                  «        yyt        |t        j                  «      rJ|j                  j                  j                  «        |j                  j                  j                  d«       yy)zInitialize the weightsg        )ÚmeanÚstdNg      ð?)r�   r   rS   ÚweightÚdataÚnormal_r`   Úinitializer_rangerå   Úzero_rH   r=   rQ   Úfill_)r_   Úmodules     r,   Ú_init_weightsz!FNetPreTrainedModel._init_weights˜  s  € ä�fœbŸi™iÔ(ð �M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSà�{‰{Ð&Ø—‘× Ñ ×&Ñ&Õ(ð 'ä˜¤§¡Ô-Ø�M‰M×Ñ×&Ñ&¨C°T·[±[×5RÑ5RÐ&ÔSØ×!Ñ!Ð-Ø—‘×"Ñ" 6×#5Ñ#5Ñ6×<Ñ<Õ>ð .ä˜¤§¡Ô-Ø�K‰K×Ñ×"Ñ"Ô$Ø�M‰M×Ñ×$Ñ$ SÕ)ð .r.   N)	rm   rn   ro   rp   r!   Úconfig_classÚbase_model_prefixÚsupports_gradient_checkpointingr
  rÁ   r.   r,   rþ   rþ   Ž  s   „ ñð
 €LØÐØ&*Ð#ó*r.   rþ   c                   ó¸   — e Zd ZU dZdZeej                     ed<   dZ	eej                     ed<   dZ
eej                     ed<   dZeeej                        ed<   y)ÚFNetForPreTrainingOutputa©  
    Output type of [`FNetForPreTraining`].

    Args:
        loss (*optional*, returned when `labels` is provided, `torch.FloatTensor` of shape `(1,)`):
            Total loss as the sum of the masked language modeling loss and the next sequence prediction
            (classification) loss.
        prediction_logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`):
            Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax).
        seq_relationship_logits (`torch.FloatTensor` of shape `(batch_size, 2)`):
            Prediction scores of the next sequence prediction (classification) head (scores of True/False continuation
            before SoftMax).
        hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`):
            Tuple of `torch.FloatTensor` (one for the output of the embeddings + one for the output of each layer) of
            shape `(batch_size, sequence_length, hidden_size)`. Hidden-states of the model at the output of each layer
            plus the initial embedding outputs.
    NÚlossÚprediction_logitsÚseq_relationship_logitsr†   )rm   rn   ro   rp   r  r   r%   ÚFloatTensorÚ__annotations__r  r  r†   r   rÁ   r.   r,   r  r  ª  sd   … ñð$ )-€Dˆ(�5×$Ñ$Ñ
%Ó,Ø59Ð�x × 1Ñ 1Ñ2Ó9Ø;?Ð˜X e×&7Ñ&7Ñ8Ó?Ø8<€M�8˜E %×"3Ñ"3Ñ4Ñ5Ô<r.   r  aG  
    This model is a PyTorch [torch.nn.Module](https://pytorch.org/docs/stable/nn.html#torch.nn.Module) sub-class. Use
    it as a regular PyTorch Module and refer to the PyTorch documentation for all matter related to general usage and
    behavior.

    Parameters:
        config ([`FNetConfig`]): Model configuration class with all the parameters of the model.
            Initializing with a config file does not load the weights associated with the model, only the
            configuration. Check out the [`~PreTrainedModel.from_pretrained`] method to load the model weights.
aÌ  
    Args:
        input_ids (`torch.LongTensor` of shape `({0})`):
            Indices of input sequence tokens in the vocabulary.

            Indices can be obtained using [`AutoTokenizer`]. See [`PreTrainedTokenizer.encode`] and
            [`PreTrainedTokenizer.__call__`] for details.

            [What are input IDs?](../glossary#input-ids)
        token_type_ids (`torch.LongTensor` of shape `({0})`, *optional*):
            Segment token indices to indicate first and second portions of the inputs. Indices are selected in `[0,
            1]`:

            - 0 corresponds to a *sentence A* token,
            - 1 corresponds to a *sentence B* token.

            [What are token type IDs?](../glossary#token-type-ids)
        position_ids (`torch.LongTensor` of shape `({0})`, *optional*):
            Indices of positions of each input sequence tokens in the position embeddings. Selected in the range `[0,
            config.max_position_embeddings - 1]`.

            [What are position IDs?](../glossary#position-ids)

        inputs_embeds (`torch.FloatTensor` of shape `({0}, hidden_size)`, *optional*):
            Optionally, instead of passing `input_ids` you can choose to directly pass an embedded representation. This
            is useful if you want more control over how to convert *input_ids* indices into associated vectors than the
            model's internal embedding lookup matrix.
        output_hidden_states (`bool`, *optional*):
            Whether or not to return the hidden states of all layers. See `hidden_states` under returned tensors for
            more detail.
        return_dict (`bool`, *optional*):
            Whether or not to return a [`~utils.ModelOutput`] instead of a plain tuple.
z^The bare FNet Model transformer outputting raw hidden-states without any specific head on top.c                   ó0  ‡ — e Zd ZdZdˆ fd„	Zd„ Zd„ Z eej                  d«      «       e
eee¬«      	 	 	 	 	 	 ddeej                      deej                      d	eej                      d
eej"                     dee   dee   deeef   fd„«       «       Zˆ xZS )Ú	FNetModelzð

    The model can behave as an encoder, following the architecture described in [FNet: Mixing Tokens with Fourier
    Transforms](https://arxiv.org/abs/2105.03824) by James Lee-Thorp, Joshua Ainslie, Ilya Eckstein, Santiago Ontanon.

    c                 óº   •— t         ‰| �  |«       || _        t        |«      | _        t        |«      | _        |rt        |«      nd | _        | j                  «        y r0   )
rF   rG   r`   r;   rk   r¹   ÚencoderrÓ   ÚpoolerÚ	post_init)r_   r`   Úadd_pooling_layerra   s      €r,   rG   zFNetModel.__init__þ  sK   ø€ Ü‰Ñ˜Ô ØˆŒä(¨Ó0ˆŒÜ" 6Ó*ˆŒá,=”j Ô(À4ˆŒð 	�‰Õr.   c                 ó.   — | j                   j                  S r0   ©rk   rL   ré   s    r,   Úget_input_embeddingszFNetModel.get_input_embeddings
  s   € Ø�‰×.Ñ.Ð.r.   c                 ó&   — || j                   _        y r0   r  )r_   Úvalues     r,   Úset_input_embeddingszFNetModel.set_input_embeddings  s   € Ø*/ˆ�‰Õ'r.   úbatch_size, sequence_length©Ú
checkpointÚoutput_typer  rf   rC   r@   rg   rÌ   rÍ   r¡   c                 ó|  — |�|n| j                   j                  }|�|n| j                   j                  }|�|�t        d«      ‚|�|j	                  «       }|\  }}	n&|�|j	                  «       d d }|\  }}	nt        d«      ‚| j                   j
                  r)|	dk  r$| j                   j                  |	k7  rt        d«      ‚|�|j                  n|j                  }
|€pt        | j                  d«      r4| j                  j                  d d …d |	…f   }|j                  ||	«      }|}n&t        j                  |t        j                  |
¬«      }| j                  ||||¬«      }| j                  |||¬	«      }|d
   }| j                   �| j!                  |«      nd }|s
||f|dd  z   S t#        |||j$                  ¬«      S )NzDYou cannot specify both input_ids and inputs_embeds at the same timerA   z5You have to specify either input_ids or inputs_embedsr{   z‹The `tpu_short_seq_length` in FNetConfig should be set equal to the sequence length being passed to the model when using TPU optimizations.rC   rc   )rf   r@   rC   rg   )rÌ   rÍ   r   r    )rÆ   Úpooler_outputr†   )r`   rÌ   Úuse_return_dictÚ
ValueErrorr\   r~   r‚   rd   re   rk   rC   rZ   r%   r[   r]   r  r  r   r†   )r_   rf   rC   r@   rg   rÌ   rÍ   rh   Ú
batch_sizer+   rd   ri   rj   Úembedding_outputÚencoder_outputsrð   r'  s                    r,   rl   zFNetModel.forward  sì  € ð  %9Ð$DÑ È$Ï+É+×JjÑJjð 	ð &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆàÐ  ]Ð%>ÜÐcÓdÐdØÐ"Ø#Ÿ.™.Ó*ˆKØ%0Ñ"ˆJ™
ØÐ&Ø'×,Ñ,Ó.¨s°Ð3ˆKØ%0Ñ"ˆJ™
äÐTÓUÐUð �K‰K×5Ò5Ø˜dÒ"Ø—‘×0Ñ0°JÒ>äð;óð ð
 &/Ð%:�×!Ò!À×@TÑ@TˆàÐ!Ü�t—‘Ð(8Ô9Ø*.¯/©/×*HÑ*HÊÈKÈZÈKÈÑ*XÐ'Ø3J×3QÑ3QÐR\Ð^hÓ3iÐ0Ø!A‘ä!&§¡¨[ÄÇ
Á
ÐSYÔ!Z�àŸ?™?ØØ%Ø)Ø'ð	 +ó 
Ðð Ÿ,™,ØØ!5Ø#ð 'ó 
ˆð
 *¨!Ñ,ˆà8<¿¹Ð8O˜Ÿ™ OÔ4ÐUYˆáØ# ]Ð3°oÀaÀbÐ6IÑIÐIä)Ø-Ø'Ø)×7Ñ7ô
ð 	
r.   )T)NNNNNN)rm   rn   ro   rp   rG   r  r!  r   ÚFNET_INPUTS_DOCSTRINGÚformatr   Ú_CHECKPOINT_FOR_DOCr   Ú_CONFIG_FOR_DOCr   r%   Ú
LongTensorr  Úboolr   rË   rl   rq   rr   s   @r,   r  r  ò  só   ø„ ñ
õ
ò/ò0ñ +Ð+@×+GÑ+GÐHeÓ+fÓgÙØ&Ø#Ø$ôð 15Ø59Ø37Ø59Ø/3Ø&*ñC
à˜E×,Ñ,Ñ-ðC
ð ! ×!1Ñ!1Ñ2ðC
ð ˜u×/Ñ/Ñ0ð	C
ð
   × 1Ñ 1Ñ2ðC
ð ' t™nðC
ð ˜d‘^ðC
ð 
ˆu�oÐ%Ñ	&òC
óó hôC
r.   r  z¨
    FNet Model with two heads on top as done during the pretraining: a `masked language modeling` head and a `next
    sentence prediction (classification)` head.
    c                   óp  ‡ — e Zd ZddgZˆ fd„Zd„ Zd„ Z eej                  d«      «       e
ee¬«      	 	 	 	 	 	 	 	 ddeej                     d	eej                     d
eej                     deej                     deej                     deej                     dee   dee   deeef   fd„«       «       Zˆ xZS )ÚFNetForPreTrainingúcls.predictions.decoder.biasúcls.predictions.decoder.weightc                 ó„   •— t         ‰| �  |«       t        |«      | _        t	        |«      | _        | j                  «        y r0   )rF   rG   r  rÿ   rú   Úclsr  r^   s     €r,   rG   zFNetForPreTraining.__init__f  s4   ø€ Ü‰Ñ˜Ô ä˜fÓ%ˆŒ	Ü'¨Ó/ˆŒð 	�‰Õr.   c                 óB   — | j                   j                  j                  S r0   ©r8  rî   rã   ré   s    r,   Úget_output_embeddingsz(FNetForPreTraining.get_output_embeddingso  ó   € Ø�x‰x×#Ñ#×+Ñ+Ð+r.   c                 ó„   — || j                   j                  _        |j                  | j                   j                  _        y r0   ©r8  rî   rã   rå   ©r_   Únew_embeddingss     r,   Úset_output_embeddingsz(FNetForPreTraining.set_output_embeddingsr  ó,   € Ø'5ˆ�‰×ÑÔ$Ø$2×$7Ñ$7ˆ�‰×ÑÕ!r.   r"  ©r%  r  rf   rC   r@   rg   ÚlabelsÚnext_sentence_labelrÌ   rÍ   r¡   c	                 óî  — |�|n| j                   j                  }| j                  ||||||¬«      }	|	dd \  }
}| j                  |
|«      \  }}d}|�u|�st	        «       } ||j                  d| j                   j                  «      |j                  d«      «      } ||j                  dd«      |j                  d«      «      }||z   }|s||f|	dd z   }|�|f|z   S |S t        ||||	j                  ¬«      S )aà  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should be in `[-100, 0, ...,
            config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-100` are ignored (masked), the
            loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`
        next_sentence_label (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the next sequence prediction (classification) loss. Input should be a sequence pair
            (see `input_ids` docstring) Indices should be in `[0, 1]`:

            - 0 indicates sequence B is a continuation of sequence A,
            - 1 indicates sequence B is a random sequence.
        kwargs (`Dict[str, any]`, *optional*, defaults to `{}`):
            Used to hide legacy arguments that have been deprecated.

        Returns:

        Example:

        ```python
        >>> from transformers import AutoTokenizer, FNetForPreTraining
        >>> import torch

        >>> tokenizer = AutoTokenizer.from_pretrained("google/fnet-base")
        >>> model = FNetForPreTraining.from_pretrained("google/fnet-base")
        >>> inputs = tokenizer("Hello, my dog is cute", return_tensors="pt")
        >>> outputs = model(**inputs)
        >>> prediction_logits = outputs.prediction_logits
        >>> seq_relationship_logits = outputs.seq_relationship_logits
        ```N©rC   r@   rg   rÌ   rÍ   rx   rA   )r  r  r  r†   )	r`   r(  rÿ   r8  r
   ÚviewrI   r  r†   )r_   rf   rC   r@   rg   rD  rE  rÌ   rÍ   r‡   rð   rÙ   rñ   rø   Ú
total_lossÚloss_fctÚmasked_lm_lossÚnext_sentence_lossr“   s                      r,   rl   zFNetForPreTraining.forwardv  s5  € ðT &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø%Ø'Ø!5Ø#ð ó 
ˆð *1°°!¨Ñ&ˆ˜Ø48·H±H¸_ÈmÓ4\Ñ1ÐÐ1àˆ
ØÐÐ"5Ð"AÜ'Ó)ˆHÙ%Ð&7×&<Ñ&<¸RÀÇÁ×AWÑAWÓ&XÐZ`×ZeÑZeÐfhÓZiÓjˆNÙ!)Ð*@×*EÑ*EÀbÈ!Ó*LÐNa×NfÑNfÐgiÓNjÓ!kÐØ'Ð*<Ñ<ˆJáØ'Ð)?Ð@À7È1È2À;ÑNˆFØ/9Ð/E�Z�M FÑ*ÐQÈ6ÐQä'ØØ/Ø$:Ø!×/Ñ/ô	
ð 	
r.   ©NNNNNNNN)rm   rn   ro   Ú_tied_weights_keysrG   r;  rA  r   r-  r.  r   r  r0  r   r%   r¥   r2  r   r   rl   rq   rr   s   @r,   r4  r4  \  s  ø„ ð 9Ð:ZÐ[Ðôò,ò8ñ +Ð+@×+GÑ+GÐHeÓ+fÓgÙÐ+CÐRaÔbð -1Ø15Ø/3Ø04Ø)-Ø6:Ø/3Ø&*ñF
à˜EŸL™LÑ)ðF
ð ! §¡Ñ.ðF
ð ˜uŸ|™|Ñ,ð	F
ð
   §¡Ñ-ðF
ð ˜Ÿ™Ñ&ðF
ð & e§l¡lÑ3ðF
ð ' t™nðF
ð ˜d‘^ðF
ð 
ˆuÐ.Ð.Ñ	/òF
ó có hôF
r.   r4  z2FNet Model with a `language modeling` head on top.c                   óR  ‡ — e Zd ZddgZˆ fd„Zd„ Zd„ Z eej                  d«      «       e
eee¬«      	 	 	 	 	 	 	 ddeej                      d	eej                      d
eej                      deej                      deej                      dee   dee   deeef   fd„«       «       Zˆ xZS )ÚFNetForMaskedLMr5  r6  c                 ó„   •— t         ‰| �  |«       t        |«      | _        t	        |«      | _        | j                  «        y r0   )rF   rG   r  rÿ   rì   r8  r  r^   s     €r,   rG   zFNetForMaskedLM.__init__Å  ó4   ø€ Ü‰Ñ˜Ô ä˜fÓ%ˆŒ	Ü" 6Ó*ˆŒð 	�‰Õr.   c                 óB   — | j                   j                  j                  S r0   r:  ré   s    r,   r;  z%FNetForMaskedLM.get_output_embeddingsÎ  r<  r.   c                 ó„   — || j                   j                  _        |j                  | j                   j                  _        y r0   r>  r?  s     r,   rA  z%FNetForMaskedLM.set_output_embeddingsÑ  rB  r.   r"  r#  rf   rC   r@   rg   rD  rÌ   rÍ   r¡   c                 ó~  — |�|n| j                   j                  }| j                  ||||||¬«      }|d   }	| j                  |	«      }
d}|�Ft	        «       } ||
j                  d| j                   j                  «      |j                  d«      «      }|s|
f|dd z   }|�|f|z   S |S t        ||
|j                  ¬«      S )a£  
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the masked language modeling loss. Indices should be in `[-100, 0, ...,
            config.vocab_size]` (see `input_ids` docstring) Tokens with indices set to `-100` are ignored (masked), the
            loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`.
        NrG  r   rA   rx   ©r  Úlogitsr†   )	r`   r(  rÿ   r8  r
   rH  rI   r   r†   )r_   rf   rC   r@   rg   rD  rÌ   rÍ   r‡   rð   rñ   rK  rJ  r“   s                 r,   rl   zFNetForMaskedLM.forwardÕ  sã   € ð, &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø%Ø'Ø!5Ø#ð ó 
ˆð " !™*ˆØ ŸH™H _Ó5ÐàˆØÐÜ'Ó)ˆHÙ%Ð&7×&<Ñ&<¸RÀÇÁ×AWÑAWÓ&XÐZ`×ZeÑZeÐfhÓZiÓjˆNáØ'Ð)¨G°A°B¨KÑ7ˆFØ3AÐ3M�^Ð%¨Ñ.ÐYÐSYÐYä >Ð:KÐ[b×[pÑ[pÔqÐqr.   ©NNNNNNN)rm   rn   ro   rN  rG   r;  rA  r   r-  r.  r   r/  r   r0  r   r%   r¥   r2  r   r   rl   rq   rr   s   @r,   rP  rP  Á  s	  ø„ à8Ð:ZÐ[Ðôò,ò8ñ +Ð+@×+GÑ+GÐHeÓ+fÓgÙØ&Ø"Ø$ôð -1Ø15Ø/3Ø04Ø)-Ø/3Ø&*ñ'rà˜EŸL™LÑ)ð'rð ! §¡Ñ.ð'rð ˜uŸ|™|Ñ,ð	'rð
   §¡Ñ-ð'rð ˜Ÿ™Ñ&ð'rð ' t™nð'rð ˜d‘^ð'rð 
ˆu�nÐ$Ñ	%ò'róó hô'rr.   rP  zJFNet Model with a `next sentence prediction (classification)` head on top.c                   ó<  ‡ — e Zd Zˆ fd„Z eej                  d«      «       eee	¬«      	 	 	 	 	 	 	 dde
ej                     de
ej                     de
ej                     de
ej                     de
ej                     d	e
e   d
e
e   deeef   fd„«       «       Zˆ xZS )ÚFNetForNextSentencePredictionc                 ó„   •— t         ‰| �  |«       t        |«      | _        t	        |«      | _        | j                  «        y r0   )rF   rG   r  rÿ   ró   r8  r  r^   s     €r,   rG   z&FNetForNextSentencePrediction.__init__
  rR  r.   r"  rC  rf   rC   r@   rg   rD  rÌ   rÍ   r¡   c                 ó´  — d|v r+t        j                  dt        «       |j                  d«      }|�|n| j                  j
                  }| j                  ||||||¬«      }	|	d   }
| j                  |
«      }d}|�2t        «       } ||j                  dd«      |j                  d«      «      }|s|f|	dd z   }|�|f|z   S |S t        |||	j                  ¬«      S )	a§  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the next sequence prediction (classification) loss. Input should be a sequence pair
            (see `input_ids` docstring). Indices should be in `[0, 1]`:

            - 0 indicates sequence B is a continuation of sequence A,
            - 1 indicates sequence B is a random sequence.

        Returns:

        Example:

        ```python
        >>> from transformers import AutoTokenizer, FNetForNextSentencePrediction
        >>> import torch

        >>> tokenizer = AutoTokenizer.from_pretrained("google/fnet-base")
        >>> model = FNetForNextSentencePrediction.from_pretrained("google/fnet-base")
        >>> prompt = "In Italy, pizza served in formal settings, such as at a restaurant, is presented unsliced."
        >>> next_sentence = "The sky is blue due to the shorter wavelength of blue light."
        >>> encoding = tokenizer(prompt, next_sentence, return_tensors="pt")
        >>> outputs = model(**encoding, labels=torch.LongTensor([1]))
        >>> logits = outputs.logits
        >>> assert logits[0, 0] < logits[0, 1]  # next sentence was random
        ```rE  zoThe `next_sentence_label` argument is deprecated and will be removed in a future version, use `labels` instead.NrG  r    rA   rx   rV  )ÚwarningsÚwarnÚFutureWarningÚpopr`   r(  rÿ   r8  r
   rH  r   r†   )r_   rf   rC   r@   rg   rD  rÌ   rÍ   Úkwargsr‡   rÙ   Úseq_relationship_scoresrL  rJ  r“   s                  r,   rl   z%FNetForNextSentencePrediction.forward  s  € ðN ! FÑ*Ü�M‰Mð%äôð
 —Z‘ZÐ 5Ó6ˆFà%0Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø%Ø'Ø!5Ø#ð ó 
ˆð   ™
ˆà"&§(¡(¨=Ó"9Ðà!ÐØÐÜ'Ó)ˆHÙ!)Ð*A×*FÑ*FÀrÈ1Ó*MÈvÏ{É{Ð[]ËÓ!_ÐáØ-Ð/°'¸!¸"°+Ñ=ˆFØ7IÐ7UÐ'Ð)¨FÑ2ÐaÐ[aÐaä*Ø#Ø*Ø!×/Ñ/ô
ð 	
r.   rX  )rm   rn   ro   rG   r   r-  r.  r   r   r0  r   r%   r¥   r2  r   r   rl   rq   rr   s   @r,   rZ  rZ    sð   ø„ ô
ñ +Ð+@×+GÑ+GÐHeÓ+fÓgÙÐ+FÐUdÔeð -1Ø15Ø/3Ø04Ø)-Ø/3Ø&*ñI
à˜EŸL™LÑ)ðI
ð ! §¡Ñ.ðI
ð ˜uŸ|™|Ñ,ð	I
ð
   §¡Ñ-ðI
ð ˜Ÿ™Ñ&ðI
ð ' t™nðI
ð ˜d‘^ðI
ð 
ˆuÐ1Ð1Ñ	2òI
ó fó hôI
r.   rZ  zœ
    FNet Model transformer with a sequence classification/regression head on top (a linear layer on top of the pooled
    output) e.g. for GLUE tasks.
    c                   ó>  ‡ — e Zd Zˆ fd„Z eej                  d«      «       eee	e
¬«      	 	 	 	 	 	 	 ddeej                     deej                     deej                     deej                     deej                     d	ee   d
ee   deee	f   fd„«       «       Zˆ xZS )ÚFNetForSequenceClassificationc                 ó,  •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        j                  |j                  «      | _        t        j                  |j                  |j                  «      | _        | j                  «        y r0   ©rF   rG   Ú
num_labelsr  rÿ   r   rU   rV   rW   rS   rJ   Ú
classifierr  r^   s     €r,   rG   z&FNetForSequenceClassification.__init__i  si   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒÜ˜fÓ%ˆŒ	ä—z‘z &×"<Ñ"<Ó=ˆŒÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒð 	�‰Õr.   r"  r#  rf   rC   r@   rg   rD  rÌ   rÍ   r¡   c                 ó$  — |�|n| j                   j                  }| j                  ||||||¬«      }|d   }	| j                  |	«      }	| j	                  |	«      }
d}|��‡| j                   j
                  €�| j                  dk(  rd| j                   _        nl| j                  dkD  rL|j                  t        j                  k(  s|j                  t        j                  k(  rd| j                   _        nd| j                   _        | j                   j
                  dk(  rIt        «       }| j                  dk(  r& ||
j                  «       |j                  «       «      }nŒ ||
|«      }n‚| j                   j
                  dk(  r=t        «       } ||
j                  d| j                  «      |j                  d«      «      }n,| j                   j
                  dk(  rt        «       } ||
|«      }|s|
f|dd z   }|�|f|z   S |S t!        ||
|j"                  ¬	«      S )
a�  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the sequence classification/regression loss. Indices should be in `[0, ...,
            config.num_labels - 1]`. If `config.num_labels == 1` a regression loss is computed (Mean-Square loss), If
            `config.num_labels > 1` a classification loss is computed (Cross-Entropy).
        NrG  r    Ú
regressionÚsingle_label_classificationÚmulti_label_classificationrA   rx   rV  )r`   r(  rÿ   rW   rh  Úproblem_typerg  rE   r%   r]   Úintr   Úsqueezer
   rH  r	   r   r†   )r_   rf   rC   r@   rg   rD  rÌ   rÍ   r‡   rÙ   rW  r  rJ  r“   s                 r,   rl   z%FNetForSequenceClassification.forwardt  sÐ  € ð, &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø%Ø'Ø!5Ø#ð ó 
ˆð   ™
ˆØŸ™ ]Ó3ˆØ—‘ Ó/ˆàˆØÑØ�{‰{×'Ñ'Ð/Ø—?‘? aÒ'Ø/;�D—K‘KÕ,Ø—_‘_ qÒ(¨f¯l©l¼e¿j¹jÒ.HÈFÏLÉLÔ\a×\eÑ\eÒLeØ/L�D—K‘KÕ,à/K�D—K‘KÔ,à�{‰{×'Ñ'¨<Ò7Ü"›9�Ø—?‘? aÒ'Ù# F§N¡NÓ$4°f·n±nÓ6FÓG‘Dá# F¨FÓ3‘DØ—‘×)Ñ)Ð-JÒJÜ+Ó-�Ù §¡¨B°·±Ó @À&Ç+Á+ÈbÃ/ÓR‘Ø—‘×)Ñ)Ð-IÒIÜ,Ó.�Ù ¨Ó/�ÙØ�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä'¨T¸&ÐPW×PeÑPeÔfÐfr.   rX  )rm   rn   ro   rG   r   r-  r.  r   r/  r   r0  r   r%   r¥   r2  r   r   rl   rq   rr   s   @r,   rd  rd  a  sô   ø„ ô	ñ +Ð+@×+GÑ+GÐHeÓ+fÓgÙØ&Ø,Ø$ôð -1Ø15Ø/3Ø04Ø)-Ø/3Ø&*ñ9gà˜EŸL™LÑ)ð9gð ! §¡Ñ.ð9gð ˜uŸ|™|Ñ,ð	9gð
   §¡Ñ-ð9gð ˜Ÿ™Ñ&ð9gð ' t™nð9gð ˜d‘^ð9gð 
ˆuÐ.Ð.Ñ	/ò9góó hô9gr.   rd  z¥
    FNet Model with a multiple choice classification head on top (a linear layer on top of the pooled output and a
    softmax) e.g. for RocStories/SWAG tasks.
    c                   ó>  ‡ — e Zd Zˆ fd„Z eej                  d«      «       eee	e
¬«      	 	 	 	 	 	 	 ddeej                     deej                     deej                     deej                     deej                     d	ee   d
ee   deee	f   fd„«       «       Zˆ xZS )ÚFNetForMultipleChoicec                 óö   •— t         ‰| �  |«       t        |«      | _        t	        j
                  |j                  «      | _        t	        j                  |j                  d«      | _
        | j                  «        y r­   )rF   rG   r  rÿ   r   rU   rV   rW   rS   rJ   rh  r  r^   s     €r,   rG   zFNetForMultipleChoice.__init__¾  sV   ø€ Ü‰Ñ˜Ô ä˜fÓ%ˆŒ	Ü—z‘z &×"<Ñ"<Ó=ˆŒÜŸ)™) F×$6Ñ$6¸Ó:ˆŒð 	�‰Õr.   z(batch_size, num_choices, sequence_lengthr#  rf   rC   r@   rg   rD  rÌ   rÍ   r¡   c                 óæ  — |�|n| j                   j                  }|�|j                  d   n|j                  d   }|�!|j                  d|j	                  d«      «      nd}|�!|j                  d|j	                  d«      «      nd}|�!|j                  d|j	                  d«      «      nd}|�1|j                  d|j	                  d«      |j	                  d«      «      nd}| j                  ||||||¬«      }	|	d   }
| j                  |
«      }
| j                  |
«      }|j                  d|«      }d}|�t        «       } |||«      }|s|f|	dd z   }|�|f|z   S |S t        |||	j                  ¬«      S )aJ  
        labels (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for computing the multiple choice classification loss. Indices should be in `[0, ...,
            num_choices-1]` where `num_choices` is the size of the second dimension of the input tensors. (See
            `input_ids` above)
        Nr    rA   éþÿÿÿrG  rx   rV  )r`   r(  r#   rH  r\   rÿ   rW   rh  r
   r   r†   )r_   rf   rC   r@   rg   rD  rÌ   rÍ   Únum_choicesr‡   rÙ   rW  Úreshaped_logitsr  rJ  r“   s                   r,   rl   zFNetForMultipleChoice.forwardÈ  s   € ð, &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆØ,5Ð,A�i—o‘o aÒ(À}×GZÑGZÐ[\ÑG]ˆà>GÐ>S�I—N‘N 2 y§~¡~°bÓ'9Ô:ÐY]ˆ	ØM[ÐMg˜×,Ñ,¨R°×1DÑ1DÀRÓ1HÔIÐmqˆØGSÐG_�|×(Ñ(¨¨\×->Ñ->¸rÓ-BÔCÐeiˆð Ð(ð ×Ñ˜r =×#5Ñ#5°bÓ#9¸=×;MÑ;MÈbÓ;QÔRàð 	ð —)‘)ØØ)Ø%Ø'Ø!5Ø#ð ó 
ˆð   ™
ˆàŸ™ ]Ó3ˆØ—‘ Ó/ˆØ Ÿ+™+ b¨+Ó6ˆàˆØÐÜ'Ó)ˆHÙ˜O¨VÓ4ˆDáØ%Ð'¨'°!°"¨+Ñ5ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä(¨d¸?ÐZa×ZoÑZoÔpÐpr.   rX  )rm   rn   ro   rG   r   r-  r.  r   r/  r   r0  r   r%   r¥   r2  r   r   rl   rq   rr   s   @r,   rq  rq  ¶  sô   ø„ ôñ +Ð+@×+GÑ+GÐHrÓ+sÓtÙØ&Ø-Ø$ôð -1Ø15Ø/3Ø04Ø)-Ø/3Ø&*ñ4qà˜EŸL™LÑ)ð4qð ! §¡Ñ.ð4qð ˜uŸ|™|Ñ,ð	4qð
   §¡Ñ-ð4qð ˜Ÿ™Ñ&ð4qð ' t™nð4qð ˜d‘^ð4qð 
ˆuÐ/Ð/Ñ	0ò4qóó uô4qr.   rq  z£
    FNet Model with a token classification head on top (a linear layer on top of the hidden-states output) e.g. for
    Named-Entity-Recognition (NER) tasks.
    c                   ó>  ‡ — e Zd Zˆ fd„Z eej                  d«      «       eee	e
¬«      	 	 	 	 	 	 	 ddeej                     deej                     deej                     deej                     deej                     d	ee   d
ee   deee	f   fd„«       «       Zˆ xZS )ÚFNetForTokenClassificationc                 ó,  •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        j                  |j                  «      | _        t        j                  |j                  |j                  «      | _        | j                  «        y r0   rf  r^   s     €r,   rG   z#FNetForTokenClassification.__init__  si   ø€ Ü‰Ñ˜Ô Ø ×+Ñ+ˆŒä˜fÓ%ˆŒ	ä—z‘z &×"<Ñ"<Ó=ˆŒÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒð 	�‰Õr.   r"  r#  rf   rC   r@   rg   rD  rÌ   rÍ   r¡   c                 óŒ  — |�|n| j                   j                  }| j                  ||||||¬«      }|d   }	| j                  |	«      }	| j	                  |	«      }
d}|�<t        «       } ||
j                  d| j                  «      |j                  d«      «      }|s|
f|dd z   }|�|f|z   S |S t        ||
|j                  ¬«      S )zÛ
        labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
            Labels for computing the token classification loss. Indices should be in `[0, ..., config.num_labels - 1]`.
        NrG  r   rA   rx   rV  )
r`   r(  rÿ   rW   rh  r
   rH  rg  r   r†   )r_   rf   rC   r@   rg   rD  rÌ   rÍ   r‡   rð   rW  r  rJ  r“   s                 r,   rl   z"FNetForTokenClassification.forward  sÝ   € ð( &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø%Ø'Ø!5Ø#ð ó 
ˆð " !™*ˆàŸ,™, Ó7ˆØ—‘ Ó1ˆàˆØÐÜ'Ó)ˆHá˜FŸK™K¨¨D¯O©OÓ<¸f¿k¹kÈ"»oÓNˆDáØ�Y ¨¨ Ñ,ˆFØ)-Ð)9�T�G˜fÑ$ÐE¸vÐEä$¨$°vÈW×MbÑMbÔcÐcr.   rX  )rm   rn   ro   rG   r   r-  r.  r   r/  r   r0  r   r%   r¥   r2  r   r   rl   rq   rr   s   @r,   rx  rx    sô   ø„ ô
ñ +Ð+@×+GÑ+GÐHeÓ+fÓgÙØ&Ø)Ø$ôð -1Ø15Ø/3Ø04Ø)-Ø/3Ø&*ñ(dà˜EŸL™LÑ)ð(dð ! §¡Ñ.ð(dð ˜uŸ|™|Ñ,ð	(dð
   §¡Ñ-ð(dð ˜Ÿ™Ñ&ð(dð ' t™nð(dð ˜d‘^ð(dð 
ˆuÐ+Ð+Ñ	,ò(dóó hô(dr.   rx  zÝ
    FNet Model with a span classification head on top for extractive question-answering tasks like SQuAD (a linear
    layers on top of the hidden-states output to compute `span start logits` and `span end logits`).
    c                   ó^  ‡ — e Zd Zˆ fd„Z eej                  d«      «       eee	e
¬«      	 	 	 	 	 	 	 	 ddeej                     deej                     deej                     deej                     deej                     d	eej                     d
ee   dee   deee	f   fd„«       «       Zˆ xZS )ÚFNetForQuestionAnsweringc                 óä   •— t         ‰| �  |«       |j                  | _        t        |«      | _        t        j                  |j                  |j                  «      | _        | j                  «        y r0   )
rF   rG   rg  r  rÿ   r   rS   rJ   Ú
qa_outputsr  r^   s     €r,   rG   z!FNetForQuestionAnswering.__init__R  sS   ø€ Ü‰Ñ˜Ô à ×+Ñ+ˆŒä˜fÓ%ˆŒ	ÜŸ)™) F×$6Ñ$6¸×8IÑ8IÓJˆŒð 	�‰Õr.   r"  r#  rf   rC   r@   rg   Ústart_positionsÚend_positionsrÌ   rÍ   r¡   c	                 ó  — |�|n| j                   j                  }| j                  ||||||¬«      }	|	d   }
| j                  |
«      }|j	                  dd¬«      \  }}|j                  d«      j                  «       }|j                  d«      j                  «       }d}|�·|�µt        |j                  «       «      dkD  r|j                  d«      }t        |j                  «       «      dkD  r|j                  d«      }|j                  d«      }|j                  d|«      }|j                  d|«      }t        |¬«      } |||«      } |||«      }||z   dz  }|s||f|	dd z   }|�|f|z   S |S t        ||||	j                  ¬	«      S )
a  
        start_positions (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for position (index) of the start of the labelled span for computing the token classification loss.
            Positions are clamped to the length of the sequence (`sequence_length`). Position outside of the sequence
            are not taken into account for computing the loss.
        end_positions (`torch.LongTensor` of shape `(batch_size,)`, *optional*):
            Labels for position (index) of the end of the labelled span for computing the token classification loss.
            Positions are clamped to the length of the sequence (`sequence_length`). Position outside of the sequence
            are not taken into account for computing the loss.
        NrG  r   r    rA   ry   )Úignore_indexrx   )r  Ústart_logitsÚ
end_logitsr†   )r`   r(  rÿ   r~  Úsplitro  Ú
contiguousÚlenr\   Úclampr
   r   r†   )r_   rf   rC   r@   rg   r  r€  rÌ   rÍ   r‡   rð   rW  rƒ  r„  rI  Úignored_indexrJ  Ú
start_lossÚend_lossr“   s                       r,   rl   z FNetForQuestionAnswering.forward]  s®  € ð6 &1Ð%<‘kÀ$Ç+Á+×B]ÑB]ˆà—)‘)ØØ)Ø%Ø'Ø!5Ø#ð ó 
ˆð " !™*ˆà—‘ Ó1ˆØ#)§<¡<°°r <Ó#:Ñ ˆ�jØ#×+Ñ+¨BÓ/×:Ñ:Ó<ˆØ×'Ñ'¨Ó+×6Ñ6Ó8ˆ
àˆ
ØÐ&¨=Ð+Dä�?×'Ñ'Ó)Ó*¨QÒ.Ø"1×"9Ñ"9¸"Ó"=�Ü�=×%Ñ%Ó'Ó(¨1Ò,Ø -× 5Ñ 5°bÓ 9�à(×-Ñ-¨aÓ0ˆMØ-×3Ñ3°A°}ÓEˆOØ)×/Ñ/°°=ÓAˆMä'°]ÔCˆHÙ! ,°Ó@ˆJÙ 
¨MÓ:ˆHØ$ xÑ/°1Ñ4ˆJáØ" JÐ/°'¸!¸"°+Ñ=ˆFØ/9Ð/E�Z�M FÑ*ÐQÈ6ÐQä+Ø¨,À:Ð]d×]rÑ]rô
ð 	
r.   rM  )rm   rn   ro   rG   r   r-  r.  r   r/  r   r0  r   r%   r¥   r2  r   r   rl   rq   rr   s   @r,   r|  r|  J  s   ø„ ô	ñ +Ð+@×+GÑ+GÐHeÓ+fÓgÙØ&Ø0Ø$ôð -1Ø15Ø/3Ø04Ø26Ø04Ø/3Ø&*ñ>
à˜EŸL™LÑ)ð>
ð ! §¡Ñ.ð>
ð ˜uŸ|™|Ñ,ð	>
ð
   §¡Ñ-ð>
ð " %§,¡,Ñ/ð>
ð   §¡Ñ-ð>
ð ' t™nð>
ð ˜d‘^ð>
ð 
ˆuÐ2Ð2Ñ	3ò>
óó hô>
r.   r|  )
rP  rq  rZ  r4  r|  rd  rx  r«   r  rþ   )Prp   r]  Údataclassesr   Ú	functoolsr   Útypingr   r   r   r%   Útorch.utils.checkpointr   Útorch.nnr	   r
   r   Úutilsr   Úscipyr   Úactivationsr   Úmodeling_outputsr   r   r   r   r   r   r   r   r   Úmodeling_utilsr   Úpytorch_utilsr   r   r   r   r   r   Úconfiguration_fnetr!   Ú
get_loggerrm   Úloggerr/  r0  r-   r1   r9   ÚModuler;   rt   r‰   r‘   r™   r§   r«   r¹   rÓ   rÛ   rà   rì   ró   rú   rþ   r  ÚFNET_START_DOCSTRINGr-  r  r4  rP  rZ  rd  rq  rx  r|  Ú__all__rÁ   r.   r,   ú<module>r�     s%  ðñ ã Ý !Ý ß )Ñ )ã Û Ý ß AÑ Aå 'ñ ÔÝå !÷
÷ 
õ 
õ .Ý 6÷õ õ +ð 
ˆ×	Ñ	˜HÓ	%€à(Ð Ø€òMò>ò
ô :�R—Y‘Yô :ôz# §	¡	ô #ôL�b—i‘iô ô
˜2Ÿ9™9ô 
ô�r—y‘yô ô �—‘ô ô�—	‘	ô ô6a�"—)‘)ô aô>�—‘ô ô  "§)¡)ô ô"*˜2Ÿ9™9ô *ô4!�b—i‘iô !ô&�b—i‘iô &ô	9˜2Ÿ9™9ô 	9ô*˜/ô *ð8 ô=˜{ó =ó ð=ð2	Ð ð Ð ñF ØdØóôc
Ð#ó c
ó	ðc
ñL ðð óô[
Ð,ó [
óð[
ñ| ÐNÐPdÓeô@rÐ)ó @ró fð@rñF ØTØóôU
Ð$7ó U
ó	ðU
ñp ðð óôKgÐ$7ó KgóðKgñ\ ðð óôEqÐ/ó EqóðEqñP ðð óô;dÐ!4ó ;dóð;dñ| ðð óôP
Ð2ó P
óðP
òf�r.   