# vLLM中Flash Attention的实现原理
from vllm.attention.backends.flash_attn import FlashAttentionBackend
from vllm.attention import AttentionMetadata

class FlashAttention:
    def __init__(self, num_heads, head_dim, block_size=64):
        self.num_heads = num_heads
        self.head_dim = head_dim
        self.block_size = block_size
    
    def forward(self, q, k, v, attn_metadata):
        """Flash Attention前向传播"""
        batch_size, seq_len, _ = q.shape
        
        # 分块计算注意力
        output = torch.zeros_like(q)
        for i in range(0, seq_len, self.block_size):
            for j in range(0, seq_len, self.block_size):
                # 计算当前块的注意力
                q_block = q[:, i:i+self.block_size]
                k_block = k[:, j:j+self.block_size]
                v_block = v[:, j:j+self.block_size]
                
                # 在线softmax计算
                attn_block = self._compute_attention_block(
                    q_block, k_block, v_block
                )
                output[:, i:i+self.block_size] += attn_block
        
        return output
    
    def _compute_attention_block(self, q, k, v):
        """计算注意力块"""
        # 计算注意力分数
        scores = torch.matmul(q, k.transpose(-2, -1))
        scores = scores / math.sqrt(self.head_dim)
        
        # 在线softmax
        attn_weights = self._online_softmax(scores)
        
        # 计算输出
        output = torch.matmul(attn_weights, v)
        return output
