如何用PyTorch将28×28的2D张量划分为7×7的16元素块?
解决方法
问题核心是torch.view()是按内存连续顺序重塑张量,而你需要的是按空间位置划分4×4的非重叠块,所以直接用view无法得到正确的分块结果。以下是两种可行的实现方式:
方法一:使用unfold分块(推荐)
unfold是PyTorch专门用于张量分块的API,适合处理这种非重叠的空间分块需求:
import torch # 模拟你的输入:28×28的图像张量 img = torch.randn(28, 28) # shape: torch.Size([28, 28]) # 第一步:沿高度方向分块,每块高度4,步长4 blocks_h = img.unfold(dimension=0, size=4, step=4) # shape: torch.Size([7, 4, 28]) # 第二步:沿宽度方向分块,每块宽度4,步长4 blocks_hw = blocks_h.unfold(dimension=2, size=4, step=4) # shape: torch.Size([7, 4, 7, 4]) # 调整维度顺序并展平每个4×4块 result = blocks_hw.permute(0, 2, 1, 3).reshape(7, 7, 16) # shape: torch.Size([7, 7, 16])
方法二:手动拆分维度
如果不想用unfold,也可以通过维度拆分和重组实现:
import torch img = torch.randn(28, 28) # 先将图像拆分为7行×4高度的块,再将每行拆分为7列×4宽度的块 split_h = img.split(4, dim=0) # 得到7个(4,28)的张量 split_hw = [row.split(4, dim=1) for row in split_h] # 得到7×7个(4,4)的张量 # 将所有块展平并重组为7×7×16的张量 result = torch.stack([torch.stack([block.flatten() for block in row], dim=0) for row in split_hw], dim=0)
为什么view不行?
28×28张量的内存存储顺序是逐行连续的(第1行所有元素→第2行所有元素→…→第28行所有元素)。img.view(7,7,16)会直接把前16个元素(第1行的前16个像素)作为第一个块,这和你需要的“4行×4列”空间块完全不符,因此无法得到预期结果。
内容的提问来源于stack exchange,提问作者Chenming Zhang
相关产品推荐
相关产品推荐

