§
    ‚ŠtjQÀ  ã                   ó¤  — d dl Z d dlZddl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  G d„ d¦  «        Z G d	„ d
ej        j        ¦  «        Zdeeee         z  eee         z  f         fd„Z G d„ dej        j        ¦  «        Z G d„ dej        j        ¦  «        Z	 	 	 	 d$dedej        dz  dej        dz  dedz  dedz  f
d„Z G d„ dej        j        ¦  «        Z G d„ dej        j        ¦  «        Z  G d„ dej        j        ¦  «        Z!	 	 d%dedej        dz  dej        dz  fd„Z"d„ Z#d efd!„Z$d"ej%        j&        j'        fd#„Z(dS )&é    Né   )ÚDynamicCacheÚDynamicLayerÚDynamicSlidingWindowLayerÚEncoderDecoderCacheÚStaticCacheÚStaticLayerÚStaticSlidingWindowLayer)ÚGenerationConfig)ÚPreTrainedModel)Úis_torch_greater_or_equalÚ"is_torch_greater_or_equal_than_2_6c                   óL   — e Zd ZdZddedefd„Zd„ Zd„ Zd	„ Zd
„ Z	d„ Z
	 dd„ZdS )ÚTorchExportableModuleForVLMa|  
    A wrapper class for exporting Vision-Language Models (VLMs) like SmolVLM2 for ExecuTorch.

    This class handles the export of three main components:
        1. Vision encoder (processes images to visual features)
        2. Connector/projector (maps visual features to text embedding space)
        3. Text decoder (generates text from combined visual and text tokens)
    é   é   Úmax_batch_sizeÚmax_cache_lenc                 óØ   — || _         || _        || _        |j        | _        |j         j        | _        |j         j        | _        |j         j        | _        d| _	        d| _
        d| _        dS )a  
        Initialize the exportable VLM module.

        Args:
            model: The VLM (e.g. SmolVLM) model instance
            max_batch_size: Maximum batch size. Always 1 for ExecuTorch
            max_cache_len: Maximum cache length for text generation
        N)Úmodelr   r   ÚconfigÚvision_modelÚvision_encoderÚ	connectorÚ
text_modelÚtext_decoderÚexported_vision_encoderÚexported_connectorÚexported_text_decoder)Úselfr   r   r   s       úb/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/integrations/executorch.pyÚ__init__z$TorchExportableModuleForVLM.__init__,   sl   € ð ˆŒ
Ø,ˆÔØ*ˆÔØ”lˆŒð $œkÔ6ˆÔØœÔ.ˆŒØ!œKÔ2ˆÔð (,ˆÔ$Ø"&ˆÔØ%)ˆÔ"Ð"Ð"ó    c                 óB  — | j                              ¦   «          t          j        ddddt          j        ¬¦  «        }dt          j        j        j        t          j        j        j        dœi}t          j                             | j         |f|d¬¦  «        | _        | j        S )	z$Export the vision encoder component.r   é   i€  ©ÚdtypeÚpixel_values)r   r%   F©ÚargsÚdynamic_shapesÚstrict)	r   ÚevalÚtorchÚrandnÚfloat32ÚexportÚDimÚAUTOr   )r    r(   r+   s      r!   Úexport_vision_encoderz1TorchExportableModuleForVLM.export_vision_encoderD   sž   € àÔ× Ò Ñ"Ô"Ð"õ ”{ 1 a¨¨c½¼ÐGÑGÔGˆð Ý”<Ô#Ô(Ý”<Ô#Ô(ðð ð
ˆõ (-¤|×':Ò':ØÔØ�Ø)Øð	 (;ñ (
ô (
ˆÔ$ð Ô+Ð+r#   c                 ó�  — | j                              ¦   «          | j        j        j        }| j        j        j        }| j        j        j        }||z  }||z  }t          j        d||t          j	        ¬¦  «        }ddt          j
        j        j        ii}t          j
         
                    | j         |f|d¬¦  «        | _        | j        S )zExport the connector component.r   r&   Úimage_hidden_statesFr)   )r   r-   r   Úvision_configÚhidden_sizeÚ
image_sizeÚ
patch_sizer.   r/   r0   r1   r2   r3   r   )r    Úvision_hidden_sizer9   r:   Úpatches_per_dimÚnum_patchesr6   r+   s           r!   Úexport_connectorz,TorchExportableModuleForVLM.export_connector\   sÈ   € àŒ×ÒÑÔÐð "œ[Ô6ÔBÐØ”[Ô.Ô9ˆ
Ø”[Ô.Ô9ˆ
Ø$¨
Ñ2ˆØ%¨Ñ7ˆÝ#œk¨!¨[Ð:LÕTYÔTaÐbÑbÔbÐð 0°!µU´\Ô5EÔ5JÐ1KÐLˆõ #(¤,×"5Ò"5ØŒNØ%Ð'Ø)Øð	 #6ñ #
ô #
ˆÔð Ô&Ð&r#   c                 ó´  — t          | j        ¬¦  «        | _        d}t          j        d|ft          j        ¬¦  «        }t          j        |t          j        ¬¦  «        }t          | j        | j	        j
        j        ¦  «        }t          j                             d|dz
  ¬¦  «        }d|id|idœ}| j                             |||d	¬
