import torch import torch.nn as nn import torch.nn.functional as F import math from typing import Optional, Tuple class MemoryEfficientAttention(nn.Module): """内存高效的多头注意力""" def __init__(self, dim: int, num_heads: int = 8, qkv_bias: bool = False, dropout: float = 0.0): super().__init__() self.dim = dim self.num_heads = num_heads self.head_dim = dim // num_heads self.scale = self.head_dim ** -0.5 self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) self.proj = nn.Linear(dim, dim) self.dropout = nn.Dropout(dropout) def forward(self, x: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: B, N, C = x.shape # 计算QKV qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) q, k, v = qkv[0], qkv[1], qkv[2] # 分块计算注意力,避免OOM chunk_size = min(32, N) attn_output = torch.zeros(B, self.num_heads, N, self.head_dim, device=x.device) for i in range(0, N, chunk_size): q_chunk = q[:, :, i:i+chunk_size, :] # 计算注意力分数 attn_scores = torch.matmul(q_chunk, k.transpose(-2, -1)) * self.scale if mask is not None: attn_scores = attn_scores + mask attn_probs = F.softmax(attn_scores, dim=-1) attn_probs = self.dropout(attn_probs) # 计算输出 attn_output[:, :, i:i+chunk_size, :] = torch.matmul(attn_probs, v) # 合并多头 attn_output = attn_output.transpose(1, 2).reshape(B, N, C) # 输出投影 output = self.proj(attn_output) output = self.dropout(output) return output class CrossAttention(nn.Module): """交叉注意力(用于文本条件)""" def __init__(self, query_dim: int, context_dim: int, num_heads: int = 8, dropout: float = 0.0): super().__init__() self.query_dim = query_dim self.context_dim = context_dim self.num_heads = num_heads self.head_dim = query_dim // num_heads self.scale = self.head_dim ** -0.5 self.to_q = nn.Linear(query_dim, query_dim) self.to_k = nn.Linear(context_dim, query_dim) self.to_v = nn.Linear(context_dim, query_dim) self.proj = nn.Linear(query_dim, query_dim) self.dropout = nn.Dropout(dropout) def forward(self, x: torch.Tensor, context: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor: B, N, C = x.shape # 计算Q q = self.to_q(x).reshape(B, N, self.num_heads, self.head_dim).transpose(1, 2) # 计算K, V k = self.to_k(context).reshape(B, -1, self.num_heads, self.head_dim).transpose(1, 2) v = self.to_v(context).reshape(B, -1, self.num_heads, self.head_dim).transpose(1, 2) # 分块计算注意力 chunk_size = min(32, N) attn_output = torch.zeros(B, self.num_heads, N, self.head_dim, device=x.device) for i in range(0, N, chunk_size): q_chunk = q[:, :, i:i+chunk_size, :] # 计算注意力分数 attn_scores = torch.matmul(q_chunk, k.transpose(-2, -1)) * self.scale if mask is not None: attn_scores = attn_scores + mask attn_probs = F.softmax(attn_scores, dim=-1) attn_probs = self.dropout(attn_probs) # 计算输出 attn_output[:, :, i:i+chunk_size, :] = torch.matmul(attn_probs, v) # 合并多头 attn_output = attn_output.transpose(1, 2).reshape(B, N, C) # 输出投影 output = self.proj(attn_output) return output class FlashAttentionWrapper(nn.Module): """FlashAttention包装器(如果可用)""" def __init__(self, dim: int, num_heads: int = 8): super().__init__() self.dim = dim self.num_heads = num_heads try: from flash_attn import flash_attn_qkvpacked_func self.use_flash = True except ImportError: self.use_flash = False if not self.use_flash: self.attention = MemoryEfficientAttention(dim, num_heads) def forward(self, x: torch.Tensor) -> torch.Tensor: if self.use_flash: return self._flash_attention(x) else: return self.attention(x) def _flash_attention(self, x: torch.Tensor) -> torch.Tensor: # FlashAttention实现 B, N, C = x.shape qkv = x.reshape(B, N, 3, self.num_heads, C // self.num_heads) qkv = qkv.permute(2, 0, 3, 1, 4) # [3, B, num_heads, N, head_dim] from flash_attn import flash_attn_qkvpacked_func output = flash_attn_qkvpacked_func(qkv) output = output.reshape(B, N, C) return output