+
    &jX                        R t ^ RIt^ RIHt ^ RIHtHt ^ RIt]P                  ! ]	4      t
. R.OtR R lt]! ^R7      R R	 l4       t ! R
 R]4      t]P                  P!                  R/ R7      R/R R ll4       t]P$                  R/R R ll4       tRRRRRR0RRRRRRRR/R R llt]P                  P!                  RR0R7      R/R R ll4       t]P$                  R/R R  ll4       tRRRRRR0RRRRRRRR/R! R" lltR# R$ lt]P                  P!                  R%/ R7      R1R& R' ll4       t]P$                  R1R( R) ll4       tR* R+ lt]P9                  ]]R,7       ]P:                  P=                  ]P>                  P@                  PB                  4       ^ R-I"H#t#H$t$H%t%H&t& ]$]&]P>                  PN                  P"                  &   ]%]&]P>                  PN                  P*                  &   ]#]&]P>                  PN                  P2                  &   R# )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
AuxRequestc                j    V ^8  d   QhR\         \        ,          R,          R\         \        ,          /# )   window_sizeNreturn)listint)formats   "q/Users/jameslopez/projects/CWCArchive/cwc-podcast/.venv/lib/python3.14/site-packages/torch/nn/attention/varlen.py__annotate__r      s'      S	D(8 T#Y     c                 d    V f   RR.p \        V 4      ^8w  d   \        R\        V 4       24      hV # )Nz$window_size must have length 2, got )len
ValueError)r   s   &r   _normalize_window_sizer      s=    2h
;1?K@P?QRSSr   )maxsizec                0    V ^8  d   QhR\         R\        /# )r   device_indexr	   )r   bool)r   s   "r   r   r      s      C D r   c                    R# )z;Cache device capability check to avoid repeated CUDA calls.F )r   s   &r   _should_use_cudnnr      s     r   c                   4   a  ] tR t^#t o RtRtV 3R ltRtV tR# )r   z
Request which auxiliary outputs to compute from varlen_attn.

Each field is a boolean indicating whether that auxiliary output should be computed.
Fc                &   < V ^8  d   Qh/ S[ ;R&   # )r   lse)r   )r   __classdict__s   "r   r   AuxRequest.__annotate__#   s      
 r   r   N)	__name__
__module____qualname____firstlineno____doc__r   __annotate_func____static_attributes____classdictcell__)r   s   @r   r   r   #   s      C  r   ztorch_attn::_varlen_attn)mutates_argsFc          !      *   V ^8  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R,          R\        R\        R	\        R
\        R,          R\
        \        ,          R,          R\        R\         P                  R,          R\         P                  R,          R\        R,          R\        \         P                  \         P                  \         P                  3,          /# r   querykeyvaluecu_seq_qcu_seq_kNmax_qmax_k	is_causalscaler   
enable_gqa	seqused_kblock_table
num_splitsr	   torchTensorr   r   floatr
   tuple)r   s   "r   r   r   .   s    V+ V+<<V+	V+ <<V+ ll	V+
 llT!V+ V+ V+ V+ 4<V+ cT!V+ V+ ||d"V+ $V+ d
V+ 5<<u||34V+r   c                D   \        V	4      p	V P                  ;'       d     \        V P                  P                  4      pV'       d   \
        P                  R4       V
'       d   \        R4      hVe   \        R4      hV	^ ,          R8w  g   V	^,          R8w  d   \        R4      hVf   Ve   \        R4      h\        P                  P                  P                  V VVRVVVVRRVR	VR
7      pV^ ,          V^,          V^,          pppMb\
        P                  R4       \        P                  P                  P                  V VVVVVVRVR	VV	^ ,          V	^,          VVVR7      w  ppp p\        P                  ! R\        P                  V P                  R7      pVVV3# )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_attnz,GQA is not supported with the cuDNN backend.Nz3num_splits is not supported with the cuDNN backend.TcuDNN backend does not support window attention. Please use Flash Attention backend.zBseqused_k/block_table is not yet supported with the cuDNN backend.T        Fr4   -Using Flash Attention backend for varlen_attn)return_debug_maskr4   window_size_leftwindow_size_rightr6   r7   r8   dtypedevicer   r   )r   is_cudar   rI   indexloginfoRuntimeErrorr:   opsaten_cudnn_attention_forward_flash_attention_forwardzerosuint64)r,   r-   r.   r/   r0   r1   r2   r3   r4   r   r5   r6   r7   r8   	use_cudnnresultoutputsoftmax_lse	rng_state_
