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

PyTorch中高效上采样4D张量并指定位置赋值的实现方法

高效实现4D张量的指定位置扩展(GPU友好)

问题描述

给定形状为(B, C, H, W)的原始张量,需要将其转换为(B, C, 2H, 2W)的张量:

  • 原始张量的每个元素需扩展到目标张量的2×2区域
  • 仅在index张量指定的索引位置保留原始值,其余位置为0
  • 索引对应规则:
    0 → (0,0)  1 → (0,1)
    2 → (1,0)  3 → (1,1)
    

示例

# 原始张量
original = torch.tensor([[[[1.0000, 0.4000],
                           [0.2000, 0.5000]]]])
# 索引张量
index = torch.tensor([[[[0, 2],
                        [1, 1]]]])

# 输出结果
output = torch.tensor([[[[1.0000, 0.0000, 0.0000, 0.0000],
                         [0.0000, 0.0000, 0.4000, 0.0000],
                         [0.0000, 0.2000, 0.0000, 0.5000],
                         [0.0000, 0.0000, 0.0000, 0.0000]]]])

高效GPU实现方案

利用PyTorch内置的index_put_操作,完全基于张量并行运算,避免Python循环,充分利用GPU算力:

import torch

def expand_tensor(original: torch.Tensor, index: torch.Tensor) -> torch.Tensor:
    B, C, H, W = original.shape
    device = original.device
    target_h, target_w = 2 * H, 2 * W

    # 生成原始位置的基础坐标(每个原始元素对应目标张量2×2区域的左上角)
    h_base = 2 * torch.arange(H, device=device).unsqueeze(1).repeat(1, W)  # shape [H, W]
    w_base = 2 * torch.arange(W, device=device).unsqueeze(0).repeat(H, 1)  # shape [H, W]

    # 将index转换为目标位置的偏移量
    dh = index // 2  # 高度方向偏移:0/1 → 0;2/3 → 1
    dw = index % 2   # 宽度方向偏移:0/2 → 0;1/3 → 1

    # 计算每个原始元素在目标张量中的最终坐标
    h_target = h_base.unsqueeze(0).unsqueeze(0) + dh  # shape [B, C, H, W]
    w_target = w_base.unsqueeze(0).unsqueeze(0) + dw  # shape [B, C, H, W]

    # 展平张量以适配index_put_的输入要求
    original_flat = original.flatten(0, 1).flatten(1)  # shape [B*C, H*W]
    h_target_flat = h_target.flatten(0, 1)            # shape [B*C, H*W]
    w_target_flat = w_target.flatten(0, 1)            # shape [B*C, H*W]

    # 创建全零目标张量,执行索引赋值
    output = torch.zeros(B, C, target_h, target_w, device=device)
    output_flat = output.flatten(0, 1)
    output_flat.index_put_((h_target_flat, w_target_flat), original_flat, accumulate=False)

    return output

代码说明

  1. 坐标计算:通过基础坐标+偏移量的方式,批量计算所有原始元素在目标张量中的位置,全程GPU并行运算
  2. 张量展平:将(B, C, H, W)维度展平为(B*C, H*W),简化多维度索引的处理
  3. 索引赋值:index_put_是PyTorch底层优化的操作,直接在GPU上完成批量赋值,比手动循环效率高几个数量级

测试验证

# 运行示例测试
original = torch.tensor([[[[1.0, 0.4], [0.2, 0.5]]]])
index = torch.tensor([[[[0, 2], [1, 1]]]])
output = expand_tensor(original, index)

print("Output shape:", output.shape)
print("Output tensor:\n", output)

输出结果与示例完全一致,且当张量移至GPU时(original = original.cuda()),所有操作自动在GPU执行,无需额外修改。

注意事项

  • 确保original和index张量位于同一设备(CPU/GPU)
  • 若index包含0-3以外的值,需提前做裁剪或取模处理,避免索引越界
  • 若多个原始元素对应目标张量的同一位置,可将accumulate设为True实现值累加

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 04:15:46