¦  «        | _        | j        S )z"Export the text decoder component.)r   r%   r   r&   Úseq_length_dim©Úmaxr   ©Ú	input_idsÚcache_positionF)rD   rE   r+   r,   )Ú%TorchExportableModuleForDecoderOnlyLMr   Úexportable_text_decoderr.   ÚzerosÚlongÚarangeÚminr   r   Útext_configÚmax_position_embeddingsr1   r2   r   )r    Ú
seq_lengthrD   rE   Úmax_seq_lengthÚseq_len_dimr+   s          r!   Úexport_text_decoderz/TorchExportableModuleForVLM.export_text_decoderu   så   € õ (MÐSWÔSdÐ'eÑ'eÔ'eˆÔ$ð ˆ
Ý”K  J µu´zÐBÑBÔBˆ	Ýœ j½¼
ÐCÑCÔCˆÝ˜TÔ/°´Ô1HÔ1`ÑaÔaˆÝ”l×&Ò&Ð'7¸^ÈaÑ=OÐ&ÑPÔPˆð ˜[Ð)Ø  +Ð.ð
ð 
ˆð
 &*Ô%A×%HÒ%HØØ)Ø)Øð	 &Iñ &
ô &
ˆÔ"ð Ô)Ð)r#   c                 óz   —  | j         di |¤Ž  | j        di |¤Ž  | j        di |¤Ž | j        | j        | j        dœS )z'Export all components of the VLM model.)r   r   r   © )r4   r>   rQ   r   r   r   )r    Úkwargss     r!   r1   z"TorchExportableModuleForVLM.export�   sn   € à"ˆÔ"Ð,Ð, VÐ,Ð,Ð,ØˆÔÐ'Ð' Ð'Ð'Ð'Ø ˆÔ Ð*Ð* 6Ð*Ð*Ð*à"Ô:ØÔ0Ø Ô6ð
ð 
ð 	
r#   c                 ó   — dS )a°  
        Simplified forward pass for inference with guaranteed non-null input_ids and cache_position.

        Args:
            pixel_values: Input images [1, channels, height, width] (optional)
            input_ids: Text token IDs [1, seq_len] (required - won't be None)
            cache_position: Cache positions [seq_len] (required - won't be None)

        Returns:
            Output with logits for text generation
        NrS   )r    r(   rD   rE   s       r!   Úforwardz#TorchExportableModuleForVLM.forward›   ó   € € € r#   Né2   Fç      ð?c                 ó   — dS )aè  
        Simplified generate method with guaranteed non-null input_ids.

        Args:
            pixel_values: Input images [1, channels, height, width] (optional)
            input_ids: Initial text tokens [1, seq_len] (required - won't be None)
            max_new_tokens: Maximum number of tokens to generate
            do_sample: Whether to use sampling or greedy decoding
            temperature: Temperature for sampling

        Returns:
            Generated sequences
        NrS   )r    r(   rD   Úmax_new_tokensÚ	do_sampleÚtemperaturerT   s          r!   Úgeneratez$TorchExportableModuleForVLM.generate¨   rW   r#   )r   r   )NNrX   FrY   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__Úintr"   r4   r>   rQ   r1   rV   r^   rS   r#   r!   r   r   "   s«   € € € € € ðð ð*ð *¨cð *Àcð *ð *ð *ð *ð0,ð ,ð ,ð0'ð 'ð 'ð2*ð *ð *ð6	
ð 	
ð 	
ðð ð ð beðð ð ð ð ð r#   r   c                   ó�  ‡ — e Zd ZdZ	 	 	 ddededz  dedz  dej        dz  ddf
ˆ fd„Z	 	 	 dd	ej	        dz  d
ej	        dz  dej	        dz  dej	        fd„Z
	 	 	 	 	 dd	ej	        dz  d
ej	        dz  dej	        dz  dedz  dedz  dej        j        fd„Ze	 	 	 	 	 	 ddej        j        dedededededededefd„¦   «         Zˆ xZS ) rF   a  
    A recipe module designed to make a `PreTrainedModel` exportable with `torch.export`,
    specifically for decoder-only LM with cache. This module ensures that the
    exported model is compatible with further lowering and execution in `ExecuTorch`.
    Nr   Ú
batch_sizer   ÚdeviceÚreturnc                 ó”  •— t          ¦   «                              ¦   «          |j                             ¦   «         }t	          |d¦  «        r	|j        du rt          d¦  «        ‚t	          |d¦  «        r*t          |dd¦  «        �t          ||||¦  «        | _	        dS t          j        d¦  «         t          ||||¦  «        | _	        dS )zõ
        Initializes the exportable module.

        Args:
            model (`PreTrainedModel`): The pretrained model to wrap.

        Raises:
            ValueError: If the model is configured with a unsupported cache implementation.
        Ú	use_cacheFz5The model must have caching enabled to be performant.Úlayer_typesÚsliding_windowNzmUsing `StaticCache` for export as `layer_types` is not specified or `sliding_window` is `null` in the config.)Úsuperr"   r   Úget_text_configÚhasattrri   Ú
ValueErrorÚgetattrÚ$TorchExportableModuleWithHybridCacher   ÚloggingÚinfoÚ$TorchExportableModuleWithStaticCache)r    r   re   r   rf   r   Ú	__class__s         €r!   r"   z.TorchExportableModuleForDecoderOnlyLM.__init__Á   sÔ   ø€ õ  	‰Œ×ÒÑÔÐà”×-Ò-Ñ/Ô/ˆå�v˜{Ñ+Ô+ð 	V¨vÔ/?À5Ð/HÐ/HÝÐTÑUÔUÐUå�6˜=Ñ)Ô)ð 	h­g°fÐ>NÐPTÑ.UÔ.UÐ.aÝ=¸eÀZÐQ^Ð`fÑgÔgˆDŒJˆJˆJõ ŒLØñô ð õ >¸eÀZÐQ^Ð`fÑgÔgˆDŒJˆJˆJr#   rD   Úinputs_embedsrE   c                 ó:   — | j                              ||¬¦  «        S )aê  
        Forward pass of the module, which is compatible with the ExecuTorch llm runner.

        Args:
            input_ids (`torch.Tensor`): Tensor representing current input token id to the module.
            inputs_embeds (`torch.Tensor`): Tensor representing current input embeddings to the module.
            cache_position (`torch.Tensor`): Tensor representing current input position in the cache.

        Returns:
            torch.Tensor: Logits output from the model.
        )rD   rv   )r   rV   )r    rD   rv   rE   s       r!   rV   z-TorchExportableModuleForDecoderOnlyLM.forwardâ   s   € ð" Œz×!Ò!¨IÀ]Ð!ÑSÔSÐSr#   r+   r,   c                 ó`  — |du |du z  st          d¦  «        ‚t          | j        d¦  «        r-t          | j        | j        j        | j        ¦  «        }|j        }n=t          | j        d¦  «        r| j        j        j        }nd}t          j        d¦  «         |�4||�|n+t          j	        |j
        d         t          j        |¬¦  «        d	œ}n3||�|n+t          j	        |j
        d
         t          j        |¬¦  «        dœ}t          j                             | j        d|||�|nd¬¦  «        }	|	S )ar  
        Export the wrapped module using `torch.export`.

        Args:
            input_ids (`Optional[torch.Tensor]`):
                Tensor representing current input token id to the module. Must specify either this or inputs_embeds.
            inputs_embeds (`Optional[torch.Tensor]`):
                Tensor representing current input embeddings to the module. Must specify either this or input_ids.
            cache_position (`Optional[torch.Tensor]`):
                Tensor representing current input position in the cache. If not provided, a default tensor will be used.
            dynamic_shapes (`Optional[dict]`):
                Dynamic shapes to use for export if specified.
            strict(`Optional[bool]`):
                Flag to instruct `torch.export` to use `dynamo`.

        Returns:
            torch.export.ExportedProgram: The exported program that can be used for inference.

        Examples:
            Export with input_ids:
            ```python
            # Prepare inputs
            input_ids = torch.tensor([[1, 2, 3]], dtype=torch.long, device=model.device)
            cache_position = torch.arange(input_ids.shape[-1], dtype=torch.long, device=model.device)

            # Export
            exported = exportable_module.export(
                input_ids=input_ids,
                cache_position=cache_position
            )
            ```

            Export with inputs_embeds:
            ```python
            # Prepare embeddings
            inputs_embeds = torch.randn(1, 3, 768, device=model.device)  # batch_size=1, seq_len=3, hidden_size=768
            cache_position = torch.arange(inputs_embeds.shape[1], dtype=torch.long, device=model.device)

            # Export
            exported = exportable_module.export(
                inputs_embeds=inputs_embeds,
                cache_position=cache_position
            )
            ```
        Nz2Need to specify either input_ids or inputs_embeds.Úbase_model_prefixr   ÚcpuzfTorchExportableModuleForDecoderOnlyLM.export Can't infer device from the model. Set to CPU by default.éÿÿÿÿ©r'   rf   rC   r   )rv   rE   rS   T©r*   rT   r+   r,   )ro   rn   r   rp   ry   rf   rr   Úwarningr.   rJ   ÚshaperI   r1   )
r    rD   rv   rE   r+   r,   ÚbaseÚmodel_deviceÚinput_kwargsÚexported_programs
             r!   r1   z,TorchExportableModuleForDecoderOnlyLM.exportõ   se  € ðj ˜TÐ! m°tÐ&;Ñ<ð 	SÝÐQÑRÔRÐRå�4”:Ð2Ñ3Ô3ð 		Ý˜4œ: t¤zÔ'CÀTÄZÑPÔPˆDØœ;ˆLˆLÝ�T”Z Ñ)Ô)ð 	Øœ:Ô+Ô2ˆLˆLà ˆLÝŒOØxñô ð ð Ð à&à!Ð-ð #1 .å”\ )¤/°"Ô"5½U¼ZÐP\Ð]Ñ]Ô]ð	ð ˆLˆLð "/à!Ð-ð #1 .å”\ -Ô"5°aÔ"8ÅÄ
