"""FlashAttention-2 backend for PaDoc's packed tree visibility pattern.""" from __future__ import annotations import torch # Transformers treats custom names containing "flash_attention" as Hub kernels. ATTN_IMPLEMENTATION = "padoc_tree_fa2" def tree_flash_attention_forward( module: torch.nn.Module, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: torch.Tensor | None, dropout: float = 0.0, scaling: float | None = None, *, tree_q_indices: torch.LongTensor, tree_kv_indices: torch.LongTensor, tree_cu_seqlens_q: torch.IntTensor, tree_cu_seqlens_kv: torch.IntTensor, tree_max_seqlen_q: int | torch.Tensor, tree_max_seqlen_kv: int | torch.Tensor, **kwargs, ) -> tuple[torch.Tensor, None]: """Run exact tree attention as a virtual FlashAttention varlen batch.""" del module, attention_mask, kwargs if tree_cu_seqlens_q.dtype != torch.int32 or tree_cu_seqlens_kv.dtype != torch.int32: raise TypeError("FlashAttention cumulative sequence lengths must be torch.int32.") try: from flash_attn import flash_attn_varlen_func except ImportError as exc: # pragma: no cover - requires a CUDA environment raise ImportError( f"{ATTN_IMPLEMENTATION} requires the optional 'flash-attn' package." ) from exc batch_size, num_heads, sequence_length, head_dim = query.shape num_kv_heads = key.shape[1] query_flat = query.transpose(1, 2).reshape(batch_size * sequence_length, num_heads, head_dim) key_flat = key.transpose(1, 2).reshape(batch_size * sequence_length, num_kv_heads, head_dim) value_flat = value.transpose(1, 2).reshape(batch_size * sequence_length, num_kv_heads, head_dim) output_packed = flash_attn_varlen_func( query_flat.index_select(0, tree_q_indices), key_flat.index_select(0, tree_kv_indices), value_flat.index_select(0, tree_kv_indices), tree_cu_seqlens_q, tree_cu_seqlens_kv, int(tree_max_seqlen_q), int(tree_max_seqlen_kv), dropout_p=dropout, softmax_scale=scaling, causal=True, ) output_flat = torch.zeros_like(query_flat) output_flat.index_copy_(0, tree_q_indices, output_packed) return output_flat.view(batch_size, sequence_length, num_heads, head_dim), None def register_tree_flash_attention() -> None: from transformers import AttentionInterface AttentionInterface.register(ATTN_IMPLEMENTATION, tree_flash_attention_forward) register_tree_flash_attention()