PyTorch张量替换:将最后维度范数为0的位置替换为另一张量对应值
PyTorch张量替换最后维度范数为0的元素
给定形状为(2000, 1, 360, 3)的PyTorch张量A,我们需要定位最后一个维度(即每个长度为3的向量)上范数为0的所有位置,并将这些位置的值替换为同形状张量B的对应元素。以下是实现方案:
核心步骤代码
import torch # 假设A和B是形状为(2000,1,360,3)的PyTorch张量 # 1. 计算最后维度的范数,生成掩码:范数为0的位置标记为True mask = torch.norm(A, dim=-1) == 0.0 # 2. 扩展掩码维度,使其与原张量维度匹配(方便广播索引) mask = mask.unsqueeze(-1) # 3. 替换元素:可以选择修改原张量,或生成新张量 # 方式一:直接修改原A A[mask] = B[mask] # 方式二:生成新张量new_A,不修改原A new_A = torch.where(mask, B, A)
示例验证
用题目给出的小维度示例测试:
# 构造示例张量A和B A = torch.tensor([[[[0, 0, 0], [1, 2, 1], [0, 1, 0]]], [[[2, 0, 0], [0, 0, 0], [1, 1, 1]]]]) B = torch.tensor([[[[0, 0, 1], [1, 1, 1], [0, 1, 0]]], [[[1, 0, 0], [0, 1, 1], [2, 1, 1]]]]) # 执行替换逻辑 mask = torch.norm(A, dim=-1) == 0.0 mask = mask.unsqueeze(-1) new_A = torch.where(mask, B, A) print(new_A)
输出结果与预期一致:
tensor([[[[0, 0, 1], [1, 2, 1], [0, 1, 0]]], [[[2, 0, 0], [0, 1, 1], [1, 1, 1]]]])
关键说明
torch.norm(A, dim=-1):对最后一个维度的每个向量计算范数,输出形状为(2000,1,360)的张量,与原张量前三维对应。unsqueeze(-1):将掩码扩展一个维度,变成(2000,1,360,1),这样可以通过广播机制匹配原张量的(2000,1,360,3)形状,确保每个向量的三个元素都能被正确选中替换。- 两种替换方式:如果不需要保留原张量,直接索引赋值效率更高;如果需要保留原数据,用
torch.where生成新张量更安全。
内容的提问来源于stack exchange,提问作者ojipadeson
相关产品推荐
相关产品推荐

