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

PyTorch中如何无循环高效对二维张量按行实现dropout

问题描述

你持有取值仅为0、1,形状为(B, I)的稀疏二维张量U:

  • 每一行对应1个用户
  • 每一列对应1个物品
  • 单元格值为1代表对应用户和物品有过交互,值为0代表无交互

需求为对该张量实现类dropout操作:逐行(用户维度)随机将该行内p%比例的1值置为0,要求不能沿B维度编写for循环,实现要高效。

注:低效循环实现思路为:先屏蔽所有0值位置,再逐行对一维张量调用PyTorch内置dropout,该方案B较大时速度极慢。

实现方案

根据需求严格程度不同,有两种全向量化无循环的实现,均比循环方案效率高1~2个数量级:

方案1:概率丢弃(和原生Dropout逻辑一致,速度最快)

如果不需要严格保证每行刚好丢弃固定比例的1,只要求每个1值独立以p%概率被置0、所有0值位置保持不变(丢弃比例的期望为p%,和PyTorch内置dropout逻辑完全对齐),可以直接用伯努利采样生成掩码,代码极简:

import torch

def prob_interaction_dropout(U: torch.Tensor, p: float, scale: bool = False) -> torch.Tensor:
    """
    逐用户对交互记录做概率dropout
    Args:
        U: 形状(B, I)的0/1二值交互张量
        p: 1值被丢弃的百分比,取值范围[0, 100]
        scale: 是否和原生dropout一样做保留值缩放,默认关闭
    """
    keep_prob = 1 - p / 100
    # 仅在原1值位置按保留概率生成伯努利掩码,0位置掩码恒为0
    mask = torch.bernoulli(U * keep_prob)
    if scale:
        mask = mask / keep_prob
    return U * mask

方案2:精确比例丢弃(严格控制每行丢弃数量)

如果要求每行必须严格丢弃p%比例的1(四舍五入取整),可以用逐行topk选保留位的方式实现,全程无循环:

import torch

def exact_ratio_interaction_dropout(U: torch.Tensor, p: float) -> torch.Tensor:
    """
    逐用户对交互记录做精确比例dropout,每行严格丢弃p%比例的1值
    Args:
        U: 形状(B, I)的0/1二值交互张量
        p: 1值被丢弃的百分比,取值范围[0, 100]
    """
    keep_prob = 1 - p / 100
    # 统计每行的1值总数
    row_1_count = U.sum(dim=1, keepdim=True)
    # 计算每行需要保留的1值数量,最小为0
    keep_num = torch.clamp(torch.round(row_1_count * keep_prob).long(), min=0).squeeze()
    # 生成随机矩阵,原0值位置填极小值,保证不会被选为保留位
    rand_mat = torch.rand_like(U, dtype=torch.float32)
    rand_mat[U == 0] = -torch.inf
    # 过滤全0行:全0行不需要选保留位
    valid_rows = keep_num > 0
    # 初始化全0掩码
    mask = torch.zeros_like(U, dtype=torch.bool)
    # 对有1值的行,选随机值最高的keep_num个位置作为保留位
    selected_idx = torch.topk(rand_mat[valid_rows], k=keep_num[valid_rows], dim=1).indices
    mask[valid_rows] = mask[valid_rows].scatter(1, selected_idx, True)
    return U * mask
性能说明

两种方案均为纯张量向量化实现,完全规避了Python层面的for循环,在GPU上运行时可以充分利用并行算力,B值越大相比循环方案的性能优势越明显。其中概率版无排序操作,速度和原生torch.nn.Dropout基本持平;精确版仅做了一次逐行topk排序,开销也远低于逐行循环。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 13:42:29