Ë
    GêñiB9  ã                   óp  — d Z ddlmZmZ ddlZddlmZ ddlmZm	Z	 ddl
mZmZmZmZ  ed«      Z e«       rdd	lmZ dd
lmZmZmZ erddlmZ ndZ e	j.                  e«      Z G d„ d«      Zdedeeeed   z  f   fd„Z	 d%dej>                  dej>                  dej>                  dej>                  e ej>                  ej>                  f   z  fd„Z!ej>                  e"z  Z#	 	 	 	 	 d&dej>                  de"dz  de e#e#f   dz  dedz  ddf
d„Z$dej>                  de"dej>                  fd„Z%	 	 	 d'dejL                  jN                  dej>                  dej>                  dej>                  d eej>                  df   d!e(dz  d"e(dz  d#ej>                  dz  de ej>                  ej>                  dz  f   fd$„Z)y)(a7  
Partially inspired by torchtune's flex attention implementation

Citation:
@software{torchtune,
  title = {torchtune: PyTorch's finetuning library},
  author = {torchtune maintainers and contributors},
  url = {https//github.com/pytorch/torchtune},
  license = {BSD-3-Clause},
  month = apr,
  year = {2024}
}
é    )ÚOptionalÚUnionN)Úversioné   )Úis_torch_flex_attn_availableÚlogging)Úget_torch_versionÚis_torch_greater_or_equalÚis_torch_less_or_equalÚis_torchdynamo_compilingz2.9.0)Ú_DEFAULT_SPARSE_BLOCK_SIZE)Ú	BlockMaskÚcreate_block_maskÚflex_attention)Ú
AuxRequestc                   óx   ‡ — e Zd ZdZdZdZdZˆ fd„Zej                  j                  d¬«      d„ «       Zd„ Zˆ xZS )ÚWrappedFlexAttentionzh
    We are doing a singleton class so that flex attention is compiled once when it's first called.
    NFc                 ó\   •— | j                   €t        ‰| �	  | «      | _         | j                   S ©N)Ú	_instanceÚsuperÚ__new__)ÚclsÚargsÚkwargsÚ	__class__s      €új/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/transformers/integrations/flex_attention.pyr   zWrappedFlexAttention.__new__D   s'   ø€ Ø�=‰=Ð ä!™G™O¨CÓ0ˆCŒMØ�}‰}Ðó    )Ú	recursivec                 óˆ  — | j                   r|| j                  k7  r§|| _        t        d«      r!t        j                  t
        d¬«      | _        nlt        j                  t        «       «      j                  dk(  r$|r"t        j                  t
        dd¬«      | _        nt        j                  t
        «      | _        d| _         yy)	z>
        Initialize or update the singleton instance.
        ú2.5.1F)Údynamicz2.6.0zmax-autotune-no-cudagraphs)r"   ÚmodeTN)Ú_is_flex_compiledÚtrainingr   ÚtorchÚcompiler   Ú_compiled_flex_attentionr   Úparser	   Úbase_version)Úselfr%   s     r   Ú__init__zWrappedFlexAttention.__init__J   s•   € ð
 ×%Ò%¨°T·]±]Ò)BØ$ˆDŒMÜ% gÔ.Ü05·±¼nÐV[Ô0\�Õ-ô —‘Ô0Ó2Ó3×@Ñ@ÀGÒKÑPXÜ05·±Ü"¨EÐ8Tô1�Õ-ô
 16·±¼nÓ0M�Ô-à%)ˆDÕ"ð *Cr   c                 ó   — | j                   S r   )r(   )r+   s    r   Ú__call__zWrappedFlexAttention.__call__`   s   € Ø×,Ñ,Ð,r   )Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   r$   r(   r   r&   ÚcompilerÚdisabler,   r.   Ú__classcell__)r   s   @r   r   r   ;   sK   ø„ ñð €IØÐØ#Ðôð ‡^�^×Ñ eÐÓ,ñ*ó -ð*ö*-r   r   Ú
return_lseÚreturnr   c                 ó<   — t         rd| rt        d¬«      iS diS d| iS )aU  
    Requests the LSE from flex_attention in a version-agnostic fashion.

    Before torch 2.9, the LSE was requested via the boolean return_lse field. However, starting with
    torch 2.9, an AuxRequest object must be passed via the aux_request field. This method conditionally
    returns the correct form based on the python version.
    Ú
