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

JAX中负采样算法的性能优化与可变数组存储方案咨询

JAX负采样算法优化与可变长度数组存储方案

一、负采样性能与逻辑优化

你的当前方案存在Python循环开销和大尺寸概率数组内存浪费的问题,针对大数据量场景(20万条样本、5万范围),可以从以下方向优化:

1. 替换Python循环为向量化操作

JAX的Python循环会触发多次独立的数组操作,效率极低。改用jax.vmap实现批量处理,配合JIT编译能大幅提升性能。

2. 重构采样逻辑,避免大尺寸概率数组

直接构造全范围概率数组在max_range=5万时完全不现实,推荐采用乱序过滤采样或拒绝采样的方式:

方案1:乱序过滤采样(稳定可控)

先将可变长度的不可接受项转换成固定形状的二维数组(用无效值填充),然后对每个样本生成乱序索引,过滤掉不可接受项后取前n_samples:

import jax.numpy as jnp
import jax

max_range = 5
n_samples = 2
true_cases = jnp.array([[1,2], [1,4], [0,5]])

# 原字典转固定长度二维数组,用-1填充空项(假设索引不会为-1)
non_acceptable_as_negatives = {
    0: jnp.array([5]), 1: jnp.array([2,4]), 2: jnp.array([]),
    3: jnp.array([]), 4: jnp.array([]), 5: jnp.array([])
}
max_non_accept_len = max(len(v) for v in non_acceptable_as_negatives.values())
non_accept_array = jnp.array([
    jnp.pad(v, (0, max_non_accept_len - len(v)), constant_values=-1)
    for k in range(max_range + 1)
])

@jax.jit
def sample_single(i, key):
    # 获取当前i的有效不可接受索引
    non_accept = non_accept_array[i][non_accept_array[i] != -1]
    # 生成全范围乱序索引
    permuted = jax.random.permutation(key, jnp.arange(max_range + 1))
    # 过滤不可接受项,取前n_samples
    valid = permuted[jnp.logical_not(jnp.isin(permuted, non_accept))]
    return valid[:n_samples]

# 批量处理所有样本
keys = jax.random.split(jax.random.PRNGKey(42), len(true_cases))
negatives = jax.vmap(sample_single)(true_cases[:, 0], keys)

方案2:拒绝采样(适合不可接受项比例低的场景)

生成候选样本后过滤不可接受项,若数量不足则补充采样,内存占用更低:

@jax.jit
def rejection_sample_single(i, key, n_samples):
    key, subkey = jax.random.split(key)
    # 生成候选(数量设为n_samples的2倍,可根据实际比例调整)
    candidates = jax.random.randint(subkey, (n_samples * 2,), 0, max_range + 1)
    non_accept = non_accept_array[i][non_accept_array[i] != -1]
    valid = candidates[jnp.logical_not(jnp.isin(candidates, non_accept))]
    
    # 递归补充不足的样本
    if valid.shape[0] < n_samples:
        remaining = n_samples - valid.shape[0]
        _, new_key = jax.random.split(key)
        additional = rejection_sample_single(i, new_key, remaining)
        return jnp.concatenate([valid, additional])
    return valid[:n_samples]

negatives = jax.vmap(rejection_sample_single, in_axes=(0, 0, None))(
    true_cases[:,0], keys, n_samples
)

二、JAX中可变长度数组的存储方案

JAX原生偏好固定形状数组,针对你的场景,推荐以下两种存储方式:

  • 固定长度二维数组填充:如上面示例,用无效值(如-1)将所有可变长度数组补全到相同长度,转换为二维数组。这种方式最容易和JAX的向量化、JIT操作兼容,实现成本最低。
  • 稀疏数组存储:如果不可接受项整体稀疏(大部分i对应空数组),可以用JAX的稀疏数组模块jax.experimental.sparse,避免填充带来的内存浪费:
    from jax.experimental.sparse import BCOO
    
    indices = []
    for i in non_acceptable_as_negatives:
        for val in non_acceptable_as_negatives[i]:
            indices.append((i, val))
    
    # 构建稀疏矩阵,标记(i, val)为不可接受对
    sparse_non_accept = BCOO(
        jnp.array(indices), 
        jnp.ones(len(indices)), 
        shape=(max_range+1, max_range+1)
    )
    
    # 查询i对应的不可接受项:sparse_non_accept[i].indices[:, 1]
    

内容的提问来源于stack exchange,提问作者Simon P.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 14:55:23