Ë
    Hêñij0  ã                  ó8  — d Z ddlmZ ddlZddlmZ ddlmZ ddlm	Z	 ddl
mZ dd	lmZmZ erdd
lmZ ddlmZ i Ze ed¬«       G d„ d«      «       «       Z G d„ d«      Z ed«      Zdd„Z	 	 	 	 	 	 	 	 dd„Zddd„Z ej4                  «       Zdd„Zdddœd„Zy) z°
Contains the logic for automatic additional output capture with our forward decorators.
This mostly describe the hooks used and the logic to make capture thread/context safe.
é    )ÚannotationsN)Ú
ContextVar)Ú	dataclass©Úwraps)ÚTYPE_CHECKINGé   )Úis_torchdynamo_compilingÚrequires)Únné   ©ÚPreTrainedModel)Útorch)Úbackendsc                  óF   — e Zd ZU dZded<   dZded<   dZded	<   dZded
<   y)ÚOutputRecordera  
    Configuration for recording outputs from a model via hooks.

    Attributes:
        target_class (Type): The class (e.g., nn.Module) to which the hook will be attached.
        index (Optional[int]): If the output is a tuple/list, optionally record only at a specific index.
        layer_name (Optional[str]): Name of the submodule to target (if needed), e.g., "transformer.layer.3.attn".
        class_name (Optional[str]): Name of the class to which the hook will be attached. Could be the suffix of class name in some cases.
    ztype[nn.Module]Útarget_classr   ÚintÚindexNú
str | NoneÚ
layer_nameÚ
class_name)Ú__name__Ú
__module__Ú__qualname__Ú__doc__Ú__annotations__r   r   r   © ó    úe/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/transformers/utils/output_capturing.pyr   r   '   s,   … ñð "Ó!Ø€Eˆ3ƒNØ!€J�
Ó!Ø!€J�
Ô!r    r   c                  ó(   — e Zd ZdZd„ Zd„ Zd„ Zd„ Zy)ÚCompileableContextVara¡  
    Convenience wrapper around a ContextVar for usage with `torch.compile`.
    This behaves exactly as a `ContextVar`, except when compilation is triggered in which case it behaves as a simple
    global variable. This is useful as `torch.compile` cannot trace the `get` method of `ContextVar`. This however means
    that the access to the underlying variable is not thread-safe when compilation is triggered.
    c                óD   — t        |d ¬«      | _        d | _        d| _        y )N)ÚdefaultF)r   Úcontext_varÚ
global_varÚ	compiling)ÚselfÚnames     r!   Ú__init__zCompileableContextVar.__init__B   s   € Ü% d°DÔ9ˆÔØˆŒØˆ�r    c                óf   — | j                   r| j                  S | j                  j                  «       S ©N)r(   r'   r&   Úget)r)   s    r!   r.   zCompileableContextVar.getG   s(   € à�>Š>Ø—?‘?Ð"à×#Ñ#×'Ñ'Ó)Ð)r    c                ój   — t        «       r|| _        d| _        y | j                  j	                  |«      S )NT)r
   r'   r(   r&   Úset)r)   Úvalues     r!   r0   zCompileableContextVar.setN   s0   € Ü#Ô%Ø#ˆDŒOØ!ˆDŒNØà×#Ñ#×'Ñ'¨Ó.Ð.r    c                ót   — | j                   s|€d | _        d| _         y | j                  j                  |«       y )NF)r(   r'   r&   Úreset)r)   Útokens     r!   r3   zCompileableContextVar.resetV   s/   € Ø�>Š>˜U˜]Ø"ˆDŒOØ"ˆD�Nà×Ñ×"Ñ" 5Õ)r    N)r   r   r   r   r+   r.   r0   r3   r   r    r!   r#   r#   :   s   „ ñòò
