即插即用系列 | CVPR 2026 | LFSB:差分双流注意力,双特征交互融合,替换传统注意力涨点! | 代码分享

发布时间:2026/7/30 22:39:45

即插即用系列 | CVPR 2026 | LFSB:差分双流注意力,双特征交互融合,替换传统注意力涨点! | 代码分享 0. 前言本文介绍了LFSB差分双流注意力其通过交替特征融合与分离策略和差分注意力机制首次在图像解耦领域实现对透射层与反射层的精准分离有效破解了深层网络中特征逐渐混淆导致的细节丢失与伪影难题。将其作为即插即用模块轻松助力CNN、YOLO、Transformer等深度学习模型精准抑制跨层干扰、增强特征独立性让模型在面对复杂光照、重叠目标或半透明物体等挑战性场景时依然能够保持清晰的边界感知与稳定的检测精度。专栏链接即插即用系列专栏链接可点击跳转免费订阅目录0. 前言1. LFSB注意力简介2. LFSB注意力原理与创新点 LFSB注意力基本原理 LFSB注意力处理流程3. 适用范围与模块效果适用范围⚡模块效果4. LFSB注意力代码实现1. LFSB注意力简介单图像反射分离Single Image Reflection Separation, SIRS旨在将混合图像解耦为透射层和反射层。现有方法在非线性混合条件下尤其是在深层解码器中由于隐式融合机制和多尺度协调不足常常出现透射-反射混淆的问题。为此我们提出 ReflexSplit一个双流框架包含三项关键创新1跨尺度门控融合Cross-scale Gated Fusion, CrGF自适应地聚合多层次的语义先验、纹理细节和解码器上下文稳定梯度流并保持特征一致性。2层融合-分离模块Layer Fusion-Separation Block, LFSB在融合与分离之间交替进行融合用于提取共享结构分离用于实现层专属的解耦。受差分变换器启发我们通过跨流减法将注意力抵消机制扩展到双流分离任务中。3课程训练Curriculum Training通过深度相关的初始化与逐轮次预热逐步增强差分分离能力。在合成与真实世界基准上的大量实验表明ReflexSplit 在感知质量和鲁棒泛化方面均达到了当前最先进的性能。原始论文https://arxiv.org/pdf/2601.17468原始代码https://github.com/wuw2135/ReflexSplit2. LFSB注意力原理与创新点 LFSB注意力基本原理ReflexSplit 的整体架构以 Layer Fusion-Separation BlockLFSB为核心模块其设计围绕差分双流注意力、窗口化分区交互与融合-分离双流机制展开构成一个高效的差分融合全局-局部特征的框架。LFSB 结构示意LFSB 的核心创新点体现在以下几个方面差分双流注意力架构通过自注意力与交叉注意力双分支协同工作。自注意力负责捕捉单流内部的空间依赖关系例如局部特征中的细节关联交叉注意力则建模双流之间的相互依赖强化局部与全局特征的互补性。两者结合使注意力覆盖更全面的多维度特征关系。跨流交互增强在自注意力和交叉注意力分支中均引入跨流投影机制使双流特征能够相互吸收有效信息如 x 融合 y、y 融合 x增强互补性避免单流特征的孤立增强。差分机制优化利用可学习的 λ 系数构建差分注意力权重attn1 − λ × attn2有效抑制双流之间的冗余信息突出各自独特且有效的关联提升特征表达的纯度。窗口化分区交互设计将特征图划分为局部窗口进行注意力计算将计算复杂度从 O((H×W)²) 降低至 O((H×W)×(ws²))ws 为窗口尺寸显著降低高分辨率特征的处理成本。同时引入窗口内相对位置编码弥补窗口分区可能带来的空间信息损失提高注意力机制的精度。融合-分离双流机制采用门控融合策略通过“通道拆分 交叉相乘”的方式实现双流的深度融合同时保持双流结构独立输出避免传统融合方式带来的特征稀释问题保留各自输入的专属特性。双流归一化与前馈网络均采用并行处理设计确保双流特征的同步增强提升交互效率。可学习残差加权与通道增强在注意力分支与前馈分支中引入可学习残差权重使模型能够自适应调节原始特征与增强特征的比例在稳定性与表达能力之间取得平衡。同时嵌入通道注意力模块强化关键通道特征抑制冗余干扰提升特征表达的针对性。 LFSB注意力处理流程LFSB 的特征处理流程模块整体遵循“预处理 → 注意力交互 → 增强融合 → 输出”的流程双输入预处理将双输入特征 x 与 y 的维度从 [b, c, h, w] 转换为 [b, h×w, c]适配注意力输入格式对双流分别进行归一化以稳定数值分布随后将特征图恢复为 [b, h, w, c] 格式推理时通过补零确保尺寸可被窗口大小整除再拆分为局部窗口。差分双流注意力交互自注意力交互将双流窗口特征沿 batch 维度拼接经跨流交互投影增强后计算单流内部的注意力权重再通过差分机制进行优化经价值加权后拆分为双流。交叉注意力交互将双流窗口特征沿序列维度拼接经跨流交互投影增强后计算双流之间的注意力权重同样经差分机制优化后拆分为双流。注意力融合将自注意力与交叉注意力的输出加权求和得到双流注意力增强后的特征还原窗口至原始尺寸并与原始特征进行残差融合。前馈增强与门控融合将特征格式转换为 [b, c, h, w]送入前馈网络。经二维层归一化、通道扩展与深度卷积增强后通过双流门控实现交叉融合与分离再经通道注意力强化关键通道最终通过 1×1 卷积恢复通道数并与原始特征进行残差融合。双输出特征输出增强后的双输入特征 out_x 与 out_y。二者在保留各自专属特性如 x 的局部细节、y 的全局结构的同时充分融入对方有效信息实现深层次互补。3. 适用范围与模块效果适用范围LFSBLayer Fusion-Separation Block适用于需要显式建模双流特征交互与解耦的视觉任务特别适合在特征层次中同时存在“共享结构”与“独立属性”需要分离的场景如图像分层、信号解耦、多源信息融合等。双流架构中的显式解耦需求LFSB 的核心在于通过交替执行“融合”与“分离”操作先通过双向投影对齐双流特征空间再通过差分注意力机制实现层间干扰抑制。这种设计适用于任何需要从混合特征中提取共享信息并分离独立成分的双流网络结构。差分注意力机制适用于抑制跨流干扰LFSB 借鉴差分变换的思想在双流之间引入 At−λArAt−λAr 的跨流减法操作有效抑制透射与反射之间的信息混淆。这一机制可推广至其他需要减少双流之间干扰的任务如多模态融合、图像去噪、特征解耦学习等。适用于深层网络中的特征保持LFSB 结合自注意力与交叉注意力分别在空间维度和序列维度上建模特征关系并通过差分操作防止深层网络中特征趋于不可分。因此LFSB 特别适用于深层网络中需要维持特征独立性与层次一致性的场景。模块化设计便于嵌入现有架构LFSB 作为即插即用的注意力模块可嵌入现有双流或双分支视觉模型中用于增强层间特征的可区分性。其结构不依赖于特定输入形式具备良好的通用性与迁移能力。⚡模块效果模块效果去反射性能和视觉效果均实现SOTA。LFSB的消融研究移除LFSB后性能和效果显著下降。内部消融说明了各组件的有效性。4. LFSB注意力代码实现以下为LFSB注意力机制的官方pytorch实现代码import torch import torch.nn as nn import torch.nn.functional as F from timm.models.layers import to_2tuple, trunc_normal_ import math from collections import OrderedDict def lambda_init_fn(depth): 动态初始化函数基于网络深度生成lambda系数预留用于层权重调制 核心公式lambda 0.8 - 0.6 * exp(-0.3 * depth) 特性深度越深lambda越接近0.8适配深层特征的融合强度 Args: depth: 网络当前层深度 Returns: lambda系数float return 0.8 - 0.6 * math.exp(-0.3 * depth) class LayerNormFunction(torch.autograd.Function): staticmethod def forward(ctx, x, weight, bias, eps): ctx.eps eps N, C, H, W x.size() # 确保 weight 和 bias 与 x 在同一设备上 weight weight.to(x.device) bias bias.to(x.device) mu x.mean(1, keepdimTrue) var (x - mu).pow(2).mean(1, keepdimTrue) y (x - mu) / (var eps).sqrt() ctx.save_for_backward(y, var, weight) y weight.view(1, C, 1, 1) * y bias.view(1, C, 1, 1) return y staticmethod def backward(ctx, grad_output): eps ctx.eps N, C, H, W grad_output.size() y, var, weight ctx.saved_tensors # 确保张量在同一设备上 grad_output grad_output.to(y.device) weight weight.to(y.device) g grad_output * weight.view(1, C, 1, 1) mean_g g.mean(dim1, keepdimTrue) mean_gy (g * y).mean(dim1, keepdimTrue) gx 1. / torch.sqrt(var eps) * (g - y * mean_gy - mean_g) return gx, (grad_output * y).sum(dim3).sum(dim2).sum(dim0), grad_output.sum(dim3).sum(dim2).sum(dim0), None class LayerNorm2d(nn.Module): def __init__(self, channels, eps1e-6): super(LayerNorm2d, self).__init__() self.register_parameter(weight, nn.Parameter(torch.ones(channels))) self.register_parameter(bias, nn.Parameter(torch.zeros(channels))) self.eps eps def forward(self, x): return LayerNormFunction.apply(x, self.weight, self.bias, self.eps) class CABlock(nn.Module): def __init__(self, channels): super(CABlock, self).__init__() self.ca nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Conv2d(channels, channels, 1), nn.Sigmoid() ) def forward(self, x): return x * self.ca(x) class DualStreamGate(nn.Module): def forward(self, x, y): # 确保输入维度匹配 if x.shape ! y.shape: raise ValueError(fShape mismatch: x{x.shape}, y{y.shape}) # 简单的门控融合机制 gate_x torch.sigmoid(x) gate_y torch.sigmoid(y) x_fused x * gate_y y_fused y * gate_x return x_fused, y_fused class DualStreamFFN(nn.Module): 双流前馈网络 def __init__(self, dim, expansion_factor2): super().__init__() hidden_dim dim * expansion_factor self.conv1 nn.Conv2d(dim, hidden_dim, 1) self.dwconv nn.Conv2d(hidden_dim, hidden_dim, 3, padding1, groupshidden_dim) self.conv2 nn.Conv2d(hidden_dim, dim, 1) self.act nn.GELU() def forward(self, x, y): # 分别处理两个流 x self.act(self.dwconv(self.conv1(x))) y self.act(self.dwconv(self.conv1(y))) x self.conv2(x) y self.conv2(y) return x, y def window_partition(x, window_size): if not isinstance(window_size, tuple): window_size (window_size, window_size) B, H, W, C x.shape x x.view(B, H // window_size[0], window_size[0], W // window_size[1], window_size[1], C) windows x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size[0], window_size[1], C) return windows def window_reverse(windows, window_size, H, W): if not isinstance(window_size, tuple): window_size (window_size, window_size) B int(windows.shape[0] / (H * W / window_size[0] / window_size[1])) x windows.view(B, H // window_size[0], W // window_size[1], window_size[0], window_size[1], -1) x x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1) return x class DifferentialDualStreamAttention(nn.Module): 差分双流注意力模块融合差分机制的双流双维度注意力自注意力交叉注意力 def __init__(self, dim, window_size, num_heads, depth1, qkv_biasTrue, qk_scaleNone, attn_drop0., proj_drop0.): super().__init__() self.dim dim self.window_size window_size self.num_heads num_heads self.head_dim dim // num_heads self.scale qk_scale or self.head_dim ** -0.5 # Differential parameters self.lambda_init lambda_init_fn(depth) self.lambda_sa nn.Parameter(torch.tensor(0.1)) self.lambda_ca nn.Parameter(torch.tensor(0.1)) # Self-Attention 分支 self.sa_qkv nn.Linear(dim, dim * 3, biasqkv_bias) # 跨流交互投影 (SA) self.sa_cross_trans_proj nn.Linear(dim, dim, biasFalse) self.sa_cross_refl_proj nn.Linear(dim, dim, biasFalse) self.sa_enhance_weight nn.Parameter(torch.tensor(0.1)) # Cross-Attention 分支 self.ca_q nn.Linear(dim, dim, biasqkv_bias) self.ca_kv nn.Linear(dim, dim * 2, biasqkv_bias) # 跨流交互投影 (CA) self.ca_cross_trans_proj nn.Linear(dim, dim, biasFalse) self.ca_cross_refl_proj nn.Linear(dim, dim, biasFalse) self.ca_enhance_weight nn.Parameter(torch.tensor(0.1)) # 相对位置编码 self.relative_position_bias_table nn.Parameter( torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)) coords_h torch.arange(self.window_size[0]) coords_w torch.arange(self.window_size[1]) coords torch.stack(torch.meshgrid([coords_h, coords_w], indexingij)) coords_flatten torch.flatten(coords, 1) relative_coords coords_flatten[:, :, None] - coords_flatten[:, None, :] relative_coords relative_coords.permute(1, 2, 0).contiguous() relative_coords[:, :, 0] self.window_size[0] - 1 relative_coords[:, :, 1] self.window_size[1] - 1 relative_coords[:, :, 0] * 2 * self.window_size[1] - 1 relative_position_index relative_coords.sum(-1) self.register_buffer(relative_position_index, relative_position_index) self.attn_drop nn.Dropout(attn_drop) self.proj_sa nn.Linear(dim, dim) self.proj_ca nn.Linear(dim, dim) self.proj_drop nn.Dropout(proj_drop) trunc_normal_(self.relative_position_bias_table, std.02) self.softmax nn.Softmax(dim-1) def forward_sa(self, x_concat): Self-Attention: 在 batch 维度 concat加入跨流交互 x_concat: [2*B, N, C] - torch.cat([x_windows, y_windows], dim0) 返回: [2*B, N, C] B2, N, C x_concat.shape B B2 // 2 # 分离 trans 和 refl x_trans x_concat[:B] x_refl x_concat[B:] # 跨流特征增强 trans_enhanced x_trans self.sa_enhance_weight * self.sa_cross_refl_proj(x_refl) refl_enhanced x_refl self.sa_enhance_weight * self.sa_cross_trans_proj(x_trans) # 计算各自的 QKV qkv_trans self.sa_qkv(trans_enhanced).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) qkv_refl self.sa_qkv(refl_enhanced).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) q_trans, k_trans, v_trans qkv_trans[0], qkv_trans[1], qkv_trans[2] q_refl, k_refl, v_refl qkv_refl[0], qkv_refl[1], qkv_refl[2] # 计算注意力 attn_trans (q_trans * self.scale) k_trans.transpose(-2, -1) attn_refl (q_refl * self.scale) k_refl.transpose(-2, -1) # 添加位置编码 relative_position_bias self.relative_position_bias_table[ self.relative_position_index.view(-1) ].view(N, N, -1).permute(2, 0, 1).contiguous() attn_trans attn_trans relative_position_bias.unsqueeze(0) attn_refl attn_refl relative_position_bias.unsqueeze(0) attn_trans self.softmax(attn_trans) attn_refl self.softmax(attn_refl) # Differential Attention lambda_sa torch.clamp(torch.sigmoid(self.lambda_sa), 0.01, 0.99) diff_attn_trans attn_trans - lambda_sa * attn_refl diff_attn_refl attn_refl - lambda_sa * attn_trans diff_attn_trans self.attn_drop(diff_attn_trans) diff_attn_refl self.attn_drop(diff_attn_refl) out_trans (diff_attn_trans v_trans).transpose(1, 2).reshape(B, N, C) out_refl (diff_attn_refl v_refl).transpose(1, 2).reshape(B, N, C) out_trans self.proj_sa(out_trans) out_refl self.proj_sa(out_refl) out_trans self.proj_drop(out_trans) out_refl self.proj_drop(out_refl) # 重新 concat 回去 return torch.cat([out_trans, out_refl], dim0) def forward_ca(self, x_concat): Cross-Attention: 在 sequence 维度 concat加入跨流交互 x_concat: [B, 2*N, C] - torch.cat([x_windows, y_windows], dim-2) 返回: [B, 2*N, C] B, N2, C x_concat.shape N N2 // 2 # 分离 trans 和 refl x_trans x_concat[:, :N, :] x_refl x_concat[:, N:, :] # 跨流特征增强 trans_enhanced x_trans self.ca_enhance_weight * self.ca_cross_refl_proj(x_refl) refl_enhanced x_refl self.ca_enhance_weight * self.ca_cross_trans_proj(x_trans) # Query 从增强后的特征计算 q_trans self.ca_q(trans_enhanced).reshape(B, N, self.num_heads, self.head_dim).permute(0, 2, 1, 3) q_refl self.ca_q(refl_enhanced).reshape(B, N, self.num_heads, self.head_dim).permute(0, 2, 1, 3) # KV 从增强后的 concat 特征计算 x_concat_enhanced torch.cat([trans_enhanced, refl_enhanced], dim-2) kv self.ca_kv(x_concat_enhanced).reshape(B, N2, 2, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) k, v kv[0], kv[1] k_trans, k_refl k[:, :, :N, :], k[:, :, N:, :] v_trans, v_refl v[:, :, :N, :], v[:, :, N:, :] # 计算注意力 attn_trans (q_trans * self.scale) k_trans.transpose(-2, -1) attn_refl (q_refl * self.scale) k_refl.transpose(-2, -1) # 添加位置编码 relative_position_bias self.relative_position_bias_table[ self.relative_position_index.view(-1) ].view(N, N, -1).permute(2, 0, 1).contiguous() attn_trans attn_trans relative_position_bias.unsqueeze(0) attn_refl attn_refl relative_position_bias.unsqueeze(0) attn_trans self.softmax(attn_trans) attn_refl self.softmax(attn_refl) # Differential Cross-Attention lambda_ca torch.clamp(torch.sigmoid(self.lambda_ca), 0.01, 0.99) diff_attn_trans attn_trans - lambda_ca * attn_refl diff_attn_refl attn_refl - lambda_ca * attn_trans diff_attn_trans self.attn_drop(diff_attn_trans) diff_attn_refl self.attn_drop(diff_attn_refl) out_trans (diff_attn_trans v_trans).transpose(1, 2).reshape(B, N, C) out_refl (diff_attn_refl v_refl).transpose(1, 2).reshape(B, N, C) out_trans self.proj_ca(out_trans) out_refl self.proj_ca(out_refl) out_trans self.proj_drop(out_trans) out_refl self.proj_drop(out_refl) # 重新 concat 回去 return torch.cat([out_trans, out_refl], dim-2) class DifferentialDualAttentionInteractiveBlock(nn.Module): 层融合-分离模块Layer Fusion-Separation Block, LFSB def __init__(self, dim, input_resolution, num_heads, window_size12, depth1, norm_layernn.LayerNorm): super().__init__() self.dim dim self.input_resolution input_resolution self.num_heads num_heads self.window_size window_size if min(self.input_resolution) self.window_size: self.window_size min(self.input_resolution) self.norm1 norm_layer(dim) # Differential 注意力 self.diff_attention DifferentialDualStreamAttention( dim, to_2tuple(self.window_size), num_heads, depth ) # 前馈网络 self.norm2 norm_layer(dim) self.ffn DualStreamFFN(dim) self.gate DualStreamGate() # 可学习权重 self.alpha nn.Parameter(torch.zeros(1)) self.beta nn.Parameter(torch.zeros(1)) def forward(self, x, y): B, C, H, W x.shape # 转换为序列形式 x_seq x.permute(0, 2, 3, 1).contiguous().view(B, H * W, C) y_seq y.permute(0, 2, 3, 1).contiguous().view(B, H * W, C) # 保存残差 x_skip, y_skip x_seq, y_seq # 归一化 x_norm self.norm1(x_seq) y_norm self.norm1(y_seq) # 转换为窗口形式 x_win x_norm.view(B, H, W, C) y_win y_norm.view(B, H, W, C) # Padding pad_l pad_t 0 pad_r (self.window_size - W % self.window_size) % self.window_size pad_b (self.window_size - H % self.window_size) % self.window_size if pad_r 0 or pad_b 0: x_win F.pad(x_win, (0, 0, pad_l, pad_r, pad_t, pad_b)) y_win F.pad(y_win, (0, 0, pad_l, pad_r, pad_t, pad_b)) _, Hp, Wp, _ x_win.shape # 窗口分割 x_windows window_partition(x_win, self.window_size).view(-1, self.window_size * self.window_size, C) y_windows window_partition(y_win, self.window_size).view(-1, self.window_size * self.window_size, C) # Self-Attention sa_out self.diff_attention.forward_sa( torch.cat([x_windows, y_windows], dim0) ) xx_windows, yy_windows sa_out.chunk(2, dim0) # Cross-Attention ca_out self.diff_attention.forward_ca( torch.cat([x_windows, y_windows], dim-2) ) xy_windows, yx_windows ca_out.chunk(2, dim-2) # 融合 x_windows (xx_windows xy_windows).view(-1, self.window_size, self.window_size, C) y_windows (yy_windows yx_windows).view(-1, self.window_size, self.window_size, C) # 窗口反转 x window_reverse(x_windows, self.window_size, Hp, Wp) y window_reverse(y_windows, self.window_size, Hp, Wp) # 移除padding if pad_r 0 or pad_b 0: x x[:, :H, :W, :].contiguous() y y[:, :H, :W, :].contiguous() # 残差连接 x x_skip x.view(B, H * W, C) * self.alpha y y_skip y.view(B, H * W, C) * self.alpha # 转换为2D形式 x x.view(B, H, W, C).permute(0, 3, 1, 2).contiguous() y y.view(B, H, W, C).permute(0, 3, 1, 2).contiguous() # 保存残差 x_skip2, y_skip2 x, y # 归一化 - 使用 LayerNorm2d 确保设备一致性 x LayerNorm2d(C)(x) y LayerNorm2d(C)(y) # 前馈网络 x_ffn, y_ffn self.ffn(x, y) # 门控融合 x_gate, y_gate self.gate(x_ffn, y_ffn) # 残差连接 x x_skip2 x_gate * self.beta y y_skip2 y_gate * self.beta return x, y if __name__ __main__: device torch.device(cuda:0 if torch.cuda.is_available() else cpu) x torch.randn(1, 64, 32, 32).to(device) y torch.randn(1, 64, 32, 32).to(device) model DifferentialDualAttentionInteractiveBlock(64, (32, 32), 8, 8, 5).to(device) out_x, out_y model(x, y) print(输入局部特征维度, x.shape) print(输入全局特征维度, y.shape) print(输出局部特征维度, out_x.shape) print(输出全局特征维度, out_y.shape)结合自己的思路可将其即插即用至任何模型做结构创新设计该模块博主已成功嵌入至YOLO26模型中可订阅博主YOLO系列算法改进或YOLO26自研改进专栏专栏链接YOLO系列算法改进专栏链接、YOLO26自研改进系列专栏

相关新闻