【发布时间】:2022-11-21 06:40:22
【问题描述】:
在 JAX 中,我希望在固定长度的数据类列表上 vmap 一个函数,例如:
import jax, chex
from flax import struct
@struct.dataclass
class EnvParams:
max_steps: int = 500
random_respawn: bool = False
def foo(params: EnvParams):
...
param_list = jnp.Array([EnvParams(max_steps=500), EnvParams(max_steps=600)])
jax.vmap(foo)(param_list)
上面的示例失败,因为无法创建自定义对象的 jnp.Array,并且 JAX 不允许在 Python 列表上进行 vmapping。我看到的唯一剩余选项是转换数据类以表示一批参数,如下所示:
@struct.dataclass
class EnvParamBatch:
max_steps: jnp.Array = jnp.array([500, 600])
random_respawn: jnp.Array = jnp.array([False, True])
def bar(params):
...
jax.vmap(bar)(EnvParamBatch())
最好使用结构容器(每个结构代表一个参数集),所以我想知道是否有任何替代方法?
注意我知道 this answer,但这不是完全相同的问题,现在可能有更好的解决方案。
【问题讨论】:
-
JAX 的
vmap不能对结构数组进行操作,但可以对数组结构进行操作,因此您的第二个解决方案是您应该与 JAX 一起使用的方法。我会添加一个答案,但您似乎已经回答了您的问题!