Ë
    3êñi\1  ã            #       óè  — d Z ddlZddlmZ ddlmZmZ ddlZ ej                  e	«      Z
ddgZdee   dz  dee   fd	„Z ed
¬«      dedefd„«       Z G d„ de«      Zej$                  j'                  di ¬«      	 	 	 d0dej(                  dej(                  dej(                  dej(                  dej(                  dededededz  dee   dz  deej(                  ej(                  ej(                  f   fd„«       Zej0                  	 	 	 d0dej(                  dej(                  dej(                  dej(                  dej(                  dededededz  dee   dz  deej(                  ej(                  ej(                  f   fd„«       Zddddœdej(                  dej(                  dej(                  dej(                  dej(                  dedededz  dedz  deeef   dej(                  eej(                  ej(                  f   z  fd„Zd ed!eed"f   d#eddfd$„Zej$                  j'                  d%i ¬«      	 	 d1d&ej(                  dej(                  dej(                  dej(                  d'ej(                  d(ej(                  dej(                  dej(                  dededed)ej(                  dedz  dee   dz  deej(                  ej(                  ej(                  f   fd*„«       Zej0                  	 	 d1d&ej(                  dej(                  dej(                  dej(                  d'ej(                  d(ej(                  dej(                  dej(                  dededed)ej(                  dedz  dee   dz  deej(                  ej(                  ej(                  f   fd+„«       Zd ed&ej(                  d,ej(                  d-ej(                  deej(                  dz  d"f   f
d.„Zej?                  ee¬/«       y)2zÊ
Variable-length attention implementation using Flash Attention.

This module provides a high-level Python interface for variable-length attention
that calls into the optimized Flash Attention kernels.
é    N)Ú	lru_cache)ÚAnyÚ
NamedTupleÚvarlen_attnÚ
AuxRequestÚwindow_sizeÚreturnc                 ó\   — | €ddg} t        | «      dk7  rt        dt        | «      › �«      ‚| S )Néÿÿÿÿé   z$window_size must have length 2, got )ÚlenÚ
ValueError)r   s    ú[/var/www/pod-logistic/pod-ai/venv/lib/python3.12/site-packages/torch/nn/attention/varlen.pyÚ_normalize_window_sizer      s=   € ØÐØ˜2�hˆä
ˆ;Ó˜1ÒÜÐ?ÄÀKÓ@PÐ?QÐRÓSÐSØÐó    é   )ÚmaxsizeÚdevice_indexc                  ó   — y)z;Cache device capability check to avoid repeated CUDA calls.F© )r   s    r   Ú_should_use_cudnnr      s   € ð r   c                   ó    — e Zd ZU dZdZeed<   y)r   z 
    Request which auxiliary outputs to compute from varlen_attn.

    Each field is a boolean indicating whether that auxiliary output should be computed.
    FÚlseN)Ú__name__Ú
__module__Ú__qualname__Ú__doc__r   ÚboolÚ__annotations__r   r   r   r   r   #   s   … ñð €CˆÔr   ztorch_attn::_varlen_attn)Úmutates_argsÚqueryÚkeyÚvalueÚcu_seq_qÚcu_seq_kÚmax_qÚmax_kÚ	is_causalÚscalec
                 óz  — t        |	«      }	| j                  xr t        | j                  j                  «      }
|
rvt
        j                  d«       |	d   dk7  s|	d   dk7  rt        d«      ‚t        j                  j                  j                  | ||d||||dd|d	|¬
«      }|d   |d   |d   }}}nWt
        j                  d«       t        j                  j                  j                  | ||||||d|d	||	d   |	d   ¬«      \  }}}}}t        j                  dt        j                  | j                  ¬«      }|||fS )zž
    Private custom op for variable-length attention.

    This is the internal implementation. Users should use the public varlen_attn function instead.
    ú#Using cuDNN backend for varlen_attnr   r   é   úTcuDNN backend does not support window attention. Please use Flash Attention backend.NTç        F©r)   é   ú-Using Flash Attention backend for varlen_attn)Úreturn_debug_maskr)   Úwindow_size_leftÚwindow_size_right©r   ©ÚdtypeÚdevice)r   Úis_cudar   r8   ÚindexÚlogÚinfoÚRuntimeErrorÚtorchÚopsÚatenÚ_cudnn_attention_forwardÚ_flash_attention_forwardÚzerosÚuint64)r!   r"   r#   r$   r%   r&   r'   r(   r)   r   Ú	use_cudnnÚresultÚoutputÚsoftmax_lseÚ	rng_stateÚ_Ú
