欢迎光临
我们一直在努力

【ICLR 2026即插即用模块】北京大学 MHLA多头线性注意力机制,适合图像分类、图像生成、视频生成、图像超分辨率、目标检测与实例分割、遥感图像分析、医学图像分析等CV任务通用,涨点起飞!

一、论文信息

MHLA的核心思想是:将传统线性注意力中的单一全局KV摘要改为多个局部KV摘要,并通过多头混合机制为不同Query动态构建专属上下文,从而在保持线性复杂度的同时恢复Softmax Attention的表达能力和选择能力。

本文目录

一、论文信息

二、论文摘要概况

三、MHLA多头线性注意力机制结构图

四、MHLA模块的作用

五、MHLA模块的原理

六、MHLA模块的优势

七、即插即用模块代码

 

论文题目:MHLA: Restoring Expressivity of Linear Attention via Token-Level Multi-Head

中文题目:MHLA :通过令牌级多头机制恢复线性注意力机制的表达能力

所属单位:北京大学

二、论文摘要概况

尽管Transformer架构在众多领域占据主导地位,但其二次自注意力机制的复杂性限制了其在大规模应用中的普及。线性注意力机制虽提供了更高效的替代方案,但直接应用时往往会导致性能下降——现有改进方案通常通过引入额外模块(如深度可分离卷积)增加计算开销,反而违背了该机制的设计初衷。本研究揭示了这些方法中存在的关键缺陷:全局上下文坍缩现象,即模型会丧失表征多样性。为解决这一问题,我们提出多头线性注意力机制(MHLA),该机制通过沿标记维度在多个独立注意力头中分别计算注意力值来保持表征多样性。实验表明,在保持线性复杂度的同时, MHLA 能有效复现softmax注意力机制的大部分表达能力;我们在多个领域验证了其有效性:在ImageNet分类任务上提升3.6%,自然语言处理任务上提升6.3%,图像生成任务上提升12.6%,视频生成任务上提升41%,且所有优化均在相同时间复杂度下实现。

图1(a)展示了我们使用 MHLA 对 SANA 模型进行微调后的生成结果;(b)对比了所提出的 MHLA 与基准方法在性能和效率方面的表现。吞吐量测试在 NVIDIA H100 Tensor Core GPU上完成。沿用前述方法,我们在表中以256×256分辨率报告了FID值;(c)展示了 MHLA 在多领域的性能表现,证明其具备强大且通用的性能优势;(d)呈现了DiT-S/2在不同设备上于4096分辨率下的吞吐量数据——所有提升均完全归功于 MHLA 算法,还可结合正交技术进一步实现更显著的加速效果。

三、MHLA多头线性注意力机制结构图

图2展示了所提出的 MHLA 与其他线性注意力机制的对比。 MHLA 在令牌维度上划分多个注意力头;通过多头混合机制, MHLA 通过将键值摘要与查询特定权重相结合来恢复查询条件依赖的选择性,从而提升令牌层面的多样性,同时保持线性复杂度。

图4(a)所示为所提出的多头线性注意力机制的总体结构;(b)当M=25时,我们分别展示了对应于模块1和模块14的初始化可学习系数矩阵的两行数据,并将这两行数据及M维维度以二维形式呈现以便更清晰地理解。

四、MHLA模块的作用

(1)解决线性注意力的表达能力不足问题

传统线性注意力为了降低计算复杂度,会将所有Token压缩成一个共享的全局特征摘要,所有Query都从同一个信息池中获取上下文。这种方式虽然高效,但会导致不同Token之间的差异性逐渐消失。MHLA通过将Token划分为多个局部组,使不同区域能够保留独立的信息表示,从而提升模型的表达能力。

(2)恢复Query相关的注意力选择能力

Softmax Attention最大的优势在于不同Query可以关注不同的Token,而传统线性注意力由于共享全局摘要,难以实现这种动态选择。MHLA通过多头混合机制,为不同Query区域构建不同的上下文表示,使模型重新具备根据当前Query动态选择重要信息的能力。

(3)增强长序列建模能力

随着序列长度增加,传统线性注意力容易出现信息稀释现象,导致长距离依赖关系难以建模。MHLA通过局部KV摘要和跨头混合机制,有效缓解长序列中的信息退化问题,提高模型对长文本、高分辨率图像和长视频序列的理解能力。

(4)提升生成任务和理解任务性能

论文将MHLA应用于图像分类、图像生成、视频生成和自然语言处理任务。实验表明,MHLA能够显著提升模型性能,同时保持与传统线性注意力几乎相同的计算效率。

图3(a)展示了 MHLA 模型与基线模型的注意力得分及注意力图可视化结果;(b)DeiT-T模型注意力得分的平均排名与熵值,表明 MHLA 模型能产生更丰富且更具聚焦性的注意力分布。

五、MHLA模块的原理

(1)Token维度多头划分

与传统Transformer在通道维度划分Attention Head不同,MHLA在Token维度进行划分。输入序列被分成多个互不重叠的Token块,每个块独立计算局部Key-Value摘要,从而避免所有Token被压缩到单一全局表示中。

(2)构建局部KV摘要

