当前位置:首页 > 技术 > 正文内容

知识蒸馏中的回顾机制:ReviewKD详解

访客 技术 2026年9月2日 2

在深度学习模型压缩技术中,知识蒸馏(Knowledge Distillation, KD)扮演着重要角色,它旨在将大型教师模型的知识迁移到小型学生模型中。然而,以往的KD方法大多集中于特征变换或要求教师与学生模型在相同层级上进行知识传递。本文介绍的Knowledge Review(KR)方法则另辟蹊径,重点研究教师与学生模型之间不同层级的连接方式,提出了一种创新的知识回顾(Knowledge Review)机制,能够有效提升学生模型的性能。

知识蒸馏回顾

知识蒸馏的概念最早可追溯到Hinton等人提出的利用教师模型与学生模型输出层(logits)的KL散度进行优化。随后,FitNets方法将蒸馏范围扩展至中间层特征,通常通过最小化学生与教师模型特征图之间的均方误差(MSE)来实现。Attention Transfer在此基础上,提出利用注意力图作为知识传递的引导。PKT将知识建模为概率分布,而Contrastive Representation Distillation (CRD) 则引入了对比学习。与上述方法不同,KR关注的是如何选择教师与学生模型之间的连接路径。

如图所示,传统的KD方法(a-c)通常使用教师与学生模型相同层级的特征进行监督。而KR方法(d)则允许学生模型的深层特征利用教师模型浅层特征进行学习,并实验证明,学生模型的深层可以有效学习到教师模型浅层的知识。

这是因为模型在不同层级所提取的知识抽象程度和学习难度是不同的。学生模型如果在早期就能从教师模型的浅层特征中获得知识,将对其整体性能大有裨益。KR将浅层知识视为"旧知识",并主张通过"回顾"来"温故知新"。因此,如何从教师模型中提取多尺度的信息成为KR的关键问题,为此,KR提出了两个核心模块:

  • Attention-based fusion (ABF): 用于特征融合。
  • Hierarchical context loss (HCL): 用于增强模型的学习能力。

Knowledge Review (KR)

形式化描述

设输入图像为 X,学生网络表示为 S,其各层输出为 (S_1, S_2, ..., S_n, S_c)。整个学生网络的输出为 Y_s。各中间层输出为 (F_s^1, ..., F_s^n)

单层知识蒸馏可以表示为:

L_SKD = D(M_s^i(F_s^i), M_t^i(F_t^i))

其中,M 是一个转换模块,用于匹配学生与教师特征图的维度,D 是衡量分布差异的损失函数。多层知识蒸馏表示为:

L_MKD = sum_{i in I} D(M_s^i(F_s^i), M_t^i(F_t^i))

上述公式表示学生与教师模型层层对应。而KR的单层表示方式为:

L_SKD_R = sum_{j=1}^{i} D(M_s^{i,j}(F_s^i), M_t^{j,i}(F_t^j))

这表示第 i 层学生网络的学习需要回顾从第 1 层到第 i 层的所有教师知识。同理,多层的KR表示为:

L_MKD_R = sum_{i in I} (sum_{j=1}^{i} D(M_s^{i,j}(F_s^i), M_t^{j,i}(F_t^j)))

融合方式设计

KR机制要求学生网络的每一层都能回顾教师网络之前的所有层。一种直观的实现方式是直接缩放学生网络最后一层的特征图以匹配教师网络的维度,并通过卷积层和插值层完成形状匹配。这种方式旨在让学生网络特征更接近教师网络。

然而,直接对所有层进行形状匹配可能会导致各阶段之间产生巨大差异,并且计算成本较高。为了提高效率和可行性,KR引入了Attention-based fusion (ABF) 模块。引入ABF后,整体蒸馏过程可以表示为:

sum_{i=j}^{n} D(F_s^i, F_t^j) approx D(U(F_s^j, ..., F_s^n), F_t^j)

ABF模块的设计如图所示(a),它采用注意力机制融合特征。中间的1x1卷积用于提取两个层级特征的联合空间注意力图,然后通过特征重标定(feature re-calibration)实现特征融合,类似于SKNet的空间注意力版本。

ABF具体实现中,首先对学生特征进行变换(conv1),然后,如果启用了注意力融合(att_conv),则会上采样教师的残差特征(y),并与变换后的学生特征(x)进行拼接。通过注意力模块(att_conv)计算注意力权重,对学生与教师特征进行加权融合。最后,通过conv2进行进一步处理,并根据需要进行插值以匹配输出形状。

HCL (Hierarchical Context Loss) 模块则对来自学生网络和教师网络的特征分别进行空间金字塔池化(SPP)处理,并使用L2距离来衡量它们之间的差异。KR认为这种方式能够捕获不同层级的语义信息,并提取不同抽象层级的信息。

实验

实验部分主要进行了消融实验,以验证不同组件的有效性。

第一个消融实验关注不同层级监督的有效性。结果表明,使用教师网络的浅层知识来监督学生网络的深层知识是有效的,能够提升模型性能。