rng_state_s                    r   Ú_varlen_attnrL   -   sW  € ô$ )¨Ó5€Kà—‘ÒGÔ"3°E·L±L×4FÑ4FÓ"G€IáÜ�‰Ð6Ô7à�q‰>˜RÒ ;¨q¡>°RÒ#7ÜØfóð ô —‘—‘×8Ñ8ØØØØØØØØØØØØØð 9ó 
ˆð  *0°©°F¸1±I¸vÀa¹y˜Y�‰ä�‰Ð@ÔAÜ/4¯y©y¯~©~×/VÑ/VØØØØØØØØØØ#ØØ(¨™^Ø)¨!™nð 0Wó 0
Ñ,ˆ�˜Y¨¨1ô  —‘Ø”E—L‘L¨¯©ô€Jð �; 
Ð*Ð*r   c
                 ó  — t        |	«      }	t        j                  | «      }
| j                  d«      }| j                  d«      }t        j                  j
                  rH|j                  d«      dz
  }t        j                  |||ft        j                  | j                  ¬«      }n2t        j                  ||ft        j                  | j                  ¬«      }t        j                  dt        j                  | j                  ¬«      }|
||fS )zç
    Fake implementation for meta tensor computation and tracing.

    Based on the 3D varlen path from meta__flash_attention_forward:
    - query shape: (total, num_heads, head_dim)
    - logsumexp shape: (num_heads, total_q)
    r   r,   r6   r5   )
r   r>   Ú
empty_likeÚsizeÚversionÚhipÚemptyÚfloatr8   rD   )r!   r"   r#   r$   r%   r&   r'   r(   r)   r   rG   Útotal_qÚ	num_headsÚ
batch_sizeÚ	logsumexprI   s                   r   Ú_varlen_attn_fakerX   s   sÍ   € ô( )¨Ó5€Kô ×Ñ˜eÓ$€Fð �j‰j˜‹m€GØ—
‘
˜1“€IÜ‡}�}×Òà—]‘] 1Ó%¨Ñ)ˆ
Ü—K‘KØ˜ EÐ*´%·+±+ÀeÇlÁlô
‰	ô —K‘KØ˜Ð ¬¯©¸E¿L¹Lô
ˆ	ô —‘˜D¬¯©¸U¿\¹\ÔJ€Ià�9˜iÐ'Ð'r   )r   r   )Ú
return_auxr)   r   rY   c                ó²   — |	dk(  }
t         j                  j                  j                  | |||||||
|t	        |	«      «
      \  }}}|�|j
                  r||fS |S )au  
    Compute variable-length attention using Flash Attention.
    This function is similar to scaled_dot_product_attention but optimized for
    variable-length sequences using cumulative sequence position tensors.

    Args:
        query (Tensor): Query tensor; shape :math:`(T_q, H, D)`
        key (Tensor): Key tensor; shape :math:`(T_k, H, D)`
        value (Tensor): Value tensor; shape :math:`(T_k, H, D)`
        cu_seq_q (Tensor): Cumulative sequence positions for queries; shape :math:`(N+1,)`
        cu_seq_k (Tensor): Cumulative sequence positions for keys/values; shape :math:`(N+1,)`
        max_q (int): Maximum query sequence length in the batch.
        max_k (int): Maximum key/value sequence length in the batch.
        return_aux (Optional[AuxRequest]): If not None and ``return_aux.lse`` is True, also returns the logsumexp tensor.
        scale (float, optional): Scaling factor for attention scores
        window_size (tuple[int, int], optional): Window size for sliding window attention as (left, right).
            Use (-1, -1) for full attention (default), (-1, 0) for causal attention,
            or (W, 0) for causal attention with sliding window of size W.

    Returns:
        output (Tensor): Output tensor from attention computation; shape :math:`(T_q, H, D)`.

        If ``return_aux`` is not None and ``return_aux.lse`` is True:
            lse (Tensor): Log-sum-exp of attention scores; shape :math:`(T_q, H)`.

    Shape legend:
        - :math:`N`: Batch size
        - :math:`T_q`: Total number of query tokens in the batch (sum of all query sequence lengths)
        - :math:`T_k`: Total number of key/value tokens in the batch (sum of all key/value sequence lengths)
        - :math:`H`: Number of attention heads
        - :math:`D`: Head dimension

    Example::

        >>> # xdoctest: +REQUIRES(env:TORCH_DOCTEST_CUDA)
        >>> batch_size, max_seq_len, embed_dim, num_heads = 2, 512, 1024, 16
        >>> head_dim = embed_dim // num_heads
        >>> seq_lengths = []
        >>> for _ in range(batch_size):
        ...     length = torch.randint(1, max_seq_len // 64 + 1, (1,)).item() * 64
        ...     seq_lengths.append(min(length, max_seq_len))
        >>> seq_lengths = torch.tensor(seq_lengths, device="cuda")
        >>> total_tokens = seq_lengths.sum().item()
        >>>
        >>> # Create packed query, key, value tensors
        >>> query = torch.randn(
        ...     total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda"
        ... )
        >>> key = torch.randn(
        ...     total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda"
        ... )
        >>> value = torch.randn(
        ...     total_tokens, num_heads, head_dim, dtype=torch.float16, device="cuda"
        ... )
        >>>
        >>> # Build cumulative sequence tensor
        >>> cu_seq = torch.zeros(batch_size + 1, device="cuda", dtype=torch.int32)
        >>> cu_seq[1:] = seq_lengths.cumsum(0)
        >>> max_len = seq_lengths.max().item()
        >>>
        >>> # Call varlen_attn
        >>> output = varlen_attn(
        ...     query, key, value, cu_seq, cu_seq, max_len, max_len
        ... )
    )r   r   )r>   r?   Ú
