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

jax.lax.scan返回的堆叠值能否回传给参数?示例列表始终为空

JAX: 将lax.scan的累积值回传为carry参数的可行方案

你的代码里这么做不可行,问题出在这两点:

  1. 你传入的result是普通Python列表,JAX的lax.scan要求carry参数必须是JAX可追踪的数组类型,Python列表无法被JAX追踪,而且迭代过程中你也没对u(即这个列表)做任何更新,所以打印始终是空。
  2. lax.scan返回的堆叠值是迭代结束后的最终结果,没法在迭代过程中回传给carry参数。

如果要在迭代过程中累积值并将其作为carry的一部分传递,正确的做法是用JAX数组维护累积状态,每次迭代更新这个数组:

from jax import lax, numpy as jnp

def cumsum(carry, el):
    current_sum, accumulated = carry
    new_sum = current_sum + el
    # 将当前计算的累积和追加到累积数组中
    new_accumulated = jnp.concatenate([accumulated, jnp.array([new_sum])])
    return (new_sum, new_accumulated), new_sum

# 初始化carry:初始和为0,累积数组初始为空JAX数组
init_carry = (0.0, jnp.array([]))
a = jnp.array([1, 2, 3, 4])

(final_carry, scan_output) = lax.scan(cumsum, init_carry, a)

final_sum, accumulated_values = final_carry
print("迭代过程中累积的所有值:", accumulated_values)
print("scan返回的堆叠输出:", scan_output)

关键说明:

  • 用jnp.array替代Python列表作为累积状态的载体,确保JAX能正确追踪状态变化。
  • 每次迭代通过jnp.concatenate生成新的累积数组,作为carry的一部分传递给下一次迭代。
  • 最终accumulated_values就是迭代过程中每次的累积结果,和scan_output的内容完全一致——因为scan_output本质就是每次返回的new_sum的堆叠结果。

如果处理的是大数据量,频繁创建新数组可能有性能损耗,这种情况下可以预先分配固定长度的数组,通过索引更新的方式维护累积状态,但lax.scan本身更适合这种顺序依赖的累积场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 13:52:38