Ë
    �…j–  ã                   óL   — d dl Z d dlZ e j                  e«      Z G d„ d«      Zy)é    Nc                   óì   — e Zd ZdZeddefd„«       Zed„ «       Zed„ «       Zede	e	e
j                        fd„«       Zedde	e
j                     d	efd
„«       Zedde	e	e
j                        fd„«       Zy)ÚPastKeyValuesHelperzEHelper functions to process past key values for encoder-decoder modelÚpresentc                 óÈ   — g }g }t        | «      D ]L  }|j                  |r
d|› �d|› �gn	d|› �d|› �g«       |j                  |r
d|› �d|› �gn	d|› �d|› �g«       ŒN ||z   S )	NÚpresent_key_self_Úpresent_value_self_Úpast_key_self_Úpast_value_self_Úpresent_key_cross_Úpresent_value_cross_Úpast_key_cross_Úpast_value_cross_)ÚrangeÚextend)Ú
num_layersr   Úpast_self_namesÚpast_cross_namesÚis        úe/root/aria/tools/markitdown-venv/lib/python3.12/site-packages/onnxruntime/transformers/past_helper.pyÚget_past_namesz"PastKeyValuesHelper.get_past_names   s³   € àˆØÐÜ�zÓ"ò 
	ˆAØ×"Ñ"áð % Q CÐ(Ð,?À¸sÐ*CÑDà& q cÐ*Ð.>¸q¸cÐ,BÐCôð
 ×#Ñ#áð & a SÐ)Ð-AÀ!ÀÐ+EÑFà'¨ sÐ+Ð/@ÀÀÐ-DÐEõð
	ð Ð!1Ñ1Ð1ó    c                 óÔ   — g }g }t        | «      D ]S  \  }}t        |«      dk(  sJ dt        |«      › �«       ‚|\  }}}}|j                  ||g«       |j                  ||g«       ŒU ||fS )a³  Split present state from grouped by layer to grouped by self/cross attention.
        Before: (past_key_self_0, past_value_self_0, past_key_cross_0, past_value_cross_0), (past_key_self_1, past_value_self_1, past_key_cross_1, past_value_cross_1), ...
        After: (past_key_self_0, past_value_self_0, past_key_self_1, past_value_self_1, ...), (past_key_cross_0, past_value_cross_0, past_key_cross_1, past_value_cross_1, ...)

        é   ú!Expected to have four items. Got ©Ú	enumerateÚlenr   )	Úpresent_key_valuesÚpresent_selfÚpresent_crossÚ_iÚpresent_layer_iÚpresent_key_selfÚpresent_value_selfÚpresent_key_crossÚpresent_value_crosss	            r   Úgroup_by_self_or_crossz*PastKeyValuesHelper.group_by_self_or_cross"   sŸ   € ð ˆØˆÜ#,Ð-?Ó#@ò 		KÑˆB�Ü�Ó'¨1Ò,ÐhÐ0QÔRUÐVeÓRfÐQgÐ.hÓhÐ,ð  ñØ Ø"Ø!Ø#à×ÑÐ!1Ð3EÐ FÔGØ× Ñ Ð"3Ð5HÐ!IÕJð		Kð ˜]Ð*Ð*r   c                 óh   ‡ ‡— t        ‰ «      d‰z  k(  sJ ‚t        ˆˆ fd„t        ‰«      D «       «      S )a©  Reorder past state from grouped by self/cross attention to grouped by layer.
        Before: past_key_self_0, past_value_self_0, past_key_self_1, past_value_self_1, ..., past_key_cross_0, past_value_cross_0, past_key_cross_1, past_value_cross_1, ...
        After: (past_key_self_0, past_value_self_0, past_key_cross_0, past_value_cross_0), (past_key_self_1, past_value_self_1, past_key_cross_1, past_value_cross_1),
        r   c              3   ó~   •K  — | ]4  }‰d |z     ‰d |z  dz      ‰d ‰z  d |z  z      ‰d ‰z  d |z  z   dz      g–— Œ6 y­w)é   é   N© )Ú.0r   r   Úpasts     €€r   ú	<genexpr>z5PastKeyValuesHelper.group_by_layer.<locals>.<genexpr>>   sf   øè ø€ ò 
