知识蒸馏中的回顾机制:ReviewKD详解
在深度学习模型压缩技术中,知识蒸馏(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