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

如何快速创建含指定比例1的PyTorch二进制张量?

优化指定占比的二进制PyTorch张量生成方法

你的问题核心在于原代码用torch.randperm处理19000×19000的超大张量时,会生成包含3.61亿元素的完整随机排列索引,这不仅内存占用爆炸,全量排序的时间开销也极高。下面提供两种更高效的实现方案:

方法一:近似占比(最快,内存占用极低)

如果允许1的数量是近似指定百分比(基于伯努利分布随机采样),直接用PyTorch内置的torch.bernoulli一步生成,这是最优解:

import torch
import time

start_time = time.time()

def create_random_binary_tensor_fast(shape, percentage, device):
    # 直接生成符合伯努利分布的张量,1的出现概率为指定百分比
    tensor = torch.bernoulli(torch.full(shape, percentage/100.0, device=device, dtype=torch.float32))
    return tensor

# 参数设置
shape = (19000, 19000)
percentage = 0.5
device = torch.device("cuda:2" if torch.cuda.is_available() else "cpu")

tensor = create_random_binary_tensor_fast(shape, percentage, device)

execution_time = time.time() - start_time
print(f"Execution Time: {execution_time:.4f} seconds")

核心优势:

  • 内存占用:无需生成全量索引,直接在目标形状上生成张量,内存仅为最终张量大小(约1.38GB,远低于原代码的19GB+)
  • 速度:内置算子经过高度优化,CUDA上几乎瞬时完成,CPU上也比原代码快一个数量级

方法二:精确占比(内存和速度均优于原代码)

如果需要精确数量的1,用torch.rand生成随机数后,通过topk筛选出指定数量的随机位置,避免生成完整随机排列:

import torch
import time

start_time = time.time()

def create_random_binary_tensor_exact(shape, percentage, device):
    size = shape[0] * shape[1]
    num_ones = int(size * (percentage / 100.0))
    
    # 生成随机数张量
    rand_tensor = torch.rand(size, device=device)
    # 获取前num_ones个最小随机数的索引(等价于随机选中num_ones个位置)
    _, indices = torch.topk(rand_tensor, num_ones, largest=False)
    
    # 创建全零张量并设置对应位置为1
    tensor = torch.zeros(size, dtype=torch.float32, device=device)
    tensor[indices] = 1
    return tensor.view(shape)

# 参数设置
shape = (19000, 19000)
percentage = 0.5
device = torch.device("cuda:2" if torch.cuda.is_available() else "cpu")

tensor = create_random_binary_tensor_exact(shape, percentage, device)

execution_time = time.time() - start_time
print(f"Execution Time: {execution_time:.4f} seconds")

核心优势:

  • 内存:topk无需存储完整随机排列,仅需保存随机数张量和索引张量,内存占用比原代码降低约50%
  • 速度:topk的时间复杂度为O(n log k),而randperm是O(n log n),当k远小于n时(比如这里0.5%占比,k=1.805e6,n=3.61e8),速度提升非常显著

原代码慢的根本原因

torch.randperm(size)会生成包含0到size-1的完整随机排列张量,对3.61亿元素做全量排序的时间和内存开销都极大。上面两种方法都避开了全量排序操作,直接针对需求生成结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 09:02:33