ð ð �Q˜‘U‘Ø�Q˜‘U˜Q‘Y‘Ø�Q˜‘^ a¨!¡eÑ+Ñ,Ø�Q˜‘^ a¨!¡eÑ+¨aÑ/Ñ0ô	ñ
ùs   ƒ:=)r   Útupler   )r.   r   s   ``r   Úgroup_by_layerz"PastKeyValuesHelper.group_by_layer7   s<   ù€ ô �4‹y˜A 
™NÒ*Ð*Ð*Üô 
ô ˜:Ó&ô
ó 
ð 	
r   Úpast_key_valuesc                 ó¬   — d}t        | «      dz  }t        t        | «      dz  «      D ])  }d|z  }|| |   | |dz      | ||z      | ||z   dz      ffz  }Œ+ |S )aè  Categorize present_key_values from self and cross attention to layer by layer.

        Reorder past state from grouped by self/cross attention to grouped by layer.
        Before: past_key_self_0, past_value_self_0, past_key_self_1, past_value_self_1, ...,
                past_key_cross_0, past_value_cross_0, past_key_cross_1, past_value_cross_1, ...
        After: (past_key_self_0, past_value_self_0, past_key_cross_0, past_value_cross_0),
                (past_key_self_1, past_value_self_1, past_key_cross_1, past_value_cross_1),

        Args:
            present_key_values: From past_key_values of a model (group by self and cross attention)

        Returns:
            past_tuples: present key and values grouped by layer.
        r,   r*   r   r+   )r   r   )r2   Úpast_tuplesÚhalf_idxr   Úidxs        r   Úback_group_by_layerz'PastKeyValuesHelper.back_group_by_layerH   sŒ   € ð  ˆÜ�Ó'¨1Ñ,ˆÜ”s˜?Ó+¨qÑ0Ó1ò 		ˆAØ�a‘%ˆCØà# CÑ(Ø# C¨!¡GÑ,Ø# H¨s¡NÑ3Ø# H¨s¡N°QÑ$6Ñ7ð	ðñ ‰Kð		ð Ðr   r   Úconcatc                 óâ   — g }g }t        | «      D ]S  \  }}t        |«      dk(  sJ dt        |«      › �«       ‚|\  }}}}	|j                  ||g«       |j                  ||	g«       ŒU |r||z   S ||fS )a˜  Categorize present_key_values into self and cross attention.

        Split present state from grouped by layer to grouped by self/cross attention.
        Before: (past_key_self_0, past_value_self_0, past_key_cross_0, past_value_cross_0),
                (past_key_self_1, past_value_self_1, past_key_cross_1, past_value_cross_1), ...
        After: (past_key_self_0, past_value_self_0, past_key_self_1, past_value_self_1, ...),
                (past_key_cross_0, past_value_cross_0, past_key_cross_1, past_value_cross_1, ...)

        Args:
            present_key_values: From past_key_values of a model (group by layer)
            concat: If concat self attention with cross attention key/value to return

        Returns:
            present_self (Tuple[torch.Tensor]): present key and values from self attention
            present_cross (Tuple[torch.Tensor]): present key and values from cross attention
        r   r   r   )
r   r8   r   r    Ú_r"   r#   r$   r%   r&   s
             r   Úgroup_by_self_and_crossz+PastKeyValuesHelper.group_by_self_and_crossf   s©   € ð$ ,.ˆØ,.ˆÜ"+Ð,>Ó"?ò 	KÑˆAˆÜ�Ó'¨1Ò,ÐhÐ0QÔRUÐVeÓRfÐQgÐ.hÓhÐ,Ø[jÑXÐÐ0Ð2CÐEXØ×ÑÐ!1Ð3EÐ FÔGØ× Ñ Ð"3Ð5HÐ!IÕJð		Kñ
 Ø -Ñ/Ð/à Ð.Ð.r   c                 óH  — g }|rt        | «      dz  n
t        | «      }|sdnd}t        |«      D ],  }|j                  d|› �d|› �fD �cg c]  }||z   ‘Œ	 c}«       Œ. t        |«      D ],  }|j                  d|› �d|› �fD �cg c]  }||z   ‘Œ	 c}«       Œ. |S c c}w c c}w )zÆProcess input names of model wrapper.

        Args:
            past_key_values: Consider `self` and `cross` past_key_values

        Returns:
            names (List[string]): input names
        r   Úpast_Úpresent_Ú	key_self_Úvalue_self_Ú
key_cross_Úvalue_cross_)r   r   r   )r2   ÚencoderÚnamesr   Úprefixr   Úss          r   Úget_input_namesz#PastKeyValuesHelper.get_input_names„   sÀ   € ð ˆÙ29”S˜Ó)¨QÒ.¼sÀ?Ó?Sˆ
Ù '‘¨ZˆÜ�zÓ"ò 	UˆAØ�L‰L°¸1¸#¨À+ÈaÈSÐ@QÐ.RÖS¨˜& 1›*ÒSÕTð	Uä�zÓ"ò 	WˆAØ�L‰L°¸A¸3Ð/?À<ÐPQÈsÐASÐ.TÖU¨˜& 1›*ÒUÕVð	Wàˆùò TùâUs   Á	B
ÂB
N)F)T)Ú__name__Ú
__module__Ú__qualname__Ú__doc__ÚstaticmethodÚboolr   r'   r1   r0   ÚtorchÚTensorr7   r;   rG   r,   r   r   r   r      sÄ   „ ÙOàñ2¨Dò 2ó ð2ð  ñ+ó ð+ð( ñ
ó ð
ð  ð¨U°5¸¿¹Ñ3FÑ-Gò ó ðð: ñ/°E¸%¿,¹,Ñ4Gð /ÐQUò /ó ð/ð: ñ¨¨u°U·\±\Ñ/BÑ)Cò ó ñr   r   )ÚloggingrN   Ú	getLoggerrH   Úloggerr   r,   r   r   ú<module>rS      s+   ðó ã à	ˆ×	Ñ	˜8Ó	$€÷Gò Gr   