torch_attnrL   Úlistr   )r!   r"   r#   r$   r%   r&   r'   rY   r)   r   r(   Úoutr   rJ   s                 r   r   r   Ÿ   sn   € ð\ ˜wÑ&€IÜ—)‘)×&Ñ&×3Ñ3ØØØØØØØØØÜˆ[Óó�K€Cˆˆað Ð *§.¢.Ø�CˆxˆØ€Jr   ÚctxÚinputs.rG   c           
      ó    — |\
  }}}}}}}	}
}}|\  }}}| j                  ||||||||«       || _        |	| _        |
| _        || _        || _        y ©N)Úsave_for_backwardr&   r'   r(   r)   r   )r^   r_   rG   r!   r"   r#   r$   r%   r&   r'   r(   r)   r   r]   r   rI   s                   r   Ú_setup_contextrc   ÿ   su   € ð 	ñØØØØØØØØØØà Ñ€Cˆˆià×Ñ˜%  e¨X°xÀÀcÈ9ÔUà€C„IØ€C„IØ€C„MØ€C„IØ!€C…Or   z!torch_attn::_varlen_attn_backwardÚgrad_outr]   r   rI   c                 óN  — t        |«      }t        j                  d|j                  ¬«      }|j                  xr t        |j                  j                  «      }|rmt        j                  d«       |d   dk7  s|d   dk7  rt        d«      ‚t        j                  j                  j                  | |||||||||	d|
|||¬«      \  }}}nYt        j                  d	«       t        j                  j                  j                  | |||||||||	d|
||||d   |d   ¬
«      \  }}}|||fS )Nr   )r8   r+   r   r,   r-   r.   r/   r1   )r)   r3   r4   )r   r>   rR   r8   r9   r   r:   r;   r<   r=   r?   r@   Ú_cudnn_attention_backwardÚ_flash_attention_backward)rd   r!   r"   r#   r]   r   r$   r%   r&   r'   r(   rI   r)   r   ÚunusedrE   ÚdqÚdkÚdvs                      r   Ú_varlen_attn_backwardrl     sD  € ô" )¨Ó5€Kä�[‰[˜ 5§<¡<Ô0€Fà—‘ÒGÔ"3°E·L±L×4FÑ4FÓ"G€IÙÜ�‰Ð6Ô7Ø�q‰>˜RÒ ;¨q¡>°RÒ#7ÜØfóð ô —Y‘Y—^‘^×=Ñ=ØØØØØØØØØØØØØØØð >ó 
‰
ˆˆB‘ô$ 	�‰Ð@ÔAÜ—Y‘Y—^‘^×=Ñ=ØØØØØØØØØØØØØØØØ(¨™^Ø)¨!™nð# >ó 
‰
ˆˆB�ð& ˆr�2ˆ:Ðr   c                 ó    — t        |«      }t        j                  |«      }t        j                  |«      }t        j                  |«      }|||fS )zF
    Fake implementation for meta tensor computation and tracing.
    )r   r>   rN   )rd   r!   r"   r#   r]   r   r$   r%   r&   r'   r(   rI   r)   r   Ú
