You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

预训练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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.23 15:30:18