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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 20:00:11