§
    ‚Štj§  ã                   óÄ  — d dl Z ddlmZ ddlmZ de j        j        de j        de j        de j        d	e j        dz  d
ede j        de j        ee	e j        f         z  de
de
ee	e
f         z  de j        dz  dee j        df         fd„Ze j        j        de j        j        de j        de j        de j        d
ede j        dee
e
f         de j        de j        fd„¦   «         ZdS )é    Né   )ÚPagedAttentionCache)Ú!lazy_import_paged_flash_attentionÚmoduleÚqÚkÚvÚattention_maskÚcacheÚcu_seq_lens_qÚcu_seq_lens_kÚmax_seqlen_qÚmax_seqlen_kÚblock_tableÚreturnc                 ó‚  — t          | j        j        ¦  «        \  }}t          | dd¦  «        sdn| j        dz
  df}|dk    rdnd}t          |t          ¦  «        r||         }|	|         }	|
�€*|                     ||| j        |d	         |d
         ¬¦  «        \  }}d|v rd| 	                    d¦  «        ini } || 
                    dd¦  «                             d¦  «                             ¦   «         |                     ¦   «         |                     ¦   «         |                     t          j        ¦  «        |                     t          j        ¦  «                             ¦   «         ||	f| j        d|dœ|¤Ž}t          |t$          ¦  «        r|d         }n%d|v r
d|d         ini }t'          | ||||||||
f	i |¤Ž}|dfS )ap  Performs the forward pass of attention with paged key-value cache. This function handles the cache updates and
    performs the attention computation. For decode-only batches (when block_table is provided), uses
    `flash_attn_with_kvcache` for fused attention + cache update. Otherwise uses `flash_attn_varlen_func`.
    See the [paged attention guide](https://huggingface.co/docs/transformers/en/paged_attention) for more details.

    Args:
        q: (1, nheads, total_q, headdim), where total_q = total number of query tokens in the batch.
        k: (1, nheads_k, total_k, headdim), where total_k = total number of key tokens in the batch.
        v: (1, nheads_k, total_k, headdim), where total_k = total number of key tokens in the batch.
        cu_seq_lens_q: (batch_size + 1,), dtype torch.int32. The cumulative sequence lengths
           of the sequences in the batch, used to index into q.
        cu_seq_lens_k: (batch_size + 1,), dtype torch.int32. The cumulative sequence lengths
           of the sequences in the batch, used to index into kv.
        max_seqlen_q: int. Maximum query sequence length in the batch.
        max_seqlen_k: int. Maximum key sequence length in the batch.
        block_table: (num_groups, batch_size, max_blocks_per_seq), dtype int32. Block table for paged KV cache.
            If provided, uses flash_attn_with_kvcache for fused attention + cache update. For each request, the block
            table is a vector of size (max_blocks_per_seq,) with indices indicating the physical location of the cache
            to read from and write to. The kernel, using the cache_seqlens for that request, knows how much cache to
            read and dispatches the read using the block table. Same for the write. If a request has fewer than
            max_blocks_per_seq blocks, the block table is padded with -1s to indicate that the block is not allocated.
    Úsliding_windowF)éÿÿÿÿr   é   r   Úfull_attentionÚsliding_attentionNÚ
read_indexÚwrite_index)Ú
key_statesÚvalue_statesÚ	layer_idxr   r   Ús_auxr   T)Úsoftmax_scaleÚcausalÚwindow_size)r   ÚconfigÚ_attn_implementationÚgetattrr   Ú
isinstanceÚdictÚupdater   ÚgetÚ	transposeÚsqueezeÚ
contiguousÚtoÚtorchÚint32ÚcloneÚscalingÚtupleÚ_paged_decode_forward)r   r   r   r	   r
   r   r   r   r   r   r   ÚkwargsÚflash_attn_varlen_funcÚflash_attn_with_kvcacher   Ú
layer_typeÚcustom_kwargsÚattn_outputÚflash_kwargss                      úc/var/www/html/CA-Chatbot/venv/lib/python3.11/site-packages/transformers/integrations/flash_paged.pyÚpaged_attention_forwardr:      s  € õJ 7XØŒÔ*ñ7ô 7Ñ3ÐÐ3õ
 &-¨VÐ5EÀuÑ%MÔ%MÐq�X�XÐTZÔTiÐlmÑTmÐopÐSq€NØ%3°xÒ%?Ð%?Ð!Ð!ÐEX€JÝ�-¥Ñ&Ô&ð 0Ø% jÔ1ˆØ# JÔ/ˆð Ñà�|Š|ØØØÔ&Ø˜lÔ+Ø˜}Ô-ð ñ 
