§
    ‚Štjå9  ã                   óü  — 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j&        j'        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j        dz  de ej        ej        dz  f         fd'„Z)dS )+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                   ó|   ‡ — e Zd ZdZdZdZdZˆ fd„Ze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                 ól   •— | j         €&t          ¦   «                              | ¦  «        | _         | j         S ©N)Ú	_instanceÚsuperÚ__new__)ÚclsÚargsÚkwargsÚ	__class__s      €úf/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/integrations/flex_attention.pyr   zWrappedFlexAttention.__new__D   s*   ø€ ØŒ=Ð å!™GœGŸOšO¨CÑ0Ô0ˆCŒMØŒ}Ðó    )Ú	recursivec                 ó€  — | j         r|| j        k    r«|| _        t          d¦  «        r!t          j        t
          d¬¦  «        | _        nkt          j        t          ¦   «         ¦  «        j
        dk    r$|r"t          j        t
          dd¬¦  «        | _        nt          j        t
          ¦  «        | _        d| _         dS dS )	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Ð)BØ$ˆDŒMÝ% gÑ.Ô.ð NÝ05´½nÐV[Ð0\Ñ0\Ô0\�Ô-Ð-õ ”Õ0Ñ2Ô2Ñ3Ô3Ô@ÀGÒKÐKÐPXÐKÝ05´Ý"¨EÐ8Tð1ñ 1ô 1�Ô-Ð-õ
 16´½nÑ0MÔ0M�Ô-à%)ˆDÔ"Ð"Ð"ð *CÐ)Br   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   ;   s“   ø€ € € € € ðð ð €IØÐØ#Ððð ð ð ð ð „^×Ò eÐÑ,Ô,ð*ð *ñ -Ô,ð*ð*-ð -ð -ð -ð -ð -ð -r   r   Ú
return_lseÚreturnr   c                 óD   — t           rd| rt          d¬¦  «        n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   s8   € õ ð LØ°jÐJ�j¨TÐ2Ñ2Ô2Ð2ÀdÐKÐKà˜*Ð%Ð%r   FÚqueryÚkeyÚvaluec                 óp   — 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`ÔF`ÐtÐ<Õ2°8Ñ<Ô<Ñ>Ô>Ð>ÕftÐØ"Ð"ØØØðð ð ð	ð ð r   TÚattention_mask_2dÚattention_chunk_sizeÚoffsetsÚ	is_causalr   c                 óv  ‡ ‡‡‡‡‡‡— ‰ j         \  }}|s|}|s|}|t          z  dz   t          z  }t          j        j                             ‰ dd||z
  f¬¦  «        Š ‰ j        }	‰                      ¦   «         Š|�@‰                     ¦   «                              d¦  «         	                    d¦  «        dz
  |z  Šˆ ˆfd„Šˆˆfd„}
ˆ ˆfd„}|s|Šn|€‰n|
Š|�>|d          
                    |	¦  «        Š|d          
                    |	¦  «        Šˆˆˆ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                 ól   •— ||k    }‰	| |f         ‰	| |f         k    }‰| |f         dk    }||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                 óV   •— ‰| |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Ë   sC   ø€ ð   	¨5Ð 0Ô1°ZÀ	È6Ð@QÔ5RÒRˆ
Ø)˜/¨)°X¸uÀfÑMÔMˆØ˜OÑ+Ð+r   c                 óZ   •— ‰| |f         ‰| |f         k    }‰| |f         dk    }||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                 ó4   •— |‰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Ô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   ˆ   sê  øøøøøøø€ ðD !2Ô 7Ñ€J�Øð #Ø"ˆ
Øð %Ø$ˆàÕ5Ñ5¸Ñ:Õ>UÑU€GÝœÔ+×/Ò/Ð0AÈÐQRÐT[Ð^hÑThÐPiÐ/ÑjÔjÐØÔ%€FØ$×*Ò*Ñ,Ô,€LàÐ'à"×(Ò(Ñ*Ô*×0Ò0°Ñ3Ô3×:Ò:¸2Ñ>Ô>ÀÑBÐH\Ñ]ˆ
ðð ð ð ð ð ð,ð ,ð ,ð ,ð ,ð ,ð	ð 	ð 	ð 	ð 	ð 	ð ð mØ"2ÐÐà5IÐ5Q / /ÐWlÐàÐØ˜1”:—=’= Ñ(Ô(ˆØ˜A”J—M’M &Ñ)Ô)ˆ	ð	Uð 	Uð 	Uð 	Uð 	Uð 	Uð 	Uð 	Uð
 +ˆåØØ
Ø
ØØØå+¨GÑ4Ô4Ð4ð	ñ 	ô 	ð 	r   Úhidden_statesÚn_repc                 ó¸   — | j         \  }}}}|dk    r| S | dd…dd…ddd…dd…f                              |||||¦  «        } |                      |||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‚   ú   s„   € ð
 2?Ô1DÑ.€EÐ  hØ�‚z€zØÐØ! ! ! ! Q Q Q¨¨a¨a¨a°°°Ð"2Ô3×:Ò:¸5ÐBUÐW\Ð^bÐdlÑmÔm€MØ× Ò  Ð(;¸eÑ(CÀTÈ8ÑTÔTÐTr   ÚmoduleÚattention_maskÚscalingÚsoftcapÚs_auxÚposition_biasc	           
      ó¸  ‡‡‡— |	                      dd¦  «        dk    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k    rTt          ||j        d         |j        d         z  ¦  «        }t          ||j        d         |j        d         z  ¦  «        }d	}|	                      d
