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

使用timm PatchEmbed后,如何反分块Transformer输出张量为图像?

解决Patch Tensor反分块(Unpatchify)问题

首先明确核心前提:你得到的torch.Size([2,77,256])张量中,77=64+1,多出来的1是ViT中添加的class token,不属于图像分块,必须先移除才能进行反分块操作。以下是完整步骤:

关键参数回顾

  • 原始图像形状:(N=2, C=4, H=64, W=64)
  • Patch尺寸:patch_size=8 → 图像每个维度的patch数:64//8=8,总patch数:8×8=64
  • Transformer输出维度:256(需映射回每个patch的像素展平维度:8×8×4=256)

完整反分块流程

1. 移除Class Token

# 去掉第一个token,得到纯图像分块张量,形状变为[2,64,256]
patch_without_cls = patch_tensor[:, 1:, :]

2. 映射回Patch像素维度

用线性层将Transformer输出的256维patch向量,映射回每个patch的像素展平维度(8×8×4=256):

import torch.nn as nn

# 定义映射层
proj_back = nn.Linear(256, patch_size * patch_size * 4)
# 得到每个patch的像素展平张量,形状[2,64,256]
patch_pixels = proj_back(patch_without_cls)

3. 重塑为Patch网格

将64个patch重新排列为8×8的网格,并恢复每个patch的空间结构:

# 先reshape为[N, H_patch, W_patch, patch_size, patch_size, C],形状[2,8,8,8,8,4]
patch_grid = patch_pixels.reshape(2, 8, 8, 8, 8, 4)

4. 拼接为目标格式图像

格式一:(N,H,W,C)

# 合并patch维度与空间维度,得到最终形状[2,64,64,4]
image_nhwc = patch_grid.permute(0, 1, 3, 2, 4, 5).reshape(2, 64, 64, 4)

格式二:(N,C,H,W)(PyTorch常用格式)

# 调整通道维度到首位,合并后得到形状[2,4,64,64]
image_nchw = patch_grid.permute(0, 5, 1, 3, 2, 4).reshape(2, 4, 64, 64)

完整代码示例

import torch
import torch.nn as nn

# 模拟输入的Transformer输出张量
patch_tensor = torch.randn(2, 77, 256)

# 固定参数
N, C, H, W = 2, 4, 64, 64
patch_size = 8
H_patch, W_patch = H // patch_size, W // patch_size

# 1. 移除class token
patch_without_cls = patch_tensor[:, 1:, :]

# 2. 映射回像素维度
proj_back = nn.Linear(256, patch_size * patch_size * C)
patch_pixels = proj_back(patch_without_cls)

# 3. 重塑为patch网格
patch_grid = patch_pixels.reshape(N, H_patch, W_patch, patch_size, patch_size, C)

# 4. 转换为目标格式
image_nhwc = patch_grid.permute(0, 1, 3, 2, 4, 5).reshape(N, H, W, C)
image_nchw = patch_grid.permute(0, 5, 1, 3, 2, 4).reshape(N, C, H, W)

print("(N,H,W,C)形状:", image_nhwc.shape)  # torch.Size([2,64,64,4])
print("(N,C,H,W)形状:", image_nchw.shape)  # torch.Size([2,4,64,64])

关键注意事项

  • 必须优先移除class token,否则会导致空间维度计算错误。
  • 如果Transformer输出维度与patch_size×patch_size×in_channels不一致,线性层的输出维度必须严格等于后者,才能正确恢复像素结构。
  • permute和reshape的顺序直接影响图像的空间对齐,需严格按照示例中的维度顺序操作。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 21:12:47