每个Token块都会生成自己的局部Key-Value表示,这些局部摘要保留了各区域内部的重要信息。相比传统线性注意力只有一个全局摘要,MHLA能够同时维护多个局部上下文表示,提高信息存储能力。

(3)多头混合机制(Multi-Head Mixing)

MHLA引入可学习的混合系数矩阵,不同Query块可以根据自身需求,对多个局部KV摘要进行加权组合。这样每个Query区域都能获得专属的上下文信息,而不是共享同一个全局表示。

(4)块级选择与Token级重加权

MHLA的信息交互分为两个阶段。首先在块级别选择重要区域,然后在选中的区域内部利用Query与Key之间的相似性进一步区分不同Token的重要程度。通过这种双层筛选机制,模型能够实现更加精准的信息聚合。

(5)保持线性计算复杂度

虽然MHLA增加了局部摘要和混合过程,但这些操作本质上仍然是矩阵乘法和线性组合,因此整体计算复杂度仍保持线性增长。相比Softmax Attention的平方复杂度,MHLA在长序列场景下具有明显效率优势。

六、MHLA模块的优势

(1)恢复注意力表达能力

论文指出,传统线性注意力的注意力矩阵秩受到严格限制,导致表达能力不足。MHLA通过多头划分显著提高注意力矩阵的秩,使模型能够学习更加丰富和多样化的特征关系。

(2)避免全局上下文坍塌

传统线性注意力随着序列长度增加,会逐渐出现注意力分布趋于均匀的问题,导致模型无法聚焦关键Token。MHLA通过多个局部摘要保持特征多样性,有效避免了这种全局上下文坍塌现象。

(3)增强注意力稀疏性

由于MHLA允许不同Query块关注不同区域,因此生成的注意力分布更加集中,能够更精准地选择与当前任务相关的信息,而不是平均关注所有Token。

(4)无需额外辅助模块

许多改进线性注意力的方法需要引入卷积模块、门控机制或混合Attention结构来弥补性能损失。MHLA仅通过Attention机制内部结构改进即可提升性能,不依赖额外模块,因此结构更加简洁。

(5)计算效率高

MHLA保持与传统线性注意力相同量级的计算复杂度和推理速度,但性能接近甚至超过Softmax Attention,实现了效率与性能的平衡。

(6)长序列场景优势明显

在视频生成、长文本建模和高分辨率图像生成等超长序列任务中,MHLA相比普通线性注意力表现出更强的稳定性和收敛能力,能够有效处理数万级Token长度的输入。

图7展示了我们经过微调的 SANA – MHLA 模型产生的更多生成结果。

七、即插即用模块代码

import math
import torch
import torch.nn as nn
from einops import rearrange

class BlockDistanceConv(nn.Module):
def __init__(
self,
num_patches_per_side=16,
patch_group_size=16,
transform="linear",
local_thres=1.5,
exp_sigma=3,
):
super().__init__()

self.num_patches_per_side = num_patches_per_side
self.patch_group_size = patch_group_size
self.transform = transform
self.local_thres = local_thres
self.exp_sigma = exp_sigma

patches_per_block_side = int(math.sqrt(patch_group_size))
self.blocks_per_side = num_patches_per_side // patches_per_block_side
self.total_blocks = self.blocks_per_side ** 2

distance_matrix = self._compute_block_distances()
weight_matrix = self._apply_transform(distance_matrix)

self.conv = nn.Conv2d(
in_channels=self.total_blocks,
out_channels=self.total_blocks,
kernel_size=1,
bias=False,
)

with torch.no_grad():
self.conv.weight.copy_(weight_matrix.unsqueeze(-1).unsqueeze(-1))

def _compute_block_distances(self):
block_centers = []

for i in range(self.blocks_per_side):
for j in range(self.blocks_per_side):
block_centers.append([i + 0.5, j + 0.5])

block_centers = torch.tensor(block_centers, dtype=torch.float32)
distance_matrix = torch.zeros(self.total_blocks, self.total_blocks)

for i in range(self.total_blocks):
for j in range(self.total_blocks):
distance_matrix[i, j] = torch.norm(
block_centers[i] – block_centers[j],
p=2,
)

return distance_matrix

def _apply_transform(self, distance_matrix):
if self.transform == "linear":
max_dist = distance_matrix.max()
mat = 1.0 – distance_matrix / max_dist
return mat / mat.sum(dim=0, keepdim=True)

if self.transform == "cos":
max_dist = distance_matrix.max()
normalized_dist = distance_matrix / max_dist * math.pi / 4
mat = torch.cos(normalized_dist)
return mat / mat.sum(dim=0, keepdim=True)

if self.transform == "exp":
mat = torch.exp(-distance_matrix / self.exp_sigma)
return mat / mat.sum(dim=0, keepdim=True)

if self.transform == "gaussian":
sigma = distance_matrix.max() / 3
mat = torch.exp(-(distance_matrix ** 2) / (2 * sigma ** 2))
return mat / mat.sum(dim=0, keepdim=True)

if self.transform == "local":
mat = (distance_matrix <= self.local_thres).float()
return mat / mat.sum(dim=0, keepdim=True)