return_auxT)ÚlseNr6   )Ú_TORCH_FLEX_USE_AUXr   )r6   s    r   Úget_flex_attention_lse_kwargsr<   d   s,   € õ Ø±jœj¨TÔ2ÐKÐKÀdÐKÐKà˜*Ð%Ð%r   ÚqueryÚkeyÚvaluec                 óX   — t        «       s t        |«      «       nt        } || ||fi |¤ŽS r   )r   r   r   )r=   r>   r?   r%   r   Úflex_attention_compileds         r   Úcompile_friendly_flex_attentionrB   r   s@   € ô G_ÔF`Ð<Ô2°8Ó<Ô>ÔftÐÙ"ØØØñð ñ	ð r   Úattention_mask_2dÚattention_chunk_sizeÚoffsetsÚ	is_causalr   c                 óD  ‡ ‡‡‡‡‡‡— ‰ j                   \  }}|s|}|s|}|t        z  dz   t        z  }t        j                  j                  j                  ‰ dd||z
  f¬«      Š ‰ j                  }	‰ j                  «       Š|�4‰j                  «       j                  d«      j                  d«      dz
  |z  Šˆ ˆfd„Šˆˆfd„}
ˆ ˆfd„}|s|Šn|€‰n|
Š|�0|d   j                  |	«      Š|d   j                  |	«      Šˆˆˆfd	„}n‰}t        ||d|||	t        d
«       ¬«      S )aG  
    IMPORTANT NOTICE: This function is deprecated in favor of using the mask primitives in `masking_utils.py`,
    and will be removed in a future version without warnings. New code should not use it. It is only kept here
    for BC for now, while models using it are being patched accordingly.

    Create a block (causal) document mask for a batch of sequences, both packed and unpacked.
    Create Block (causal) logic and passing it into :func:`torch.nn.attention.flex_attention.create_block_mask`.
    The resultant BlockMask is a compressed representation of the full (causal) block
    mask. BlockMask is essential for performant computation of flex attention.
    See: https://pytorch.org/blog/flexattention/

    Args:
        attention_mask_2d (torch.Tensor): Attention mask for packed and padded sequences
        of shape (batch_size, total_seq_len). e.g.

        For unpacked sequence:
        [[1, 1, 1, 1, 0, 0, 0],
         [1, 1, 1, 1, 1, 0, 0]]

        For packed sequence:
        [[1, 1, 1, 2, 2, 2, 0],
         [1, 1, 2, 2, 2, 3, 3]]

    Returns:
        BlockMask
    é   r   )r?   ÚpadNéÿÿÿÿc                 óT   •— ||k\  }‰	| |f   ‰	| |f   k(  }‰| |f   dkD  }||z  |z  }|S )zü
        Defines the logic of a block causal mask by combining both a standard causal mask
        and a block diagonal document mask.
        See :func:`~torchtune.modules.attention_utils.create_block_causal_mask`
        for an illustration.
        r   © )
Ú	batch_idxÚhead_idxÚq_idxÚkv_idxÚcausal_maskÚdocument_maskÚpadding_maskÚ
final_maskrC   Údocument_idss
           €€r   Úcausal_mask_modz4make_flex_block_causal_mask.<locals>.causal_mask_mod¾   sV   ø€ ð ˜v‘oˆØ$ Y°Ð%5Ñ6¸,ÀyÐRXÐGXÑ:YÑYˆØ(¨°EÐ)9Ñ:¸QÑ>ˆØ  <Ñ/°-Ñ?ˆ
ØÐr   c                 óB   •— ‰| |f   ‰| |f   k(  } ‰| |||«      }||z  S )zU
        Combines the chunk mask with the causal mask for chunked attention.
        rL   )rM   rN   rO   rP   Ú
chunk_maskÚcausal_doc_maskrV   Ú
chunk_idxss         €€r   Úchunk_causal_mask_modz:make_flex_block_causal_mask.<locals>.chunk_causal_mask_modË   s>   ø€ ð   	¨5Ð 0Ñ1°ZÀ	È6Ð@QÑ5RÑRˆ
Ù)¨)°X¸uÀfÓMˆØ˜OÑ+Ð+r   c                 óD   •— ‰| |f   ‰| |f   k(  }‰| |f   dkD  }||z  }|S )zp
        Utilizes default attention mask to enable encoder and encoder-decoder
        attention masks.
        r   rL   )	rM   rN   rO   rP   rR   rS   rT   rC   rU   s	          €€r   Údefault_mask_modz5make_flex_block_causal_mask.<locals>.default_mask_modÓ   sH   ø€ ð
 % Y°Ð%5Ñ6¸,ÀyÐRXÐGXÑ:YÑYˆà(¨°FÐ):Ñ;¸aÑ?ˆØ! MÑ1ˆ
ØÐr   c                 ó.   •— |‰z   }|‰z   } ‰| |||«      S r   rL   )	rM   rN   rO   rP   Úoffset_qÚ	offset_kvÚ	kv_offsetÚmask_mod_maybe_combinedÚq_offsets	         €€€r   Úmask_modz-make_flex_block_causal_mask.<locals>.mask_modç   s(   ø€ Ø˜xÑ'ˆHØ Ñ*ˆIÙ*¨9°hÀÈ)ÓTÐTr   r!   )rd   ÚBÚHÚQ_LENÚKV_LENÚdeviceÚ_compile)ÚshapeÚflex_default_block_sizer&   ÚnnÚ
functionalrI   ri   ÚcloneÚfill_ÚcumsumÚtor   r   )rC   rD   Úquery_lengthÚ
key_lengthrE   rF   Ú
batch_sizeÚtotal_seq_lenÚpad_lenri   r[   r]   rd   rV   rZ   rU   ra   rb   rc   s   `            @@@@@@r   Úmake_flex_block_causal_maskrx   ˆ   sC  þ€ ðD !2× 7Ñ 7Ñ€J�ÙØ"ˆ
ÙØ$ˆàÔ5Ñ5¸Ñ:Ô>UÑU€GÜŸ™×+Ñ+×/Ñ/Ð0AÈÐQRÐT[Ð^hÑThÐPiÐ/ÓjÐØ×%Ñ%€FØ$×*Ñ*Ó,€LàÐ'à"×(Ñ(Ó*×0Ñ0°Ó3×:Ñ:¸2Ó>ÀÑBÐH\Ñ]ˆ
õõ,õ	ñ Ø"2Ñà5IÐ5Q¡/ÐWlÐàÐØ˜1‘:—=‘= Ó(ˆØ˜A‘J—M‘M &Ó)ˆ	÷	Uð
 +ˆäØØ
Ø
ØØØä+¨GÓ4Ð4ô	ð 	r   Úhidden_statesÚn_repc                 óª   — | j                   \  }}}}|dk(  r| S | dd…dd…ddd…dd…f   j                  |||||«      } | j                  |||z  ||«      S )zÔ
    This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
    num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
    rH   N)rk   ÚexpandÚreshape)ry   rz   ÚbatchÚnum_key_value_headsÚslenÚhead_dims         r   Ú	repeat_kvr‚   ú   so   € ð
 2?×1DÑ1DÑ.€EÐ  hØ�‚zØÐØ!¢!¢Q¨ªa²Ð"2Ñ3×:Ñ:¸5ÐBUÐW\Ð^bÐdlÓm€MØ× Ñ  Ð(;¸eÑ(CÀTÈ8ÓTÐTr   ÚmoduleÚattention_maskÚscalingÚsoftcapÚs_auxc           
      óŽ  ‡‡— |j                  dd«      dkD  rt        d«      ‚d }	d Št        |t        «      r|}	n|Š‰�‰d d …d d …d d …d |j                  d   …f   Šˆˆfd„}
