一、论文信息本文目录一、论文信息二、论文摘要概况三、DHOGSA 模块结构图四、DHOGSA 模块的作用五、DHOGSA 模块的原理六、DHOGSA 模块的优势七、即插即用模块代码论文题目Gradient as Conditions: Rethinking HOG for All-in-one Image Restoration中文题目梯度作为条件重新思考HOG算法在全功能图像复原中的应用论文链接https://arxiv.org/abs/2504.09377所属单位北京大学计算机科学学院二、论文摘要概况全功能图像恢复AIR旨在通过利用具有信息量的退化条件来指导恢复过程从而在一个统一的模型中处理多种不同的退化问题。然而现有方法往往依赖于隐式学习的先验知识这可能会导致特征表示相互干扰并在复杂或未见过的场景下影响性能。作为经典的梯度表征方法方向梯度直方图HOG展现出强大的跨不同退化类型的判别能力使其成为AIR领域一个强大且可解释的先验模型。基于这一洞察我们提出了HOGformer——一个基于Transformer的模型它整合了可学习的HOG特征以实现退化感知的图像恢复。HOGformer的核心是动态HOG感知自注意力DHOGSA机制该机制能够根据HOG描述符编码的退化特定线索自适应地建模长程空间依赖关系。为进一步适应AIR中退化的异质性我们提出了动态交互前馈DIFF模块该模块促进通道与空间之间的相互作用从而能够在各种退化条件下实现鲁棒的特征转换。此外我们还引入了HOG损失函数以显式增强结构保真度和边缘锐度。我们在包括恶劣天气和自然退化在内的多种基准数据集上进行了广泛的实验结果表明HOGformer不仅实现了最先进的性能还能很好地泛化至复杂的现实场景。不同退化条件下HOG特征分布的可视化结果。(a) 不同天气条件下的示例图像及其对应的HOG特征可视化结果(b) 基于每种退化类型随机选取的100张图像展示五种自然退化Zheng et al.2024及三种恶劣天气退化Sun et al.2024对应的HOG特征。针对多种降解条件雨-雾-雪的对比分析上图Histformer与本方法在 PSNR 和 SSIM 指标上的对比结果下图代表性示例的可视化结果。三、DHOGSA 模块结构图模块结构如图 1 所示论文 Fig. 2 原图即 HOGformer 整体架构与 HOG Transformer Block其中 DHOGSA 为 HOGTB 的核心自注意力模块。图 1 论文 Fig. 2 原图HOGformer 整体架构与 HOGTB含 DHOGSA 与 DIFF 模块结构图3HOG引导机制(a) 神经元级分类(b) 像素级分类(c) HOG特征提取过程模块内部组件与数据流1.HOG 特征提取用 Sobel 算子求方向梯度 gx、gy计算幅值 m√(gx²gy²) 与方向 o⌊(atan2(gy,gx)π)/(2π/N_bin)⌋并按 N_bin9 个方向分箱得到 HOG 描述子。2.LDRConv局部动态范围卷积将输入 F 拆分为 F1、F2仅对 F1 做 patch 级 HOG 排序并用可学习 bin 级 HOG 先验调制F1Sort_patch(F1)HOG_θ(F1)再与 F2 拼接后经 1×1 逐点卷积和 3×3 深度卷积。3.HOG 引导像素级排序基于信息丰富的 V 特征生成排序索引 idxSort(reshape(o(V)·m(V)))并用同一索引一致地排序 Q、K、V聚合受同种退化影响的远距离像素。4.双分支 HOG 重整形自注意力BHOGRbin-wise固定大小分箱捕获雾等大尺度结构与 FHOGRfrequency-wise按 HOG 值聚类捕获雨线等细尺度纹理两个并行注意力分支输出按 Hadamard 积融合FA_B·R_B(V) ⊙ A_F·R_F(V)再按索引 scatter 还原空间顺序。5.与 DIFF 组成 HOGTBF_l F_{l-1} DHOGSA(LN(F_{l-1}))F_l F_l DIFF(LN(F_l))DHOGSA 负责退化感知的全局空间聚合DIFF 负责通道-空间交互的特征精炼。四、DHOGSA 模块的作用1.退化感知的自注意力利用 HOG 梯度先验调整注意力权重使模型自适应地强调退化敏感区域缓解传统固定窗口/通道式注意力对非均匀退化建模不足的问题。2.全局与局部长程依赖建模像素级排序聚合同退化像素的长程依赖BHOGR 捕获大尺度结构伪影如雾的均匀分布FHOGR 捕获细尺度重复退化如雨线。3.增强注意力输入LDRConv 以动态范围重组局部特征在保持全局空间结构一致性的同时提升退化敏感度为自注意力提供更强的特征基础。4.作为 HOGformer 每个 HOGTB 的核心模块与 DIFF 配合实现一体化图像复原在恶劣天气与自然退化等多个基准上取得 SOTA 性能。五、DHOGSA 模块的原理给定输入特征 F ∈ R^(C×H×W)DHOGSA 按以下步骤计算对应论文式 2-41.HOG 特征提取式 2用 Sobel 滤波器计算方向梯度 gx、gy得幅值 m√(gx²gy²) 与方向 o⌊(atan2(gy,gx)π)/(2π/N_bin)⌋将方向量化到 N_bin 个 bin。2.LDRConv式 3拆分 F 为 F1、F2对 F1 做 patch 级 HOG 排序并叠加可学习 bin 级 HOG 先验 HOG_θ(F1)拼接后经 1×1 逐点卷积与 3×3 深度卷积F Conv3×3^d(Conv1×1^p(Concat(Sort_patch(F1)HOG_θ(F1), F2)))。3.像素级排序与统一 Gather式 4idx Sort(reshape(o(V)·m(V)))对 Q、K、V 按同一索引 Gather 排序使同退化像素在注意力中彼此邻近。4.双分支注意力与还原BHOGRR_B与 FHOGRR_F两种重整形分支分别计算注意力Q/K 经 L2 归一化、温度缩放与 softmax-1输出 A_B·R_B(V) 与 A_F·R_F(V) 做 Hadamard 积最后按 idx scatter 还原到原始空间位置。六、DHOGSA 模块的优势1.可解释的退化先验HOG 梯度特征天然判别不同退化类型比 prompt 条件与灰度直方图更可靠模糊图与清晰图的直方图差异极小而梯度分布差异显著。2.轻量高性能HOGformer-L 在 AllWeather 五任务去雪/去雨/去雾/去雨滴平均 PSNR 达 33.99 dB刷新 SOTA超越 Histoformer33.67 dB且参数量显著低于同类 AIR 方法约 1/10。3.双路径建模patch 级 HOG 排序保持全局结构一致像素级排序聚合同退化长程依赖BHOGR 与 FHOGR 双分支兼顾粗粒度大尺度结构与细粒度局部纹理。4.即插即用模块可替换任意 Transformer 的自注意力部分独立实现仅依赖 PyTorch便于嵌入复原、检测、分割等网络。5.泛化性强在真实复杂退化数据集Practical、RealBlur、HIDE 等上表现出良好的跨退化泛化能力。七、即插即用模块代码# -*- coding: utf-8 -*- # DHOGSA (HOGformer, AAAI 2026, arXiv:2504.09377) | 参考官方开源实现 github.com/Fire-friend/HOGformer import torch import torch.nn as nn import torch.nn.functional as F class LayerNorm(nn.Module): def __init__(self, dim): super().__init__() def forward(self, x): b, c, h, w x.shape x x.flatten(2).transpose(1, 2) mu x.mean(-1, keepdimTrue) sigma x.var(-1, keepdimTrue, unbiasedFalse) x (x - mu) / torch.sqrt(sigma 1e-5) return x.transpose(1, 2).reshape(b, c, h, w) class ElementScale(nn.Module): def __init__(self, dims, init_value0.): super().__init__() self.scale nn.Parameter(init_value * torch.ones((1, dims, 1, 1))) def forward(self, x): return x * self.scale class FFN_DIFF(nn.Module): Dynamic Interaction Feed-Forward通道-空间交互增强对不同退化的适应性 def __init__(self, dim, ffn_expansion_factor2.667, biasFalse): super().__init__() hidden int(dim * ffn_expansion_factor) self.sigma ElementScale(hidden // 4, init_value1e-5) self.decompose nn.Conv2d(hidden // 4, 1, 1) self.decompose_act nn.GELU() self.project_in nn.Conv2d(dim, hidden * 2, 1, biasbias) self.dwconv_5 nn.Conv2d(hidden // 4, hidden // 4, 5, padding2, groupshidden // 4, biasbias) self.dwconv_dilated2_1 nn.Conv2d(hidden // 4, hidden // 4, 3, padding2, groupshidden // 4, biasbias, dilation2) self.p_unshuffle nn.PixelUnshuffle(2) self.p_shuffle nn.PixelShuffle(2) self.project_out nn.Conv2d(hidden, dim, 1, biasbias) def feat_decompose(self, x): return x self.sigma(x - self.decompose_act(self.decompose(x))) def forward(self, x): x self.p_shuffle(self.project_in(x)) x1, x2 x.chunk(2, dim1) x F.mish(self.dwconv_dilated2_1(x2)) * self.dwconv_5(x1) x self.p_unshuffle(self.feat_decompose(x)) return self.project_out(x) class DHOGSA(nn.Module): Dynamic HOG-aware Self-AttentionHOG 梯度先验引导的动态自注意力 def __init__(self, dim, num_heads, biasFalse, patch_size8, n_bins9): super().__init__() self.factor num_heads self.num_heads num_heads self.temperature nn.Parameter(torch.ones(num_heads, 1, 1)) self.qkv nn.Conv2d(dim, dim * 5, 1, biasbias) self.qkv_dwconv nn.Conv2d(dim * 5, dim * 5, 3, padding1, groupsdim * 5, biasbias) self.project_out nn.Conv2d(dim, dim, 1, biasbias) self.bin_proj nn.Conv2d(n_bins, dim // 2, 1, biasbias) self.patch_size patch_size self.n_bins n_bins sobel_x torch.tensor([[-1, 0, 1], [-2, 0, 2], [-1, 0, 1]], dtypetorch.float32).view(1, 1, 3, 3) sobel_y torch.tensor([[-1, -2, -1], [0, 0, 0], [1, 2, 1]], dtypetorch.float32).view(1, 1, 3, 3) self.register_buffer(sobel_x, sobel_x.repeat(dim, 1, 1, 1)) self.register_buffer(sobel_y, sobel_y.repeat(dim, 1, 1, 1)) def pad(self, x, factor): hw x.shape[-1] t_pad [0, 0] if hw % factor 0 else [0, (hw // factor 1) * factor - hw] return F.pad(x, t_pad, constant, 0), t_pad def unpad(self, x, t_pad): hw x.shape[-1] return x[:, :, t_pad[0]:hw - t_pad[1]] def softmax_1(self, x, dim-1): logit x.exp() return logit / (logit.sum(dim, keepdimTrue) 1) def reshape_attn(self, q, k, v, ifBox): b, c q.shape[:2] q, t_pad self.pad(q, self.factor) k, _ self.pad(k, self.factor) v, _ self.pad(v, self.factor) hw q.shape[-1] // self.factor heads, c_ch self.num_heads, c // self.num_heads if ifBox: # b (head c) (factor hw) - b head (c factor) hw q q.view(b, heads, c_ch, self.factor, hw).reshape(b, heads, c_ch * self.factor, hw) k k.view(b, heads, c_ch, self.factor, hw).reshape(b, heads, c_ch * self.factor, hw) v v.view(b, heads, c_ch, self.factor, hw).reshape(b, heads, c_ch * self.factor, hw) else: # b (head c) (hw factor) - b head (c factor) hw q q.view(b, heads, c_ch, hw, self.factor).permute(0, 1, 2, 4, 3).reshape(b, heads, c_ch * self.factor, hw) k k.view(b, heads, c_ch, hw, self.factor).permute(0, 1, 2, 4, 3).reshape(b, heads, c_ch * self.factor, hw) v v.view(b, heads, c_ch, hw, self.factor).permute(0, 1, 2, 4, 3).reshape(b, heads, c_ch * self.factor, hw) q F.normalize(q, dim-1) k F.normalize(k, dim-1) attn (q k.transpose(-2, -1)) * self.temperature attn self.softmax_1(attn, dim-1) out attn v if ifBox: out out.reshape(b, heads, c_ch, self.factor, hw).reshape(b, heads * c_ch, self.factor * hw) else: out out.reshape(b, heads, c_ch, self.factor, hw).permute(0, 1, 2, 4, 3).reshape(b, heads * c_ch, hw * self.factor) return self.unpad(out, t_pad) def split_into_patches(self, x): b, c, h, w x.shape pad_h (self.patch_size - h % self.patch_size) % self.patch_size pad_w (self.patch_size - w % self.patch_size) % self.patch_size if pad_h 0 or pad_w 0: x F.pad(x, (0, pad_w, 0, pad_h)) p self.patch_size n_h, n_w (h pad_h) // p, (w pad_w) // p patches x.view(b, c, n_h, p, n_w, p).permute(0, 2, 4, 1, 3, 5).reshape(b, n_h * n_w, c, p * p) return patches, (b, c, h, w, pad_h, pad_w, n_h, n_w) def merge_patches(self, patches, shape_info): b, c, h, w, pad_h, pad_w, n_h, n_w shape_info p self.patch_size x patches.view(b, n_h, n_w, c, p, p).permute(0, 3, 1, 4, 2, 5).reshape(b, c, n_h * p, n_w * p) if pad_h 0 or pad_w 0: x x[:, :, :h, :w] return x def apply_hog_to_patch(self, x_half): b, c, h, w x_half.shape gx F.conv2d(x_half, self.sobel_x[:c], padding1, groupsc) gy F.conv2d(x_half, self.sobel_y[:c], padding1, groupsc) magnitude torch.sqrt(gx ** 2 gy ** 2 1e-6) orientation torch.atan2(gy, gx) orientation_bin ((orientation torch.pi) / (2 * torch.pi) * self.n_bins).long() % self.n_bins patches_x, shape_info self.split_into_patches(x_half) patches_mag, _ self.split_into_patches(magnitude) patches_ori, _ self.split_into_patches(orientation_bin.float()) b, n_patches, c, patch_pixels patches_x.shape sort_values torch.zeros_like(patches_x) hog_features torch.zeros(b, n_patches, self.n_bins, devicex_half.device) for i in range(self.n_bins): bin_mask (patches_ori i).float() bin_magnitude patches_mag * bin_mask sort_values bin_magnitude * (i 1) hog_features[..., i] bin_magnitude.mean(dim[-1, -2]) hog_features hog_features / (hog_features.sum(dim-1, keepdimTrue) 1e-8) _, sort_indices sort_values.sum(dim2, keepdimTrue).expand_as(patches_x).sort(dim-1) patches_x_sorted torch.gather(patches_x, -1, sort_indices) return self.merge_patches(patches_x_sorted, shape_info), sort_indices, hog_features, shape_info def forward(self, x): b, c, h, w x.shape half_c c // 2 x_half x[:, :half_c] x_half_processed, idx_patch, hog_features, shape_info self.apply_hog_to_patch(x_half) b, n_patches, n_bins hog_features.shape n_h, n_w shape_info[-2], shape_info[-1] hog_map hog_features.permute(0, 2, 1).view(b, n_bins, n_h, n_w).contiguous() hog_map self.bin_proj(hog_map) hog_map F.interpolate(hog_map, size(h, w), modebilinear) x torch.cat((x_half_processed hog_map, x[:, half_c:]), dim1) qkv self.qkv_dwconv(self.qkv(x)) q1, k1, q2, k2, v qkv.chunk(5, dim1) gx F.conv2d(v, self.sobel_x[:c], padding1, groupsc) gy F.conv2d(v, self.sobel_y[:c], padding1, groupsc) magnitude torch.sqrt(gx ** 2 gy ** 2 1e-6).view(b, c, -1) orientation torch.atan2(gy, gx).view(b, c, -1) weighted_magnitude magnitude * (orientation torch.pi) / (2 * torch.pi) _, idx weighted_magnitude.sum(dim1).sort(dim-1) idx idx.unsqueeze(1).expand(b, c, -1) v torch.gather(v.view(b, c, -1), 2, idx) q1 torch.gather(q1.view(b, c, -1), 2, idx) k1 torch.gather(k1.view(b, c, -1), 2, idx) q2 torch.gather(q2.view(b, c, -1), 2, idx) k2 torch.gather(k2.view(b, c, -1), 2, idx) out1 self.reshape_attn(q1, k1, v, True) out2 self.reshape_attn(q2, k2, v, False) out1 torch.scatter(out1, 2, idx, out1).view(b, c, h, w) out2 torch.scatter(out2, 2, idx, out2).view(b, c, h, w) out self.project_out(out1 * out2) out_replace out[:, :half_c] patches_out, shape_info self.split_into_patches(out_replace) patches_out torch.scatter(patches_out, -1, idx_patch, patches_out) out_replace self.merge_patches(patches_out, shape_info) out[:, :half_c] out_replace return out if __name__ __main__: input torch.randn(2, 32, 128, 128) model DHOGSA(dim32, num_heads4, biasFalse) print(model) print(CSDN:AI魔改博士) output model(input) print(DHOGSA input_size:, input.size()) print(DHOGSA output_size:, output.size())