第二个消融实验验证了ABF和HCL模块的作用。实验结果表明,ABF和HCL模块都能在不同程度上提升模型性能,其中ABF模块的效果尤为显著。

代码实现

ABF实现


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

class ABF(nn.Module):
    def __init__(self, in_channel, mid_channel, out_channel, fuse):
        super(ABF, self).__init__()
        # 变换学生特征
        self.conv1 = nn.Sequential(
            nn.Conv2d(in_channel, mid_channel, kernel_size=1, bias=False),
            nn.BatchNorm2d(mid_channel),
        )
        # 输出层
        self.conv2 = nn.Sequential(
            nn.Conv2d(mid_channel, out_channel, kernel_size=3, stride=1, padding=1, bias=False),
            nn.BatchNorm2d(out_channel),
        )
        # 注意力融合模块
        if fuse:
            self.att_conv = nn.Sequential(
                nn.Conv2d(mid_channel * 2, 2, kernel_size=1),
                nn.Sigmoid(),
            )
        else:
            self.att_conv = None
        
        # 初始化权重
        nn.init.kaiming_uniform_(self.conv1[0].weight, a=1)
        nn.init.kaiming_uniform_(self.conv2[0].weight, a=1)

    def forward(self, x, y=None, teacher_shape=None, student_out_shape=None):
        # x: 学生特征, y: 教师特征 (用于融合)
        n, _, h, w = x.shape
        
        # 变换学生特征
        transformed_x = self.conv1(x)
        
        fused_feature = transformed_x
        if self.att_conv is not None and y is not None and teacher_shape is not None:
            # 上采样教师特征以匹配形状
            upsampled_y = F.interpolate(y, (teacher_shape, teacher_shape), mode="nearest")
            # 特征融合
            combined_features = torch.cat([transformed_x, upsampled_y], dim=1)
            attention_weights = self.att_conv(combined_features)
            
            # 应用注意力权重
            # attention_weights[:, 0] for transformed_x, attention_weights[:, 1] for upsampled_y
            fused_feature = (transformed_x * attention_weights[:, 0].view(n, 1, h, w) + 
                             upsampled_y * attention_weights[:, 1].view(n, 1, h, w))
        
        # 调整最终输出形状
        if fused_feature.shape[-1] != student_out_shape:
            fused_feature = F.interpolate(fused_feature, (student_out_shape, student_out_shape), mode="nearest")
            
        # 输出层处理
        output_feature = self.conv2(fused_feature)
        
        return output_feature, fused_feature # 返回处理后的特征和融合后的特征

    

HCL实现


def hierarchical_context_loss(student_features, teacher_features):
    # student_features 和 teacher_features 都是列表,包含各个 stage 的特征
    total_loss = 0.0
    for fs, ft in zip(student_features, teacher_features):
        # 计算初始 MSE Loss
        loss = F.mse_loss(fs, ft, reduction='mean')
        
        # 渐进式池化并计算 Loss
        scale_factor = 1.0 # 用于加权
        total_scale = 1.0
        for pool_size in [4, 2, 1]:
            if pool_size >= fs.shape[-2]: # 确保池化尺寸不大于特征图尺寸
                continue
            
            pooled_fs = F.adaptive_avg_pool2d(fs, (pool_size, pool_size))
            pooled_ft = F.adaptive_avg_pool2d(ft, (pool_size, pool_size))
            
            scale_factor /= 2.0 # 减小后续池化层的权重
            loss += F.mse_loss(pooled_fs, pooled_ft, reduction='mean') * scale_factor
            total_scale += scale_factor
            
        loss = loss / total_scale # 归一化总 Loss
        total_loss += loss
        
    return total_loss

    

ReviewKD实现


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

# 假设 ABF 类已定义

