如何使用JAX vmap高效计算重要性采样估计
问题描述
我有一段计算强化学习离线策略重要性采样估计的代码,核心是处理一个自定义Episode类实例的数组,每个Episode包含四个浮点数组属性(observations、actions、rewards、action_probs)。函数需要遍历每个episode计算单个浮点数结果,最终返回结果数组。
原始未优化代码:
def IS_estimate(model, theta, episodes): """ Calculate the unweighted importance sampling estimate for each episode in episodes. Return as an array, one element per episode """ # episodes is an array of custom Python class instances gamma = 1.0 result = np.zeros(len(episodes)) for ii, ep in enumerate(episodes): obs = ep.observations # 1D array of floats actions = ep.actions # 1D array of floats rewards = ep.rewards # 1D array of floats action_probs = ep.action_probs # 1D array of floats pi_news = np.zeros(len(obs)) for jj in range(len(obs)): pi_news[jj] = model.get_prob_this_action(obs[jj],actions[jj]) pi_ratio_prod = np.prod(pi_news / action_probs) weighted_return = weighted_sum_gamma(rewards, gamma) result[ii] = pi_ratio_prod * weighted_return return np.array(result)
我已经用jax.vmap消除了生成pi_news的内层循环,优化后的代码如下:
def IS_estimate(model, theta, episodes): """ Calculate the unweighted importance sampling estimate for each episode in episodes. Return as an array, one element per episode """ # episodes is an array of custom Python class instances gamma = 1.0 result = np.zeros(len(episodes)) for ii, ep in enumerate(episodes): obs = ep.observations # 1D array of floats actions = ep.actions # 1D array of floats rewards = ep.rewards # 1D array of floats action_probs = ep.action_probs # 1D array of floats vmapped_get_prob_this_action = vmap(model.get_prob_this_action,in_axes=(0,0)) pi_news = vmapped_get_prob_this_action(obs,actions) pi_ratio_prod = np.prod(pi_news / action_probs) weighted_return = weighted_sum_gamma(rewards, gamma) result[ii] = pi_ratio_prod * weighted_return return np.array(result)
但JAX不支持直接对自定义Episode对象使用vmap,所以外部循环还没法用vmap优化,想问问怎么实现对外部循环的vmap?
解决方案
要对外部循环用vmap,核心思路是把自定义Episode对象的属性转换成JAX能处理的数组结构,而非直接传递Episode实例数组。具体步骤如下:
1. 批量提取所有episode的属性
先将所有episode的四个属性分别提取出来,组成批量数组:
import jax.numpy as jnp # 提取所有episode的属性,每个属性转为JAX数组 obs_batch = jnp.array([ep.observations for ep in episodes]) actions_batch = jnp.array([ep.actions for ep in episodes]) rewards_batch = jnp.array([ep.rewards for ep in episodes]) action_probs_batch = jnp.array([ep.action_probs for ep in episodes])
注意:若不同episode的属性数组长度不一致,JAX数组无法直接处理,需先做padding补全,统一子数组长度,同时记录每个episode的真实长度,后续计算时忽略padding部分。
2. 编写单episode计算的纯函数
将针对单个episode的计算逻辑抽成JAX兼容的纯函数(无副作用、仅依赖输入):
def compute_single_episode(model, gamma, obs, actions, rewards, action_probs): # 用vmap处理单episode内的概率计算 vmapped_get_prob = vmap(model.get_prob_this_action, in_axes=(0, 0)) pi_news = vmapped_get_prob(obs, actions) pi_ratio_prod = jnp.prod(pi_news / action_probs) weighted_return = weighted_sum_gamma(rewards, gamma) return pi_ratio_prod * weighted_return
3. 对单episode函数做vmap,处理批量输入
将批量属性作为输入,用vmap实现批量计算:
def IS_estimate(model, theta, episodes): gamma = 1.0 # 提取批量属性 obs_batch = jnp.array([ep.observations for ep in episodes]) actions_batch = jnp.array([ep.actions for ep in episodes]) rewards_batch = jnp.array([ep.rewards for ep in episodes]) action_probs_batch = jnp.array([ep.action_probs for ep in episodes]) # 定义vmap后的批量计算函数 # in_axes说明:model和gamma是全局参数,不参与vmap;四个属性数组在第0维上批量处理 vmapped_compute = vmap(compute_single_episode, in_axes=(None, None, 0, 0, 0, 0)) # 执行批量计算 result = vmapped_compute(model, gamma, obs_batch, actions_batch, rewards_batch, action_probs_batch) return result
处理变长episode的情况
若不同episode的属性数组长度不一致,需先做padding补全:
import jax.numpy as jnp # 找到最长的episode长度 max_len = max(len(ep.observations) for ep in episodes) def pad_array(arr, max_len): return jnp.pad(arr, (0, max_len - len(arr)), mode='constant') # 批量提取并补全 obs_batch = jnp.array([pad_array(ep.observations, max_len) for ep in episodes]) actions_batch = jnp.array([pad_array(ep.actions, max_len) for ep in episodes]) rewards_batch = jnp.array([pad_array(ep.rewards, max_len) for ep in episodes]) action_probs_batch = jnp.array([pad_array(ep.action_probs, max_len) for ep in episodes]) # 记录每个episode的真实长度 lengths = jnp.array([len(ep.observations) for ep in episodes])
随后修改单episode函数,用长度mask忽略padding部分:
def compute_single_episode(model, gamma, obs, actions, rewards, action_probs, length): # 只取真实长度内的元素 obs = obs[:length] actions = actions[:length] rewards = rewards[:length] action_probs = action_probs[:length] vmapped_get_prob = vmap(model.get_prob_this_action, in_axes=(0, 0)) pi_news = vmapped_get_prob(obs, actions) pi_ratio_prod = jnp.prod(pi_news / action_probs) weighted_return = weighted_sum_gamma(rewards, gamma) return pi_ratio_prod * weighted_return
对应的vmap调用也要加上length参数:
vmapped_compute = vmap(compute_single_episode, in_axes=(None, None, 0, 0, 0, 0, 0)) result = vmapped_compute(model, gamma, obs_batch, actions_batch, rewards_batch, action_probs_batch, lengths)
内容的提问来源于stack exchange,提问作者marvin
相关产品推荐
相关产品推荐

