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

PyTorch中如何用index张量在dim=1维度替换source张量?

解决PyTorch中scatter方法的维度匹配与替换问题

问题分析

你需要利用index张量(形状[20, 3001])在dim=1维度上,将source张量(形状[20, 3, 3001])的对应值替换到目标张量的指定位置。核心问题是scatter方法要求索引张量与源张量的维度匹配,需要对index做维度扩展。

解决方案代码

import torch

# 模拟你的输入张量
index = torch.randint(0, 3, (20, 3001))  # 生成0/1/2的索引,与你的index形状一致
source = torch.randn(20, 3, 3001)        # 模拟你的source张量

# 初始化目标张量(与source同形状,初始值可按需调整)
target = torch.zeros_like(source)

# 关键步骤:将index扩展一个维度,使其形状变为[20, 1, 3001],与source的dim=1维度匹配
expanded_index = index.unsqueeze(1)

# 从source中取出index指定通道的值,再用scatter_替换到target的对应位置
selected_values = source.gather(dim=1, index=expanded_index)
target.scatter_(dim=1, index=expanded_index, src=selected_values)

代码解释

  1. 维度扩展:index.unsqueeze(1)将原形状[20, 3001]变为[20, 1, 3001],这样才能与source的[20, 3, 3001]在dim=1维度上匹配,满足scatter的维度要求。
  2. 提取对应值:source.gather(dim=1, index=expanded_index)会根据index取出每个(batch, time_step)对应的通道值,得到形状[20, 1, 3001]的张量。
  3. 执行替换:target.scatter_会将提取出的值,放到target中dim=1维度对应index指定的位置,其他位置保持初始值(这里是0)。

验证结果

可以通过以下代码验证替换是否正确:

# 随机选一个batch和time_step检查
batch_idx = 5
time_step = 1000
selected_channel = index[batch_idx, time_step]

# 检查target对应位置的值是否与source一致
assert torch.allclose(target[batch_idx, selected_channel, time_step], source[batch_idx, selected_channel, time_step])
# 检查其他通道是否保持初始值0
for channel in range(3):
    if channel != selected_channel:
        assert target[batch_idx, channel, time_step] == 0

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 16:21:00