如何利用torch.topk返回的索引将原张量对应位置元素置零?
Zero Out Top k Elements in PyTorch Tensor
Step 1: Get Indices of Top k Elements
First, use torch.topk to retrieve the indices of the k largest elements in your tensor. You can ignore the returned values by using an underscore placeholder:
import torch x = torch.arange(1., 6.) k = 3 _, top_indices = torch.topk(x, k)
Step 2: Zero Out Target Positions
Use PyTorch's native advanced indexing to set the elements at these indices to zero. You have two options depending on whether you want to modify the original tensor or create a new one:
Option 1: In-place Modification
Modify the original tensor directly (saves memory if you don't need the original data):
x[top_indices] = 0.0 print(x) # Output: tensor([1., 2., 0., 0., 0.])
Option 2: Create a New Tensor
Keep the original tensor intact by cloning it first:
x = torch.arange(1., 6.) new_x = x.clone() new_x[top_indices] = 0.0 print(new_x) # Output: tensor([1., 2., 0., 0., 0.]) print(x) # Original remains unchanged: tensor([1., 2., 3., 4., 5.])
Efficiency Notes
- This method leverages PyTorch's optimized native indexing, which runs efficiently even on large tensors.
- For multi-dimensional/batched tensors, you can extend this approach using
torch.scatter_if needed, but the above code works seamlessly for 1D tensors as in your example.
内容的提问来源于stack exchange,提问作者SanadMarji
相关产品推荐
相关产品推荐