ÐS_Ð`Ñ`Ô`ð	ð ˆLõ !œ<×.Ò.ØŒJØØØ)Ø#Ð/�6�6°Tð /ñ 
ô 
Ðð  Ðr#   é   FrY   rX   rz   rƒ   Úpromptr[   r\   r]   Útop_kÚtop_pc	                 óz  — |                       ¦   «         }	 ||d¬¦  «        j                             |¦  «        }
|
                     ¦   «         }d}t	          |
j        d         ¦  «        D ]G}|
dd…||dz   …f         }t          j        |gt          j        |¬¦  «        } |	||¬¦  «        }|dz  }ŒHt	          |¦  «        D �]Ý}|dd…dd…f         }t          j        |gt          j        |¬¦  «        } |	||¬¦  «        }|�r|dk    r||z  }n|}|dk    r7|t          j	        ||¦  «        d         d	         k     }t          d
¦  «        ||<   |dk     rœt          j        |d¬¦  «        \  }}t          j        t          j        |d¬¦  «        d¬¦  «        }||k    }|ddd…f                              ¦   «         |ddd…f<   d|d<   |                     d||¦  «        }t          d
¦  «        ||<   t          j        |d¬¦  «        }t          j        |d¬¦  «        }n|                     dd¬¦  «        }|                     ¦   «         dk    r|                     d¦  «        }t          j        ||gd¬¦  «        }|dz  }|                     ¦   «         |j        k    r n�Œß|                     |d         d¬¦  «        S )a   
        Generate a sequence of tokens using an exported program.

        Args:
            exported_program (`torch.export.ExportedProgram`): The exported model being used for generate.
            tokenizer: The tokenizer to use.
            prompt (str): The input prompt.
            max_new_tokens (int): Maximum number of new tokens to generate.
            do_sample (bool): Whether to use sampling or greedy decoding.
            temperature (float): The temperature for sampling.
            top_k (int): The number of highest probability tokens to keep for top-k sampling.
            top_p (float): The cumulative probability for nucleus sampling.
            device (str): The device to use.

        Returns:
            str: The generated text.
        Úpt)Úreturn_tensorsr   r   Nr|   rC   r{   ).r{   Nz-infrY   T)Ú
descending©Údim.).r   )Únum_samples)r�   Úkeepdimr   )Úskip_special_tokens)ÚmodulerD   ÚtoÚcloneÚranger   r.   ÚtensorrI   ÚtopkÚfloatÚsortÚcumsumÚsoftmaxÚscatterÚmultinomialÚargmaxr�   ÚsqueezeÚcatÚitemÚeos_token_idÚdecode)rƒ   Ú	tokenizerr…   r[   r\   r]   r†   r‡   rf   Úexported_modulerD   Úgenerated_idsÚcurr_positionÚiÚcurr_input_idsÚcurr_cache_positionÚ_ÚoutputsÚlogitsÚindices_to_removeÚsorted_logitsÚsorted_indicesÚcumulative_probsÚsorted_indices_to_removeÚprobsÚnext_token_ids                             r!   r^   z.TorchExportableModuleForDecoderOnlyLM.generateQ  s,  € ð< +×1Ò1Ñ3Ô3ˆð �I˜f°TÐ:Ñ:Ô:ÔD×GÒGÈÑOÔOˆ	ð "ŸšÑ)Ô)ˆð ˆÝ�y” qÔ)Ñ*Ô*ð 	ð 	ˆAà& q q q¨!¨a°!©e¨) |Ô4ˆNÝ"'¤,°¨ÅeÄjÐY_Ð"`Ñ"`Ô"`Ðð  �¨.ÐI\Ð]Ñ]Ô]ˆAØ˜QÑˆMˆMõ �~Ñ&Ô&ð 5	ñ 5	ˆAà*¨1¨1¨1¨b¨c¨c¨6Ô2ˆNÝ"'¤,°¨ÅeÄjÐY_Ð"`Ñ"`Ô"`Ðð &�o°ÐObÐcÑcÔcˆGð ñ  Eà ’?�?Ø$ {Ñ2�F�Fà$�Fð ˜1’9�9Ø(.µ´¸FÀEÑ1JÔ1JÈ1Ô1MÈmÔ1\Ò(\Ð%Ý05°f±´�FÐ,Ñ-ð ˜3’;�;Ý49´J¸vÐRVÐ4WÑ4WÔ4WÑ1�M >Ý',¤|µE´MÀ-ÐUWÐ4XÑ4XÔ4XÐ^`Ð'aÑ'aÔ'aÐ$ð 0@À%Ò/GÐ,à8PÐQTÐVYÐWYÐVYÐQYÔ8Z×8`Ò8`Ñ8bÔ8bÐ,¨S°!°"°"¨WÑ5Ø78Ð,¨VÑ4ð )A×(HÒ(HÈÈ^Ð]uÑ(vÔ(vÐ%Ý05°f±´�FÐ,Ñ-õ œ f°"Ð5Ñ5Ô5�Ý %Ô 1°%ÀQÐ GÑ GÔ G��ð !(§¢°2¸t Ñ DÔ D�ð × Ò Ñ"Ô" QÒ&Ð&Ø -× 5Ò 5°bÑ 9Ô 9�õ "œI }°mÐ&DÈ"ÐMÑMÔMˆMØ˜QÑˆMð ×!Ò!Ñ#Ô# yÔ'=Ò=Ð=Ø�ñ >ð ×Ò ¨aÔ 0ÀdÐÑKÔKÐKr#   ©NNN)NNNNN)r„   FrY   rX   rY   rz   )r_   r`   ra   rb   r   rc   r.   rf   r"   ÚTensorrV   ÚdictÚboolr1   ÚExportedProgramÚstaticmethodÚstrr—   r^   Ú__classcell__©ru   s   @r!   rF   rF   º   sY  ø€ € € € € ðð ð "&Ø$(Ø&*ðhð hàðhð ˜$‘Jðhð ˜T‘zð	hð
 ”˜tÑ#ðhð 
ðhð hð hð hð hð hðF *.Ø-1Ø.2ð	Tð Tà”< $Ñ&ðTð ”| dÑ*ðTð œ tÑ+ð	Tð
 
ŒðTð Tð Tð Tð* *.Ø-1Ø.2Ø&*Ø"ðZ ð Z à”< $Ñ&ðZ ð ”| dÑ*ðZ ð œ tÑ+ð	Z ð
 ˜t™ðZ ð �t‘ðZ ð 
ŒÔ	%ðZ ð Z ð Z ð Z ðx ð
 !ØØ ØØØðiLð iLØœ,Ô6ðiLð ðiLð ð	iLð
 ðiLð ðiLð ðiLð ðiLð ðiLð 
ðiLð iLð iLñ „\ðiLð iLð iLð iLð iLr#   rF   rg   c                 ó  ‡ — t          ‰ d¦  «        rCˆ fd„‰ j        d‰ j         …         D ¦   «         }ˆ fd„‰ j        d‰ j         …         D ¦   «         }n4t          ‰ d‰ j        ‰ j        z  ¦  «        }t          ‰ d‰ j        ¦  «        }||fS )zuReturns a tuple `(num_heads, head_dim)` containing either 2 ints, or a list of int with the value for each
    layer.Úglobal_head_dimc                 ó8   •— g | ]}|d k    r‰j         n‰j        ‘ŒS ©Úfull_attention)r¾   Úhead_dim©Ú.0Úlayerr   s     €r!   ú
<listcomp>z#get_head_shapes.<locals>.<listcomp>Ã  s>   ø€ ð 
ð 
ð 
àð ',Ð/?Ò&?Ð&?ˆFÔ"Ð"ÀVÄ_ð
ð 
ð 
r#   Nc                 óF   •— g | ]}|d k    r‰j         r‰j        n‰j        ‘ŒS rÀ   )Úattention_k_eq_vÚnum_global_key_value_headsÚnum_key_value_headsrÃ   s     €r!   rÆ   z#get_head_shapes.<locals>.<listcomp>Ç  sM   ø€ ð 
ð 
ð 
ð ð Ð(Ò(Ð(¨VÔ-DÐ(ð Ô-Ð-àÔ+ð
ð 
ð 
r#   rÂ   rÊ   )rn   rj   Únum_kv_shared_layersrp   r8   Únum_attention_heads)r   rÂ   Ú	num_headss   `  r!   Úget_head_shapesrÎ   ¾  sÓ   ø€ õ ˆvÐ(Ñ)Ô)ð Wð
ð 
ð 
ð 
àÔ+Ð,J¨vÔ/JÐ.JÐ,JÔKð
ñ 
ô 
ˆð
ð 
ð 
ð 
ð  Ô+Ð,J¨vÔ/JÐ.JÐ,JÔKð	
ñ 
ô 
ˆ	ˆ	õ ˜6 :¨vÔ/AÀVÔE_Ñ/_Ñ`Ô`ˆÝ˜FÐ$9¸6Ô;UÑVÔVˆ	à�hÐÐr#   c                   óø   ‡ — e Zd ZdZ	 	 	 ddededz  dedz  dej        dz  ddf
ˆ fd„Z	 	 	 dd	ej	        dz  d
ej
        dz  dej
        dz  fd„Zedej        j        dej
        dedej
        fd„¦   «         Zˆ xZS )rt   aÒ  
    A recipe module designed to make a `PreTrainedModel` exportable with `torch.export`,
    specifically for decoder-only LM to `StaticCache`. This module ensures that the
    exported model is compatible with further lowering and execution in `ExecuTorch`.

    Note:
        This class is specifically designed to support export process using `torch.export`
        in a way that ensures the model can be further lowered and run efficiently in `ExecuTorch`.
    Nr   re   r   rf   rg   c                 óX  •— t          ¦   «                              ¦   «          |j                             ¦   «         }|j        }|€t          d¦  «        ‚|j        st          d¦  «        ‚|j        dk    rt          d¦  «        ‚|j        €i n|j        }|€'| 	                    dd¦  «        }|€t          d¦  «        ‚|€'| 	                    dd¦  «        }|€t          d	¦  «        ‚|€| 	                    d
