向量化JAX内层循环:支持反向自动微分且避免冗余计算
JAX嵌套循环优化问题
我在JAX中实现了嵌套循环,内层循环有两种实现方案:
- 第一种使用
jax.lax.fori_loop:优势是仅执行获取结果所需的最少操作,但存在两大缺点:1)无法并行化;2)嵌套循环通常不支持反向自动微分(如官方文档所述)。 - 第二种采用并行广播实现:支持反向自动微分且可并行,但会对大于内层循环计数器的索引执行大量冗余操作。
示例代码
# CHOOSE FUNCTION def comb(n, k): # 实现组合数comb(n,k),输入必须为整数。这里的取整操作可能多余,但保留以确保结果为整数 return jnp.round(jnp.exp(jax.scipy.special.gammaln(n+1) - jax.scipy.special.gammaln(k+1) - jax.scipy.special.gammaln(n-k+1))) # FIRST INNER LOOP IMPLEMENTATION: FORI_LOOP def inner_loop_fori(n, Aks, Bks): init_conv = jnp.zeros(Aks.shape[1:], dtype=Aks.dtype) return jax.lax.fori_loop(0, n, update_inner_loop, (init_conv, n, Aks, Bks))[0] # 循环前n个元素 def update_inner_loop(k, val): conv, n, Aks, Bks = val conv = conv + comb(n-1, k) * Aks[k] @ Bks[(n-1)-k] return conv, n, Aks, Bks # SECOND INNER LOOP IMPLEMENTATION: BROADCAST def multiply_along_axis(vec, arr, axis): return jnp.swapaxes(vec * jnp.swapaxes(arr, axis, -1), -1, axis) def inner_loop_bcst(n, Aks, Bks): Bks = jnp.roll(Bks, len(Bks)-n, axis=0) Bks = jnp.flip(Bks, axis=0) coeffs = comb(n-1, jnp.arange(len(Aks))) # 索引>=n时系数为0 conv = jnp.sum(multiply_along_axis(coeffs, Aks @ Bks, axis=0), axis=0) # 能否避免索引>=n的冗余操作? return conv # TOGGLE INNER LOOP IMPLEMENTATION def inner_loop(n, Aks, Bks): conv = inner_loop_bcst(n, Aks, Bks) # conv = inner_loop_fori(n, Aks, Bks) # 切换内层循环实现方式 return conv # OUTER LOOP def outer_loop(n, Aks, Bks): return jax.lax.scan(scan_func, (Aks, Bks), jnp.arange(n))[0][0] def scan_func(carry, k): Aks, Bks = carry Aks = Aks.at[k+1].set(inner_loop(k+1, Aks, Bks)) # 每一步理论上只需要操作内层循环的k+1个元素 return (Aks, Bks), k+1 # 需要求梯度的目标函数 def main(t): # 初始化数组 Aks = t * jnp.ones((5,1,1)) Bks = t**2 * jnp.ones((5,1,1)) n = len(Aks) return outer_loop(n, Aks, Bks)
使用示例
main(2.) # 评估函数,两种实现都能正常运行 jax.jacrev(main)(2.) # 反向模式微分,使用fori_loop实现会失败
提问
请问是否可以修改上述代码,实现鱼与熊掌兼得:内层循环既能并行化、支持反向自动微分,又不会执行冗余计算?
我强烈怀疑答案是否定的,因为这必然涉及对数组进行动态大小的切片,而任何条件操作都可能默认使用lax.select,会执行所有分支。但或许在这个特定场景中我忽略了某些显而易见的方法。
附:我可以用标准Python循环替代lax.scan,但更希望保留lax.scan。
解决方案:切片+向量化计算消除冗余
你可以通过提前切片截取有效数据结合向量化运算的方式,同时满足并行化、反向微分支持和无冗余计算的需求,核心思路是只对参与计算的前n个元素操作,完全规避无效索引的运算。
优化后的内层循环实现
def inner_loop_optimized(n, Aks, Bks): # 仅截取前n个有效元素,截断后续冗余部分 A_slice = Aks[:n] B_slice = Bks[:n] # 翻转B的前n个元素,匹配原逻辑中的(n-1)-k索引映射 B_flipped = jnp.flip(B_slice, axis=0) # 仅计算0到n-1索引对应的组合数 coeffs = comb(n-1, jnp.arange(n)) # 向量化计算:系数乘对应矩阵积后求和,维度适配通过[:, None, None]实现 conv = jnp.sum(coeffs[:, None, None] * (A_slice @ B_flipped), axis=0) return conv
替换内层循环
将原inner_loop函数替换为优化版本:
def inner_loop(n, Aks, Bks): conv = inner_loop_optimized(n, Aks, Bks) return conv
方案优势
- 无冗余计算:通过
Aks[:n]和Bks[:n]直接限定运算范围,后续所有操作都只针对有效数据,彻底消除原广播方案中对>=n索引的无效运算。 - 支持并行化:所有运算都是向量化的,JAX可自动进行并行调度。
- 兼容反向微分:未使用
fori_loop这类不支持反向传播的循环结构,所有操作都属于JAX可微分的范畴。
验证
main(2.) # 正常执行 jax.jacrev(main)(2.) # 反向微分正常运行
内容的提问来源于stack exchange,提问作者Ben
相关产品推荐
相关产品推荐