rng_state_s   &&&&&&&&&&&&&&       r   _varlen_attnr]   -   s   , )5KGG"3ELL4F4F"GI67MNN!TUUq>R;q>R#7f   K$; T  88 9 
  *0F1IvayYY@A/4yy~~/V/V#(^)!n#!! 0W 0
,Y1& ELLJ ;
**r   c          !      *   V ^8  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R,          R\        R\        R	\        R
\        R,          R\
        \        ,          R,          R\        R\         P                  R,          R\         P                  R,          R\        R,          R\        \         P                  \         P                  \         P                  3,          /# r+   r9   )r   s   "r   r   r      s     .( .(<<.(	.( <<.( ll	.(
 llT!.( .( .( .( 4<.( cT!.( .( ||d".( $.( d
.( 5<<u||34.(r   c                   \        V	4      p	\        P                  ! V 4      pV P                  ^ 4      pV P                  ^4      p\        P                  ! VV3\        P
                  V P                  R7      p\        P                  P                  '       d   \        P                  P                  4       pV\        P                  P                  P                  8X  dM   VP                  ^ 4      ^,
          p\        P                  ! VVV3\        P
                  V P                  R7      p\        P                  ! R\        P                  V P                  R7      pVVV3# )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)
rG   rJ   )r   r:   
empty_likesizeemptyr<   rI   versionhip_C_get_rocm_fa_preferred_backend_ROCmFABackendAOTritonrU   )r,   r-   r.   r/   r0   r1   r2   r3   r4   r   r5   r6   r7   r8   rX   total_q	num_heads	logsumexp	preferred
batch_sizerZ   s   &&&&&&&&&&&&&&       r   _varlen_attn_fakern      s    0 )5K e$F jjmG

1I	GEKKI }}HH;;=	//888!q)A-JY.ekk%,,I DU\\JI9i''r   
return_auxr4   r   r5   r6   r7   r8   c          !      B   V ^8  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R,          R\        R\        R	\        R,          R
\        R,          R\
        \        \        3,          R\        R\         P                  R,          R\         P                  R,          R\        R,          R\         P                  \
        \         P                  \         P                  3,          ,          /# )r   r,   r-   r.   r/   r0   Nr1   r2   ro   r4   r   r5   r6   r7   r8   r	   r:   r;   r   r   r<   r=   r   )r   s   "r   r   r      s
    V V<<V	V <<V ll	V
 llT!V V V T!V 4<V sCxV V ||d"V $V d
V  \\E%,,455!Vr   c                  V P                  ^4      pVe   VP                  ^4      MVP                  ^4      pV
'       g   W8w  d   \        RV RV R24      hV
'       d    W,          ^ 8w  d   \        RV RV R24      hV	R8H  p\        P                  P                  P                  V VVVVVVVV\        V	4      V
VVV4      w  pppVe   VP                  '       d   VV3# V# )a-  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_q, D)`
    key (Tensor): Key tensor; shape :math:`(T_k, H_{kv}, D)`, or
        :math:`(\text{total\_pages}, \text{page\_size}, H_{kv}, D)` when ``block_table`` is provided.
    value (Tensor): Value tensor; shape :math:`(T_k, H_{kv}, D)`, or
        :math:`(\text{total\_pages}, \text{page\_size}, H_{kv}, D)` when ``block_table`` is provided.
    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.
    enable_gqa (bool): If set to True, enables Grouped Query Attention (GQA)
        and allows key/value to have fewer heads than query.
        Each KV head is shared by a group of :math:`H_q / H_{kv}` query heads,
        so :math:`H_q` must be divisible by :math:`H_{kv}`.
        Default is False.
    seqused_k (Tensor, optional): Number of valid KV tokens per batch element; shape :math:`(N,)`.
        When set, only the first ``seqused_k[i]`` tokens in the key/value sequence for batch
        element *i* participate in attention. Useful for KV-cache decoding where the cache slot
        is larger than the actual sequence. Inference-only (not supported in backward).
    block_table (Tensor, optional): Block table for paged KV cache; shape
        :math:`(N, \text{max\_pages\_per\_seq})`, dtype ``int32``.
        Requires ``seqused_k``. Inference-only (not supported in backward).

        When ``block_table`` is provided, ``key`` and ``value`` are a "pool" of
        pages of tokens of KV data and the pages belong to any sequence/order.
        The ``block_table`` is what maps each sequence's logical chunks
        back to physical pages in this pool.

        ``seqused_k[i]`` tells the kernel how many tokens in sequence *i* are
        actually valid, since the last page is typically only partially filled.
    num_splits (int, optional): Number of splits for split-KV. Set to ``1``
        to disable split-KV which enables batch invariance. Split-KV
        parallelizes the key/value sequence dimension across multiple thread
        blocks and combines partial results. The split decision depends
        on ``max_k`` (the longest sequence in the batch), so different batch
        compositions can change the reduction order and produce different
        floating-point results for the same sequence. When this is disabled,
        bitwise identical outputs are guaranteed for a given sequence
        regardless of what other sequences are in the batch, at the
        cost of lower GPU utilization when there are few queries. When
        ``None`` (default), the kernel chooses automatically.

Returns:
    output (Tensor): Output tensor from attention computation; shape :math:`(T_q, H_q, 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_q)`.

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_q`: Number of query attention heads
    - :math:`H_{kv}`: Number of key/value attention heads (equal to :math:`H_q` unless GQA is enabled)
    - :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
    ... )
GExpect query and key/value to have the same number of heads but got Hq=	 and Hkv=&. Try setting enable_gqa=True for GQA.MExpect number of query heads to be a multiple of kv heads for GQA but got Hq=.r       )ra   r   r:   rP   
torch_attnr]   r
   r   )r,   r-   r.   r/   r0   r1   r2   ro   r4   r   r5   r6   r7   r8   num_heads_qnum_heads_kr3   outr   r[   s   &&&&&&&$$$$$$$      r   varlen_attnr~      s	   j **Q-K!,!8#((1+chhqkK+4%i} =34
 	

 k/14%i}A?
 	

 w&I))&&33[KCa  *...CxJr   ztorch_attn::_varlen_attn_outr}   c                    V ^8  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R,          R\        R	\        R
\        R\        R,          R\
        \        ,          R,          R\        R\         P                  R,          R\         P                  R,          R\        R,          R\         P                  /# r   r}   r,   r-   r.   r/   r0   Nr1   r2   r3   r4   r   r5   r6   r7   r8   r	   r:   r;   r   r   r<   r
   )r   s   "r   r   r   S  s     2 2	2<<2 
2 <<	2
 ll2 llT!2 2 2 2 4<2 cT!2 2 ||d"2 $2 d
2  \\!2r   c                b   \        V
4      p
VP                  ;'       d     \        VP                  P                  4      pV'       d   \        R4      h\        P                  R4       \        P                  P                  P                  V VVVVVVVRVRV	V
^ ,          V
^,          VVVR7      pV# )z
Private custom op for variable-length attention with pre-allocated output.
Same as _varlen_attn but writes the attention output into the provided out tensor.
z+cuDNN backend does not support out variant.z1Using Flash Attention backend for varlen_attn_outrA   F)r4   rE   rF   r6   r7   r8   )r   rK   r   rI   rL   rO   rM   rN   r:   rP   rQ   +_flash_attention_forward_no_dropout_inplace)r}   r,   r-   r.   r/   r0   r1   r2   r3   r4   r   r5   r6   r7   r8   rV   rY   s   &&&&&&&&&&&&&&&  r   _varlen_attn_outr   R  s    , )5KGG"3ELL4F4F"GIHIIHH@A))..LL$Q%a.# M K( r   c                    V ^8  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R,          R\        R	\        R
\        R\        R,          R\
        \        ,          R,          R\        R\         P                  R,          R\         P                  R,          R\        R,          R\         P                  /# r   r   )r   s   "r   r   r     s     " "	"<<" 
" <<	"
 ll" llT!" " " " 4<" cT!" " ||d"" $" d
"  \\!"r   c                   VP                  ^ 4      pVP                  ^4      p\        P                  ! VV3\        P                  VP                  R7      p\        P
                  P                  '       d   \        P                  P                  4       pV\        P                  P                  P                  8X  dM   VP                  ^ 4      ^,
          p\        P                  ! VVV3\        P                  VP                  R7      pV# )>
Fake implementation for meta tensor computation and tracing.
rG   )ra   r:   rb   r<   rI   rc   rd   re   rf   rg   rh   )r}   r,   r-   r.   r/   r0   r1   r2   r3   r4   r   r5   r6   r7   r8   ri   rj   rk   rl   rm   s   &&&&&&&&&&&&&&&     r   _varlen_attn_out_faker     s    * jjmG

1I	GEKKI }}HH;;=	//888!q)A-JY.ekk%,,I r   c          #      b   V ^8  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R,          R\        R	\        R
\        R,          R\        R,          R\
        \        \        3,          R\        R\         P                  R,          R\         P                  R,          R\        R,          R\         P                  \
        \         P                  \         P                  3,          ,          /# )r   r}   r,   r-   r.   r/   r0   Nr1   r2   ro   r4   r   r5   r6   r7   r8   r	   rq   )r   s   "r   r   r     s    : :	:<<: 
: <<	:
 ll: llT!: : : T!: 4<: sCx: : ||d": $:  d
!:" \\E%,,455#:r   c                  VP                  ^4      pVe   VP                  ^4      MVP                  ^4      pV'       g   VV8w  d   \        RV RV R24      hV'       d!   VV,          ^ 8w  d   \        RV RV R24      hV
R8H  p\        P                  P                  P                  V VVVVVVVVV	\        V
4      VVVV4      pVe   VP                  '       d   V V3# V # )zCompute variable-length attention using Flash Attention with a pre-allocated output tensor.

Same as :func:`varlen_attn` but writes the attention output into the provided ``out`` tensor
instead of allocating a new one.

rs   rt   ru   rv   rw   rx   )ra   r   r:   rP   rz   r   r
   r   )r}   r,   r-   r.   r/   r0   r1   r2   ro   r4   r   r5   r6   r7   r8   r{   r|   r3   r   s   &&&&&&&&$$$$$$$    r   varlen_attn_outr     s	   0 **Q-K!,!8#((1+chhqkK+4%i} =34
 	

 kK/14%i}A?
 	

 w&I
))


/
/[C" *...CxJr   c                Z    V ^8  d   QhR\         R\        \         R3,          R\         RR/# )r   ctxinputs.rX   r	   N)r   r=   )r   s   "r   r   r     s0     " " "U38_ "c "d "r   c                     Vw  ppppppp	p
ppppppVw  pppVe   \        R4      hVe   \        R4      hV P                  W4WVVVVV4       Wn        Wn        Wn        Wn        Wn        R # )Nz)seqused_k is an inference-only parameter.z+block_table is an inference-only parameter.)rO   save_for_backwardr1   r2   r3   r4   r   )r   r   rX   r,   r-   r.   r/   r0   r1   r2   r3   r4   r   r5   r6   r7   r8   r}   r   rZ   s   &&&                 r   _setup_contextr     s      	 CiFGGHII%exc9UIIMI!Or   z!torch_attn::_varlen_attn_backwardc          !         V ^8  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R	\        R
\        R\        R\         P                  R\        R,          R\
        \        ,          R,          R\        \         P                  \         P                  \         P                  3,          /# r   grad_outr,   r-   r.   r}   r   r/   r0   r1   r2   r3   rZ   r4   Nr   r	   r9   )r   s   "r   r   r     s     A AllA<<A 
A <<	A
 
A 
A llA llA A A A ||A 4<A cT!A 5<<u||34Ar   c                    \        V4      p\        P                  ! ^ VP                  R7      pVP                  ;'       d     \        VP                  P                  4      pV'       dz   \        P                  R4       V^ ,          R8w  g   V^,          R8w  d   \        R4      h\        P                  P                  P                  V VVVVVVVVV	RV
VVVR7      w  pppMa\        P                  R4       \        P                  P                  P                  V VVVVVVVVV	RV
VVVV^ ,          V^,          R7      w  pppVVV3# )	ry   )rI   r?   r@   rA   rB   rC   )r4   rE   rF   r   )r   r:   rb   rI   rK   r   rL   rM   rN   rO   rP   rQ   _cudnn_attention_backward_flash_attention_backward)r   r,   r-   r.   r}   r   r/   r0   r1   r2   r3   rZ   r4   r   unusedrV   dqdkdvs   &&&&&&&&&&&&&&     r   _varlen_attn_backwardr     sG   " )5K[[5<<0FGG"3ELL4F4F"GI67q>R;q>R#7f  YY^^== > 

B$ 	@AYY^^==(^)!n# > 

B& r2:r   c          !         V ^8  d   QhR\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R\         P                  R	\        R
\        R\        R\         P                  R\        R,          R\
        \        ,          R,          R\        \         P                  \         P                  \         P                  3,          /# r   r9   )r   s   "r   r   r   R  s     , ,ll,<<, 
, <<	,
 
, 
, ll, ll, , , , ||, 4<, cT!, 5<<u||34,r   c                    \        V4      p\        P                  ! V4      p\        P                  ! V4      p\        P                  ! V4      pWV3# )r   )r   r:   r`   )r   r,   r-   r.   r}   r   r/   r0   r1   r2   r3   rZ   r4   r   
grad_querygrad_key
grad_values   &&&&&&&&&&&&&&   r   _varlen_attn_backward_faker   Q  sI    ( )5K!!%(J$H!!%(J++r   c                    V ^8  d   QhR\         R\        P                  R\        P                  R\        P                  R\        \        P                  R,          R3,          /# )r   r   r   grad_lsegrad_rngr	   N.)r   r:   r;   r=   )r   s   "r   r   r   n  sS     1 1	11051HM1
5<<$#$1r   c                 2   V P                   w  rErgrrV P                  pV P                  pV P                  pV P                  pV P
                  p\        P                  P                  P                  VVVVV	V
VVVVVVVV4      w  ppp^pVVV.RV,          O5# )   )N)
saved_tensorsr1   r2   r3   r4   r   r:   rP   rz   r   )r   r   r   r   r,   r-   r.   r/   r0   r}   r   rZ   r1   r2   r3   r4   r   r   r   r   
num_paramss   &&&&                 r   	_backwardr   n  s     BEARAR>EIIEIIEIIIE//K%%;;JBB$ JB0'J.00r   )setup_context)_varlen_attn_backward_flop_varlen_attn_forward_flop_varlen_attn_out_flopflop_registry)r~   r   r   )FNNFNNN)r   r   )NN)(r%   logging	functoolsr   typingr   r   r:   	getLoggerr!   rM   __all__r   r   r   library	custom_opr]   register_fakern   r~   r   r   r   r   r   r   r   register_autograd_dynamodisallow_in_graphrP   rQ   r   torch.utils.flop_counterr   r   r   r   rz   r   r   r   <module>r      so     "  !
: 1 
  3"EV+ FV+r .( .(bV %)V V $,V V &*V (,V "Vr 7ugN2 O2j "  "J: %): : $,: : &*: (,:  "!:z"B <2NA OAH $$, %,81B   y  G   	IINN>>  4Meii""// 07Leii""33 4<Veii""88 9r   