如何以torch.take风格将PyTorch张量中索引指向的值置零?
解决方案
PyTorch 中没有直接对应 torch.take 的 torch.set_value 方法,但可以通过两种简洁方式实现你要的功能:
方法1:使用 torch.put_ 原地修改
torch.put_ 是 torch.take 的逆操作,支持用相同逻辑的索引直接原地修改张量元素,无需手动展平:
import torch import torch.nn as nn input1 = torch.randn(1, 1, 6, 6) m = nn.MaxPool2d(2, 2, return_indices=True) val, indx = m(input1) # 原地将indx指向的元素置零 input1.put_(indx, torch.zeros_like(val))
方法2:展平后直接索引赋值
如果需要非原地操作,可以先展平张量修改,再恢复原形状:
# 非原地版本,生成新张量 input_flat = input1.flatten() input_flat[indx.flatten()] = 0 input_modified = input_flat.view(input1.shape) # 原地版本,直接修改原张量 input1.flatten()[indx.flatten()] = 0
自定义池化层封装
把逻辑封装成可复用的自定义层:
class MaxPoolDropout(nn.Module): def __init__(self, kernel_size, stride=None): super().__init__() self.max_pool = nn.MaxPool2d(kernel_size, stride, return_indices=True) def forward(self, x): _, indx = self.max_pool(x) # 克隆输入避免修改原张量,若允许原地修改可直接操作x x_out = x.clone() x_out.put_(indx, torch.zeros_like(indx, dtype=x.dtype)) return x_out
内容的提问来源于stack exchange,提问作者FortCpp
相关产品推荐
相关产品推荐

