如何在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
相关产品推荐
相关产品推荐

