预训练CNN推理阶段跳过零乘法操作的实现方法咨询
不用重训练!推理时跳过CNN零运算的简易实现
针对MNIST预训练CNN的推理需求,下面给出纯PyTorch层面的简易修改方案,不用碰C++代码,也不需要重新训练模型,专门用来观察不同稀疏度下的推理时间差异。
一、全连接层(FC)的零运算跳过
全连接层本质是矩阵乘法,输入里的零元素乘权重结果还是零,完全可以跳过计算。我们直接继承PyTorch的nn.Linear类,重写前向传播逻辑,只处理非零元素:
import torch import torch.nn as nn class SparseLinear(nn.Linear): def forward(self, x): # x的形状: (batch_size, 输入特征数) batch_size = x.size(0) # 找出每个样本里的非零元素索引 non_zero_indices = x.nonzero(as_tuple=True) if len(non_zero_indices[0]) == 0: # 输入全是零,直接返回偏置 return self.bias.expand(batch_size, -1) # 提取非零的输入值和对应的权重列 x_nonzero = x[non_zero_indices] w_selected = self.weight[:, non_zero_indices[1]].T # 初始化输出并累加非零元素的贡献 output = torch.zeros(batch_size, self.out_features, device=x.device) output.index_add_(0, non_zero_indices[0], x_nonzero.unsqueeze(1) * w_selected) # 加上偏置 output += self.bias return output
使用方法
把预训练模型里的普通全连接层替换成这个稀疏版本,再加载原权重即可:
# 假设你的原模型是这样的 class OriginalCNN(nn.Module): def __init__(self): super().__init__() self.conv1 = nn.Conv2d(1, 32, kernel_size=3) self.conv2 = nn.Conv2d(32, 64, kernel_size=3) self.pool = nn.MaxPool2d(2) self.fc1 = nn.Linear(64*12*12, 128) self.fc2 = nn.Linear(128, 10) def forward(self, x): x = torch.relu(self.conv1(x)) x = self.pool(torch.relu(self.conv2(x))) x = x.flatten(1) x = torch.relu(self.fc1(x)) x = self.fc2(x) return x # 加载预训练模型 model = OriginalCNN() model.load_state_dict(torch.load('pretrained_mnist_cnn.pth')) # 替换全连接层为稀疏版本 model.fc1 = SparseLinear(model.fc1.in_features, model.fc1.out_features) model.fc2 = SparseLinear(model.fc2.in_features, model.fc2.out_features) # 复制原模型的权重和偏置 model.fc1.load_state_dict(model.fc1.state_dict()) model.fc2.load_state_dict(model.fc2.state_dict()) # 切换到推理模式 model.eval()
二、卷积层(Conv)的零运算跳过
MNIST输入是单通道图像,背景全为零,只有数字区域有非零值。我们可以先提取输入里的非零像素,只计算这些像素对输出特征图的贡献,跳过零像素的无效运算:
class SparseConv2d(nn.Conv2d): def forward(self, x): # x的形状: (batch_size, 输入通道数, 高, 宽) batch_size, in_channels, h, w = x.shape kernel_h, kernel_w = self.kernel_size stride_h, stride_w = self.stride padding_h, padding_w = self.padding # 计算输出特征图的尺寸 out_h = (h + 2*padding_h - kernel_h) // stride_h + 1 out_w = (w + 2*padding_w - kernel_w) // stride_w + 1 # 初始化输出为全零 output = torch.zeros(batch_size, self.out_channels, out_h, out_w, device=x.device) # 逐个处理每个样本 for b in range(batch_size): img = x[b, 0] # 找出当前图像的非零像素位置 y_pos, x_pos = img.nonzero(as_tuple=True) if len(y_pos) == 0: continue # 遍历每个非零像素,计算它对输出的贡献 for y, x_p in zip(y_pos, x_pos): # 计算卷积核覆盖的输入区域(考虑padding) start_y = y - padding_h start_x = x_p - padding_w # 遍历卷积核的每个位置 for ky in range(kernel_h): for kx in range(kernel_w): # 计算输入中的实际位置 input_y = start_y + ky input_x = start_x + kx # 检查是否在原始输入的有效范围内 if 0 <= input_y < h and 0 <= input_x < w: # 计算对应的输出特征图位置 out_y = (input_y - start_y) // stride_h out_x = (input_x - start_x) // stride_w if 0 <= out_y < out_h and 0 <= out_x < out_w: # 获取卷积核对应位置的权重 weight_val = self.weight[:, 0, ky, kx] # 累加像素值乘权重的结果 output[b, :, out_y, out_x] += img[y, x_p] * weight_val # 加上偏置(如果有) if self.bias is not None: output += self.bias.view(1, -1, 1, 1) return output
使用方法
同样替换原模型的卷积层,加载原权重:
# 替换卷积层为稀疏版本 model.conv1 = SparseConv2d(model.conv1.in_channels, model.conv1.out_channels, kernel_size=model.conv1.kernel_size, stride=model.conv1.stride, padding=model.conv1.padding) model.conv2 = SparseConv2d(model.conv2.in_channels, model.conv2.out_channels, kernel_size=model.conv2.kernel_size, stride=model.conv2.stride, padding=model.conv2.padding) # 复制原模型的权重和偏置 model.conv1.load_state_dict(model.conv1.state_dict()) model.conv2.load_state_dict(model.conv2.state_dict())
三、测试说明
- 这个实现是简易版本,没做底层优化,核心是快速验证稀疏度和推理时间的关系。卷积层用了嵌套循环,Python层面有一定开销,但输入越稀疏,省的运算量越多,整体时间还是会比原模型少。
- 不需要重新训练,替换层后直接加载原权重即可,因为参数维度和原模型完全一致。
- 你可以手动生成不同稀疏比例的MNIST输入(比如随机把数字区域的部分像素置零),测试不同情况下的推理时间差异。
内容的提问来源于stack exchange,提问作者Rezajh
相关产品推荐
相关产品推荐

