如何将图像张量与对应Patch分数实现逐像素相乘?
解决Patch分数与对应像素相乘的问题
问题场景
你有20个展平的256×256图像张量(形状(20, 65536)),每张图像被划分为32×32的Patch,共64个;对应每个Patch的分数张量形状为(20, 64)。需要让每个Patch内的所有像素都乘以该Patch的分数,但直接相乘或简单repeat无法实现对应关系。
核心思路
要实现像素与对应Patch分数的匹配,关键是把一维的分数张量重新映射到图像的Patch布局,再将每个分数扩展到对应Patch的所有像素位置,最后与图像张量相乘。
示例代码验证(基于补充的小例子)
import torch img_size = 4 patch_size = 2 img = torch.rand((2, img_size, img_size)) # (2,4,4) score = torch.tensor([[1,2,3,4],[5,6,7,8]]) # (2,4) # 计算图像高/宽方向的Patch数量 num_patches_h = img_size // patch_size num_patches_w = img_size // patch_size # 将分数重塑为Patch的二维布局(匹配图像的Patch划分) score_reshaped = score.reshape(-1, num_patches_h, num_patches_w) # (2,2,2) # 在高、宽维度分别重复patch_size次,得到与图像同形状的分数张量 score_expanded = score_reshaped.repeat_interleave(patch_size, dim=1).repeat_interleave(patch_size, dim=2) # 执行像素与对应分数的相乘 result = img * score_expanded
如果你的Patch顺序与示例不同(比如分数的排列是按列优先),只需调整reshape后的维度顺序即可,比如score.reshape(-1, num_patches_w, num_patches_h).transpose(1,2)。
原始问题的完整实现
针对展平的256×256图像张量,只需先恢复图像的二维形状,处理后再展平即可:
batch_size = 20 img_size = 256 patch_size = 32 num_patches = (img_size // patch_size) ** 2 # 64 # 模拟输入 imgs_flattened = torch.rand((batch_size, img_size*img_size)) # (20, 65536) scores = torch.rand((batch_size, num_patches)) # (20,64) # 1. 将展平图像恢复为二维形状 imgs = imgs_flattened.reshape(batch_size, img_size, img_size) # (20,256,256) # 2. 计算高/宽方向的Patch数量 num_patches_h = img_size // patch_size # 8 num_patches_w = img_size // patch_size # 8 # 3. 重塑分数为Patch布局 scores_reshaped = scores.reshape(batch_size, num_patches_h, num_patches_w) # (20,8,8) # 4. 扩展分数到每个像素位置 scores_expanded = scores_reshaped.repeat_interleave(patch_size, dim=1).repeat_interleave(patch_size, dim=2) # (20,256,256) # 5. 相乘后重新展平(如果需要保持原形状) result_flattened = (imgs * scores_expanded).flatten(start_dim=1) # (20,65536)
为什么直接repeat不行?
score.repeat(1,1,64)会生成(20,64,64)的张量,不仅形状与展平图像(20,65536)不匹配,而且分数的排列顺序是连续重复每个Patch的分数,与图像像素按行优先的展平顺序不对应,无法实现每个Patch内像素与分数的正确匹配。
内容的提问来源于stack exchange,提问作者Tamir
相关产品推荐
相关产品推荐