*ò/ó*r    r#   Úoutput_collectorc                ó6   ‡‡— ˆˆfd„}| j                  |«       y)zaInstall the forward hook needed to capture the output described by `key` and `index` in `module`.c                ó6  •— t         j                  «       }|�‰|j                  «       vry ‰dk(  r(t        |‰   «      dk(  r|‰   j	                  |d   «       t        |t        «      s|‰   j	                  |«       y |‰   �|‰   j	                  |‰   «       y y )NÚhidden_statesr   )Ú_active_collectorr.   ÚkeysÚlenÚappendÚ
isinstanceÚtuple)ÚmoduleÚargsÚoutputÚcollected_outputsr   Úkeys       €€r!   Úoutput_capturing_hookz<install_output_capturing_hook.<locals>.output_capturing_hooke   s    ø€ ä-×1Ñ1Ó3ÐàÐ$¨Ð3D×3IÑ3IÓ3KÑ(KØà�/Ò!¤cÐ*;¸CÑ*@Ó&AÀQÒ&FØ˜cÑ"×)Ñ)¨$¨q©'Ô2Ü˜&¤%Ô(Ø˜cÑ"×)Ñ)¨&Õ1Ø�E‰]Ð&Ø˜cÑ"×)Ñ)¨&°©-Õ8ð 'r    N)Úregister_forward_hook)r?   rC   r   rD   s    `` r!   Úinstall_output_capturing_hookrF   b   s   ù€ õ9ð × Ñ Ð!6Õ7r    c                ó°  — ddl m} | j                  «       D ]6  \  }}t        ||«      st	        ||› d|› �|«       Œ%t        ||› d|› �¬«       Œ8 |D ]‚  \  }}|j                  �t        | |j                  «      s)|j                  €Œ5|j                  |j                  «      sŒQ|j                  �|j                  |vrŒlt        | ||j                  «       Œ„ y)aÖ  
    Recursively install all output capturing hooks on all submodules of `parent_module`.
    Note that we need to use this recursive approach instead of simply iterating over all modules, because we want
    to respect the `capture_tasks` of all individual submodels (`PreTrainedModel` instances) in the graph. That is, once
    we reach a submodel in the graph, its children should use this submodel's `capture_tasks`, but other parts of the graph
    should not.
    r   r   ú.)ÚprefixN)Úmodeling_utilsr   Únamed_childrenr=   Úrecursively_install_hooksÚ"install_all_output_capturing_hooksr   r   Úendswithr   rF   r   )Úparent_moduleÚmodule_nameÚcapture_tasksr   r*   r?   rC   Úspecss           r!   rL   rL   v   sÜ   € õ 1ð &×4Ñ4Ó6ò W‰ˆˆfä˜& /Ô2Ü% f°°¸Q¸t¸fÐ.EÀ}ÕUô /¨vÀÀÈQÈtÈfÐ>UÖVðWð $ò K‰
ˆˆUà×ÑÐ*¬z¸-È×I[ÑI[Ô/\Ø×ÑÑ(¨[×-AÑ-AÀ%×BRÑBRÕ-Sà×ÑÐ+°×0@Ñ0@ÈÑ0SØÜ)¨-¸¸e¿k¹kÕJñKr    c                óÆ  — t         j                  t        | j                  «      «      xs i }g }|j	                  «       D ]€  \  }}t        |t        «      s|g}|D ]c  }t        |t        «      s>d|v rdnd}t        |t        «      sdn|}t        |t        «      s|nd}	t        |	||¬«      }|j                  ||f«       Œe Œ‚ |�|nd}t        | ||«       t        | dd«       y)	zÙ
    Install the output recording hooks on all the modules in `model`. This will take care of correctly dispatching
    the `_can_record_outputs` property of each individual submodels in case of composite models.
    r8   r   r	   N)r   r   r   Ú Ú!_output_capturing_hooks_installedT)Ú_CAN_RECORD_REGISTRYr.   ÚstrÚ	__class__Úitemsr=   Úlistr   r<   rL   Úsetattr)