raise ValueError(f"Unknown transform: {self.transform}")

def forward(self, x):
return self.conv(x)

def get_weight_matrix(self):
return self.conv.weight.data.squeeze(-1).squeeze(-1)

class MHLA(nn.Module):
def __init__(
self,
dim,
heads=8,
dim_head=None,
dropout=0.1,
fixed_weight_value=None,
qk_norm=False,
transform="cos",
**kwargs,
):
super().__init__()

if dim_head is None:
dim_head = dim // heads

inner_dim = dim_head * heads

self.num_heads = heads
self.head_dim = dim_head
self.scale = dim_head ** -0.5
self.eps = kwargs.get("eps", 1e-6)

self.norm = nn.LayerNorm(dim)

qkv_bias = kwargs.get("qkv_bias", False)
self.to_qkv = nn.Linear(dim, inner_dim * 3, bias=qkv_bias)

self.q_norm = nn.RMSNorm(dim) if qk_norm else nn.Identity()
self.k_norm = nn.RMSNorm(dim) if qk_norm else nn.Identity()

self.lepe = nn.Conv2d(dim, dim, kernel_size=5, stride=1, padding=2, groups=dim)

self.window_size = kwargs.get("window_size", 49)
self.window_len = int(self.window_size ** 0.5)

self.embed_len = kwargs.get("embed_len", 196)
self.num_pieces = self.embed_len // self.window_size
self.pieces_len = int(self.num_pieces ** 0.5)

local_thres = kwargs.get("local_thres", 1.5)
exp_sigma = kwargs.get("exp_sigma", 3)

self.piece_attn = BlockDistanceConv(
num_patches_per_side=int(self.embed_len ** 0.5),
patch_group_size=self.window_size,
transform=transform,
local_thres=local_thres,
exp_sigma=exp_sigma,
)

self.to_out = nn.Sequential(
nn.Linear(inner_dim, dim),
nn.Dropout(dropout),
)

if fixed_weight_value is not None:
self._init_weights_with_fixed_value(fixed_weight_value)

def _init_weights_with_fixed_value(self, value):
for name, param in self.named_parameters():
if "weight" in name:
nn.init.constant_(param, value)
elif "bias" in name and param is not None:
nn.init.zeros_(param)

nn.init.constant_(self.to_qkv.weight, value)

for module in self.to_out:
if isinstance(module, nn.Linear):
nn.init.constant_(module.weight, value)
if module.bias is not None:
nn.init.zeros_(module.bias)

@staticmethod
def init_to_value(model, value=1.0):
for name, param in model.named_parameters():
if "weight" in name:
nn.init.constant_(param, value)
elif "bias" in name and param is not None:
nn.init.zeros_(param)

return model

def _process_qkv_impl(self, q, k, v, b, n, h, d):
q = self.q_norm(q)
k = self.k_norm(k)

q = torch.relu(q) + self.eps
k = torch.relu(k) + self.eps

q, k, v = map(
lambda t: rearrange(
t,
"b n w (h d) -> (b h) n w d",
h=h,
d=d,
),
(q, k, v),
)

k = k.transpose(-2, -1)

return q, k, v

def _mlp_lepe(self, x):
q, k, v = self.to_qkv(x).chunk(3, dim=-1)

lepe = self.lepe(
rearrange(
v,
"b (h w) (p1 p2) d -> b d (h p1) (w p2)",
h=self.pieces_len,
w=self.pieces_len,
p1=self.window_len,
p2=self.window_len,
)
)

lepe = rearrange(
lepe,
"b d (h p1) (w p2) -> b (h w) (p1 p2) d",
h=self.pieces_len,
w=self.pieces_len,
p1=self.window_len,
p2=self.window_len,
)

return q, k, v, lepe

def forward(self, x):
x = self.norm(x)

b, n, w, c = x.shape
h = self.num_heads
d = self.head_dim

q, k, v, lepe = self._mlp_lepe(x)
q, k, v = self._process_qkv_impl(q, k, v, b, n, h, d)

kv = torch.matmul(k, v)
kv = self.piece_attn(kv)

k_sum = k.sum(dim=-1, keepdim=True)
normalizer = self.piece_attn(torch.matmul(q, k_sum)) + self.eps

out = torch.matmul(q, kv) / normalizer

out = rearrange(
out,
"(b h) n w d -> b n w (h d)",
b=b,
h=self.num_heads,
)

out = out + lepe
return self.to_out(out)

if __name__ == "__main__":
input = torch.randn(2, 16, 16,64)
model = MHLA(
dim=64,
heads=8,
embed_len=256,
window_size=16,
qkv_bias=False,
transform="cos",
)
print(model)
print("CSDN:AI魔改博士")
output = model(input)
print('MHLA input_size:', input.size())
print('MHLA output_size:', output.size())

赞(0)
未经允许不得转载:171主机测评 » 【ICLR 2026即插即用模块】北京大学 MHLA多头线性注意力机制,适合图像分类、图像生成、视频生成、图像超分辨率、目标检测与实例分割、遥感图像分析、医学图像分析等CV任务通用,涨点起飞!
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址