如何在PyTorch中为张量创建基于行全零条件的二进制掩码
PyTorch实现行全0判断的二进制Mask
实现思路
对输入的n×m张量逐行判断是否所有元素均为0,将判断结果取反后转换为二进制数值,最后调整形状为n×1的mask张量。
代码实现
import torch # 示例输入:3×4的张量 input_tensor = torch.tensor([ [0, 0, 0, 0], # 全0行,对应mask为0 [1, 0, 2, 0], # 非全0行,对应mask为1 [0, 0, 0, 1] # 非全0行,对应mask为1 ]) # 1. 判断每行是否全为0(得到形状为(n,)的布尔张量) all_zero_rows = torch.all(input_tensor == 0, dim=1) # 2. 取反并转换为数值型,再扩展为n×1形状 mask = (~all_zero_rows).float().unsqueeze(1) print(mask)
输出结果
tensor([[0.], [1.], [1.]])
补充说明
- 若需要整数类型的mask(如
long型),将.float()替换为.long()即可。 dim=1指定沿列维度(即每行)进行判断;unsqueeze(1)用于将一维张量扩展为二维的n×1格式。
内容的提问来源于stack exchange,提问作者EEQ
相关产品推荐
相关产品推荐