¦  «        }|j        j        dk    }|s|�t          d¦  «        ‚t          |||f||
|||| j	        dœt          |¦  «        ¤Ž}|rèt          r|\  }}|j        }n|\  }}|                     |j        ¦  «        }|�²|j        \  }}}}|                     dddd¦  «                             |||d¦  «        }|                     d¦  «        }t%          j        t%          j        ||gd¬¦  «        dd¬¦  «        }t%          j        ||z
  ¦  «        }||z  }|                     |j        ¦  «        }n|}d }|                     dd¦  «                             ¦   «         }||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                 ó    •— ‰�‰t          j        | ‰z  ¦  «        z  } ‰�| ‰|         d         |         |         z   } ‰�| ‰||||f         z   } | S )Nr   )r&   Útanh)ÚscorerM   rN   rO   rP   rˆ   Ú
score_maskr†   s        €€€r   Ú	score_modz)flex_attention_forward.<locals>.score_mod"  sj   ø€ ØÐØ�eœj¨°©Ñ9Ô9Ñ9ˆEØÐ!Ø˜J yÔ1°!Ô4°UÔ;¸FÔCÑCˆEØÐ$Ø˜M¨)°X¸uÀfÐ*LÔMÑMˆ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�   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 1 1 a a a¨¨3¬9°R¬=¨Ð 8Ô9ˆ
ð
ð 
ð 
ð 
ð 
ð 
ð 
ð €JØ!œK¨œNÐð 	Ð!6¸Ñ!:Ñ;ÀÒAÐAÝ˜˜Uœ[¨œ^¨s¬y¸¬|Ñ;Ñ<Ô<ˆÝ˜% ¤¨Q¤°5´;¸q´>Ñ!AÑBÔBˆØˆ
à—Z’ZÐ 0Ñ1Ô1€Nà”Ô" eÒ+€Jàð 
˜%Ð+ÝØvñ
ô 
ð 	
õ <ØØØðð ØØØØ%ð ”ðð õ (¨
Ñ
3Ô
3ðð Ðð  ð õ ð 	:Ø$9Ñ!Ð˜cØ”'ˆCˆCà$9Ñ!Ð˜cð �fŠf�U”[Ñ!Ô!ˆàÐà2BÔ2HÑ/ˆJ˜	 9¨aØ—J’J˜q " a¨Ñ+Ô+×2Ò2°:¸yÈ)ÐUVÑWÔWˆEð
 Ÿ=š=¨Ñ,Ô,ˆLÝ œ?­5¬9°lÀEÐ5JÐPRÐ+SÑ+SÔ+SÐY[ÐeiÐjÑjÔjˆLõ "œI l°\Ñ&AÑBÔBˆMØ/°-Ñ?ÐØ/×2Ò2°5´;Ñ?Ô?Ðøà0ÐØˆà'×1Ò1°!°QÑ7Ô7×BÒBÑDÔDÐØ˜SÐ Ð r   )F)NNNNT)NNNN)*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Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9Ð 9ðð ð ð ð ð ð ð ð ð ð ð ð 0Ð/°Ñ8Ô8Ð ð  ÐÑ!Ô!ð ØgÐgÐgÐgÐgÐgØ^Ð^Ð^Ð^Ð^Ð^Ð^Ð^Ð^Ð^àð Ø@Ð@Ð@Ð@Ð@Ð@Ð@àˆ
ð 
ˆÔ	˜HÑ	%Ô	%€ð&-ð &-ð &-ð &-ð &-ñ &-ô &-ð &-ðR&¨dð &°t¸CÀÈÐQ]ÔH^ÑA^Ð<^Ô7_ð &ð &ð &ð &ð$ ð	ð ØŒ<ðà	Œðð Œ<ðð „\�E˜%œ,¨¬Ð4Ô5Ñ5ðð ð ð ð$ 
Œ˜Ñ	€ð (,ØØØ,0Ø!ðoð oØ”|ðoà ™*ðoð
 �6˜6�>Ô" TÑ)ðoð �d‰{ðoð ðoð oð oð oðd	U˜Uœ\ð 	U°#ð 	U¸%¼,ð 	Uð 	Uð 	Uð 	Uð$ !Ø Ø!%Ø)-ðj!ð j!ØŒHŒOðj!àŒ<ðj!ð 
Œðj!ð Œ<ð	j!ð
 ˜%œ,¨Ð3Ô4ðj!ð �T‰\ðj!ð �T‰\ðj!ð Œ<˜$Ñðj!ð ”< $Ñ&ðj!ð ˆ5Œ<˜œ¨Ñ,Ð,Ô-ðj!ð j!ð j!ð j!ð j!ð j!r   