ô 
‰ˆˆ1ð ;BÀVÐ:KÐ:K˜ &§*¢*¨WÑ"5Ô"5Ð6Ð6ÐQSˆØ,Ð,Ø�KŠK˜˜1ÑÔ×%Ò% aÑ(Ô(×3Ò3Ñ5Ô5Ø�LŠL‰NŒNØ�LŠL‰NŒNØ×Ò�Uœ[Ñ)Ô)Ø×Ò�Uœ[Ñ)Ô)×/Ò/Ñ1Ô1ØØð
ð !œ.ØØ&ð
ð 
ð ð
ð 
ˆõ �k¥5Ñ)Ô)ð 	)Ø% aœ.ˆKøð 6=ÀÐ5FÐ5F˜ ¨¤Ð1Ð1ÈBˆÝ+Ø�A�q˜!˜U M°>ÐCZÐ\gð
ð 
Økwð
ð 
ˆð ˜ÐÐó    r   c	                 óX  — |j         | j                 \  }
}|j        |                              d|j        |j        |j        ¦  «        }|j        |                              d|j        |j        |j        ¦  «        }|                     dddd¦  «         	                    ¦   «         }|                     dddd¦  «         	                    ¦   «         }|                     dddd¦  «         	                    ¦   «         }| 
                    d¦  «        }|d|dz   …         |d|…         z
  dz
                       t          j        ¦  «        }||
         |	|                     |¦  «        <    |d	||||||| j        d|dœ	|	¤Ž}t!          |t"          ¦  «        r|d         }|                     d¦  «        S )
zaDecode fast path using flash_attn_with_kvcache. Disabled because FA3 has issue with tracing this.r   r   r   r   é   NT)	r   Úk_cacheÚv_cacher   r	   Úcache_seqlensr   r   r    © )Úlayer_index_to_group_indicesr   Ú	key_cacheÚviewÚ
block_sizeÚnum_key_value_headsÚhead_dimÚvalue_cacheÚpermuter*   Úsizer+   r,   r-   Úget_block_table_keyr/   r$   r0   r)   )r   r   r   r	   r   r   r   r4   r   r8   Ú	group_idxÚlayer_idx_in_groupr>   r?   Ú
batch_sizer@   r7   s                    r9   r1   r1   [   s¼  € ð %*Ô$FÀvÔGWÔ$XÑ!€IÐ!àŒoÐ0Ô1×6Ò6°r¸5Ô;KÈUÔMfÐhmÔhvÑwÔw€GØÔÐ 2Ô3×8Ò8Ø
ˆEÔ˜eÔ7¸¼ñô €Gð 	
�	Š	�!�Q˜˜1ÑÔ×(Ò(Ñ*Ô*€AØ	�	Š	�!�Q˜˜1ÑÔ×(Ò(Ñ*Ô*€AØ	�	Š	�!�Q˜˜1ÑÔ×(Ò(Ñ*Ô*€Að —’˜‘”€JØ" 1 z°A¡~Ð#5Ô6¸À{È
À{Ô9SÑSÐVWÑW×[Ò[Õ\aÔ\gÑhÔh€MàGRÐS\ÔG]€L�×*Ò*Ð+BÑCÔCÑDà)Ð)ð Ø
ØØØ
Ø
Ø#Ø”nØØ"ðð ð ðð €Kõ �+�uÑ%Ô%ð %Ø! !”nˆà×Ò˜qÑ!Ô!Ð!r;   )r,   Úgeneration.continuous_batchingr   Úmodeling_flash_attention_utilsr   ÚnnÚModuleÚTensorr%   ÚstrÚintr0   r:   ÚcompilerÚdisabler1   rA   r;   r9   ú<module>rX      s¯  ðØ €€€à @Ð @Ð @Ð @Ð @Ð @Ø NÐ NÐ NÐ NÐ NÐ NðQØŒHŒOðQà„|ðQð „|ðQð „|ð	Qð
 ”L 4Ñ'ðQð ðQð ”<ðQð ”< $ s¨E¬LÐ'8Ô"9Ñ9ðQð ðQð ˜˜S #˜XœÑ&ðQð ” Ñ$ðQð ˆ5Œ<˜ÐÔðQð Qð Qð Qðh „Ôð."ØŒHŒOð."à„|ð."ð „|ð."ð „|ð	."ð
 ð."ð ”<ð."ð ˜#˜s˜(”Oð."ð ”ð."ð „\ð."ð ."ð ."ñ Ôð."ð ."ð ."r;   