Ë
    îÍ:j‰5  ã                   óH  — 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  e«       rddlmZ dd	lmZmZmZ  e	j(                  e«      Z G d
„ d«      Z	 d!dej0                  dej0                  dej0                  deej0                  eej0                  ej0                  f   f   fd„Zeej0                  ef   Z	 	 	 	 	 d"dej0                  dee   deeeef      dee   ddf
d„Zdej0                  dedej0                  fd„Z	 	 	 	 d#dej@                  jB                  dej0                  dej0                  dej0                  deej0                  df   dee"   dee"   deej0                     deej0                     deej0                  eej0                     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)Ú_torch_versionÚis_torch_less_or_equalÚis_torchdynamo_compiling)Ú_DEFAULT_SPARSE_BLOCK_SIZE)Ú	BlockMaskÚcreate_block_maskÚflex_attentionc                   ó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      €ú}/home/mcse/projects/srt_converter/srt-converter-venv/lib/python3.12/site-packages/transformers/integrations/flex_attention.pyr   zWrappedFlexAttention.__new__7   s'   ø€ Ø�=‰=Ð ä!™G™O¨CÓ0ˆCŒMØ�}‰}Ðó    )Ú	recursivec                 ó€  — | j                   r|| j                  k7  r£|| _        t        d«      r!t        j                  t
        d¬«      | _        nht        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__=   s�   € ð
 ×%Ò%¨°T·]±]Ò)BØ$ˆDŒMÜ% gÔ.Ü05·±¼nÐV[Ô0\�Õ-ô —‘œ~Ó.×;Ñ;¸wÒFÉ8Ü05·±Ü"¨EÐ8Tô1�Õ-ô
 16·±¼nÓ0M�Ô-à%)ˆDÕ"ð *Cr   c                 ó   — | j                   S r   )r&   )r)   s    r   Ú__call__zWrappedFlexAttention.__call__S   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   ÚqueryÚkeyÚvalueÚreturnc                 óX   — t        «       s t        |«      «       nt        } || ||fi |¤ŽS r   )r   r   r   )r4   r5   r6   r#   r   Úflex_attention_compileds         r   Úcompile_friendly_flex_attentionr:   W   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   )r6   Ú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_maskr;   Ú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.
        rD   )rE   rF   rG   rH   Ú
chunk_maskÚcausal_doc_maskrN   Ú
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   rD   )	rE   rF   rG   rH   rJ   rK   rL   r;   rM   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   rD   )	rE   rF   rG   rH   Ú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   )r\   ÚBÚHÚQ_LENÚKV_LENÚdeviceÚ_compile)ÚshapeÚflex_default_block_sizer$   ÚnnÚ
functionalrA   ra   ÚcloneÚfill_ÚcumsumÚtor   r
   )r;   r<   Úquery_lengthÚ
key_lengthr=   r>   Ú
batch_sizeÚtotal_seq_lenÚpad_lenra   rS   rU   r\   rN   rR   rM   rY   rZ   r[   s   `            @@@@@@r   Úmake_flex_block_causal_maskrp   m   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)
    r@   N)rc   ÚexpandÚreshape)rq   rr   ÚbatchÚnum_key_value_headsÚslenÚhead_dims         r   Ú	repeat_kvrz   ß   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Ú	head_maskÚs_auxc	                 óH  ‡‡‡— ‰�t         j                  d«       |	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        |||||
||||| j                  ¬«
      }|r·|\  }}|j                  |j                  «      }|�™|j                  \  }}}}|j                  d	dd	d	«      j                  |||d	«      }|j!                  d«      }t#        j$                  t#        j&                  ||gd¬«      dd¬«      }t#        j(                  ||z
  «      }||z  }n|}d }|j+                  d	d«      j-                  «       }||fS )Nzm`flex_attention` does not support `head_mask`. Please set your attention to `eager` if you want this feature.Ú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                 óŽ   •— ‰�‰t        j                  | ‰z  «      z  } ‰�| ‰|   d   |   |   z   } ‰�| ‰|   |   d   d   z   } | S )Nr   )r$   Útanh)ÚscorerE   rF   rG   rH   r   Ú
score_maskr~   s        €€€r   Ú	score_modz)flex_attention_forward.<locals>.score_mod  so   ø€ ØÐØœeŸj™j¨°©Ó9Ñ9ˆEØÐ!Ø˜J yÑ1°!Ñ4°UÑ;¸FÑCÑCˆEØÐ Ø˜I iÑ0°Ñ:¸1Ñ=¸aÑ@Ñ@ˆEð ˆr   Tr@   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‰   Ú
return_lser#   rB   )Údim)r�   Úkeepdimr   )ÚloggerÚwarning_onceÚgetÚ
ValueErrorÚ
isinstancer   rc   rz   ra   Útyper:   r#   rj   ÚdtypeÚviewrt   Ú	unsqueezer$   Ú	logsumexpÚcatÚexpÚ	transposeÚ
contiguous)r{   r4   r5   r6   r|   r}   r~   r   r€   r   r‹   rˆ   rŒ   Únum_local_query_headsr‰   rŽ   Úflex_attention_outputÚattention_outputÚlserm   Ú	num_headsÚ	seq_len_qÚ_ÚsinksÚlse_expandedÚcombined_lseÚrenorm_factorr‡   s         ``                   @r   Úflex_attention_forwardrª   ë   sT  ú€ ð ÐÜ×ÑØ{ô	
ð ‡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ó
ð 	
ô <ØØØØØØØØ%ð Ø—‘ôÐñ Ø 5ÑÐ˜#Ø�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Ø/°-Ñ?Ñà0ÐØˆà'×1Ñ1°!°QÓ7×BÑBÓDÐØ˜SÐ Ð r   )F)NNNNT)NNNN)$r0   Útypingr   r   r$   Ú	packagingr   Úutilsr   r   Úutils.import_utilsr	   r
   r   Ú!torch.nn.attention.flex_attentionr   rd   r   r   r   Ú
get_loggerr-   r‘   r   ÚTensorÚtupler:   ÚintÚOffsetÚboolrp   rz   re   ÚModuleÚfloatrª   rD   r   r   ú<module>r¸      s  ðñ÷8 #ã Ý ç 9ß aÑ añ  Ô!Ýgß^Ñ^ð 
ˆ×	Ñ	˜HÓ	%€÷&-ñ &-ðZ ñ	Ø�<‰<ðà	�‰ðð �<‰<ðð ˆ5�<‰<˜˜uŸ|™|¨U¯\©\Ð9Ñ:Ð:Ñ;óð$ 
ˆu�|‰|˜SÐ Ñ	!€ð +/ØØØ/3Ø $ñoØ—|‘|ðoà" 3™-ðoð
 �e˜F F˜NÑ+Ñ,ðoð ˜‰~ðoð óoðd	U˜UŸ\™\ð 	U°#ð 	U¸%¿,¹,ó 	Uð$  $Ø#Ø(,Ø$(ñe!Ø�H‰H�O‰Oðe!à�<‰<ðe!ð 
�‰ðe!ð �<‰<ð	e!ð
 ˜%Ÿ,™,¨Ð3Ñ4ðe!ð �e‰_ðe!ð �e‰_ðe!ð ˜Ÿ™Ñ%ðe!ð �E—L‘LÑ!ðe!ð ˆ5�<‰<˜ %§,¡,Ñ/Ð/Ñ0ôe!r   