ÚmodelrI   Úcapture_flagsrQ   rC   Úlayer_specsrR   r   r   r   s
             r!   rM   rM   –   sä   € ô )×,Ñ,¬S°·±Ó-AÓBÒHÀb€Mà€MØ)×/Ñ/Ó1ò 	/Ñˆˆ[Ü˜+¤tÔ,Ø&˜-ˆKØ ò 	/ˆEÜ˜e¤^Ô4Ø,°Ñ3™¸�Ü)3°E¼3Ô)?™TÀU�
Ü,6°u¼cÔ,B™uÈ�Ü&°LÈÐZdÔe�Ø× Ñ  # u Õ.ñ	/ð	/ð Ð)‰V¨r€FÜ˜e V¨]Ô;äˆEÐ6¸Õ=r    c                óš   — t        | dd«      ryt        5  t        | dd«      r
	 ddd«       yt        | «       ddd«       y# 1 sw Y   yxY w)zà
    Check if the model already has output capturing hooks installed, and install them if it is not already the
    case.
    Note that this is thread-safe, in case 2 (or more) threads want to install them concurrently.
    rU   FN)ÚgetattrÚ_hook_installation_lockrM   )r\   s    r!   Úmaybe_install_capturing_hooksrb   ¶   sR   € ô ˆuÐ9¸5ÔAØä	 ñ 2ô �5Ð=¸uÔEØ÷	2ð 2ô 	+¨5Ô1÷2÷ 2ñ 2ús   •A­AÁA
T)Útie_last_hidden_statesc               ó&   ‡— ˆfd„}| � || «      S |S )aÿ  
    Decorator to intercept specific layer outputs through hooks. The hooks are installed only once and lazily,
    the first time output capture is requested with the `output_xxx` kwargs/config.
    The implementation is fully context/thread safe, except when using `torch.compile`, as dynamo is unable to trace
    through `ContextVar` methods.

    Args:
        tie_last_hidden_states (`bool`, *optional*, defaults to `True`):
            Whether to overwrite `out.hidden_states[-1]` with the `out.last_hidden_state`.
            This is true for all language models and should be toggled off only if
            `out.hidden_states[-1]` has to be the hidden state before last layer norm, which
            is needed for some vision models (e.g. CLIP, SigLIP)
    c                ó2   •‡ — t        ‰ «      ˆ ˆfd„«       }|S )Nc                óx  •— |j                  dt        | j                  dd«      «      }t        j	                  t        | j                  «      «      xs i }|D �ci c]3  }d|› �|j	                  d|› �t        | j                  d|› �d«      «      “Œ5 }}d|v r*|j	                  dt        | j                  dd«      «      |d<   d|v r*|j	                  dt        | j                  dd«      «      |d	<   |j                  «       D ��ci c]  \  }}|sŒ	|j                  dd
