如何将pygame手绘数字调整为匹配MNIST格式用于TensorFlow识别
解决方案
你遇到的问题核心是当前直接输出的28x28手绘图像和MNIST训练集的预处理规则不匹配,只要增加一步图像预处理步骤即可解决,具体流程如下:
- 从手绘的28x28原始图像中提取数字的最小外接矩形,即所有灰度值大于0的像素的上下左右边界
- 将矩形内的数字内容等比例缩放到20x20像素,和MNIST训练集的数字尺寸保持一致
- 将20x20的数字图像居中放置在28x28的纯黑背景画布上,得到和MNIST格式完全匹配的输入
你可以直接用下面的代码修改原逻辑即可:
首先安装依赖:pip install opencv-python
import cv2 def preprocess_image(raw_image): # raw_image为从grid获取的28x28灰度数组,值范围0-1 img = (raw_image * 255).astype(np.uint8) # 提取非零像素边界 coords = cv2.findNonZero(img) if coords is None: return raw_image x, y, w, h = cv2.boundingRect(coords) # 裁剪数字区域 digit = img[y:y+h, x:x+w] # 等比例缩放到20x20 max_side = max(w, h) scale = 20 / max_side resized_digit = cv2.resize(digit, (int(w*scale), int(h*scale)), interpolation=cv2.INTER_AREA) # 创建28x28黑底画布 new_img = np.zeros((28,28), dtype=np.uint8) # 计算居中偏移量 start_x = (28 - resized_digit.shape[1]) // 2 start_y = (28 - resized_digit.shape[0]) // 2 # 粘贴数字到居中位置 new_img[start_y:start_y+resized_digit.shape[0], start_x:start_x+resized_digit.shape[1]] = resized_digit # 转回0-1浮点格式 return new_img / 255.0
替换原有的get_data函数:
def get_data(grid): image = [] for r in range(len(grid)): image.append([]) for spot in grid[r]: image[r].append(spot.get_gray_scale()) raw_img = np.array(image, dtype=np.float64) / 255 processed_img = preprocess_image(raw_img) return [processed_img]
如果不想引入opencv依赖,也可以用numpy手动计算像素边界和实现缩放逻辑,效果一致。修改后输入完全匹配MNIST数据集格式,识别准确率会大幅提升。
内容的提问来源于stack exchange,提问作者xxSirFartAlotxx
相关产品推荐
相关产品推荐

