ó
    pyüi@#  ã                   ór  • S SK Jr  S SKJr  S SKJrJr  S SKJr  S SK	r	S SK
Jr  S SKJr  SS	KJrJrJrJr   " S
 S5      r\ " S S5      5       rS\S\4S jrS\S\S\S\4S jrS'S\S\S\S\4S jjrS\S\S\S\4S jr S(S\	R6                  S\\   S\\   S\SS4
S  jjrS!\S"\S#\S$\S%\S\\   4S& jrg))é    )ÚOrderedDict)Ú	dataclass)ÚceilÚlog2)ÚAnyN)ÚPretrainedConfig)ÚContinuousBatchingConfigé   )ÚFutureRequestStateÚRequestStateÚRequestStatusÚloggerc                   óÜ   • \ rS rSrSrS\SS4S jrSS jrS\\S	4   S\	R                  R                  S-  4S
 jrSS\SS4S jjrS\\S	4   S\	R                  R                  SS4S jrSrg)ÚCudaGraphBufferé   z>A fixed-size dict for CUDA graphs with LRU eviction when full.Úmax_sizeÚreturnNc                 óV   • US::  a  [        SU 35      eXl        [        5       U l        g )Nr   z#max_size must be positive, but got )Ú
ValueErrorr   r   Ú_storage)Úselfr   s     Ún/home/mande/repo/quber/.venv/lib/python3.13/site-packages/transformers/generation/continuous_batching/utils.pyÚ__init__ÚCudaGraphBuffer.__init__   s*   € Ø�q‹=ÜÐBÀ8À*ÐMÓNÐNØ ŒÜLWËMˆ�ó    c                 óT   • U R                   nSU l         U R                  SS9  Xl         g )Nr
   T)Úsilent)r   Úplan_for_new_graph)r   Úoriginal_max_sizes     r   Ú__del__ÚCudaGraphBuffer.__del__$   s)   € Ø ŸM™MÐØˆŒØ×Ñ tÐÑ,Ø)�r   Úkey.c                 óx   • U R                   R                  U5      nUb  U R                   R                  U5        U$ ©N)r   ÚgetÚmove_to_end©r   r"   Úgraphs      r   Ú	get_graphÚCudaGraphBuffer.get_graph*   s3   € Ø—‘×!Ñ! #Ó&ˆØÑØ�M‰M×%Ñ% cÔ*Øˆr   r   c                 ó.  • [        U R                  5      U R                  :¼  ar  U R                  R                  SS9u  p#U(       d  [        R
                  " SU< 35        UR                  5         [        U R                  5      U R                  :¼  a  Mq  g g )NF)Úlastz!Evicting graph for evicted_key = )Úlenr   r   Úpopitemr   ÚinfoÚreset)r   r   Úevicted_keyÚevicted_graphs       r   r   Ú"CudaGraphBuffer.plan_for_new_graph0   sl   € Ü�$—-‘-Ó  D§M¡MÓ1Ø)-¯©×)>Ñ)>ÀEÐ)>Ð)JÑ&ˆKÞÜ—’Ð@°+Ñ1AÐBÔCØ×ÑÔ!ô	 �$—-‘-Ó  D§M¡M×1r   r(   c                 ó@   • U R                  5         X R                  U'   g r$   )r   r   r'   s      r   Ú	set_graphÚCudaGraphBuffer.set_graph7   s   € à×ÑÔ!Ø"�‰�cÒr   )r   r   )r   N)F)Ú__name__Ú
__module__Ú__qualname__Ú__firstlineno__Ú__doc__Úintr   r    ÚtupleÚtorchÚcudaÚ	CUDAGraphr)   Úboolr   r5   Ú__static_attributes__© r   r   r   r      s‰   † ÙHðZ ð Z¨ô Zô*ð˜U 3¨ 8™_ð °·±×1EÑ1EÈÑ1Lô ñ"¨ð "¸$õ "ð#˜U 3¨ 8™_ð #°U·Z±Z×5IÑ5Ið #Èd÷ #r   r   c                   ó@   • \ rS rSr% SrSr\\S'   Sr\\S'   S	S jr	Sr
g)
ÚWorkloadHintsé=   zRA tiny dataclass containing hints to help choose good continuous batching defaultsr   Úmax_prompt_lengthÚmax_generated_lengthNc                 óø   • U R                   (       ai  U R                  (       aW  UR                  cI  U R                   U R                  -   n[        [	        X!R
                  -  5      5      S-   nX3S-  -   Ul        gggg)z*Resolves the config using the given hints.Nr
   é   )rG   rH   Úmax_blocks_per_requestr<   r   Ú
block_size)r   Ú	cb_configÚmax_sequence_lengthÚblocks_per_requests       r   Úresolve_using_hintsÚ!WorkloadHints.resolve_using_hintsE   sw   € ð ×!×! d×&?×&?Ø×/Ñ/Ñ7Ø&*×&<Ñ&<¸t×?XÑ?XÑ&XÐ#Ü%(¬Ð.A×DXÑDXÑ.XÓ)YÓ%ZÐ]^Ñ%^Ð"Ø3EÐ^_ÑI_Ñ3`�	Õ0ð 8ð '@Ð!r   rC   )rM   r	   r   N)r7   r8   r9   r:   r;   rG   r<   Ú__annotations__rH   rP   rB   rC   r   r   rE   rE   =   s!   ‡ á\àÐ�sÓØ !Ð˜#Ó!÷ar   rE   Úconfigr   c                 ó    • U R                   S;   $ )z:Checks if attention mask is needed for the given (config).)zpaged|eagerz
paged|sdpa)Ú_attn_implementation)rS   s    r   Úattn_mask_is_neededrV   O   s   € à×&Ñ&Ð*GÑGÐGr   ÚsizeÚinterval_sizeÚ	max_valuec                 óX   • US::  a  U$ U S:”  a  [        X-  5      U-  OUn[        X25      $ )zQReturn the smallest multiple of (interval_size) >= (size), capped at (max_value).r   )r   Úmin)rW   rX   rY   Úpaddeds       r   Úpad_to_intervalr]   T   s5   € à˜ÓØÐØ;?À!»8ŒT�$Ñ&Ó'¨-Ò7È€FÜˆvÓ!Ð!r   ÚvalueÚ	min_valuec                 ó„   • [        U [        SU5      5      n S[        [        [        U 5      5      5      -  n[	        X15      $ )z�Return the smallest power of 2 >= (value), capped at (max_value). If a minimum value is provided, the value is at
least padded to that value.r
   rJ   )Úmaxr<   r   r   r[   )r^   rY   r_   r\   s       r   Úpad_to_pow2rb   \   s:   € ô �”s˜1˜iÓ(Ó)€EØ”#”dœ4 ›;Ó'Ó(Ñ(€FÜˆvÓ!Ð!r   ÚxÚ	divide_byÚalign_toc                 óV   • [        [        X-  5      5      n X-  (       a	  XX-  -
  -  n U $ r$   )r<   r   )rc   rd   re   s      r   Úaligned_dividerg   d   s,   € ÜŒD�‘ÓÓ €AØ‡|Ø	˜™Ñ&Ñ&ˆØ€Hr   Úattention_maskÚcumulative_seqlens_qÚcumulative_seqlens_kÚsliding_windowc                 ó*  • [         R                  " U R                  5      R                  n[	        [        U5      S-
  5       HÎ  nXS-      X   -
  nX%S-      X%   -
  nXg:  a  US:¼  a  Xv-
  S-   nOSn[        X   XS-      5      n	[        X%   X%S-      5      n
[         R                  " U SXš4   R                  UU R                  U R                  S9n[         R                  " X¸S9nUS:”  a  Xv-
  U-
  nU[         R                  " X½S9-  nXÀSXš4'   MÐ     g)u~  Builds an attention mask inplace using the cumulative seqlens of the query and key. If given a sliding window, it
will also apply a sliding window mask on top. The attention mask is not boolean, it uses zeroes and -inf (or its
equivalent) so it's more of an attention score bias tensor.
The attention mask is a block-diagonal matrix, with each block an attention mask for a single query-key pair.
Each of those block is built from a causal mask and, if there is a sliding window, a sliding window mask.

An example is represented below, with seqlen_k = 8, seqlen_q = 4 and sliding_window = 6:

CAUSAL MASK:

       â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–‘ â–‘ â–‘
       â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–‘ â–‘
       â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–‘
       â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ

SLIDING WINDOW MASK:
     â”Œâ”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€ seqlen_k - seqlen_q - sliding_window = 8 - 4 - 6 = -2 offset to the left
   <â”€â”´â”€>
 â–‘ â–ˆ | â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ
 â–‘ â–‘ | â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ
 â–‘ â–‘ | â–‘ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ
 â–‘ â–‘ | â–‘ â–‘ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ

ATTENTION MASK (sum of causal and sliding window masks):

       â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–‘ â–‘ â–‘
       â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–‘ â–‘
       â–‘ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–‘
       â–‘ â–‘ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ

Another example with seqlen_k = 5, seqlen_q = 3 and sliding_window = 2:

CAUSAL MASK:

       â–ˆ â–ˆ â–ˆ â–‘ â–‘
       â–ˆ â–ˆ â–ˆ â–ˆ â–‘
       â–ˆ â–ˆ â–ˆ â–ˆ â–ˆ

SLIDING WINDOW MASK:
     â”Œâ”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€â”€ seqlen_k - seqlen_q - sliding_window = 5 - 3 - 2 = 0 offset to the left
    <â”´>
     | â–‘ â–ˆ â–ˆ â–ˆ â–ˆ
     | â–‘ â–‘ â–ˆ â–ˆ â–ˆ
     | â–‘ â–‘ â–‘ â–ˆ â–ˆ

ATTENTION MASK (sum of causal and sliding window masks):

       â–‘ â–ˆ â–ˆ â–‘ â–‘
       â–‘ â–‘ â–ˆ â–ˆ â–‘
       â–‘ â–‘ â–‘ â–ˆ â–ˆ

r
   .)ÚdtypeÚdevice)ÚdiagonalN)r>   Úfinform   r[   Úranger-   ÚsliceÚfullÚshapern   ÚtriuÚtril)rh   ri   rj   rk   r_   ÚiÚseqlen_qÚseqlen_kÚcausal_diagonalÚquery_rangeÚ	key_rangeÚ	minus_infÚmaskedÚsliding_diagonals                 r   Úbuild_attention_maskr€   k   s-  € ôt —’˜N×0Ñ0Ó1×5Ñ5€IÜ”3Ð+Ó,¨qÑ0Ö1ˆØ'¨A©Ñ.Ð1EÑ1HÑHˆØ'¨A©Ñ.Ð1EÑ1HÑHˆØÓ 8¨q£=Ø&Ñ1°AÑ5‰OàˆOÜÐ0Ñ3Ð5IÈaÉ%Ñ5PÓQˆÜÐ.Ñ1Ð3GÈAÉÑ3NÓOˆ	ä—J’JØ˜3 Ð6Ñ7×=Ñ=ØØ ×&Ñ&Ø!×(Ñ(ñ	
ˆ	ô —’˜IÑ@ˆà˜AÓØ'Ñ2°^ÑCÐØ”e—j’j ÑFÑFˆFà6<�s˜KÐ2Ó3ò- 2r   ÚnumÚstatusÚnum_query_tokensÚnum_cache_tokensÚcachec           
      ó|  • [        U 5       Vs/ s H  nSUR                   SU S3PM     nnX#-   n[        XtR                  -  5      n/ n	U Hg  n
[	        U
S/U-  SS9nXl        S/U-  Ul        X;l        UR                  X‹R                  S5      nUc  U	s  $ U	R                  [        USSUS95        Mi     U	$ s  snf )	zQAn utility function to create a list of FutureRequestStates for the warmup of CB.Ú	__warmup_Ú_Ú__r   r
   )Ú
request_idÚinitial_tokensÚmax_new_tokensT)Úhas_new_tokenÚcomplete_blocksÚquery_length)rq   Únamer   rL   r   Ú_statusÚtokens_to_processÚposition_offsetÚallocate_blocksrŠ   Úappendr   )r�   r‚   rƒ   r„   r…   rw   Úrequest_idsÚtotal_tokensÚblocks_neededÚfuture_statesÚreq_idÚstateÚ	allocateds                r   Úcreate_warmup_future_statesr�   ¿   sÚ   € ô =BÀ#¼JÓGºJ°q�Y˜vŸ{™{˜m¨1¨Q¨C¨rÓ2¹J€KÐGØ#Ñ6€LÜ˜×(8Ñ(8Ñ8Ó9€Mà€MÛˆÜ¨À¸sÀ\Ñ?QÐbcÑdˆØŒØ#$ #Ð(8Ñ"8ˆÔØ 0Ôà×)Ñ)¨-×9IÑ9IÈ1ÓMˆ	ØÑØ Ò Ø×ÑÜ˜u°DÈ!ÐZjÑkö	
ñ ð Ðùò# Hs   ŽB9)r   )r
   )Úcollectionsr   Údataclassesr   Úmathr   r   Útypingr   r>   Ú transformers.configuration_utilsr   Ú+transformers.generation.configuration_utilsr	   Úrequestsr   r   r   r   r   rE   rA   rV   r<   r]   rb   rg   ÚTensorÚlistr€   r�   rC   r   r   Ú<module>r§      sQ  ðõ $Ý !ß Ý ã å =Ý Pç MÓ M÷#ñ #ðD ÷að aó ðað"HÐ 0ð H°Tô Hð
"˜#ð "¨cð "¸cð "Àcô "ñ"�sð " sð "°sð "À3õ "ð�cð  cð °Sð ¸Sô ð ñ	Q=Ø—L‘LðQ=à˜s™)ðQ=ð ˜s™)ðQ=ð ð	Q=ð
 
õQ=ðhØ	ðàðð ðð ð	ð
 ðð 
Ð
Ñõr   