【问题标题】:Jax - vmap over batch of dataclassesJax - 对一批数据类进行 vmap
【发布时间】: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 一起使用的方法。我会添加一个答案,但您似乎已经回答了您的问题!

标签: python jax flax


【解决方案1】:

vmap 无法处理对象列表,只能处理包含数组的单个对象。这是一个例子:

import typing
import jax
import jax.numpy as jnp

class EnvParams(typing.NamedTuple):
    max_steps: int = 500
    random_respawn: bool = False

param_array = EnvParams(
    max_steps=jnp.array([500, 600]),
    random_respawn=jnp.array([False, False]))
vmap_param_array = jax.vmap(lambda x: x)(param_array)

大多数时候最好使用上述方法,这样对象就可以存储在 GPU / TPU 内存中而不是 CPU 中......但是如果你真的必须在 CPU 上的列表/数组之间进行转换,这里有一个例子:

def list_to_array(list):
    cls = type(list[0])
    return cls(**{k: jnp.array([getattr(v, k) for v in list]) for k in cls._fields})

def array_to_list(array):
    cls = type(array)
    size = len(getattr(array, cls._fields[0]))
    return [cls(**{k: v(getattr(array, k)[i]) for k, v in cls._field_types.items()}) for i in range(size)]

param_list = [EnvParams(max_steps=500), EnvParams(max_steps=600)]
param_array = list_to_array(param_list)
vmap_param_array = jax.vmap(lambda x: x)(param_array)
vmap_param_list = array_to_list(vmap_param_array)

【讨论】:

    猜你喜欢
    • 2021-06-07
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2021-09-20
    • 1970-01-01
    • 2021-11-05
    • 1970-01-01
    • 2021-08-11
    相关资源
    最近更新 更多