Lumina_Dev_Legacy / src /models /attention.py
TAI Research
Initial commit: Lumina_Dev_Legacy (archived)
29691f6
Raw
History Blame Contribute Delete
5.28 kB
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