|j        ¦  «        }|| _        t          ||¬¦  «        | _        t          | j        j        ¦  «        D ]6\  }}	t#          |	t$          ¦  «        rt'          |¦  «        | j        j        |<   Œ7t)          |¦  «        \  }
}| j        j        }| j                             ||
|||¦  «         t          | j        j        ¦  «        D ]e\  }}	|                      d|› �|	j        d¬¦  «         |                      d|› �|	j        d¬¦  «         |                      d|› �|	j        d¬¦  «         ŒfdS )a…  
        Initializes the wrapper module with the pretrained model.

        Args:
            model (`PreTrainedModel`): The pretrained model to wrap. The model must have caching
                enabled and use a 'static' caching implementation.
            batch_size (`Optional[int]`): The batch size of the model. If not provided, we check if a value can be found
                in `generation_config.cache_config` and otherwise we raise a ValueError.
            max_cache_len (`Optional[int]`): The maximum cache length for generation. Same mechanism as `batch_size` if
                not provided.
            device (`Optional[torch.device]`): The device to use. If not provided, we check if a value can be found
                in `generation_config.cache_config` and otherwise we use `model.device` (no error is raised).

        Raises:
            AssertionError: If the pretrained model does not have caching enabled or if it does
            not use a 'static' caching implementation in `model.generation_config`.
            ValueError: If `batch_size` or `max_cache_len` is not provided, either as an argument or in `cache_config`.
        NúvThe model must have a generation config to be exported with static caching. Please set `generation_config` in `model`.zvThe model must have caching enabled to be exported with static caching. Please set `generation_config.use_cache=True`.Ústaticz–The model must use a 'static' caching implementation to be exported with static caching. Please set `generation_config.cache_implementation='static'`.re   úFbatch_size must be provided, either as an argument or in cache_config.r   úImax_cache_len must be provided, either as an argument or in cache_config.rf   )r   r   Ú
key_cache_F©Ú
persistentÚvalue_cache_Úcumulative_length_)rl   r"   r   rm   Úgeneration_configÚAssertionErrorri   Úcache_implementationÚcache_configÚgetro   rf   r   r   Ústatic_cacheÚ	enumerateÚlayersÚ
isinstancer
   r	   rÎ   r'   Úearly_initializationÚregister_bufferÚkeysÚvaluesÚcumulative_length©r    r   re   r   rf   r   rÚ   rÝ   r§   rÅ   rÍ   rÂ   r'   ru   s                €r!   r"   z-TorchExportableModuleWithStaticCache.__init__ß  s�  ø€ õ2 	‰Œ×ÒÑÔÐà”×-Ò-Ñ/Ô/ˆØ!Ô3Ðð Ð$Ý ð=ñô ð ð !Ô*ð 	Ý ðAñô ð ð Ô1°XÒ=Ð=Ý ðPñô ð ð
 /Ô;ÐC�r�rÐIZÔIgˆð ÐØ%×)Ò)¨,¸Ñ=Ô=ˆJØÐ!Ý Ð!iÑjÔjÐjØÐ Ø(×,Ò,¨_¸dÑCÔCˆMØÐ$Ý Ð!lÑmÔmÐmàˆ>Ø!×%Ò% h°´Ñ=Ô=ˆFð ˆŒ
Ý'°mÈFÐSÑSÔSˆÔõ " $Ô"3Ô":Ñ;Ô;ð 	Ið 	I‰HˆAˆuÝ˜%Õ!9Ñ:Ô:ð IÝ.9¸-Ñ.HÔ.H�Ô!Ô(¨Ñ+øÝ-¨fÑ5Ô5Ñˆ	�8Ø”
Ô ˆàÔ×.Ò.¨z¸9ÀhÐPUÐW]Ñ^Ô^Ð^õ " $Ô"3Ô":Ñ;Ô;ð 	fð 	f‰HˆAˆuØ× Ò Ð!1¨aÐ!1Ð!1°5´:È%Ð ÑPÔPÐPØ× Ò Ð!3°Ð!3Ð!3°U´\ÈeÐ ÑTÔTÐTØ× Ò Ð!9°aÐ!9Ð!9¸5Ô;RÐ_dÐ ÑeÔeÐeÐeð	fð 	fr#   rD   rv   rE   c                 óÞ   — | j         j        D ]"}|j                             |d         ¦  «         Œ#| j         }|                      ||d|d¬¦  «        }t          |d¦  «        r|j        S |j        S )a8  
        Forward pass of the module, which is compatible with the ExecuTorch runtime.

        Args:
            input_ids (`torch.Tensor`): Tensor representing current input token id to the module.
            inputs_embeds (`torch.Tensor`): Tensor representing current input embeddings to the module.
            cache_position (`torch.Tensor`): Tensor representing current input position in the cache.

        Returns:
            torch.Tensor: Logits output from the model.

        This forward adapter serves two primary purposes:

        1. **Making the Model `torch.export`-Compatible**:
            The adapter hides unsupported objects, such as the `Cache`, from the graph inputs and outputs,
            enabling the model to be exportable using `torch.export` without encountering issues.

        2. **Ensuring Compatibility with `ExecuTorch` runtime**:
            The adapter matches the model's forward signature with that in `executorch/extension/llm/runner`,
            ensuring that the exported model can be executed in `ExecuTorch` out-of-the-box.
        r   NT©rD   rv   Úattention_maskÚpast_key_valuesri   r¬   )rß   rá   rç   Úcopy_r   rn   r¬   Úlast_hidden_state)r    rD   rv   rE   rÅ   rì   Úoutss          r!   rV   z,TorchExportableModuleWithStaticCache.forward0  sŽ   € ð< Ô&Ô-ð 	=ð 	=ˆEØÔ#×)Ò)¨.¸Ô*;Ñ<Ô<Ð<Ð<àÔ+ˆà�zŠzØØ'ØØ+Øð ñ 
ô 
ˆõ �4˜Ñ"Ô"ð 	*à”;Ðð Ô)Ð)r#   rƒ   Úprompt_token_idsr[   c           	      óÐ  — |j         }|j        d         }||z   }|                      ¦   «         D ]9\  }}|                     d¦  «        r|j        d         }t	          ||¦  «        } nŒ:g }	t          t	          ||¦  «        ¦  «        D ]�}
|                      ¦   «                              |dd…|
|
dz   …f         t          j	        |
gt          j
        |¬¦  «        ¬¦  «        }|	                     |d         |
                              ¦   «         ¦  «         ŒŽt          j        |dd…ddd…f         d¬	¦  «                             ¦   «         }|	                     |¦  «         t          |	¦  «        |k     rÔ|                      ¦   «                              t          j	        |ggt          j
        |¬¦  «        t          j	        t          |	¦  «        gt          j
        |¬¦  «        ¬¦  «        }t          j        |dd…ddd…f         d¬	¦  «                             ¦   «         }|	                     |¦  «         t          |	¦  «        |k     °Ôt          j	        |	gt          j
        |¬¦  «        S )
aá  
        Generate a sequence of tokens using an exported program.

        This util function is designed to test exported models by simulating the generation process.
        It processes the input prompt tokens sequentially (no parallel prefill).
        This generate function is not intended to replace the original `generate` method, and the support
        for leveraging the original `generate` is potentially planned!

        Args:
            exported_program (`torch.export.ExportedProgram`): The exported program generated via `torch.export`.
            prompt_token_ids (`torch.Tensor`): Tensor representing the input prompt token IDs.
            max_new_tokens (`int`): Maximum number of new tokens to generate. Note that the total generation
                length is limited by both `max_new_tokens` and the model's cache size.

        Returns:
            torch.Tensor: A tensor containing the generated sequence of token IDs, including the original prompt tokens.
        r{   Ú	key_cacher   Nr   r|   rC   r   rŒ   )rf   r   Únamed_buffersÚ
