| 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 = 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]
|
|
|
|
|
| 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 = self.to_q(x).reshape(B, N, self.num_heads, self.head_dim).transpose(1, 2)
|
|
|
|
|
| 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:
|
|
|
| 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)
|
|
|
| from flash_attn import flash_attn_qkvpacked_func
|
| output = flash_attn_qkvpacked_func(qkv)
|
| output = output.reshape(B, N, C)
|
|
|
| return output |