«      g “Œ }}}t        |«      dkD  rt        | «       t        j                  |«      }		  ‰| g|¢­i |¤Ž}
t        j                  |	«       |D �]  }|dk(  r€‰snkt        |
d«      r*||   d d ||<   ||   j                  |
j                   «       n5t        |
d«      r)||   d d ||<   ||   j                  |
j"                  «       t%        ||   «      |
|<   Œ‰|dk(  rht'        ||   t(        «      rCt        ||   «      dk(  r2t%        ||   dd d…   «      |
|<   t%        ||   dd d…   «      |
d|z   <   Œät%        ||   «      |
|<   Œöt%        ||   «      |
|<   �Œ	 |du r|
j+                  «       }
|
S c c}w c c}}w # t        j                  |	«       w xY w)NÚreturn_dictTÚoutput_FÚcross_attentionsÚoutput_attentionsÚoutput_cross_attentionsÚmask_decoder_attentionsÚoutput_mask_decoder_attentionsrT   r   r8   Úvision_hidden_stateséÿÿÿÿÚlast_hidden_stateÚ
attentionsr   r	   Úcross_)Úpopr`   ÚconfigrV   r.   rW   rX   rY   Úreplacer;   rb   r9   r0   r3   Úhasattrr<   rn   rp   r>   r=   rZ   Úto_tuple)r)   r@   Úkwargsrg   Úcapturable_flagsÚkÚrecordable_keysÚvrB   Úoutput_tokenÚoutputsrC   Úfuncrc   s               €€r!   Úwrapperz4capture_outputs.<locals>.wrapped_fn.<locals>.wrapperÙ   s  ø€ ð !Ÿ*™* ]´G¸D¿K¹KÈÐX\Ó4]Ó^ˆKô  4×7Ñ7¼¸D¿N¹NÓ8KÓLÒRÐPRÐð *öàð ˜!˜�˜vŸz™z¨G°A°3¨-¼ÀÇÁÐPWÐXYÐWZÈmÐ]bÓ9cÓdÑdðˆOð ð
 "Ð%5Ñ5Ø=C¿Z¹ZØ'¬°·±Ð>QÐSXÓ)Yó>�Ð 9Ñ:ð )Ð,<Ñ<ØDJÇJÁJØ'¬°·±Ð>QÐSXÓ)YóE�Ð @ÑAð KZ×J_ÑJ_ÓJa× gÁ$À!ÀQÒef §¡¨9°bÓ!9¸2Ñ!=Ð gÐÑ gäÐ$Ó%¨Ò)Ü-¨dÔ3ä,×0Ñ0Ð1BÓCˆLð6Ù˜tÐ5 dÒ5¨fÑ5�ô "×'Ñ'¨Ô5ð )ó A�Ø˜/Ò)Ù1ØÜ  Ð*@ÔAØ1BÀ3Ñ1GÈÈÐ1LÐ)¨#Ñ.Ø)¨#Ñ.×5Ñ5°g×6RÑ6RÕSÜ  Ð*=Ô>Ø1BÀ3Ñ1GÈÈÐ1LÐ)¨#Ñ.Ø)¨#Ñ.×5Ñ5°g×6OÑ6OÔPä#(Ð):¸3Ñ)?Ó#@�G˜C’LØ˜LÒ(ä!Ð"2°3Ñ"7¼Ô>Ä3ÐGWÐX[ÑG\ÓC]ÐabÒCbÜ',Ð->¸sÑ-CÀAÀDÀqÀDÑ-IÓ'J˜ ™Ü27Ð8IÈ#Ñ8NÈqÈtÐRSÈtÑ8TÓ2U˜ ¨3¡Ò/ä',Ð->¸sÑ-CÓ'D˜ šä#(Ð):¸3Ñ)?Ó#@�G˜C“Lð)Að, ˜eÑ#Ø!×*Ñ*Ó,�àˆNùòoùó !høô "×'Ñ'¨Õ5ús   Á8JÄ
JÄJÅJ" Ê"J9r   )r   r€   rc   s   ` €r!   Ú
wrapped_fnz#capture_outputs.<locals>.wrapped_fnØ   s!   ù€ Ü	ˆt‹ô=	ó 
ð=	ð~ ˆr    r   )r   rc   r�   s    ` r!   Úcapture_outputsr‚   É   s#   ø€ ôAðF ÐÙ˜$ÓÐØÐr    )r?   ú	nn.ModulerC   rW   r   r   ÚreturnÚNone)rO   rƒ   rP   rW   rQ   z list[tuple[str, OutputRecorder]]r„   r…   r-   )r\   r   rI   r   r„   r…   )r\   r   r„   r…   )r   Ú
__future__r   Ú	threadingÚcontextvarsr   Údataclassesr   Ú	functoolsr   Útypingr   Úimport_utilsr
   r   r   r   rJ   r   rV   r   r#   r9   rF   rL   rM   ÚLockra   rb   r‚   r   r    r!   ú<module>rŽ      sÌ   ðñõ
 #ã Ý "Ý !Ý Ý  ç <ñ Ýå0ð Ð ð Ù	�:Ô÷"ð "ó ó ð"÷"!*ñ !*ñJ *Ð*<Ó=Ð ó8ð(KØðKØ+.ðKØ?_ðKà	óKô@>ð: )˜)Ÿ.™.Ó*Ð ó2ð&T¸õ Tr    