startswithrK   r”   r‘   rV   r.   r•   rI   Úappendr    r�   Úlen)rƒ   rð   r[   rf   Úprompt_token_lenÚmax_generation_lengthÚbuffer_nameÚbufferr   Úresponse_tokensÚ	input_posÚresultÚcurrent_tokens                r!   r^   z-TorchExportableModuleWithStaticCache.generatea  sr  € ð. "Ô(ˆØ+Ô1°"Ô5ÐØ 0°>Ñ AÐØ#3×#AÒ#AÑ#CÔ#Cð 	ð 	ÑˆK˜Ø×%Ò% kÑ2Ô2ð Ø &¤¨Q¤�Ý(+Ð,AÀ=Ñ(QÔ(QÐ%Ø�ðð
 ˆÝ�sÐ#8Ð:JÑKÔKÑLÔLð 	Jð 	JˆIØ%×,Ò,Ñ.Ô.×6Ò6Ø*¨1¨1¨1¨i¸)Àa¹-Ð.GÐ+GÔHÝ$œ|¨Y¨K½u¼zÐRXÐYÑYÔYð 7ñ ô ˆFð ×"Ò"Ð#3°AÔ#6°yÔ#A×#FÒ#FÑ#HÔ#HÑIÔIÐIÐIåœ V¨A¨A¨A¨r°1°1°1¨HÔ%5¸2Ð>Ñ>Ô>×CÒCÑEÔEˆØ×Ò˜}Ñ-Ô-Ð-å�/Ñ"Ô"Ð%:Ò:Ð:Ø%×,Ò,Ñ.Ô.×6Ò6Ýœ,¨¨Ð'8ÅÄ
ÐSYÐZÑZÔZÝ$œ|­S°Ñ-AÔ-AÐ,BÍ%Ì*Ð]cÐdÑdÔdð 7ñ ô ˆFõ "œL¨°°°°2°q°q°q°Ô)9¸rÐBÑBÔB×GÒGÑIÔIˆMØ×"Ò" =Ñ1Ô1Ð1õ �/Ñ"Ô"Ð%:Ò:Ð:õ Œ|˜_Ð-µU´ZÈÐOÑOÔOÐOr#   r´   )r_   r`   ra   rb   r   rc   r.   rf   r"   Ú
LongTensorrµ   rV   r¹   r1   r¸   r^   r»   r¼   s   @r!   rt   rt   Ô  s_  ø€ € € € € ðð ð "&Ø$(Ø&*ðOfð OfàðOfð ˜$‘JðOfð ˜T‘zð	Ofð
 ”˜tÑ#ðOfð 
