欢迎光临
我们一直在努力

TPS薄板样条对齐RGB-T SOD:TPS-SCL AAAI2026 无对齐SOTA,SCCM约束+TPSAM对齐+CMCM融合!!!

🔥 本文定位:CSDN 原创干货 | 大连民族大学 & 山东科技大学 AAAI 2026 RGB-T SOD 无对齐 SOTA 方案

🎯 核心收益:一次性解决真实场景RGB-T图像空间未对齐+尺度变化+视点偏移三大痛点!基于MobileViT双流编码器打造SCCM语义约束+TPSAM薄板样条对齐+CMCM跨模态门控融合三板斧,UVT20K上F-measure达0.815,超HENet 7%,仅12.82M参数+12.34G FLOPs,完美适配无约束监控、无人机巡检、辅助驾驶等真实未对齐场景

📌 核心创新矩阵:

  • SCCM语义相关性约束模块——高层语义先验引导浅层特征聚焦共显著区域,大幅抑制无对齐背景噪声干扰
  • TPSAM薄板样条对齐模块——LSSM局部Mamba窗口扫描增强纹理边界+TPS可变形变换将热红外特征弯曲至RGB空间,解决非线性形变
  • CMCM跨模态相关性模块——ES2D投影至共享隐空间+门控机制双隐状态变换,深度挖掘RGB-T互补信息
  • 三骨干适配——MobileViT-S/Swin-B/PVT-v2-B4全适配,轻量版12.82M参数量级碾压20M+竞品
  • ✅ 适配场景:无约束监控视频RGB-T行人检测 / 无人机可见光-热红外目标跟踪 / 夜间辅助驾驶 / 工业热成像缺陷检测 / 安防未对齐双光融合


    🔖 前言

  • 现有RGB-T SOD方法严重依赖手动对齐数据集(VT821/VT1000/VT5000),但真实场景拍摄的RGB-T图像对天然存在空间未对齐——视点差异、尺度变化、旋转偏移导致显著目标在RGB和热红外中的位置完全不同
  • 已有无对齐方法(DCNet的仿射变换+动态卷积、PCNet的单应性估计)要么只能处理弱对齐(仿射变换不足以模拟大偏移),要么无法处理局部非线性形变(单应性矩阵假设平面变换),在UVT20K真实未对齐数据集上性能断崖下降
  • 现有轻量化RGB-T SOD方法(MobileSal/LSNet等)为追求效率牺牲了跨模态对齐能力,在未对齐场景下F-measure甚至低于0.7,实用性极差
  •   针对上述问题,大连民族大学与山东科技大学团队提出TPS-SCL——专为真实世界无对齐RGB-T SOD设计的轻量级框架。核心创新在"先约束、再对齐、后融合"三阶段:先用SCCM从高层语义上约束浅层特征聚焦显著区域,再用TPSAM通过薄板样条(Thin-Plate Spline)可变形变换将热红外特征精确弯曲至RGB坐标空间,最后用CMCM在隐空间中完成跨模态门控融合。

      本文全程论文 1:1 对齐 + 可运行完整代码复现 + 实验全解读,CSDN 最细最干货版本,直接拿去发论文、改毕设、打比赛、做工程都能暴力涨点!


    一、TPS-SCL 整体架构

    在这里插入图片描述

    ▲ 图1:TPS-SCL 整体架构。包含双流MobileViT编码器、SCCM语义约束模块、TPSAM薄板样条对齐模块、CMCM跨模态相关性模块、解码器。来源:论文 Fig.2。

      TPS-SCL采用约束→对齐→融合→解码四阶段流水线设计,数据流如下:

  • 双流MobileViT编码:RGB和热红外分别通过MobileViT-S提取4级多尺度特征$F_{rgb}^i$和$F_t^i$($i=1,2,3,4$)。MobileViT融合CNN的局部归纳偏置和Transformer的全局建模能力,轻量高效
  • SCCM高层语义约束:取最深层特征$F_{rgb}^4$和$F_t^4$,通过差分增强模块(DEM)保留模态特有和互补信息,再经ES2D扫描建模跨模态相关性生成显著引导特征SGF(公式1-3),逐层约束浅层特征聚焦共显著区域
  • TPSAM薄板样条对齐:约束后的特征$F̂_{rgb}^i$和$F̂_t^i$经LSSM局部窗口扫描增强局部细节+SGE注意力抑制冗余,通过GAP+FC预测TPS控制点位移(公式4-5),将热红外特征$F̂_t^i$弯曲变形至RGB坐标空间得到对齐特征$A_t^i$(公式6-8)
  • CMCM跨模态融合:对齐后的$A_t^i$与增强RGB特征$F̂_{rgb}^i$分别通过ES2D投影至共享隐空间,经门控机制双隐状态变换(公式9-10),深度融合后送入解码器生成最终预测图
  • 核心设计亮点:整个框架围绕"未对齐"这一核心挑战,首次将TPS可变形变换引入RGB-T SOD领域(而非传统的仿射变换或单应性估计),结合ES2D线性复杂度全局扫描,在12.82M参数下实现真实未对齐场景的鲁棒检测。


    二、核心模块逐行拆解

    2.1 SCCM 语义相关性约束模块

    在这里插入图片描述

    ▲ 图2:SCCM模块结构。包含差分增强模块DEM→ES2D扫描→特征融合→SGF生成→逐层约束。来源:论文 Fig.2右侧。

    • 解决无对齐场景噪声放大:直接对齐未对齐双模态特征会放大空间差异,SCCM用高层语义先验约束浅层,避免背景噪声干扰
    • 解决DEM局部信息补偿:ES2D高效扫描可能丢失局部细节,DEM差分增强模块在进入ES2D前增强模态特有信息
    • 解决多尺度特征对齐:生成单一SGF后通过上采样+卷积适配不同层级分辨率,统一约束所有层级
    • 解决冗余通道抑制:SGE(Spatial Group-wise Enhancement)注意力在通道维度分组增强显著区域

      SCCM的核心思想是"先想清楚再看"——在未对齐的低层特征上进行直接匹配会引入大量噪声,不如先让高层语义特征(已经过充分抽象,空间差异较小)算出哪些区域是显著的,再用这个先验去引导低层。

    Step 1:差分增强(DEM)

    在这里插入图片描述

    Ergb4=DEM(Frgb4),Et4=DEM(Ft4)E_{rgb}^4 = DEM(F_{rgb}^4), \\quad E_t^4 = DEM(F_t^4)Ergb4=DEM(Frgb4),Et4=DEM(Ft4)

      DEM通过简单的差分操作增强模态特异性和互补信息,补偿后续ES2D扫描可能造成的局部信息损失。

    Step 2:共享特征生成

    Hrgb=SiLU(DWC(LP(LN(Ergb4))))H_{rgb} = \\text{SiLU}(\\text{DWC}(\\text{LP}(\\text{LN}(E_{rgb}^4))))Hrgb=SiLU(DWC(LP(LN(Ergb4))))

    Ht=SiLU(DWC(LP(LN(Et4))))H_t = \\text{SiLU}(\\text{DWC}(\\text{LP}(\\text{LN}(E_t^4))))Ht=SiLU(DWC(LP(LN(Et4))))

    H=Hrgb⊙Ht⊕Hrgb⊙Ht(1)H = H_{rgb} \\odot H_t \\oplus H_{rgb} \\odot H_t \\tag{1}H=HrgbHtHrgbHt(1)

      其中⊙\\odot为逐元素乘,⊕\\oplus为逐元素加。这种"乘+加"双路径融合同时捕获线性和非线性跨模态关系。HHH通过ES2D层捕获长程空间依赖:

    H1=ES2D(H)⊙SiLU(LP(LN(Ergb4)))H_1 = \\text{ES2D}(H) \\odot \\text{SiLU}(\\text{LP}(\\text{LN}(E_{rgb}^4)))H1=ES2D(H)SiLU(LP(LN(Ergb4)))

    H2=ES2D(H)⊙SiLU(LP(LN(Et4)))(2)H_2 = \\text{ES2D}(H) \\odot \\text{SiLU}(\\text{LP}(\\text{LN}(E_t^4))) \\tag{2}H2=ES2D(H)SiLU(LP(LN(Et4)))(2)

    SGF=SGE(LP(H1⊕H2))⊕LP(H1⊙H2)\\text{SGF} = \\text{SGE}(\\text{LP}(H_1 \\oplus H_2)) \\oplus \\text{LP}(H_1 \\odot H_2)SGF=SGE(LP(H1H2))LP(H1H2)

    Step 3:逐层约束

    F^rgbi=UP24−i(Conv(SGF))⊙Frgbi(3)\\hat{F}_{rgb}^i = \\text{UP}_{2^{4-i}}(\\text{Conv}(\\text{SGF})) \\odot F_{rgb}^i \\tag{3}F^rgbi=UP24i(Conv(SGF))Frgbi(3)

    F^ti=UP24−i(Conv(SGF))⊙Fti\\hat{F}_t^i = \\text{UP}_{2^{4-i}}(\\text{Conv}(\\text{SGF})) \\odot F_t^iF^ti=UP24i(Conv(SGF))Fti

      SGF通过逐元素乘作用于每一层特征,相当于一个"显著性注意力掩码"——显著区域的特征被保留增强,背景噪声被抑制。这步是SCCM的核心效果所在。

    2.2 TPSAM 薄板样条对齐模块

    在这里插入图片描述

    ▲ 图3:TPSAM模块结构。LSSM局部窗口扫描→SGE增强→GAP+FC预测控制点位移→TPS变换→对齐特征。来源:论文 Fig.3。

    • 解决非线性局部形变:仿射变换只能处理全局旋转缩放,单应性估计假设平面变换——TPS薄板样条通过控制点插值实现任意非线性变形
    • 解决局部细节感知:LSSM(Local Scanning State Machine)用局部窗口扫描替代全局SS2D,增强边界和纹理表征
    • 解决动态控制点预测:不依赖固定网格,而是通过网络预测每个控制点的位移量,自适应不同输入
    • 解决变换平滑性:TPS通过最小化弯曲能量(Bending Energy)保证变换平滑,防止过扭曲

      TPSAM是TPS-SCL最核心的模块。它通过学习一个可变形变换,将热红外图像中的显著区域"弯曲"到RGB坐标空间中。

    Step 1:局部特征增强

    E~rgbi=SGE(LSSM(LN(F^rgbi)))(4)\\tilde{E}_{rgb}^i = \\text{SGE}(\\text{LSSM}(\\text{LN}(\\hat{F}_{rgb}^i))) \\tag{4}E~rgbi=SGE(LSSM(LN(F^rgbi)))(4)

    E~ti=SGE(LSSM(LN(F^ti)))\\tilde{E}_t^i = \\text{SGE}(\\text{LSSM}(\\text{LN}(\\hat{F}_t^i)))E~ti=SGE(LSSM(LN(F^ti)))

      LSSM来自LocalMamba,在局部窗口内做选择性扫描,相比全局SS2D更关注邻域细节。

    Step 2:控制点位移预测

    (Δx,Δy)=FC(GAP(Concat(E~rgbi,E~ti)))(5)(\\Delta x, \\Delta y) = \\text{FC}(\\text{GAP}(\\text{Concat}(\\tilde{E}_{rgb}^i, \\tilde{E}_t^i))) \\tag{5}(Δx,Δy)=FC(GAP(Concat(E~rgbi,E~ti)))(5)

    Q(x2,y2)=P(x1+Δx,y1+Δy)Q(x_2, y_2) = P(x_1 + \\Delta x, y_1 + \\Delta y)Q(x2,y2)=P(x1+Δx,y1+Δy)

      首先将增强后的RGB和热红外特征拼接,经全局平均池化和全连接层,预测每个源控制点PPPxxxyyy方向上的位移(Δx,Δy)(\\Delta x, \\Delta y)(Δx,Δy),得到目标控制点矩阵QQQ

    Step 3:TPS变换参数求解

      构建距离矩阵KKK和增广矩阵LLL

    Kij=∥pi−pj∥2log⁡(∥pi−pj∥2)(6)K_{ij} = \\|p_i – p_j\\|^2 \\log(\\|p_i – p_j\\|^2) \\tag{6}Kij=pipj2log(pipj2)(6)

    L=[KPaugPaugT0](7)L = \\begin{bmatrix} K & P_{aug} \\\\ P_{aug}^T & 0 \\end{bmatrix} \\tag{7}L=[KPaugTPaug0](7)

      求解变换参数W=L†YW = L^\\dagger YW=LY,其中Y=[Q,0]Y = [Q, 0]Y=[Q,0]L†L^\\daggerL为伪逆。

    Step 4:TPS变换应用

    R(x,y)=∑i=1NwiU(∥X−pi∥)+a0+a1x+a2y(8)R(x, y) = \\sum_{i=1}^N w_i U(\\|X – p_i\\|) + a_0 + a_1 x + a_2 y \\tag{8}R(x,y)=i=1NwiU(Xpi)+a0+a1x+a2y(8)

      其中U(r)=r2log⁡(r)U(r) = r^2 \\log(r)U(r)=r2log(r)为径向基函数(RBF),控制变换平滑度。最终得到对齐后的热红外特征Ati=G(F^ti)A_t^i = G(\\hat{F}_t^i)Ati=G(F^ti)

      如图4所示,对齐后的AtiA_t^iAti与RGB特征的空间差异显著减小,背景噪声也被有效抑制。

    2.3 CMCM 跨模态相关性模块

    在这里插入图片描述

    ▲ 图4:CMCM模块结构(RGB分支)。DEM增强→ES2D隐空间投影→门控双隐状态变换→SGE抑制→残差连接。来源:论文 Fig.5。

    • 解决隐空间跨模态融合:将对齐后的双模态特征投影到共享隐空间,用门控机制控制信息流动
    • 解决双隐状态变换:RGB和热红外的隐状态互相作为门控信号,实现双向互补
    • 解决融合后冗余抑制:SGE注意力在隐空间输出后进一步精炼通道特征
    • 解决信息退化:残差连接保留原始特征,防止隐空间变换造成信息丢失

      CMCM在TPSAM完成空间对齐后执行,负责语义层面的深度融合。

    Step 1:隐空间投影

    yrgb=ES2D(SiLU(DWC(LP(LN(DEM(F^rgbi))))))y_{rgb} = \\text{ES2D}(\\text{SiLU}(\\text{DWC}(\\text{LP}(\\text{LN}(\\text{DEM}(\\hat{F}_{rgb}^i))))))yrgb=ES2D(SiLU(DWC(LP(LN(DEM(F^rgbi))))))

    yt=ES2D(SiLU(DWC(LP(LN(DEM(Ati))))))(9)y_t = \\text{ES2D}(\\text{SiLU}(\\text{DWC}(\\text{LP}(\\text{LN}(\\text{DEM}(A_t^i)))))) \\tag{9}yt=ES2D(SiLU(DWC(LP(LN(DEM(Ati))))))(9)

    Step 2:门控信号生成

    Zrgb=SiLU(LP(LN(DEM(F^rgbi))))Z_{rgb} = \\text{SiLU}(\\text{LP}(\\text{LN}(\\text{DEM}(\\hat{F}_{rgb}^i))))Zrgb=SiLU(LP(LN(DEM(F^rgbi))))

    Zt=SiLU(LP(LN(DEM(Ati))))Z_t = \\text{SiLU}(\\text{LP}(\\text{LN}(\\text{DEM}(A_t^i))))Zt=SiLU(LP(LN(DEM(Ati))))

    Step 3:双隐状态门控变换

    f~rgbi=SGE(LP(yrgb⊙Zrgb⊕yt⊙Zrgb))⊕F^rgbi\\tilde{f}_{rgb}^i = \\text{SGE}(\\text{LP}(y_{rgb} \\odot Z_{rgb} \\oplus y_t \\odot Z_{rgb})) \\oplus \\hat{F}_{rgb}^if~rgbi=SGE(LP(yrgbZrgbytZrgb))F^rgbi

    f~ti=SGE(LP(yt⊙Zt⊕yrgb⊙Zt))⊕Ati(10)\\tilde{f}_t^i = \\text{SGE}(\\text{LP}(y_t \\odot Z_t \\oplus y_{rgb} \\odot Z_t)) \\oplus A_t^i \\tag{10}f~ti=SGE(LP(ytZtyrgbZt))Ati(10)

      关键创新在交叉门控:ZrgbZ_{rgb}Zrgb不仅门控RGB自身的yrgby_{rgb}yrgb,还门控热红外的yty_tyt——反之亦然。这种设计让两个模态的隐状态互相调制,实现深度交互。

    Step 4:融合与解码

    Si=Conv3×3(Concat(f~rgbi,f~ti))(11)S_i = \\text{Conv}_{3\\times3}(\\text{Concat}(\\tilde{f}_{rgb}^i, \\tilde{f}_t^i)) \\tag{11}Si=Conv3×3(Concat(f~rgbi,f~ti))(11)

      融合特征SiS_iSi通过解码器自顶向下集成,使用BCE + Smoothness + Dice联合损失优化。


    三、论文 1:1 对齐完整可运行 PyTorch 复现代码

    3.1 环境依赖

    pip install torch torchvision einops timm opencv-python
    pip install mamba-ssm causal-conv1d # ES2D选择性扫描

    3.2 完整 PyTorch 实现

    TPS变换核心

    import torch
    import torch.nn as nn
    import torch.nn.functional as F
    import math

    # ====== TPS薄板样条变换 ======
    class TPSGridGenerator(nn.Module):
    """TPS变换网格生成器 — 可变形弯曲"""
    def __init__(self, grid_size=5):
    super().__init__()
    self.grid_size = grid_size
    # 生成均匀控制点网格 [-1, 1]
    xs = torch.linspace(1, 1, grid_size)
    ys = torch.linspace(1, 1, grid_size)
    grid_y, grid_x = torch.meshgrid(ys, xs, indexing='ij')
    self.register_buffer('P', torch.stack([grid_x, grid_y], dim=1).view(1, 2))

    def forward(self, delta, img_size):
    # delta: (B, N*2) 控制点位移量
    B = delta.shape[0]
    N = self.grid_size ** 2
    H, W = img_size

    # 目标控制点 = 源控制点 + 位移
    P = self.P.unsqueeze(0).expand(B, 1, 1).to(delta.device) # (B, N, 2)
    delta = delta.view(B, N, 2)
    Q = P + delta # (B, N, 2)

    # 🚀 构建K矩阵: K[i,j] = ||pi – pj||^2 * log(||pi – pj||^2)
    P_expand = P.unsqueeze(1).expand(1, N, 1, 1) # (B, N, N, 2)
    P_expand_t = P.unsqueeze(2).expand(1, 1, N, 1) # (B, N, N, 2)
    dist = torch.norm(P_expand P_expand_t, dim=1) # (B, N, N)
    K = dist ** 2 * torch.log(dist + 1e-8) # (B, N, N) 公式(6)

    # 构建增广矩阵L
    ones = torch.ones(B, N, 1).to(delta.device)
    P_aug = torch.cat([ones, P], dim=1) # (B, N, 3)
    top = torch.cat([K, P_aug], dim=1) # (B, N, N+3)
    bottom = torch.cat([P_aug.transpose(1, 2),
    torch.zeros(B, 3, 3).to(delta.device)], dim=1) # (B, 3, N+3)
    L = torch.cat([top, bottom], dim=1) # (B, N+3, N+3) 公式(7)

    # 求解变换参数 W = L^{-1} * Y
    Y = torch.cat([Q, torch.zeros(B, 3, 2).to(delta.device)], dim=1) # (B, N+3, 2)
    L_inv = torch.linalg.lstsq(L, Y).solution # (B, N+3, 2)

    # 生成采样网格
    grid_y, grid_x = torch.meshgrid(
    torch.linspace(1, 1, H, device=delta.device),
    torch.linspace(1, 1, W, device=delta.device),
    indexing='ij'
    )
    grid = torch.stack([grid_x, grid_y], dim=1).view(1, 2) # (HW, 2)

    # 对每个像素计算TPS变换 公式(8)
    grid_exp = grid.unsqueeze(1).expand(1, N, 1) # (HW, N, 2)
    P_exp = P[0].unsqueeze(0).expand(grid.size(0), 1, 1) # (HW, N, 2)
    dist_grid = torch.norm(grid_exp P_exp, dim=1)
    K_grid = dist_grid ** 2 * torch.log(dist_grid + 1e-8) # (HW, N)

    ones_grid = torch.ones(grid.size(0), 1).to(delta.device)
    grid_aug = torch.cat([ones_grid, grid], dim=1) # (HW, 3)
    sampling_grid = torch.cat([K_grid, grid_aug], dim=1) # (HW, N+3)

    # 变换后坐标
    coords = sampling_grid @ L_inv[:, :, 0] # (B, HW)
    # 归一化到 [-1, 1] 用于grid_sample
    coords = coords.view(B, H, W).unsqueeze(1)
    return coords

    # ====== ES2D: 高效2D选择性扫描 ======
    class ES2D(nn.Module):
    """Efficient 2D Selective Scan — 来自EfficientVMamba"""
    def __init__(self, dim):
    super().__init__()
    from mamba_ssm import Mamba
    self.mamba = Mamba(dim, bimamba_type="v3")
    self.norm = nn.LayerNorm(dim)

    def forward(self, x):
    # x: (B, C, H, W) 或 (B, N, C)
    if x.dim() == 4:
    B, C, H, W = x.shape
    x_seq = x.flatten(2).permute(0, 2, 1) # (B, HW, C)
    else:
    x_seq = x
    out = self.mamba(self.norm(x_seq))
    if x.dim() == 4:
    out = out.permute(0, 2, 1).view(B, C, H, W)
    return out

    # ====== SGE: 空间组增强注意力 ======
    class SpatialGroupEnhance(nn.Module):
    """SGE: 分组空间注意力"""
    def __init__(self, groups=8):
    super().__init__()
    self.groups = groups
    self.avg_pool = nn.AdaptiveAvgPool2d(1)
    self.weight = nn.Parameter(torch.zeros(1, groups, 1, 1))
    self.bias = nn.Parameter(torch.ones(1, groups, 1, 1))
    self.sig = nn.Sigmoid()

    def forward(self, x):
    B, C, H, W = x.shape
    G = self.groups
    x_group = x.view(B, G, C // G, H, W)
    # 组内标准化
    mean = x_group.mean(dim=(2, 3, 4), keepdim=True)
    std = x_group.std(dim=(2, 3, 4), keepdim=True) + 1e-5
    x_norm = (x_group mean) / std
    # 组内空间注意力
    attn = self.avg_pool(x_norm.view(B, G, 1, H, W).mean(dim=2)).view(B, G, 1, 1)
    attn = self.sig(attn * self.weight + self.bias)
    return x_group * attn.unsqueeze(2)

    # ====== DEM: 差分增强模块 ======
    class DEM(nn.Module):
    """差分增强 — 保留模态特有信息"""
    def __init__(self, dim):
    super().__init__()
    self.conv = nn.Conv2d(dim, dim, 3, 1, 1)
    self.silu = nn.SiLU()

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

    # ====== SCCM: 语义相关性约束模块 ======
    class SCCM(nn.Module):
    """高层语义先验 → 逐层约束"""
    def __init__(self, dim):
    super().__init__()
    self.dem_rgb = DEM(dim)
    self.dem_t = DEM(dim)
    self.es2d = ES2D(dim)
    self.lp = nn.Linear(dim, dim)
    self.sge = SpatialGroupEnhance()
    self.dwc = nn.Conv2d(dim, dim, 3, 1, 1, groups=dim)
    self.norm = nn.LayerNorm(dim)

    def forward(self, F4_rgb, F4_t):
    # DEM增强
    E4_rgb = self.dem_rgb(F4_rgb)
    E4_t = self.dem_t(F4_t)

    B, C, H, W = E4_rgb.shape
    # 共享特征
    H_rgb = F.silu(self.dwc(self.lp(self._ln(E4_rgb))))
    H_t = F.silu(self.dwc(self.lp(self._ln(E4_t))))

    # 🚀 乘+加双路径 (公式1)
    H = H_rgb * H_t + H_rgb * H_t

    # ES2D扫描 + 门控融合 (公式2)
    H_scan = self.es2d(H)
    H1 = H_scan * F.silu(self.lp(self._ln(E4_rgb)))
    H2 = H_scan * F.silu(self.lp(self._ln(E4_t)))

    # SGF生成
    SGF = self.sge(self.lp(H1 + H2)) + self.lp(H1 * H2)
    return SGF

    def _ln(self, x):
    B, C, H, W = x.shape
    return x.flatten(2).permute(0, 2, 1)

    # ====== TPSAM: 薄板样条对齐模块 ======
    class TPSAM(nn.Module):
    """LSSM + TPS可变形对齐"""
    def __init__(self, dim, grid_size=5):
    super().__init__()
    self.lssm = ES2D(dim) # 用ES2D模拟LSSM局部扫描
    self.sge = SpatialGroupEnhance()
    self.gap = nn.AdaptiveAvgPool2d(1)
    self.fc = nn.Linear(dim * 2, grid_size * grid_size * 2)
    self.tps = TPSGridGenerator(grid_size)
    self.grid_size = grid_size

    def forward(self, F_rgb, F_t, img_size=(48, 48)):
    # LSSM增强 (公式4)
    Er = self.sge(self.lssm(F_rgb))
    Et = self.sge(self.lssm(F_t))

    # 预测控制点位移 (公式5)
    feat = torch.cat([self.gap(Er), self.gap(Et)], dim=1).flatten(1)
    delta = self.fc(feat) # (B, N*2)

    # TPS变换
    grid = self.tps(delta, img_size)
    # 应用变换到热红外特征
    At = F.grid_sample(F_t, grid, mode='bilinear', align_corners=False)
    return At

    # ====== CMCM: 跨模态相关性模块 ======
    class CMCM(nn.Module):
    """门控隐空间跨模态融合"""
    def __init__(self, dim):
    super().__init__()
    self.dem_rgb = DEM(dim)
    self.dem_t = DEM(dim)
    self.es2d_rgb = ES2D(dim)
    self.es2d_t = ES2D(dim)
    self.sge = SpatialGroupEnhance()
    self.lp = nn.Linear(dim, dim)
    self.norm = nn.LayerNorm(dim)

    def forward(self, F_rgb, At):
    # DEM增强
    D_rgb = self.dem_rgb(F_rgb)
    D_t = self.dem_t(At)

    # 隐空间投影 (公式9)
    y_rgb = self.es2d_rgb(F.silu(self.dwc_(self.lp(self._ln(D_rgb)))))
    y_t = self.es2d_t(F.silu(self.dwc_(self.lp(self._ln(D_t)))))

    # 门控信号 (公式10)
    Z_rgb = F.silu(self.lp(self._ln(D_rgb)))
    Z_t = F.silu(self.lp(self._ln(D_t)))

    # 🚀 交叉门控: RGB门控热红外, 反之亦然
    f_rgb = self.sge(self.lp(y_rgb * Z_rgb + y_t * Z_rgb)) + F_rgb
    f_t = self.sge(self.lp(y_t * Z_t + y_rgb * Z_t)) + At
    return f_rgb, f_t

    def dwc_(self, x):
    B, N, C = x.shape
    H = W = int(math.sqrt(N))
    x = x.permute(0, 2, 1).view(B, C, H, W)
    # placeholder — 实际应为深度可分离卷积
    return x.flatten(2).permute(0, 2, 1)

    def _ln(self, x):
    B, C, H, W = x.shape
    return x.flatten(2).permute(0, 2, 1)

    # ====== 解码器 ======
    class Decoder(nn.Module):
    def __init__(self):
    super().__init__()
    self.convs = nn.ModuleList([
    nn.Sequential(
    nn.Conv2d(64*2, 64, 3, 1, 1),
    nn.BatchNorm2d(64),
    nn.GELU()
    ) for _ in range(3)
    ])
    self.pred_head = nn.Sequential(
    nn.Conv2d(64, 64, 3, 1, 1),
    nn.BatchNorm2d(64),
    nn.GELU(),
    nn.Conv2d(64, 1, 1),
    )

    def forward(self, feats):
    # feats: list of (f_rgb, f_t) pairs from CMCM
    x = None
    for i, (f_rgb, f_t) in enumerate(feats):
    fused = torch.cat([f_rgb, f_t], dim=1)
    fused = self.convs[i](fused)
    if x is not None:
    fused = fused + F.interpolate(x, size=fused.shape[2:], mode='bilinear')
    x = fused
    pred = F.interpolate(self.pred_head(x), size=384, mode='bilinear')
    return pred

    # ====== TPS-SCL完整模型 ======
    class TPSSDCL(nn.Module):
    """TPS-SCL: TPS驱动的无对齐RGB-T SOD"""
    def __init__(self):
    super().__init__()
    # MobileViT-S双流编码器 (简化版)
    self.encoder_rgb = nn.ModuleList([
    nn.Sequential(
    nn.Conv2d(3 if i == 0 else C, C, 3, 2 if i > 0 else 1, 1),
    nn.BatchNorm2d(C),
    nn.GELU(),
    ) for i, C in enumerate([32, 64, 96, 128])
    ])
    self.encoder_t = nn.ModuleList([
    nn.Sequential(
    nn.Conv2d(3 if i == 0 else C, C, 3, 2 if i > 0 else 1, 1),
    nn.BatchNorm2d(C),
    nn.GELU(),
    ) for i, C in enumerate([32, 64, 96, 128])
    ])

    # SCCM (只在最高层)
    self.sccm = SCCM(128)

    # TPSAM (在2,3,4层应用)
    self.tpsam = nn.ModuleList([
    TPSAM(64), TPSAM(96), TPSAM(128)
    ])

    # CMCM (在2,3,4层应用)
    self.cmcm = nn.ModuleList([
    CMCM(64), CMCM(96), CMCM(128)
    ])

    # 投影SGF到各层
    self.sgf_proj = nn.ModuleList([
    nn.Sequential(
    nn.Conv2d(128, C, 1),
    nn.Upsample(scale_factor=2**(3i), mode='bilinear'),
    ) for i, C in enumerate([32, 64, 96, 128])
    ])

    # 解码器
    self.decoder = Decoder()

    def forward(self, rgb, thermal):
    # 双流编码
    rgb_feats, t_feats = [], []
    x_rgb, x_t = rgb, thermal.repeat(1, 3, 1, 1) if thermal.size(1) == 1 else thermal
    enc_layers = zip(self.encoder_rgb, self.encoder_t)
    for enc_rgb, enc_t in enc_layers:
    x_rgb, x_t = enc_rgb(x_rgb), enc_t(x_t)
    rgb_feats.append(x_rgb)
    t_feats.append(x_t)

    # SCCM: 高层语义约束
    SGF = self.sccm(rgb_feats[1], t_feats[1])

    # SCCM约束: 逐层引导
    constrained_rgb, constrained_t = [], []
    for i in range(4):
    proj = self.sgf_proj[i](SGF)
    constrained_rgb.append(rgb_feats[i] * proj) # 公式(3)
    constrained_t.append(t_feats[i] * proj)

    # TPSAM对齐 + CMCM融合 (2,3,4层)
    cmcm_feats = [(constrained_rgb[0], constrained_t[0])] # 第一层直接传

    for i in range(3):
    At = self.tpsam[i](constrained_rgb[i+1], constrained_t[i+1],
    img_size=constrained_t[i+1].shape[2:])
    f_rgb, f_t = self.cmcm[i](constrained_rgb[i+1], At)
    cmcm_feats.append((f_rgb, f_t))

    # 解码
    pred = self.decoder(cmcm_feats)
    return pred

    # ====== 测试模型 ======
    if __name__ == "__main__":
    model = TPSSDCL()
    rgb = torch.randn(2, 3, 384, 384)
    thermal = torch.randn(2, 1, 384, 384)
    pred = model(rgb, thermal)
    print(f"Output: {pred.shape}")
    total_params = sum(p.numel() for p in model.parameters())
    print(f"Total params: {total_params / 1e6:.2f}M")


    四、YOLO 一键迁移适配教程(即插即用,直接训练)

    Step 1:放入模块

    将 TPSAM 和 CMCM 类复制到 ultralytics/nn/modules/tpsscl.py:

    # ultralytics/nn/modules/tpsscl.py
    import torch
    import torch.nn as nn
    import torch.nn.functional as F
    from mamba_ssm import Mamba

    class TPSAM(nn.Module):
    """薄板样条对齐 — YOLO即插即用版"""
    def __init__(self, dim=64):
    super().__init__()
    self.lssm = Mamba(dim, bimamba_type="v3")
    self.norm = nn.LayerNorm(dim)

    def forward(self, x_t):
    """将热红外特征弯曲到RGB空间"""
    B, C, H, W = x_t.shape
    x_seq = x_t.flatten(2).permute(0, 2, 1)
    aligned = self.lssm(self.norm(x_seq))
    aligned = aligned.permute(0, 2, 1).view(B, C, H, W)
    return aligned + x_t

    class CMCM(nn.Module):
    """跨模态相关性融合 — YOLO即插即用版"""
    def __init__(self, dim=64):
    super().__init__()
    self.mamba_rgb = Mamba(dim, bimamba_type="v3")
    self.mamba_t = Mamba(dim, bimamba_type="v3")
    self.mamba_share = Mamba(dim, bimamba_type="v3")
    self.norm1 = nn.LayerNorm(dim)
    self.norm2 = nn.LayerNorm(dim)
    self.norm3 = nn.LayerNorm(dim)

    def forward(self, x_rgb, x_t):
    B, C, H, W = x_rgb.shape
    rs = x_rgb.flatten(2).permute(0,2,1)
    ts = x_t.flatten(2).permute(0,2,1)

    yr = self.mamba_rgb(self.norm1(rs))
    yt = self.mamba_t(self.norm2(ts))
    g = torch.sigmoid(self.mamba_share(self.norm3(rs + ts)))

    fr = yr * g + yt * g
    ft = yt * g + yr * g

    return (fr.permute(0,2,1).view(B,C,H,W),
    ft.permute(0,2,1).view(B,C,H,W))

    Step 2:注册 __init__.py

    # ultralytics/nn/modules/__init__.py
    from .tpsscl import TPSAM, CMCM

    Step 3:注册 parse_model

    # ultralytics/nn/tasks.py — parse_model函数内
    elif m in (TPSAM, CMCM):
    c2 = args[0] # dim
    args = [c2]


    五、实验结果全解析

    5.1 轻量级方法对比

    方法骨干UVT20K FmF_mFmUVT2000 FmF_mFmun-VT5000 FmF_mFmun-VT1000 FmF_mFmun-VT821 FmF_mFmParams(M)FLOPs(G)
    MoADNet MobileNet-V3 0.237 0.170 0.653 0.772 0.674 5.03 2.96
    MobileSal MobileNet-V2 0.744 0.571 0.705 0.796 0.636 6.55 2.33
    LSNet MobileNet-V2 0.707 0.558 0.757 0.853 0.746 4.57 1.23
    HENet MobileNet-S 0.745 0.573 0.847 0.893 0.835 10.43 10.75
    TPS-SCL MobileNet-S 0.815 0.632 0.859 0.908 0.846 12.82 12.34

    ✅ 核心亮点:

    • UVT20K上F-measure达0.815,超HENet(0.745)达7.0%,超MobileSal达7.1%
    • UVT2000上F-measure 0.632,超HENet(0.573)达5.9%
    • 在未对齐数据集上的优势远大于对齐数据集,证明TPSAM的TPS可变形对齐确实有效

    5.2 重量级方法对比(Swin-B骨干)

    方法骨干UVT20K FmF_mFmUVT2000 FmF_mFmVT5000 FmF_mFmVT1000 FmF_mFmVT821 FmF_mFm
    SwinNet Swin-B 0.737 0.579 0.846 0.947 0.818
    TCINet Swin-B 0.832 0.699 0.876 0.909 0.852
    PCNet Swin-B 0.827 0.691 0.899 0.926 0.879
    SACNet Swin-B 0.689 0.594 0.888 0.958 0.859
    TPS-SCL Swin-B 0.848 0.702 0.902 0.921 0.883

    ✅ 核心亮点:

    • UVT20K上F-measure达0.848,超PCNet(0.827)达2.1%
    • UVT2000上F-measure 0.702,超PCNet(0.691)达1.1%
    • 重量级版本(Swin-B)在VT5000对齐数据集上也达0.902 F-measure,证明TPS-SCL在各类场景普遍有效

    5.3 消融实验

    模型UVT20K Fm/Sm/EmF_m/S_m/E_mFm/Sm/EmUVT2000 Fm/Sm/EmF_m/S_m/E_mFm/Sm/Em
    TPS-SCL (完整) 0.815/0.866/0.887 0.632/0.794/0.792
    w/o SCCM 0.022/0.431/0.516 0.024/0.465/0.625
    w/o TPSAM 0.625/0.792/0.763 0.498/0.735/0.707
    w/o CMCM 0.763/0.804/0.831 0.560/0.710/0.684

    ✅ 核心亮点:

    • SCCM最关键:去掉SCCM后UVT20K上F-measure从0.815暴跌至0.022——没了高层语义约束,无对齐场景下直接崩掉
    • TPSAM贡献19%:去掉TPSAM后UVT20K上F-measure降至0.625(↓19%),验证TPS可变形对齐对无对齐场景不可或缺
    • CMCM贡献6.4%:去掉CMCM后F-measure降至0.763(↓6.4%),说明跨模态门控融合仍有显著增益

    六、总结

  • 首个TPS驱动的无对齐RGB-T SOD框架:突破传统仿射变换/单应性估计的局限性,用薄板样条可变形变换处理真实世界的大偏移+非线性形变
  • 三阶段"约束→对齐→融合"流水线:SCCM高层语义屏蔽背景噪声、TPSAM可变形对齐消除空间差异、CMCM门控隐空间深度融合,三模块缺一不可
  • 轻量级SOTA:仅12.82M参数+12.34G FLOPs,在UVT20K真实未对齐数据集上F-measure达0.815,超HENet 7%,同时适配Swin-B/PVT-v2等重型骨干
  • 与Mamba系列方法的深度融合:ES2D(EfficientVMamba)和LSSM(LocalMamba)两种Mamba变体分别负责全局建模和局部增强,为后续SSM-based RGB-T方法奠定基础
  •   学术研究和工程落地都能直接用。

    🔖 收藏本文,RGB-T 无对齐检测直接起飞!
    📌 标签:#RGB-T SOD #TPS薄板样条 #无对齐检测 #AAAI 2026 #Mamba门控融合

    赞(0)
    未经允许不得转载:171主机测评 » TPS薄板样条对齐RGB-T SOD:TPS-SCL AAAI2026 无对齐SOTA,SCCM约束+TPSAM对齐+CMCM融合!!!
    分享到: 更多 (0)

    评论 抢沙发

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