You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.17 23:31:19