ðOfð Ofð Ofð Ofð Ofð Ofðf .2Ø-1Ø.2ð	/*ð /*àÔ# dÑ*ð/*ð ”| dÑ*ð/*ð œ tÑ+ð	/*ð /*ð /*ð /*ðb ð2PØœ,Ô6ð2Pàœ,ð2Pð ð2Pð 
Œð	2Pð 2Pð 2Pñ „\ð2Pð 2Pð 2Pð 2Pð 2Pr#   rt   c                   ó¶   ‡ — e Zd ZdZ	 	 	 ddededz  dedz  dej        dz  ddf
ˆ fd„Z	 	 	 dd	ej	        dz  d
ej
        dz  dej
        dz  dej
        fd„Zˆ xZS )rq   a  
    A recipe module designed to make a `PreTrainedModel` exportable with `torch.export`,
    specifically for decoder-only LM to hybrid `StaticCache`. This module ensures that the
    exported model is compatible with further lowering and execution in `ExecuTorch`.
    Nr   re   r   rf   rg   c                 ó$  •— t          ¦   «                              ¦   «          || _        |j                             ¦   «         }|j        }|€t          d¦  «        ‚|j        st          d¦  «        ‚|j        €i n|j        }|€'| 	                    dd¦  «        }|€t          d¦  «        ‚|€'| 	                    dd¦  «        }|€t          d¦  «        ‚|€| 	                    d|j        ¦  «        }t          ||¬	¦  «        | _        t          | j        j        ¦  «        D ]6\  }}	t!          |	t"          ¦  «        rt%          |¦  «        | j        j        |<   Œ7t'          |¦  «        \  }
}| j        j        }| j                             ||
|||¦  «         t          | j        j        ¦  «        D ]e\  }}	|                      d
|› �|	j        d¬¦  «         |                      d|› �|	j        d¬¦  «         |                      d|› �|	j        d¬¦  «         ŒfdS )aÃ  
        Initializes the exportable module.

        Args:
            model (`PreTrainedModel`): The pretrained model to wrap.
            batch_size (`Optional[int]`): The batch size of the model. If not provided, we check if a value can be found
                in `generation_config.cache_config` and otherwise we raise a ValueError.
            max_cache_len (`Optional[int]`): The maximum cache length for generation. Same mechanism as `batch_size` if
                not provided.
            device (`Optional[torch.device]`): The device to use. If not provided, we check if a value can be found
                in `generation_config.cache_config` and otherwise we use `model.device` (no error is raised).
        Raises:
            AssertionError: If the model doesn't have the expected configuration for hybrid StaticCache.
            ValueError: If `batch_size` or `max_cache_len` is not provided, either as an argument or in `cache_config`.
        NrÑ   z Model must have caching enabled.re   rÓ   r   rÔ   rf   ©r   r   rÕ   FrÖ   rØ   rÙ   )rl   r"   r   r   rm   rÚ   rÛ   ri   rÝ   rÞ   ro   rf   r   Úcacherà   rá   râ   r
   r	   rÎ   r'   rã   rä   rå   ræ   rç   rè   s                €r!   r"   z-TorchExportableModuleWithHybridCache.__init__ž  sS  ø€ õ, 	‰Œ×ÒÑÔÐØˆŒ
Ø”×-Ò-Ñ/Ô/ˆØ!Ô3Ðð Ð$Ý ð=ñô ð ð Ôð 	EÝ Ð!CÑDÔDÐDà.Ô;ÐC�r�rÐIZÔIgˆàÐØ%×)Ò)¨,¸Ñ=Ô=ˆJØÐ!Ý Ð!iÑjÔjÐjØÐ Ø(×,Ò,¨_¸dÑCÔCˆMØÐ$Ý Ð!lÑmÔmÐmàˆ>Ø!×%Ò% h°´Ñ=Ô=ˆFõ !¨¸mÐLÑLÔLˆŒ
õ " $¤*Ô"3Ñ4Ô4ð 	Bð 	B‰HˆAˆuÝ˜%Õ!9Ñ:Ô:ð BÝ'2°=Ñ'AÔ'A�”
Ô! !Ñ$øÝ-¨fÑ5Ô5Ñˆ	�8Ø”
Ô ˆàŒ
×'Ò'¨
°I¸xÈÐPVÑWÔWÐWõ " $¤*Ô"3Ñ4Ô4ð 	fð 	f‰HˆAˆuØ× Ò Ð!1¨aÐ!1Ð!1°5´:È%Ð ÑPÔPÐPØ× Ò Ð!3°Ð!3Ð!3°U´\ÈeÐ ÑTÔTÐTØ× Ò Ð!9°aÐ!9Ð!9¸5Ô;RÐ_dÐ ÑeÔeÐeÐeð	fð 	fr#   rD   rv   rE   c                 ó¬   — | j         j        D ]"}|j                             |d         ¦  «         Œ#|                      ||d| j         d¬¦  «        }|j        S )aô  
        Forward pass of the module, which is compatible with the ExecuTorch llm runner.

        Args:
            input_ids (`torch.Tensor`): Tensor representing current input token id to the module.
            inputs_embeds (`Optional[torch.Tensor]`): Tensor representing current input embeddings to the module.
            cache_position (`torch.Tensor`): Tensor representing current input position in the cache.

        Returns:
            torch.Tensor: Logits output from the model.
        r   NTrê   )r  rá   rç   rí   r   r¬   )r    rD   rv   rE   rÅ   r«   s         r!   rV   z,TorchExportableModuleWithHybridCache.forwardâ  sl   € ð( ”ZÔ&ð 	=ð 	=ˆEØÔ#×)Ò)¨.¸Ô*;Ñ<Ô<Ð<Ð<ð —*’*ØØ'ØØ œJØð ñ 
ô 
ˆð Œ~Ðr#   r´   )r_   r`   ra   rb   r   rc   r.   rf   r"   rÿ   rµ   rV   r»   r¼   s   @r!   rq   rq   —  s  ø€ € € € € ðð ð "&Ø$(Ø&*ðBfð BfàðBfð ˜$‘JðBfð ˜T‘zð	Bfð
 ”˜tÑ#ðBfð 
ðBfð Bfð Bfð Bfð Bfð BfðL .2Ø-1Ø.2ð	!ð !àÔ# dÑ*ð!ð ”| dÑ*ð!ð œ tÑ+ð	!ð
 
Œð!ð !ð !ð !ð !ð !ð !ð !r#   rq   r   Úexample_input_idsÚexample_cache_positionr+   r,   c                 ó0  — ddl } |j        ¦   «         5  |�|n |j        dgg|j        | j        ¬¦  «        }|�|n |j        dg|j        | j        ¬¦  «        }t          d¦  «        r4|j                             t          | ¦  «        d||dœ||�|nd¬	¦  «        }n`|�t          j	        d
¦  «         |�t          j	        d¦  «         |j        j
                             t          | ¦  «        d||dœdd¬¦  «        }|cddd¦  «         S # 1 swxY w Y   dS )aî  
    Convert a `PreTrainedModel` into an exportable module and export it using `torch.export`,
    ensuring the exported model is compatible with `ExecuTorch`.

    Args:
        model (`PreTrainedModel`): The pretrained model to be exported.
        example_input_ids (`Optional[torch.Tensor]`): Example input token id used by `torch.export`.
        example_cache_position (`Optional[torch.Tensor]`): Example current cache position used by `torch.export`.
        dynamic_shapes(`Optional[dict]`): Dynamic shapes used by `torch.export`.
        strict(`Optional[bool]`): Flag to instruct `torch.export` to use `dynamo`.

    Returns:
        Exported program (`torch.export.ExportedProgram`): The exported program generated via `torch.export`.
    r   Nr   r|   z2.6.0rS   rC   Tr}   zWDynamic shapes spec will be ignored by convert_and_export_with_cache for torch < 2.6.0.zSThe strict flag will be ignored by convert_and_export_with_cache for torch < 2.6.0.F)r*   rT   Úpre_dispatchr,   )Útorch.export._traceÚno_gradr•   rI   rf   r   r1   rt   rr   r~   Ú_traceÚ_export)r   r  r  r+   r,   r.   rƒ   s          r!   Úconvert_and_export_with_cacher    sÁ  € ð, ÐÐÐà	ˆŒ‰Œð ' ð ' ð !Ð,ð Ðà�” ˜s˜e¨5¬:¸e¼lÐKÑKÔKð 	ð &Ð1ð #Ð"à�”˜q˜c¨¬¸E¼LÐIÑIÔIð 	õ % WÑ-Ô-ð 	Ø$œ|×2Ò2Ý4°UÑ;Ô;ØØ%6ÐJ`ÐaÐaØ-Ø!'Ð!3�v�v¸ð  3ñ  ô  ÐÐð Ð)Ý”Ømñô ð ð Ð!Ý”Ð uÑvÔvÐvð
  %œ|Ô2×:Ò:Ý4°UÑ;Ô;ØØ%6ÐJ`ÐaÐaØ"Øð  ;ñ  ô  Ðð  ðO' ð ' ð ' ð ' ñ ' ô ' ð ' ð ' ð ' ð ' ð ' ð ' øøøð ' ð ' ð ' ð ' ð ' ð ' s   ”C*DÄDÄDc                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )Ú Seq2SeqLMEncoderExportableModulez·
    A wrapper module designed to make a Seq2Seq LM encoder exportable with `torch.export`.
    This module ensures that the exported encoder model is compatible with ExecuTorch.
    c                 óV   •— t          ¦   «                              ¦   «          || _        d S ©N)rl   r"   Úencoder)r    Úencoder_modelru   s     €r!   r"   z)Seq2SeqLMEncoderExportableModule.__init__N  s$   ø€ Ý‰Œ×ÒÑÔÐØ$ˆŒˆˆr#   c                 ó8   — |                       |¬¦  «        j        S )N)rD   )r  rî   )r    rD   s     r!   rV   z(Seq2SeqLMEncoderExportableModule.forwardR  s   € Ø�|Š| iˆ|Ñ0Ô0ÔBÐBr#   ©r_   r`   ra   rb   r"   rV   r»   r¼   s   @r!   r  r  H  sX   ø€ € € € € ðð ð
%ð %ð %ð %ð %ðCð Cð Cð Cð Cð Cð Cr#   r  c                   ó(   ‡ — e Zd ZdZˆ fd„Zd„ Zˆ xZS )Ú/Seq2SeqLMDecoderExportableModuleWithStaticCachezÚ
    A wrapper module designed to make a Seq2Seq LM decoder exportable with `torch.export`,
    specifically for use with static caching. This module ensures the exported decoder
    is compatible with ExecuTorch.
    c                 ó¾  •— t          ¦   «                              ¦   «          |                     ¦   «         | _        |j        | _        |j        | _        t          |                     ¦   «         ¦  «        j        }t          | j        |¬¦  «        | _
        t          | j
        j        ¦  «        D ]6\  }}t          |t          ¦  «        rt          |¦  «        | j
        j        |<   Œ7t!          | j        ¦  «        \  }}| j
                             |||t$          j        |¦  «         t)          | j
        t+          | j        ¬¦  «        ¦  «        | _        t/          ¦   «          t          | j
        j        ¦  «        D ]e\  }}|                      d|› �|j        d¬¦  «         |                      d|› �|j        d¬¦  «         |                      d|› �|j        d¬¦  «         Œfd S )Nr  ©r   rÕ   FrÖ   rØ   rÙ   )rl   r"   Úget_decoderÚdecoderÚlm_headr   ÚnextÚ
parametersrf   r   rß   rà   rá   râ   r
   r	   rÎ   rã   r.   r0   r   r   r  Ú%register_dynamic_cache_export_supporträ   rå   ræ   rç   )
r    r   Úmax_static_cache_lengthre   r�   r§   rÅ   rÍ   rÂ   ru   s
            €r!   r"   z8Seq2SeqLMDecoderExportableModuleWithStaticCache.__init__]  sÚ  ø€ Ý‰Œ×ÒÑÔÐð ×(Ò(Ñ*Ô*ˆŒØ”}ˆŒØ”lˆŒõ ˜E×,Ò,Ñ.Ô.Ñ/Ô/Ô6ˆõ (¨t¬{ÐJaÐbÑbÔbˆÔõ " $Ô"3Ô":Ñ;Ô;ð 	Sð 	S‰HˆAˆuÝ˜%Õ!9Ñ:Ô:ð SÝ.9Ð:QÑ.RÔ.R�Ô!Ô(¨Ñ+øÝ-¨d¬kÑ:Ô:Ñˆ	�8ØÔ×.Ò.¨z¸9ÀhÕPUÔP]Ð_kÑlÔlÐlÝ(¨Ô):½LÐPTÔP[Ð<\Ñ<\Ô<\Ñ]Ô]ˆŒ
å-Ñ/Ô/Ð/õ " $Ô"3Ô":Ñ;Ô;ð 	fð 	f‰HˆAˆuØ× Ò Ð!1¨aÐ!1Ð!1°5´:È%Ð ÑPÔPÐPØ× Ò Ð!3°Ð!3Ð!3°U´\ÈeÐ ÑTÔTÐTØ× Ò Ð!9°aÐ!9Ð!9¸5Ô;RÐ_dÐ ÑeÔeÐeÐeð	fð 	fr#   c                 óÖ   — | j         j        D ]"}|j                             |d         ¦  «         Œ#|                      ||| j        d¬¦  «        }|                      |d         ¦  «        }|S )Nr   T)rD   Úencoder_hidden_statesrì   ri   )rß   rá   rç   rí   r  r  r  )r    Údecoder_input_idsr"  rE   rÅ   r«   Ú	lm_logitss          r!   rV   z7Seq2SeqLMDecoderExportableModuleWithStaticCache.forward|  s}   € ð Ô&Ô-ð 	=ð 	=ˆEØÔ#×)Ò)¨.¸Ô*;Ñ<Ô<Ð<Ð<ð —,’,Ø'Ø"7Ø œJØð	 ñ 
ô 
ˆð —L’L ¨¤Ñ,Ô,ˆ	àÐr#   r  r¼   s   @r!   r  r  V  sV   ø€ € € € € ðð ðfð fð fð fð fð>ð ð ð ð ð ð r#   r  c                   ó<   ‡ — e Zd Z	 dˆ fd„	Zd„ Zd„ Zdd	„Zd
„ Zˆ xZS )ÚSeq2SeqLMExportableModuler   é   rÒ   r   c                 ó  •— t          ¦   «                              ¦   «          || _        |                     ¦   «         | _        |j        | _        || _        t          d||||dœ|j        j	        ¬¦  «        | _        d | _
        d | _        d S )NT)re   r   )ri   Ú
max_lengthrÜ   rÝ   r¡   )rl   r"   Ú
full_modelÚget_encoderr  r   Úmax_hidden_seq_lengthr   rÚ   r¡   Úexported_encoderÚexported_decoder)r    r   re   r,  rÜ   Úmax_cache_lengthru   s         €r!   r"   z"Seq2SeqLMExportableModule.__init__’  s™   ø€ õ 	‰Œ×ÒÑÔÐàˆŒØ×(Ò(Ñ*Ô*ˆŒØ”lˆŒØ%:ˆÔ"Ý!1ØØ'Ø!5à(Ø!1ðð ð Ô0Ô=ð	"
ñ 	"
ô 	"
ˆÔð !%ˆÔØ $ˆÔÐÐr#   c                 ó~  — t          | j        ¦  «                             | j        j        ¦  «                             ¦   «         }t          j                             d| j	        ¬¦  «        }t          j
        ¦   «         5  t          j                             ||fdd|iid¬¦  «        }d d d ¦  «         n# 1 swxY w Y   |S )NÚencoder_seq_lengthrA   rD   r   T©r+   r,   )r  r  r’   r*  rf   r-   r.   r1   r2   r,  r
  )r    Úencoder_input_idsÚwrapped_encoderrP   r-  s        r!   Ú_export_encoderz)Seq2SeqLMExportableModule._export_encoder¨  sò   € Ý:¸4¼<ÑHÔH×KÒKÈDÌOÔLbÑcÔc×hÒhÑjÔjˆõ ”l×&Ò&Ð';ÀÔA[Ð&Ñ\Ô\ˆõ Œ]‰_Œ_ð 	ð 	Ý$œ|×2Ò2ØÐ"3Ð!5À{ÐUVÐXcÐTdÐFeÐnrð  3ñ  ô  Ðð	ð 	ð 	ñ 	ô 	ð 	ð 	ð 	ð 	ð 	ð 	øøøð 	ð 	ð 	ð 	ð
  Ðs   Á=)B2Â2B6Â9B6c           	      ó‚  — | j         j        }t          | j         | j        j                             d¦  «        | j        j                             d¦  «        ¬¦  «                             |¦  «                             ¦   «         }|                     |¦  «        }|                     |¦  «        }|                     |¦  «        }t          j	         
                    d| j        ¬¦  «        }t          j        ¦   «         5  t          j	         	                    ||||fd d|id dœd¬	¦  «        }d d d ¦  «         n# 1 swxY w Y   |S )
Nr   re   )r   r   re   Úencoder_hidden_seq_lengthrA   r   )r#  r"  rE   Tr2  )r*  rf   r  rÚ   rÝ   rÞ   r’   r-   r.   r1   r2   r,  r
  )r    r#  r"  rE   Útarget_deviceÚwrapped_decoderÚencoder_seq_len_dimr.  s           r!   Ú_export_decoderz)Seq2SeqLMExportableModule._export_decoder¶  s‚  € ØœÔ.ˆå;Ø”oØ(,Ô(>Ô(K×(OÒ(OÐP_Ñ(`Ô(`ØÔ1Ô>×BÒBÀ<ÑPÔPðñ ô ÷
 ŠR�ÑÔßŠT‰VŒVð 	ð .×0Ò0°Ñ?Ô?ÐØ 5× 8Ò 8¸Ñ GÔ GÐØ'×*Ò*¨=Ñ9Ô9ˆõ $œl×.Ò.Ð/JÐPTÔPjÐ.ÑkÔkÐõ Œ]‰_Œ_ð 
	ð 
	Ý$œ|×2Ò2ØØ"Ð$9¸>ÐJà)-Ø./Ð1DÐ-EØ&*ð ð  ð
 ð  3ñ 	 ô 	 Ðð
	ð 
	ð 
	ñ 
	ô 
	ð 
	ð 
	ð 
	ð 
	ð 
	ð 
	øøøð 
	ð 
	ð 
	ð 
	ð  Ðs   Ã;-D4Ä4D8Ä;D8Nc                 ó  — | j         j        }|�|n t          j        dt          j        |¬¦  «        }|�|n"t          j        dggt          j        |¬¦  «        }|�|n!t          j        dgt          j        |¬¦  «        }|�|nJt          j        | j        j         	                    d¦  «        d| j
        j        ft          j        |¬¦  «        }	|                      |¦  «        | _        |                      ||	|¦  «        | _        | S )N)r   é
   r|   r   re   r=  )r*  rf   r.   ÚonesrI   r•   rH   rÚ   rÝ   rÞ   r   Úd_modelr0   r5  r-  r;  r.  )
r    r3  r#  r"  rE   rf   Úexample_encoder_input_idsÚexample_decoder_input_idsr  Úexample_encoder_hidden_statess
             r!   r1   z Seq2SeqLMExportableModule.exportÙ  s&  € Ø”Ô'ˆð !Ð,ð Ðå”˜G­5¬:¸fÐEÑEÔEð 	"ð !Ð,ð Ðå” ˜s˜e­5¬:¸fÐEÑEÔEð 	"ð -Ð8ˆNˆN½e¼lÈAÈ3ÕV[ÔV`ÐioÐ>pÑ>pÔ>pð 	ð
 %Ð0ð "Ð!å”ØÔ'Ô4×8Ò8¸ÑFÔFÈÈDÌKÔL_Ð`Ý”mØðñ ô ð 	&ð !%× 4Ò 4Ð5NÑ OÔ OˆÔØ $× 4Ò 4Ø%Ð'DÐF\ñ!
ô !
ˆÔð
 ˆr#   c                 óø  — t          j        ¦   «         5  | j        j        }|j        |k    r|                     |¦  «        } | j                             ¦   «         |¦  «        }t          j        dggt           j        |¬¦  «        }dg}t          |dz
  ¦  «        D ]Å} | j
                             ¦   «         ||t          j        |gt           j        |¬¦  «        ¦  «        }t          j        |d d …dd d …f         d¬¦  «                             ¦   «         }	|                     |	¦  «         t          j        |	ggt           j        |¬¦  «        }|	| j        j        k    r nŒÆ|cd d d ¦  «         S # 1 swxY w Y   d S )Nr   r|   r   r{   rŒ   )r.   r
  r*  rf   r’   r-  r‘   r•   rI   r”   r.  r�   r    rõ   rÚ   r¡   )
r    rð   r[   r�   Úencoder_outputr#  r¥   r§   r¬   Ú
next_tokens
             r!   r^   z"Seq2SeqLMExportableModule.generateù  sÏ  € ÝŒ]‰_Œ_ð  	!ð  	!Øœ?Ô1ˆLð  Ô&¨,Ò6Ð6Ø#3×#6Ò#6°|Ñ#DÔ#DÐ ð <˜TÔ2×9Ò9Ñ;Ô;Ð<LÑMÔMˆNõ !&¤¨q¨c¨U½%¼*È\Ð ZÑ ZÔ ZÐØ˜CˆMõ ˜>¨AÑ-Ñ.Ô.ð ð �à7˜Ô.×5Ò5Ñ7Ô7Ø% ~µu´|ÀQÀCÍuÌzÐbnÐ7oÑ7oÔ7oñô �õ
 #œ\¨&°°°°B¸¸¸°Ô*:ÀÐCÑCÔC×HÒHÑJÔJ�
Ø×$Ò$ ZÑ0Ô0Ð0õ %*¤L°:°,°ÅuÄzÐZfÐ$gÑ$gÔ$gÐ!ð  Ô!7Ô!DÒDÐDØ�Eð Eð !ðA 	!ð  	!ð  	!ð  	!ñ  	!ô  	!ð  	!ð  	!ð  	!ð  	!ð  	!ð  	!øøøð  	!ð  	!ð  	!ð  	!ð  	!ð  	!s   ”EE/Å/E3Å6E3)r   r'  rÒ   r   ©NNNN)	r_   r`   ra   r"   r5  r;  r1   r^   r»   r¼   s   @r!   r&  r&  ‘  sƒ   ø€ € € € € àosð%ð %ð %ð %ð %ð %ð, ð  ð  ð! ð ! ð ! ðFð ð ð ð@!!ð !!ð !!ð !!ð !!ð !!ð !!r#   r&  Úexample_attention_maskc           
      óò   — t          ¦   «          t          j        ¦   «         5  t          j                             | d||t	          | j        ¬¦  «        ddœd¬¦  «        }|cddd¦  «         S # 1 swxY w Y   dS )a  
    Export a model with DynamicCache using `torch.export`, ensuring the exported model is compatible with `ExecuTorch`.

    Args:
        model (`PreTrainedModel`): The pretrained model to be exported.
        example_input_ids (`Optional[torch.Tensor]`): Example input token id used by `torch.export`.
        example_attention_mask (`Optional[torch.Tensor]`): Example attention mask used by `torch.export`.

    Returns:
        Exported program (`torch.export.ExportedProgram`): The exported program generated via `torch.export`.
    rS   r  T)rD   rë   rì   ri   F)r,   N)r  r.   r
  r1   r   r   )r   r  rG  rƒ   s       r!   Úexport_with_dynamic_cacherI    sÐ   € õ" *Ñ+Ô+Ð+å	Œ‰Œð  ð  Ý œ<×.Ò.ØØà.Ø"8Ý#/°u´|Ð#DÑ#DÔ#DØ!ð	ð ð ð /ñ 

ô 

Ðð  ð ð  ð  ð  ñ  ô  ð  ð  ð  ð  ð  ð  øøøð  ð  ð  ð  ð  ð  s   ¢=A,Á,A0Á3A0c                  óN  — 	 t           j        j                             t          d„ t
          t          j        › dt          j        › �d„ ¬¦  «         t           j        j         	                    t          d„ ¦  «         dS # t          $ r} dt          | ¦  «        vr‚ Y d} ~ dS d} ~ ww xY w)z>
    Utilities for `DynamicCache` <> torch.export support
    c                 ód   — t           j        j                             t	          | ¦  «        ¦  «        S r  )r.   ÚutilsÚ_pytreeÚ_dict_flattenÚ_get_cache_dict©Údynamic_caches    r!   ú<lambda>z7register_dynamic_cache_export_support.<locals>.<lambda>G  s"   € ¥%¤+Ô"5×"CÒ"CÅOÐTaÑDbÔDbÑ"cÔ"c€ r#   ú.c                 ód   — t           j        j                             t	          | ¦  «        ¦  «        S r  )r.   rL  rM  Ú_dict_flatten_with_keysrO  rP  s    r!   rR  z7register_dynamic_cache_export_support.<locals>.<lambda>J  s&   € µu´{Ô7J×7bÒ7bÝ Ñ.Ô.ñ8ô 8€ r#   )Úserialized_type_nameÚflatten_with_keys_fnc                 óf   — t           j        j                             t	          | ¦  «        |¦  «        S r  )r.   ÚfxrM  Ú_dict_flatten_specrO  )r  Úspecs     r!   rR  z7register_dynamic_cache_export_support.<locals>.<lambda>Q  s%   € ¥¤Ô 0× CÒ CÅOÐTYÑDZÔDZÐ\`Ñ aÔ a€ r#   z!already registered as pytree nodeN)r.   rL  rM  Úregister_pytree_noder   Ú_unflatten_dynamic_cacher`   r_   rY  Úregister_pytree_flatten_specro   rº   )Úes    r!   r  r  ?  sÑ   € ð
ÝŒÔ×0Ò0ÝØcÐcÝ$Ý$0Ô$;Ð!UÐ!U½lÔ>SÐ!UÐ!Uð"ð "ð 	1ñ 	
ô 	
ð 	
õ 	ŒÔ×5Ò5ÝØaÐañ	
ô 	
ð 	
ð 	
ð 	
øõ
 ð ð ð Ø.µc¸!±f´fÐ<Ð<Øð =Ð<Ð<Ð<Ð<Ð<øøøøðøøøs   ‚A9A= Á=
B$ÂBÂB$r  c                 óØ   — t          d„ | j        D ¦   «         ¦  «        rt          d¦  «        ‚t          st	          j        d¦  «         d„ | j        D ¦   «         d„ | j        D ¦   «         dœS )z9Convert cache to dictionary format for pytree operations.c              3   óP   K  — | ]!}t          |t          t          f¦  «         V — Œ"d S r  )râ   r   r   ©rÄ   rÅ   s     r!   ú	<genexpr>z"_get_cache_dict.<locals>.<genexpr>[  s6   è è € Ð
fÐ
fÐPU�z˜%¥,Õ0IÐ!JÑKÔKÐKÐ
fÐ
fÐ
fÐ
fÐ
fÐ
fr#   zFThis pytree flattening function should only be applied to DynamicCachez[DynamicCache + torch.export is tested on torch 2.6.0+ and may not work on earlier versions.c                 ó*   — g | ]}|j         ®	|j         ‘ŒS r  )rå   rb  s     r!   rÆ   z#_get_cache_dict.<locals>.<listcomp>b  s!   € ÐUÐUÐU U¸e¼jÐ>T�e”jÐ>TÐ>TÐ>Tr#   c                 ó*   — g | ]}|j         ®	|j         ‘ŒS r  )ræ   rb  s     r!   rÆ   z#_get_cache_dict.<locals>.<listcomp>c  s!   € Ð[Ð[Ð[¨À%Ä,ÐBZ˜œÐBZÐBZÐBZr#   )rò   Úvalue_cache)Úanyrá   ÚRuntimeErrorr   rr   r~   )r  s    r!   rO  rO  Y  s‡   € å
Ð
fÐ
fÐY^ÔYeÐ
fÑ
fÔ
fÑfÔfð eÝÐcÑdÔdÐdå-ð wÝŒÐuÑvÔvÐvð VÐU¨e¬lÐUÑUÔUØ[Ð[°%´,Ð[Ñ[Ô[ðð ð r#   Úcontextc                 óÚ  — t           j        j                             | |¦  «        }t	          ¦   «         }|                     dg ¦  «        }|                     dg ¦  «        }t          t          t          |¦  «        t          |¦  «        ¦  «        ¦  «        D ]S}|t          |¦  «        k     r||         nd }|t          |¦  «        k     r||         nd }| 	                    |||¦  «         ŒT|S )Nrò   rf  )
r.   rL  rM  Ú_dict_unflattenr   rÞ   r”   rB   rö   Úupdate)	ræ   ri  Ú
dictionaryr  Úkey_listÚ
value_listÚidxÚkeyÚvalues	            r!   r]  r]  g  sÐ   € Ý”Ô$×4Ò4°V¸WÑEÔE€JÝ‰NŒN€Eà�~Š~˜k¨2Ñ.Ô.€HØ—’ ¨rÑ2Ô2€JÝ•S�˜X™œ­¨J©¬Ñ8Ô8Ñ9Ô9ð &ð &ˆØ"¥S¨¡]¤]Ò2Ð2ˆh�sŒmˆm¸ˆØ#&­¨Z©¬Ò#8Ð#8�
˜3”�¸dˆØ�Š�S˜% Ñ%Ô%Ð%Ð%Ø€Lr#   rF  )NN))rr   r.   Úcache_utilsr   r   r   r   r   r	   r
   Úgeneration.configuration_utilsr   Úmodeling_utilsr   Úpytorch_utilsr   r   r   ÚnnÚModulerF   Útuplerc   ÚlistrÎ   rt   rq   rµ   r¶   r·   r  r  r  r&  rI  r  rO  rL  rM  ÚContextr]  rS   r#   r!   ú<module>r|     s…  ðð €€€à €€€ðð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð ð >Ð =Ð =Ð =Ð =Ð =Ø ,Ð ,Ð ,Ð ,Ð ,Ð ,ðð ð ð ð ð ð ð ðUð Uð Uð Uð Uñ Uô Uð UðpALð ALð ALð ALð AL¨E¬H¬Oñ ALô ALð ALðH˜u S¨4°¬9¡_°c¸DÀ¼I±oÐ%EÔFð ð ð ð ð,@Pð @Pð @Pð @Pð @P¨5¬8¬?ñ @Pô @Pð @PðFlð lð lð lð l¨5¬8¬?ñ lô lð lðb .2Ø26Ø"&Øð? ð ? Øð? à”| dÑ*ð? ð "œL¨4Ñ/ð? ð ˜4‘Kð	? ð
 �4‰Kð? ð ? ð ? ð ? ðDCð Cð Cð Cð C u¤x¤ñ Cô Cð Cð8ð 8ð 8ð 8ð 8°e´h´oñ 8ô 8ð 8ðvI!ð I!ð I!ð I!ð I! ¤¤ñ I!ô I!ð I!ð\ .2Ø26ð ð  Øð à”| dÑ*ð ð "œL¨4Ñ/ð ð  ð  ð  ðDð ð ð4˜<ð ð ð ð ð
¨e¬kÔ.AÔ.Ið 
ð 
ð 
ð 
ð 
ð 
r#   