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

如何在Jax中实现可vmap的动态范围求和并支持反向自动微分?

Jax动态kmax求和:vmap并行与反向微分实现方案

问题背景

需要在Jax中实现如下求和函数,通过vmap实现并行化,同时支持对输入x的反向模式自动微分:

def f(x,kmax):
  return sum ([x**k for k in range(1,kmax+1)])

注:此为简化示例,实际场景无闭合求和公式可用。

当前遇到的障碍:

  • 动态kmax下,jax.lax.fori_loop不支持反向微分;
  • jax.lax.scan要求输入为静态形状数组,否则触发ConcretizationTypeError;
  • Python原生range在vmap中使用会抛出TracerIntegerConversionError。

目标功能代码(运行报错):

import jax

def f(x,kmax):
  return sum ([x**k for k in range(1,kmax+1)])

fmap = jax.vmap(f,in_axes=(None,-1))

x = 3.
kmaxes = jax.numpy.array([1,2,3])

print(fmap(x,kmaxes))

fmap_sum = lambda k,kmaxes:jax.numpy.sum(fmap(k,kmaxes))

print(fmap_sum(x,kmaxes))
print(jax.grad(fmap_sum)(x,kmaxes))

报错位置在range(1,kmax+1),触发TracerIntegerConversionError。期望效果等价于以下纯Python循环代码,但需保留vmap的并行加速:

import jax

def f(x,kmax):
  return sum ([x**k for k in range(1,kmax+1)])

def fmap(x,kmaxes):
  return [f(x,kmax) for kmax in kmaxes]

x = 3.
kmaxes = jax.numpy.array([1,2,3])

print(fmap(x,kmaxes))

def fmap_sum(x,kmaxes):
  return sum(fmap(x,kmaxes))

print(fmap_sum(x,kmaxes))
print(jax.grad(fmap_sum)(x,kmaxes))

解决方案:静态序列+掩码实现动态求和

核心思路:先构造覆盖所有kmax最大值的静态k序列,预计算所有可能的x^k项,再通过布尔掩码筛选出每个kmax对应的有效项并求和。这种方式既满足Jax对静态形状的要求,又支持动态kmax的反向微分,同时兼容vmap并行。

高效实现代码

import jax
import jax.numpy as jnp

def f_dynamic(x, kmax):
    # 获取所有kmax中的最大值(静态整数,用于构造序列)
    max_k = kmax if isinstance(kmax, int) else kmax.max().item()
    # 生成1到max_k的静态序列
    k_seq = jnp.arange(1, max_k + 1)
    # 预计算所有x^k项
    terms = x ** k_seq
    # 生成掩码:保留k <= kmax的项
    mask = k_seq <= kmax
    # 对有效项求和
    return jnp.sum(terms * mask)

# 用vmap包装,针对kmaxes的最后一维并行处理
fmap = jax.vmap(f_dynamic, in_axes=(None, -1))

# 测试验证
x = 3.
kmaxes = jnp.array([1, 2, 3])

# 前向计算
print(fmap(x, kmaxes))  # 输出: [ 3. 12. 39.]

# 求和与梯度计算
fmap_sum = lambda x, ks: jnp.sum(fmap(x, ks))
print(fmap_sum(x, kmaxes))  # 输出: 54.
print(jax.grad(fmap_sum)(x, kmaxes))  # 输出: 14.0

关键细节说明

  1. 静态序列构造:通过max_k(静态整数)生成完整的k序列,满足Jax对输入形状静态可知的要求;
  2. 掩码筛选:利用jnp.arange(1, max_k+1) <= kmax生成布尔数组,过滤超出当前kmax的无效项;
  3. 反向微分兼容性:直接对带掩码的项求和,确保反向传播时仅计算有效项的梯度,避免无效项干扰;
  4. vmap并行支持:函数输入kmax为Jax数组,vmap可直接对其并行处理,完全保留并行加速效果。

内容的提问来源于stack exchange,提问作者Ian Holmes

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 22:15:32