padoc-document-parser / padoc /tree_flash_attention.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
414b4fe verified
Raw
History Blame Contribute Delete
2.56 kB
"""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()