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.
相关产品推荐
相关产品推荐

