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

