如何为MNIST数据集添加随机位置偏移,避免数字平移出帧?
解决MNIST数据集随机平移且数字不越界的问题
方案1:将数字放到更大画布上随机平移(简单直接)
如果你的需求是把原28x28的MNIST数字放到更大尺寸的画布(比如56x56)中随机平移,且数字完全处于画布内,可以自定义一个PyTorch Transform类,把你的手动逻辑封装进去,直接应用到整个数据集:
自定义Transform类
import random import numpy as np from PIL import Image from torchvision import datasets, transforms class RandomTranslateOnLargeCanvas: def __init__(self, canvas_size=(56, 56), img_size=(28, 28)): self.canvas_h, self.canvas_w = canvas_size self.img_h, self.img_w = img_size # 计算最大偏移量,确保数字不会超出画布 self.max_h_offset = self.canvas_h - self.img_h self.max_w_offset = self.canvas_w - self.img_w def __call__(self, img): # 将PIL图像转为numpy数组 img_np = np.array(img) # 创建空白画布 canvas = np.zeros((self.canvas_h, self.canvas_w), dtype=img_np.dtype) # 生成随机偏移坐标 h_offset = random.randint(0, self.max_h_offset) w_offset = random.randint(0, self.max_w_offset) # 将数字图像放到画布的对应位置 canvas[h_offset:h_offset+self.img_h, w_offset:w_offset+self.img_w] = img_np # 转回PIL图像供后续Transform处理 return Image.fromarray(canvas)
加载数据集并应用变换
# 构建Transform链 transform = transforms.Compose([ RandomTranslateOnLargeCanvas(canvas_size=(56,56), img_size=(28,28)), transforms.ToTensor(), # 可选:添加MNIST的标准化 transforms.Normalize((0.1307,), (0.3081,)) ]) # 加载带变换的MNIST数据集 train_dataset = datasets.MNIST( root='./data', train=True, download=True, transform=transform )
方案2:在原28x28画布内平移且数字不越界
如果需要保持原28x28的画布尺寸,同时让数字随机平移且不超出画面,需要先检测数字的实际边界,再计算允许的平移范围:
自定义Transform类
import random import numpy as np from PIL import Image from torchvision import datasets, transforms class RandomTranslateWithinFrame: def __call__(self, img): img_np = np.array(img) # 找到数字非零区域的边界(去除空白区域) non_zero_coords = np.where(img_np > 0) min_h, max_h = non_zero_coords[0].min(), non_zero_coords[0].max() min_w, max_w = non_zero_coords[1].min(), non_zero_coords[1].max() # 计算允许的平移范围:保证平移后数字不会超出28x28画布 max_up_shift = min_h # 向上最多平移min_h像素(避免顶部出界) max_down_shift = 27 - max_h # 向下最多平移(27-max_h)像素 max_left_shift = min_w # 向左最多平移min_w像素 max_right_shift = 27 - max_w # 向右最多平移(27-max_w)像素 # 生成随机平移量 dh = random.randint(-max_up_shift, max_down_shift) dw = random.randint(-max_left_shift, max_right_shift) # 构建平移仿射矩阵 affine_matrix = np.array([[1, 0, dw], [0, 1, dh]], dtype=np.float32) # 应用变换,空白区域填充黑色(0) return img.transform(img.size, Image.AFFINE, affine_matrix.flatten(), fillcolor=0)
加载数据集并应用变换
transform = transforms.Compose([ RandomTranslateWithinFrame(), transforms.ToTensor(), transforms.Normalize((0.1307,), (0.3081,)) ]) train_dataset = datasets.MNIST( root='./data', train=True, download=True, transform=transform )
为什么你的手动代码无法直接应用?
你之前的代码只处理了单个样本,而PyTorch的Dataset需要将变换逻辑封装成可调用的Transform类,这样在加载每个样本时会自动执行变换操作,上面的两个方案就是把你的手动逻辑包装成了符合PyTorch要求的Transform结构。
内容的提问来源于stack exchange,提问作者Dankrupt
相关产品推荐
相关产品推荐