grad_queryÚgrad_keyÚ
grad_values                    r   Ú_varlen_attn_backward_fakerq   \  sK   € ô( )¨Ó5€Kä×!Ñ! %Ó(€JÜ×Ñ Ó$€HÜ×!Ñ! %Ó(€Jà�x Ð+Ð+r   Úgrad_lseÚgrad_rngc                 ó0  — | j                   \  }}}}}}	}
}| j                  }| j                  }| j                  }| j                  }| j
                  }t        j                  j                  j                  |||||	|
||||||||«      \  }}}|||d d d d d d d f
S ra   )
Úsaved_tensorsr&   r'   r(   r)   r   r>   r?   r[   rl   )r^   rd   rr   rs   r!   r"   r#   r$   r%   r]   r   rI   r&   r'   r(   r)   r   ri   rj   rk   s                       r   Ú	_backwardrv   y  s¶   € ð BE×ARÑARÑ>€Eˆ3��x ¨3°°Yà�I‰I€EØ�I‰I€EØ—‘€IØ�I‰I€EØ—/‘/€Kä—‘×%Ñ%×;Ñ;ØØØØØØØØØØØØØØó�J€BˆˆBð  ˆr�2�t˜T 4¨¨t°T¸4Ð?Ð?r   )Úsetup_context)FNN)NN) r   ÚloggingÚ	functoolsr   Útypingr   r   r>   Ú	getLoggerr   r;   Ú__all__r\   Úintr   r   r   r   ÚlibraryÚ	custom_opÚTensorrS   ÚtuplerL   Úregister_fakerX   r   rc   rl   rq   rv   Úregister_autogradr   r   r   ú<module>r„      s«  ðñó Ý ß "ã ð €g×Ñ˜Ó!€à˜,Ð
'€ð¨¨S©	°DÑ(8ð ¸TÀ#¹Yó ñ �1Ôð Cð ¨Dò ó ðô
�ô ð ‡�×ÑÐ3À"ÐÓEð ØØ$(ñB+Ø�<‰<ðB+à	�‰ðB+ð �<‰<ðB+ð �l‰lð	B+ð
 �l‰lðB+ð ðB+ð ðB+ð ðB+ð �4‰<ðB+ð �c‘˜TÑ!ðB+ð ˆ5�<‰<˜Ÿ™ u§|¡|Ð3Ñ4òB+ó FðB+ðJ ×Ñð ØØ$(ñ((Ø�<‰<ð((à	�‰ð((ð �<‰<ð((ð �l‰lð	((ð
 �l‰lð((ð ð((ð ð((ð ð((ð �4‰<ð((ð �c‘˜TÑ!ð((ð ˆ5�<‰<˜Ÿ™ u§|¡|Ð3Ñ4ò((ó ð((ðh %)ØØ#+ò]Ø�<‰<ð]à	�‰ð]ð �<‰<ð]ð �l‰lð	]ð
 �l‰lð]ð ð]ð ð]ð ˜TÑ!ð]ð �4‰<ð]ð �s˜C�x‘ð]ð ‡\�\�E˜%Ÿ,™,¨¯©Ð4Ñ5Ñ5ó]ð@"˜ð " U¨3°¨8¡_ð "¸cð "Àdó "ð0 ‡�×ÑÐ<È2ÐÓNð Ø$(ñAØ�l‰lðAà�<‰<ðAð 
�‰ðAð �<‰<ð	Að
 
�‰ðAð 
�‰ðAð �l‰lðAð �l‰lðAð ðAð ðAð ðAð �|‰|ðAð �4‰<ðAð �c‘˜TÑ!ðAð ˆ5�<‰<˜Ÿ™ u§|¡|Ð3Ñ4òAó OðAðH ×$Ñ$ð Ø$(ñ,Ø�l‰lð,à�<‰<ð,ð 
�‰ð,ð �<‰<ð	,ð
 
�‰ð,ð 
�‰ð,ð �l‰lð,ð �l‰lð,ð ð,ð ð,ð ð,ð �|‰|ð,ð �4‰<ð,ð �c‘˜TÑ!ð,ð ˆ5�<‰<˜Ÿ™ u§|¡|Ð3Ñ4ò,ó %ð,ð8@Ø	ð@ØŸ™ð@Ø05·±ð@ØHMÏÉð@à
ˆ5�<‰<˜$Ñ Ð#Ñ$ó@ð< × Ñ ˜y¸Ð Õ Gr   