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

如何对MNIST数据集每行应用卷积核并保存全量处理结果?

问题分析与解决

你的代码只输出最后一行结果,核心原因是没有存储每次循环处理后的图像数据——每次循环都在处理新图像,但前一次的结果没被保存,循环结束后自然只剩最后一张的处理结果。另外代码还有几个细节错误需要修正:

具体问题点&修复

  1. 无结果存储容器:没创建列表/数组来存所有处理后的行数据,每次处理完就丢了前面的结果。
  2. 变量名错误:定义了conv_image但赋值给了未定义的sharpened,实际运行会报错。
  3. 循环索引错误:i初始为0,先执行i+=1,导致从i=1开始处理,漏掉了图像顶部第一行像素。
  4. 未还原CSV格式:处理后的图像没展平为一维,也没加回标签,无法保存成原CSV的每行结构。

修正后的完整代码

import numpy as np
import matplotlib.pyplot as plt

# 读取数据集
test_data_file = open("mnist_test.csv", 'r')      
test_data_list = test_data_file.readlines()    
test_data_file.close() 

# 用来存所有处理后的行数据
processed_data = []

# 提前定义锐化核,不用每次循环重复创建
sharpen_kernel = np.array([
    [0, -1, 0],
    [-1, 5, -1],
    [0, -1, 0]])
kernel_size = 3

for record in test_data_list:
    # 拆分一行数据
    all_values = record.split(',')
    label = int(all_values[0])
    pixel_data = np.asfarray(all_values[1:])
    
    # 转成28×28图像
    original_img = pixel_data.reshape((28,28))
    # 初始化卷积结果图像
    sharpened_img = np.zeros((28,28))
    
    # 遍历图像做卷积(用for循环比while更不容易出错)
    for i in range(28 - kernel_size + 1):
        for j in range(28 - kernel_size + 1):
            # 取3×3的子图像
            sub_img = original_img[i:i+kernel_size, j:j+kernel_size]
            # 计算卷积值
            conv_val = np.sum(sub_img * sharpen_kernel)
            # 把结果放在子图像的中心位置
            sharpened_img[i+1, j+1] = conv_val
    
    # 处理边缘像素:直接用原图像的边缘(因为3×3核碰不到最外层1像素)
    sharpened_img[0, :] = original_img[0, :]
    sharpened_img[-1, :] = original_img[-1, :]
    sharpened_img[:, 0] = original_img[:, 0]
    sharpened_img[:, -1] = original_img[:, -1]
    
    # 把处理后的图像转成CSV行格式:标签+展平的像素
    flattened_img = sharpened_img.flatten()
    # 转成字符串列表,方便拼接成CSV行
    row_data = [str(label)] + [str(val) for val in flattened_img]
    processed_data.append(','.join(row_data))

# 保存到新CSV文件
with open("mnist_test_sharpened.csv", 'w') as out_file:
    out_file.write('\n'.join(processed_data))

# 测试绘图:显示前5张处理后的图像
plt.figure(figsize=(10, 2))
for idx in range(5):
    row = processed_data[idx].split(',')
    img = np.asfarray(row[1:]).reshape((28,28))
    plt.subplot(1,5,idx+1)
    plt.imshow(img, cmap='gray')
    plt.title(f"Label: {row[0]}")
    plt.axis('off')
plt.show()

关键说明

  • 结果存储:用processed_data列表存每一行的处理结果,循环结束后一次性写入文件,效率更高。
  • 卷积简化:用np.sum(sub_img * sharpen_kernel)替代reshape+dot,逻辑一样但代码更简洁。
  • 边缘处理:3×3卷积核没法处理图像最外层1像素,这里直接保留原图像边缘,你也可以换成零填充等方式。
  • 绘图改进:循环显示多张图像,方便验证所有处理结果是否正确。

内容的提问来源于stack exchange,提问作者Moe

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 14:56:30