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

如何用NumPy实现高效的广播式负采样优化?

高效负采样优化方案

原代码问题分析

原代码的核心瓶颈在于:

  • 循环遍历30000个正样本,每次都生成接近50000长度的候选数组(通过np.concatenate拼接两个范围数组),内存开销大且重复计算多
  • 每个采样操作都基于大数组,耗时严重

优化思路

利用索引映射的方式避免生成候选数组:每个正样本pos的候选集是除自身外的所有整数,我们可以将这些候选集映射到0~total-2的索引范围,采样后再反向映射回原整数范围,全程通过numpy向量化操作完成,无Python循环。

优化后代码

import numpy as np

total = 50000
positive_samples = np.random.choice(np.arange(total), size=(30000,), replace=False)

def get_negative_samples(positive_samples, total, num_negatives=100):
    n_pos = len(positive_samples)
    # 在0~total-2范围内采样,每行num_negatives个不重复元素
    rand_indices = np.random.choice(total - 1, size=(n_pos, num_negatives), replace=False)
    # 映射回原范围:索引≥pos时加1,跳过当前正样本
    negatives = rand_indices.copy()
    # 利用广播实现逐行比较
    mask = rand_indices >= positive_samples[:, np.newaxis]
    negatives[mask] += 1
    return negatives

为什么这个方法更快?

  • 无循环开销:所有操作都是numpy向量化计算,比Python循环快几个数量级
  • 无冗余内存占用:不需要生成任何大尺寸的候选数组,仅存储采样的索引和最终结果
  • 采样效率高:直接在缩小的索引范围内采样,避免了对大数组的遍历筛选

验证逻辑正确性

假设total=10,正样本pos=5,候选集为[0,1,2,3,4,6,7,8,9]:

  1. 采样范围是0~8(共9个元素,对应候选集的索引)
  2. 若采样到索引5,因为5≥5,映射为5+1=6,对应候选集中的6
  3. 若采样到索引3,因为3<5,直接映射为3,对应候选集中的3
  4. 所有采样结果自动避开了正样本5,且保证每行元素不重复

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 21:30:54