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

如何以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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 02:45:33