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

PyTorch使用nn.Unfold提取重叠patch报im2col_out_cpu不支持Byte错误

报错诱因

核心报错RuntimeError: "im2col_out_cpu" not implemented for 'Byte'的直接原因是传入nn.Unfold的图像张量为torch.uint8(即Byte类型),PyTorch CPU端的im2col(Unfold底层算子)不支持8位无符号整型运算,仅支持浮点类数据类型(float32、float64等)。
你代码中虽然写了img = img.to(torch.float32),但存在两个逻辑问题:

  • 旧版本PyTorch/torchvision匹配环境下(你使用的是Python3.7+PyTorch1.7左右版本),transforms.Resize接收uint8类型张量输入时,插值计算完成后会自动将输出转回uint8类型,若类型转换语句位置不对或数值范围不匹配,容易出现类型不符合算子要求的问题。
  • 后续可视化逻辑存在错误:Unfold输出的是展平后的patch向量,你直接索引取值拿到的是单个像素值而非完整图像块,同时16×16的子图网格和实际输出的patch数量不匹配,即使解决类型问题也会出现索引越界、显示异常的问题。
修复方案

按以下步骤调整代码即可解决问题:

  1. 调整张量类型转换逻辑,在传入Unfold前明确将张量转为float32类型,同时将像素值从0-255区间归一化到0-1区间,匹配后续可视化的输入要求。
  2. 对Unfold输出的展平patch做形状重构,还原为(通道数, 块高, 块宽)的标准图像张量格式。
  3. 根据实际输出的patch行列数设置子图网格,避免索引越界。

修复后的可运行代码如下:

import torch
import torch.nn as nn
from torchvision import io, transforms
from torchvision.transforms.functional import to_pil_image
import matplotlib.pyplot as plt

IMG_SIZE = 112
image_size = 112
patch_size = 28
ac_patch_size = 12
pad = 4

# 读取并预处理图像
resize = transforms.Resize((IMG_SIZE, IMG_SIZE))
img = resize(io.read_image("Adam_Brody_233.png"))
# 先转float32再归一化到0-1区间,满足Unfold算子和可视化要求
img = img.to(torch.float32) / 255.0
img = img.unsqueeze(0)  # 新增batch维度,形状变为(1, 3, 112, 112)

# 生成重叠patch
soft_split = nn.Unfold(
    kernel_size=(ac_patch_size, ac_patch_size),
    stride=(patch_size, patch_size),
    padding=(pad, pad)
)
patches = soft_split(img)  # 输出形状: (1, 3*12*12, L),L为patch总数
# 计算输出patch的行列数
h_out = (image_size + 2*pad - ac_patch_size) // patch_size + 1
w_out = (image_size + 2*pad - ac_patch_size) // patch_size + 1
# 重构形状为(1, h_out, w_out, 3, ac_patch_size, ac_patch_size)
patches = patches.transpose(1, 2).reshape(1, h_out, w_out, 3, ac_patch_size, ac_patch_size)

# 可视化
fig, ax = plt.subplots(h_out, w_out)
for i in range(h_out):
    for j in range(w_out):
        sub_img = patches[0, i, j]  # 取出单块patch,形状(3,12,12)
        ax[i][j].imshow(to_pil_image(sub_img))
        ax[i][j].axis('off')

plt.tight_layout()
plt.show()

注:如果需要更高密度的重叠patch,适当调小stride参数即可,当前参数下输出patch为4×4共16个,若要得到16×16的patch网格,将stride调整为8左右即可。

内容的提问来源于stack exchange,提问作者Khawar Islam

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 08:54:23