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

如何为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 22:41:10