如何对小批量CIFAR-10图像分块、打乱重组?循环为何无限运行?
循环无法终止原因及CIFAR-10图像分块、打乱、重组实现方案
一、循环无法终止的原因
- 代码中
for images, labels in trainloader:会完整遍历整个CIFAR-10训练集(50000张样本),batch_size设为4的话需要执行12500次迭代。如果没有添加任何输出或终止逻辑,程序会持续处理数据,看起来像是“循环无法终止”。 - 若仅为测试功能,无需遍历全部数据集,需添加终止条件(比如只处理前N个batch或前N张图)。
二、正确实现图像分块、打乱与重组的步骤
CIFAR-10图像尺寸为32×32×3,要实现3×3分块的话,由于32无法被3整除,需先对图像做padding(比如补成33×33);如果是做非重叠的2×2分块(16×16),可直接用patch_size=16、stride=16。下面以**非重叠2×2分块(适配32×32尺寸)**为例,给出完整实现流程:
核心步骤
- 将PyTorch张量格式的图像转换为
patchify兼容的numpy数组(通道维度从[C,H,W]转回[H,W,C])。 - 使用
patchify分割图像为小块,展平块列表后随机打乱。 - 将打乱后的块恢复为原分块结构,用
unpatchify重组图像。 - 可视化原图与重组后的图像。
完整代码示例
import numpy as np from patchify import patchify, unpatchify import matplotlib.pyplot as plt import torch import torchvision import torchvision.transforms as transforms # 数据预处理:保留numpy转换的可逆性 transform = transforms.Compose( [transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))]) batch_size = 4 trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform) trainloader = torch.utils.data.DataLoader(trainset, batch_size=batch_size, shuffle=True, num_workers=2) # 设定分块参数:适配32×32图像的非重叠分块 patch_size = 16 # 每个小块的尺寸 stride = 16 # 步长=块尺寸,保证无重叠 # 仅处理第一个batch做测试(避免遍历全数据集) for images, labels in trainloader: # 取batch中的第一张图演示 img_tensor = images[0] # 张量转numpy:从[C,H,W]转为[H,W,C],并反归一化回到[0,1]范围 img = img_tensor.permute(1,2,0).numpy() img = (img * 0.5) + 0.5 # 反归一化 # 1. 分块:返回形状为(2,2,1,16,16,3),对应(行块数,列块数,1,块高,块宽,通道) patches = patchify(img, (patch_size, patch_size, 3), stride=stride) # 展平块:从(2,2,1,16,16,3)转为(4,16,16,3) patches_flat = patches.reshape(-1, patch_size, patch_size, 3) # 2. 随机打乱块顺序 np.random.shuffle(patches_flat) # 3. 重组块:恢复为原分块结构(2,2,1,16,16,3) patches_shuffled = patches_flat.reshape(patches.shape) # 用unpatchify重组图像 img_shuffled = unpatchify(patches_shuffled, img.shape) # 4. 绘图展示 plt.figure(figsize=(8,4)) plt.subplot(121) plt.imshow(img) plt.title('原图') plt.axis('off') plt.subplot(122) plt.imshow(img_shuffled) plt.title('打乱分块重组图') plt.axis('off') plt.show() # 处理完第一个batch后终止循环 break
若要实现3×3分块(适配32×32图像)
由于32无法被3整除,需先对图像做padding,比如补成33×33:
# 在转numpy后添加padding img_padded = np.pad(img, ((0,1),(0,1),(0,0)), mode='constant') # 补1行1列 patch_size = 11 # 33//3=11 patches = patchify(img_padded, (patch_size, patch_size, 3), stride=patch_size) # 后续打乱、重组步骤同上,最后重组后裁剪回32×32:img_shuffled = img_shuffled[:32,:32,:]
内容的提问来源于stack exchange,提问作者linkho
相关产品推荐
相关产品推荐

