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

如何在PyTorch张量的指定维度中随机将k个元素设为特定值?

在PyTorch张量指定维度随机将k个元素设为特定值

假设我们有一个形状为[2, 3, 5]的PyTorch张量:

import torch

tensor = torch.tensor([[[0.8823, 0.9150, 0.3829, 0.9593, 0.3904],
                        [0.6009, 0.2566, 0.7936, 0.9408, 0.1332],
                        [0.9346, 0.5936, 0.8694, 0.5677, 0.7411]],

                       [[0.4294, 0.8854, 0.5739, 0.2666, 0.6274],
                        [0.2696, 0.4414, 0.2969, 0.8317, 0.1053],
                        [0.2695, 0.3588, 0.1994, 0.5472, 0.0062]]])

当k=2时,需要在指定维度(如dim=2)中,为每个子张量随机选择k个元素并设置为特定值(如5),得到类似如下结果:

[[[0.8823, 0.9150, 0.3829, 5, 5],
  [0.6009, 0.2566, 5, 0.9408, 5],
  [5, 5, 0.8694, 0.5677, 0.7411]],

 [[5, 0.8854, 0.5739, 5, 0.6274],
  [5, 0.4414, 0.2969, 5, 0.1053],
  [0.2695, 0.3588, 5, 0.5472, 5]]]

实现方案

每个子张量独立随机选k个元素

这个方案会让每个目标维度下的子张量(比如示例中每个长度为5的向量)都拥有独立的随机索引,完全匹配需求:

import torch

def set_k_random_elements_per_group(tensor, dim, k, value):
    # 获取目标维度的长度
    dim_length = tensor.size(dim)
    # 计算需要生成多少组随机索引(其他维度的元素总数)
    num_groups = tensor.numel() // dim_length
    
    # 为每组生成k个不重复的随机索引,拼接后调整形状匹配原张量结构
    rand_indices = torch.cat(
        [torch.randperm(dim_length, device=tensor.device)[:k].unsqueeze(0) 
         for _ in range(num_groups)],
        dim=0
    ).view(tensor.shape[:dim] + (k,) + tensor.shape[dim+1:])
    
    # 构造其他维度的索引网格,确保每个位置都能对应到随机索引
    other_dim_indices = [torch.arange(s, device=tensor.device) for s in tensor.shape[:dim] + tensor.shape[dim+1:]]
    grid = torch.meshgrid(*other_dim_indices, indexing='ij')
    
    # 扩展网格维度以适配k的维度
    expanded_grid = [g.unsqueeze(-1) for g in grid]
    
    # 组合所有索引,完成赋值
    tensor[tuple(expanded_grid) + (rand_indices,)] = value
    return tensor

# 测试代码
tensor = torch.tensor([[[0.8823, 0.9150, 0.3829, 0.9593, 0.3904],
                        [0.6009, 0.2566, 0.7936, 0.9408, 0.1332],
                        [0.9346, 0.5936, 0.8694, 0.5677, 0.7411]],

                       [[0.4294, 0.8854, 0.5739, 0.2666, 0.6274],
                        [0.2696, 0.4414, 0.2969, 0.8317, 0.1053],
                        [0.2695, 0.3588, 0.1994, 0.5472, 0.0062]]])

# 注意使用clone()避免修改原张量
result = set_k_random_elements_per_group(tensor.clone(), dim=2, k=2, value=5)
print(result)

代码说明

  • torch.randperm(dim_length):生成目标维度内的随机排列,取前k个索引保证不重复。
  • torch.meshgrid:构造其他维度的索引网格,确保每个子张量都能对应到自己的随机索引组。
  • 索引组合:将扩展后的网格索引与随机索引组合,利用PyTorch的高级索引实现批量赋值。

可选:全局统一随机索引

如果希望所有子张量使用同一组随机索引(比如所有长度为5的向量都替换相同的2个位置),可以使用以下简化版本:

def set_k_random_elements_global(tensor, dim, k, value):
    dim_length = tensor.size(dim)
    rand_indices = torch.randperm(dim_length, device=tensor.device)[:k]
    
    # 构造其他维度的索引网格
    other_dim_indices = [torch.arange(s, device=tensor.device) for s in tensor.shape[:dim] + tensor.shape[dim+1:]]
    grid = torch.meshgrid(*other_dim_indices, indexing='ij')
    expanded_grid = [g.unsqueeze(dim) for g in grid]
    
    tensor[tuple(expanded_grid) + (rand_indices,)] = value
    return tensor

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 23:45:45