class ReviewKD(nn.Module):
    def __init__(self, student_model, teacher_in_channels, student_out_channels, 
                 student_feature_shapes, student_output_shapes=None):
        super(ReviewKD, self).__init__()
        self.student_model = student_model
        self.student_feature_shapes = student_feature_shapes
        # 如果未提供学生输出形状,则默认与特征形状相同
        self.student_output_shapes = student_feature_shapes if student_output_shapes is None else student_output_shapes

        # 初始化ABF模块列表,注意反转顺序以匹配从深层到浅层的处理
        self.abf_modules = nn.ModuleList()
        # 确定中间通道数,可以根据实际情况调整
        mid_channel = min(512, teacher_in_channels[-1]) 
        
        for idx, in_channel in enumerate(teacher_in_channels):
            # fuse 参数设置为 True,表示在非最后一层 ABF 模块中进行融合
            fuse_flag = idx < len(teacher_in_channels) - 1
            self.abf_modules.append(ABF(in_channel, mid_channel, student_out_channels[idx], fuse=fuse_flag))
        
        # 反转列表,使得 ABF 模块按从深层到浅层的顺序应用
        self.abf_modules = self.abf_modules[::-1]
        
        # 移动模型到 GPU
        self.to('cuda')

    def forward(self, input_tensor):
        # 获取学生模型的特征图和最终 logits
        # 假设 student_model 支持 is_feat 参数返回所有中间特征
        student_features, logits = self.student_model(input_tensor, is_feat=True)
        
        # 反转学生特征列表,使其从浅层到深层排列
        student_features_reversed = student_features[::-1]
        
        results = [] # 存储 KR 蒸馏的特征
        
        # 应用第一个 ABF 模块 (对应学生最深层)
        # 传入学生特征、教师特征(这里暂时用学生反转后的第一个特征,实际应为教师特征)、形状等
        # 注意:这里的教师特征输入需要根据实际的 KR 实现进行调整,通常是从教师模型对应层获取
        # 示例中,我们暂时使用学生特征作为占位符,实际应替换为教师特征
        current_teacher_feature = student_features_reversed[0] # 占位符,实际应为教师模型的浅层特征
        
        # ABF(in_channel, mid_channel, out_channel, fuse)
        # 对应学生模型的深层特征
        processed_feature, fused_residual = self.abf_modules[0](
            x=student_features_reversed[0], 
            y=None, # 教师特征 Y 暂不传入,实际根据 KR 实现补充
            teacher_shape=self.student_feature_shapes[0], # 教师特征形状,这里用学生第一个特征形状作为占位
            student_out_shape=self.student_output_shapes[0]
        )
        results.append(processed_feature) # 记录处理后的特征

        # 迭代应用剩余的 ABF 模块
        for i in range(1, len(self.abf_modules)):
            current_student_feature = student_features_reversed[i]
            # 这里的教师特征也需要根据实际 KR 实现进行调整
            # 假设 KR 使用的是教师模型 i-1 层特征与融合残差进行计算
            # 示例中,我们使用 fused_residual 作为教师特征的占位符
            
            processed_feature, fused_residual = self.abf_modules[i](
                x=current_student_feature, 
                y=fused_residual, # 使用上一层 ABF 的融合残差作为教师特征输入
                teacher_shape=self.student_feature_shapes[i], # 教师特征形状,这里用学生第 i 个特征形状作为占位
                student_out_shape=self.student_output_shapes[i]
            )
            results.insert(0, processed_feature) # 插入到结果列表的开头

        # 返回蒸馏后的特征列表和最终 logits
        return results, logits

    

相关文章

Linux crontab 详解

1) crontab 是什么cron 是 Linux 的定时任务守护进程;crontab 是用来编辑/查看“按时间周期执行命令”的表(cron table)。常见两类:用户 crontab:每个用户一份(crontab -e 编辑)系统级 crontab / cron.d:可指定执行用户(/etc/crontab、/etc/cron.d/*)2) crontab 时间...

富文本里可以允许的 HTML 属性

一、所有标签默认允许的安全属性(极少)class        (可选)id           (通常建议禁用)title️ 注意:id 容易被滥用做锚点注入,很多系统直接禁用class 允许的话最好只允许固定前缀(如 editor-*)二、a 标签允许属性<a href="" t...

Mac 安装 Node.js 指南

方法一:通过官网安装包(最简单,适合初学者)如果你只是想快速安装并开始使用,这是最直接的方法。访问 Node.js 官网。页面会显示两个版本:LTS (Recommended For Most Users):长期支持版,最稳定。建议选这个。Current:最新特性版,包含最新功能但可能不够稳定。下载 .pkg 安装包并运行。按照安装向导点击“下一步”即可完成。方法二:使用 Homebrew 安装(...

Dom\HTML_NO_DEFAULT_NS 的副作用:自动加闭合标签

在使用Dom\HTMLDocument时,Dom\HTML_NO_DEFAULT_NS 将禁止在解析过程中设置元素的命名空间, 此设置是为了与DOMDocument向后兼容而存在的。当使用它时,已知的一个副作用就是:自动加闭合标签例如 </img> 为什么会这样?当你使用:Dom\HTML_NO_DEFAULT_NS文档会变成 无命名空间模式,此时内部更接近 XML...

Laravel 事件和监听器创建

在 Laravel 中,使用 Artisan 命令创建 Events(事件) 和 Listeners(监听器) 是非常高效的。你可以通过以下几种方式来实现:1. 手动创建单个 Event如果你只想创建一个事件类,可以使用 make:event 命令:Bashphp artisan make:event UserRegistered执行后,文件将生成在 app/Even...

自定义域名解析神器 dnsmasq

什么是 dnsmasq?dnsmasq 是一个轻量级、功能强大的网络服务工具,专为小型和中等规模网络设计。它是一个综合的网络基础设施解决方案[1]。dnsmasq 能做什么?功能说明应用场景DNS 转发与缓存将 DNS 查询转发到上游服务器(ISP、Google DNS 等),并在本地缓存结果加快 DNS 查询速度,减少外部 DNS 流量本地 DNS解析本地网络设备的主机名,无需编辑&n...

发表评论

访客

◎欢迎参与讨论,请在这里发表您的看法和观点。