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

如何优化列表随机负采样(规避指定元素)的代码性能?

优化负样本生成速度的方案

原代码性能瓶颈分析

原代码耗时的核心原因是循环内重复执行全量数据过滤:每次循环都对50万行的all_pairs执行all_pairs.loc[all_pairs['TARGET_ID'] != target_id]操作,单次过滤是O(n)复杂度,10240次循环累计带来巨大的性能开销。此外,每次都提取values数组,重复操作进一步拖慢速度。


优化方案一:针对原代码逻辑(负样本为其他TARGET_ID)

预处理减少重复计算

先提取所有唯一的TARGET_ID,避免每次过滤全量数据:

import pandas as pd
import numpy as np
import time

num_neg_samples = 4

# 样本数据生成
rng = np.random.default_rng()
all_pairs = pd.DataFrame({
    'TARGET_ID': rng.choice(1_000_000, size=500_000, replace=True),
    'CONTEXT_ID': rng.choice(1_000_000, size=500_000, replace=True)
})
pair_sample = all_pairs.sample(10240)

# 预处理:获取所有唯一的TARGET_ID(去重后候选集更小)
unique_targets = all_pairs['TARGET_ID'].unique()

start = time.time()
targets = np.repeat(pair_sample['TARGET_ID'].values, num_neg_samples + 1)
negids = []

for target_id in pair_sample['TARGET_ID']:
    # 快速生成排除当前target_id的候选集(比过滤全量数据快100+倍)
    candidates = np.setdiff1d(unique_targets, [target_id], assume_unique=True)
    neg_samples = rng.choice(candidates, size=num_neg_samples)
    negids.append(neg_samples)

print(time.time() - start, 'seconds')

进一步矢量化优化(几乎无循环)

通过批量生成随机候选,再替换冲突值,彻底避免逐行过滤:

start = time.time()
targets_np = pair_sample['TARGET_ID'].values.reshape(-1, 1)
# 批量生成所有负样本候选
neg_candidates = rng.choice(unique_targets, size=(len(pair_sample), num_neg_samples), replace=True)
# 找到与当前target_id重复的位置
mask = neg_candidates == targets_np

# 处理冲突:替换重复的值
for i in range(len(pair_sample)):
    conflict_count = mask[i].sum()
    if conflict_count == 0:
        continue
    # 生成新的候选,直到没有重复
    new_samples = rng.choice(unique_targets, size=conflict_count)
    while (new_samples == targets_np[i]).any():
        bad_mask = new_samples == targets_np[i]
        new_samples[bad_mask] = rng.choice(unique_targets, size=bad_mask.sum())
    neg_candidates[i][mask[i]] = new_samples

negids = neg_candidates.tolist()
print(time.time() - start, 'seconds')

优化方案二:针对实际业务需求(负样本为不存在的物品对)

如果你的负样本定义是**(TARGET_ID, CONTEXT_ID) 未在原列表中出现**,可以通过哈希集合快速校验,避免全量查询:

start = time.time()
# 预处理:将所有正样本对转为哈希集合,O(1)查询
positive_pairs = set(zip(all_pairs['TARGET_ID'].values, all_pairs['CONTEXT_ID'].values))
unique_contexts = all_pairs['CONTEXT_ID'].unique()

neg_contexts = []
for target_id in pair_sample['TARGET_ID']:
    valid_neg = []
    # 一次生成双倍候选,减少循环次数
    candidates = rng.choice(unique_contexts, size=num_neg_samples * 2)
    # 过滤掉已存在的正样本对
    valid_neg = [c for c in candidates if (target_id, c) not in positive_pairs]
    # 补充不足的数量
    while len(valid_neg) < num_neg_samples:
        c = rng.choice(unique_contexts)
        if (target_id, c) not in positive_pairs:
            valid_neg.append(c)
    neg_contexts.append(valid_neg[:num_neg_samples])

print(time.time() - start, 'seconds')

额外优化建议

  1. 减少数据复制:尽量使用numpy数组而非Pandas Series进行操作,避免不必要的类型转换。
  2. 预先生成全量候选集:如果训练过程中all_pairs不变化,可将预处理的unique_targets、positive_pairs等保存为文件,避免每次训练重复预处理。
  3. 并行生成:如果样本量极大,可使用multiprocessing模块并行生成负样本,但需注意随机种子的一致性。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.12 07:07:01