d}|j                  d   }||dz
  z  dk7  rTt        ||j                  d   |j                  d   z  «      }t        ||j                  d   |j                  d   z  «      }d	}|j                  d
«      }|j                  j                  dk7  }|s|�t        d«      ‚t        |||f|
|	|||| j                  dœt        |«      ¤Ž}|rêt        r|\  }}|j                  }n|\  }}|j                  |j                  «      }|�´|j                  \  }}}}|j                  dddd«      j!                  |||d«      }|j#                  d«      }t%        j&                  t%        j(                  ||gd¬«      dd¬«      }t%        j*                  ||z
  «      }||z  }|j                  |j                  «      }n|}d }|j-                  dd«      j/                  «       }||fS )NÚdropoutg        r   z›`flex_attention` does not support `dropout`. Please use it with inference only (`model.eval()`) or turn off the attention dropout in the respective config.éþÿÿÿc                 óh   •— ‰�‰t        j                  | ‰z  «      z  } ‰�| ‰|   d   |   |   z   } | S )Nr   )r&   Útanh)ÚscorerM   rN   rO   rP   Ú
score_maskr†   s        €€r   Ú	score_modz)flex_attention_forward.<locals>.score_mod!  sK   ø€ ØÐØœeŸj™j¨°©Ó9Ñ9ˆEØÐ!Ø˜J yÑ1°!Ñ4°UÑ;¸FÑCÑCˆEð ˆr   TrH   FÚkernel_optionsÚcpuzhAttention sinks cannot be run on CPU with flex attention. Please switch to a different device, e.g. CUDA)r�   Ú
block_maskÚ
enable_gqaÚscaler�   r%   rJ   )Údim)r•   Úkeepdimr   )ÚgetÚ
ValueErrorÚ
isinstancer   rk   r‚   ri   ÚtyperB   r%   r<   r;   r:   rr   ÚdtypeÚviewr|   Ú	unsqueezer&   Ú	logsumexpÚcatÚexpÚ	transposeÚ
contiguous)rƒ   r=   r>   r?   r„   r…   r†   r‡   r   r’   r�   r“   Únum_local_query_headsr�   r6   Úflex_attention_outputÚattention_outputÚauxr:   ru   Ú	num_headsÚ	seq_len_qÚ_ÚsinksÚlse_expandedÚcombined_lseÚrenorm_factorrŽ   s         `                    @r   Úflex_attention_forwardr®     s~  ù€ ð ‡z�z�)˜SÓ! AÒ%Üðaó
ð 	
ð
 €JØ€JÜ�.¤)Ô,Ø#‰
à#ˆ
àÐØ¢¢1¢a¨¨3¯9©9°R©=¨Ð 8Ñ9ˆ
õð €JØ!ŸK™K¨™NÐð 	Ð!6¸Ñ!:Ñ;ÀÒAÜ˜˜UŸ[™[¨™^¨s¯y©y¸©|Ñ;Ó<ˆÜ˜% §¡¨Q¡°5·;±;¸q±>Ñ!AÓBˆØˆ
à—Z‘ZÐ 0Ó1€Nà—‘×"Ñ" eÑ+€Já˜%Ð+ÜØvó
ð 	
ô <ØØØðð ØØØØ%ð —‘ñô (¨
Ó
3ñÐñ  õ Ø$9Ñ!Ð˜cØ—'‘'‰Cà$9Ñ!Ð˜cð �f‰f�U—[‘[Ó!ˆàÐà2B×2HÑ2HÑ/ˆJ˜	 9¨aØ—J‘J˜q " a¨Ó+×2Ñ2°:¸yÈ)ÐUVÓWˆEð
 Ÿ=™=¨Ó,ˆLÜ Ÿ?™?¬5¯9©9°lÀEÐ5JÐPRÔ+SÐY[ÐeiÔjˆLô "ŸI™I l°\Ñ&AÓBˆMØ/°-Ñ?ÐØ/×2Ñ2°5·;±;Ó?Ñà0ÐØˆà'×1Ñ1°!°QÓ7×BÑBÓDÐØ˜SÐ Ð r   )F)NNNNT)NNN)*r2   Útypingr   r   r&   Ú	packagingr   Úutilsr   r   Úutils.import_utilsr	   r
   r   r   r;   Ú!torch.nn.attention.flex_attentionr   rl   r   r   r   r   Ú
get_loggerr/   Úloggerr   ÚboolÚdictÚstrr<   ÚTensorÚtuplerB   ÚintÚOffsetrx   r‚   rm   ÚModuleÚfloatr®   rL   r   r   ú<module>r¿      s<  ðñ÷8 #ã Ý ç 9÷ó ñ 0°Ó8Ð ñ  Ô!Ýgß^Ñ^áÞ@àˆ
ð 
ˆ×	Ñ	˜HÓ	%€÷&-ñ &-ðR&¨dð &°t¸CÀÈÐQ]ÑH^ÑA^Ð<^Ñ7_ó &ð$ ñ	Ø�<‰<ðà	�‰ðð �<‰<ðð ‡\�\�E˜%Ÿ,™,¨¯©Ð4Ñ5Ñ5óð$ 
�‰˜Ñ	€ð (,ØØØ,0Ø!ñoØ—|‘|ðoà ™*ðoð
 �6˜6�>Ñ" TÑ)ðoð �d‰{ðoð óoðd	U˜UŸ\™\ð 	U°#ð 	U¸%¿,¹,ó 	Uð$ !Ø Ø!%ñg!Ø�H‰H�O‰Oðg!à�<‰<ðg!ð 
�‰ðg!ð �<‰<ð	g!ð
 ˜%Ÿ,™,¨Ð3Ñ4ðg!ð �T‰\ðg!ð �T‰\ðg!ð �<‰<˜$Ñðg!ð ˆ5�<‰<˜Ÿ™¨Ñ,Ð,Ñ-ôg!r   