在JAX中实现非规则/非均匀数组行为的推荐方案
在JAX中处理非规则/非均匀数据的最优方案探讨
针对具有固有维度的非规则/非均匀数据,在JAX中实现数组行为与方法的推荐方案是什么?
目前有两种主流思路:
- 转为均匀数组并使用掩码(make homogeneous and use a mask)
- 扁平化数据并实现自定义方法(如广播、归约等)
显然第一种方案实现开销更低(验证和测试工作也更少),但存在内存复杂度问题。当内存优先级极高(需避免数组分布式存储)时,有没有比第二种方案更优的替代方案,既能利用JAX高度优化的数组方法,又能解决内存问题?
编辑补充:
以下是一个包含稀疏性的具体实现示例:
import jax as jx import jax.numpy as jnp jx.config.update("jax_enable_x64", True) # 问题相关静态变量 n_vars = 3 # 变量集数量 n_smps = 10 # 集合元素最大数量 p_smps = 0.2 # 问题稀疏度 # 每个集合包含不同数量的元素(例如二项分布随机生成) n_lvls = jx.random.bernoulli( jx.random.PRNGKey(0), p_smps, (n_vars, n_smps) ).sum(axis=1, dtype='i4') # 衍生量依赖于常数系数(例如均匀随机生成) a_vars = jx.random.uniform(jx.random.PRNGKey(1), (n_vars, ), dtype='f8') b_vars = jx.random.uniform(jx.random.PRNGKey(2), (n_vars, ), dtype='f8') b_vars = 10.0*b_vars c_vars = jx.random.uniform(jx.random.PRNGKey(3), (n_vars, ), dtype='f8') c_vars = 2.0*c_vars
该问题本质上仅需7个元素的状态即可表示。以下是方案1的一种实现:
### 转为均匀数组并使用掩码 ### # 定义层级索引数组 i_smps = jnp.arange(n_smps, dtype='i4') mask = n_lvls[:,None]>i_smps[None,:] # 生成满足单位公理的初始状态 x_vars = 1.0/(1.0+n_lvls[:,None]*i_smps[None,:]).astype('f8') x_vars = jnp.where(mask, x_vars, 0.0) x_vars = x_vars/x_vars.sum() # 生成系数张量 P_vars = a_vars[:,None]+b_vars[:,None]*i_smps[None,:] P_vars = jnp.where(mask, P_vars, 0.0) # 计算标量矩 scalar_moment = (x_vars*c_vars[:,None]).sum() # >>> DeviceArray(0.66574861, dtype=float64) # 计算转移张量 trans_tens = (P_vars[:,:,None,None]-P_vars[None,None,:,:]) trans_tens = trans_tens*x_vars[None,None,:,:]*x_vars[:,:,None,None] trans_tens.sum(axis=(2,3)) # >>> DeviceArray([[-0.37032842, 0.16153429, 0.22063015, 0.24335933, ...
转为均匀数组后,状态规模增至30个元素,且计算衍生量时会涉及大量零值乘法操作。
内容的提问来源于stack exchange,提问作者DavidJ
相关产品推荐
相关产品推荐

