如何用JAX并行运行树形结构状态的马尔可夫链?
用JAX vmap并行运行多初始状态的马尔可夫链(树形字典状态)
JAX的vmap天然支持PyTree结构(包括你用Haiku生成的嵌套字典状态),不需要额外将字典转换成数组,只要保证批量状态的PyTree结构对齐、每个叶子节点带有批量维度即可,具体实现步骤如下:
1. 构造批量初始状态
将多个独立的初始状态合并为一个带批量维度的PyTree,每个叶子节点(比如weights、bias)的第0维对应链的索引。可以用jax.tree_map批量堆叠单个状态的叶子节点:
import jax import jax.numpy as jnp import haiku as hk # 单条链的状态初始化函数 def init_single_chain(key): key1, key2 = jax.random.split(key) return { 'layer1': { 'weights': jax.random.normal(key1, (32, 64)), 'bias': jax.random.normal(key1, (64,)), }, 'layer2': { 'weights': jax.random.normal(key2, (64, 10)), 'bias': jax.random.normal(key2, (10,)), }, } # 生成3条独立链的初始状态 key = jax.random.PRNGKey(42) keys = jax.random.split(key, num=3) single_states = [init_single_chain(k) for k in keys] # 合并为批量状态:每个叶子节点堆叠成[batch_size, ...]的shape batch_state = jax.tree_map(lambda *xs: jnp.stack(xs), *single_states) # 此时layer1/weights的shape为(3, 32, 64),对应3条链的权重
2. 用vmap封装step函数
通过vmap的in_axes参数指定哪些输入需要映射批量维度:
- 对于
state,指定in_axes=0表示映射每个叶子节点的第0轴(批量轴) - 对于随机密钥
key,需要生成对应批量数的密钥,同样映射第0轴 - 对于共享的
X_train,指定in_axes=None表示所有链复用同一数据
示例代码:
# 原状态推进函数(假设你已实现advance逻辑) def advance(state, key, X_train): # 示例:给权重添加小噪声模拟状态转移 new_state = jax.tree_map( lambda param: param + jax.random.normal(key, param.shape) * 0.01, state ) return new_state def step(state: dict, key, X_train) -> dict: return advance(state, key, X_train) # 封装批量step函数 batch_step = jax.vmap(step, in_axes=(0, 0, None)) # 准备批量随机密钥(每条链对应一个独立密钥) batch_keys = jax.random.split(key, num=3) # 训练数据(所有链共享) X_train = jax.random.normal(key, (1000, 32)) # 运行一次批量状态推进 new_batch_state = batch_step(batch_state, batch_keys, X_train) # 返回的new_batch_state结构与batch_state完全一致,每个叶子节点保持[3, ...]的shape
3. 批量多步推进(可选)
如果需要运行多步马尔可夫链,可以结合jax.lax.scan和vmap,实现高效的批量多步迭代:
def batch_multistep(batch_state, step_batch_keys, X_train, num_steps): def scan_fn(current_state, step_keys): next_state = batch_step(current_state, step_keys, X_train) return next_state, next_state # 返回当前状态用于累积记录 final_state, all_states = jax.lax.scan( scan_fn, init=batch_state, xs=step_batch_keys ) return final_state, all_states # 生成10步的批量密钥:每一步对应3条链的独立密钥 step_keys = jax.random.split(key, num=10) step_batch_keys = jnp.array([jax.random.split(k, 3) for k in step_keys]) # 运行10步批量推进 final_batch_state, all_batch_states = batch_multistep(batch_state, step_batch_keys, X_train, 10) # all_states的shape为[10, 3, ...],记录了每一步所有链的状态
注意事项
- 所有初始状态的PyTree结构必须完全一致(比如不能有的链有额外层),否则
vmap会抛出结构不匹配的错误 - 如果每条链需要使用不同的训练数据子集,可以给
X_train添加批量维度,同时将in_axes调整为(0, 0, 0)
内容的提问来源于stack exchange,提问作者Hylke
相关产品推荐
相关产品推荐

