Team Ai
Modelpublic

KwangHwi/quantization

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes7downloads
llama_bidirectional_model.cpython-312.pyc107 linesDownload Raw Back to __pycache__
1�

2͂Xjw#���dZddlZddlZddlmZmZddlmZddlm	Z	ddl3mZmZddl
mZeje�Z	ddlmZd	Zej0ej2�j4Zej0ej8�j4ZdevZd
evZGd�de	�Z Gd�de�Z!y#e$rdd4lmZdZY�zwxYw)a5Bidirectional Llama model for embedding tasks.6 7This module provides a modified LlamaModel that uses bidirectional (non-causal)8attention, suitable for generating embeddings where each token should attend9to all other tokens in the sequence.10 11Supports transformers version 4.44 and above with a unified forward() implementation.12 13Version compatibility notes:14    - transformers 4.47: Setting _attn_implementation in __init__ had no effect due to15      attention initialization order16    - transformers 4.48+: Attention refactor (transformers#35235) activated the17      _attn_implementation setting, which defaulted to "eager" instead of "sdpa"18    - transformers < 4.53: LlamaModel has _update_causal_mask method that can be overridden19    - transformers 4.53+: _update_causal_mask removed; masking moved to masking_utils module,20      necessitating a full forward() override for custom attention masks21    - transformers < 4.54: Decoder layer returns tuple, uses past_key_value (singular)22    - transformers 4.54-4.55: Decoder layer returns tensor, uses past_key_value (singular)23    - transformers 4.56+: Decoder layer returns tensor, uses past_key_values (plural),24      DynamicCache accepts config parameter25    - transformers 5.0+: Has native create_bidirectional_mask in masking_utils26�N)�Cache�DynamicCache)�BaseModelOutputWithPast)�LlamaConfig)�LlamaDecoderLayer�27LlamaModel)�logging)�create_bidirectional_maskT)�_prepare_4d_attention_maskF�past_key_values�configc�8��eZdZdZdZ	ddededdf�fd�
Z�xZS)	�LlamaBidirectionalConfigzPConfiguration for LlamaBidirectionalModel with pooling and temperature settings.�
llama_bidirec�pooling�temperature�returnNc�@��||_||_t�|�di|��y)a28        Initialize bidirectional Llama configuration.29 30        Args:31            pooling: Pooling strategy for embeddings ("avg", "cls", "last", etc.)32            temperature: Temperature scaling for embeddings33            **kwargs: Additional arguments passed to LlamaConfig34        N�)rr�super�__init__)�selfrr�kwargs�	__class__s    ���/home/namlt/Desktop/Huyhq21/LLMs/embedding_model/quantization/model/model_vai_embed_1B_multiSFT_v6.1_86.32/output_merge_multiSFT_v6_1/llama_bidirectional_model.pyrz!LlamaBidirectionalConfig.__init__?s$������&���
���"�6�"�)�avgg�?)	�__name__�35__module__�__qualname__�__doc__�36model_type�str�floatr�
__classcell__�rs@rrr:s2���Z� �J�:=�
#��
#�16�
#�	
�
#�
#rrc�R��eZdZdZeZdeddf�fd�Zdejdejdzdejdzfd�Z37							dd	ejdzdejdzd38ejdzdedzdejdzd
ejdzdedzdefd�Z�xZS)�LlamaBidirectionalModela�39    LlamaModel modified to use bidirectional (non-causal) attention.40 41    In standard Llama, each token can only attend to previous tokens (causal attention).42    This model removes that restriction, allowing each token to attend to all tokens43    in the sequence, which is useful for embedding tasks.44 45    The key modifications are:46        1. Setting is_causal=False on all attention layers47        2. Using a bidirectional attention mask instead of causal mask48    r
rNc�h��t�|�|�|jD]}d|j_�y)NF)rr�layers�	self_attn�	is_causal)rr
�layerrs   �rrz LlamaBidirectionalModel.__init__^s*���
���� ��[�[�E�(-�E�O�O�%�!r�input_embeds�attention_maskc���|�ytrt|j||��St|jdd�dk(r|dk(j	�}|r|SdSt||j�S)a�49        Create bidirectional attention mask.50 51        Args:52            input_embeds: Input embeddings tensor of shape (batch_size, seq_len, hidden_size)53            attention_mask: Optional 2D attention mask of shape (batch_size, seq_len)54                where 1 indicates tokens to attend to and 0 indicates masked tokens55 56        Returns:57            4D attention mask suitable for the attention implementation, or None58            if no masking is needed59        N)r
r.r/�_attn_implementation�flash_attention_2r)�_HAS_NATIVE_BIDIRECTIONAL_MASKr60r
�getattr�anyr�dtype)rr.r/�has_masked_tokenss    r�_create_bidirectional_maskz2LlamaBidirectionalModel._create_bidirectional_maskcsw��"�!��)�,��{�{�)�-��
��4�;�;� 6��=�AT�T�!/�1�!4� 9� 9� ;��%6�>�@�D�@�)�.�,�:L�:L�M�Mr�	input_ids�position_idsr�
inputs_embeds�cache_position�	use_cachec��|du|duzrtd��|�|j|�}|r)|�'trt|j��}n61t�}|�F|�|j�nd}	t
j|	|	|jdz|j��}|�|jd�}|j||�}62|}|j||�}|63||||d�}
tr||
d<n||
d	<|jd|jjD]#}||fi|
��}t!|t"�r|d}�"|}�%|j%|�}t'||�64�S)a�65        Forward pass with bidirectional attention.66 67        Args:68            input_ids: Input token IDs of shape (batch_size, seq_len)69            attention_mask: Attention mask of shape (batch_size, seq_len)70            position_ids: Position IDs for rotary embeddings71            past_key_values: Cached key/value states for incremental decoding72            inputs_embeds: Pre-computed input embeddings (alternative to input_ids)73            cache_position: Position indices for cache updates74            use_cache: Whether to return cached key/value states75            **kwargs: Additional arguments passed to decoder layers76 77        Returns:78            BaseModelOutputWithPast containing last_hidden_state and past_key_values79        Nz:You must specify exactly one of input_ids or inputs_embeds)r
r�)�device)r/r:r=r<�position_embeddingsr�past_key_value)�last_hidden_stater)�80ValueError�embed_tokens�_DYNAMIC_CACHE_ACCEPTS_CONFIGrr
�get_seq_length�torch�arange�shaper@�	unsqueezer8�81rotary_emb�_USE_PLURAL_CACHE_PARAMr*�num_hidden_layers�82isinstance�tuple�normr)rr9r/r:rr;r<r=r�past_seen_tokens�bidirectional_mask�
hidden_statesrA�layer_kwargs�
decoder_layer�
layer_outputss                r�forwardzLlamaBidirectionalModel.forward�s���6
���-�t�";�<��L��
�� � �-�-�i�8�M���0�,�".�d�k�k�"B��".�.���!�4C�4O��.�.�0�UV�
�#�\�\� � �=�#6�#6�q�#9�9�$�+�+��N���)�3�3�A�6�L�!�<�<��>�83��&�
�"�o�o�m�\�J��841�(�"�,�#6�85��#�.=�L�*�+�-<�L�)�*�!�[�[�)H�4�;�;�+H�+H�I�M�)�-�H�<�H�M��-��/� -�a� 0�
� -�
�J��	�	�-�0�
�&�+�+�86�	87r)NNNNNNN)rrr r!r�config_classrrrH�Tensorr8�88LongTensorr�FloatTensor�boolrrXr%r&s@rr(r(Os���89�,�L�.�{�.�t�.�90#N��l�l�#N����t�+�#N�91����	�	#N�N.2�.2�04�(,�26�26�!%�Z92��#�#�d�*�Z93����t�+�Z94��&�&��-�	Z95�96���Z97��(�(�4�/�
Z98��(�(�4�/�Z99��$�;�Z100�101!�Z102rr()"r!�inspectrH�transformers.cache_utilsrr�transformers.modeling_outputsr�-transformers.models.llama.configuration_llamar�(transformers.models.llama.modeling_llamarr�transformers.utilsr	�103get_loggerr�logger�transformers.masking_utilsr104r3�ImportError�%transformers.modeling_attn_mask_utilsr�	signaturerX�105parameters�_decoder_forward_paramsr�_dynamic_cache_init_paramsrMrFrr(rrr�<module>rms����0��8�A�E�R�&�	��	�	�H�	%��+�D�%)�"�,�'�+�+�,=�,E�,E�F�Q�Q��.�W�.�.�|�/D�/D�E�P�P��,�/F�F�� (�,F� F��#�{�#�*S106�j�S107��I�+�P�%*�"�+�s�B3�3
C�C