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

jax.lax.fori_loop触发ConcretizationTypeError的解决方法

问题描述

以下JAX循环代码在step函数中使用Python内置min比较参数:

import jax

def step(timestep: int, order: int = 4) -> int:
    order = min(timestep + 1, order)
    return order

num_steps = 10
order = 100
order = jax.lax.fori_loop(0, num_steps, step, order)

运行后触发jax._src.errors.ConcretizationTypeError,完整错误堆栈如下:

WARNING:jax._src.lib.xla_bridge:未找到GPU/TPU,退回到CPU运行。(设置TF_CPP_MIN_LOG_LEVEL=0并重新运行以获取更多信息。)
---------------------------------------------------------------------------
UnfilteredStackTrace                      Traceback (most recent call last)
<ipython-input-4-9ec280f437cb> in <module>
      2 order = 100
----> 3 order = jax.lax.fori_loop(0, num_steps, step, order)

16 frames
/usr/local/lib/python3.8/dist-packages/jax/_src/traceback_util.py in reraise_with_filtered_traceback(*args, **kwargs)
    161     try:
--> 162       return fun(*args, **kwargs)
    163     except Exception as e:

/usr/local/lib/python3.8/dist-packages/jax/_src/lax/control_flow/loops.py in fori_loop(lower, upper, body_fun, init_val)
   1691 
--> 1692     (_, result), _ = scan(_fori_scan_body_fun(body_fun), (lower_, init_val),
   1693                           None, length=upper_ - lower_)

/usr/local/lib/python3.8/dist-packages/jax/_src/traceback_util.py in reraise_with_filtered_traceback(*args, **kwargs)
    161     try:
--> 162       return fun(*args, **kwargs)
    163     except Exception as e:

/usr/local/lib/python3.8/dist-packages/jax/_src/lax/control_flow/loops.py in scan(f, init, xs, length, reverse, unroll)
    258   # necessary, a second time with modified init values.
--> 259   init_flat, carry_avals, carry_avals_out, init_tree, *rest = _create_jaxpr(init)
    260   new_init_flat, changed = _promote_weak_typed_inputs(init_flat, carry_avals, carry_avals_out)

/usr/local/lib/python3.8/dist-packages/jax/_src/lax/control_flow/loops.py in _create_jaxpr(init)
    244     carry_avals = tuple(_map(_abstractify, init_flat))
--> 245     jaxpr, consts, out_tree = _initial_style_jaxpr(
    246         f, in_tree, (*carry_avals, *x_avals), "scan")

/usr/local/lib/python3.8/dist-packages/jax/_src/lax/control_flow/common.py in _initial_style_jaxpr(fun, in_tree, in_avals, primitive_name)
     59                          primitive_name: Optional[str] = None):
--> 60   jaxpr, consts, out_tree = _initial_style_open_jaxpr(
     61       fun, in_tree, in_avals, primitive_name)

/usr/local/lib/python3.8/dist-packages/jax/_src/lax/control_flow/common.py in _initial_style_open_jaxpr(fun, in_tree, in_avals, primitive_name)
     53   debug = pe.debug_info(fun, in_tree, False, primitive_name or "<unknown>")
--> 54   jaxpr, _, consts = pe.trace_to_jaxpr_dynamic(wrapped_fun, in_avals, debug)
     55   return jaxpr, consts, out_tree()

/usr/local/lib/python3.8/dist-packages/jax/_src/profiler.py in wrapper(*args, **kwargs)
    313     with TraceAnnotation(name, **decorator_kwargs):
--> 314       return func(*args, **kwargs)
    315     return wrapper

/usr/local/lib/python3.8/dist-packages/jax/interpreters/partial_eval.py in trace_to_jaxpr_dynamic(fun, in_avals, debug_info, keep_inputs)
   1980     main.jaxpr_stack = ()  # type: ignore
--> 1981     jaxpr, out_avals, consts = trace_to_subjaxpr_dynamic(
   1982       fun, main, in_avals, keep_inputs=keep_inputs, debug_info=debug_info)

/usr/local/lib/python3.8/dist-packages/jax/interpreters/partial_eval.py in trace_to_subjaxpr_dynamic(fun, main, in_avals, keep_inputs, debug_info)
   1997     in_tracers_ = [t for t, keep in zip(in_tracers, keep_inputs) if keep]
--> 1998     ans = fun.call_wrapped(*in_tracers_)
   1999     out_tracers = map(trace.full_raise, ans)

/usr/local/lib/python3.8/dist-packages/jax/linear_util.py in call_wrapped(self, *args, **kwargs)
    166     try:
--> 167       ans = self.f(*args, **dict(self.params, **kwargs))
    168     except:

/usr/local/lib/python3.8/dist-packages/jax/_src/lax/control_flow/loops.py in scanned_fun(loop_carry, _)
   1607     i, x = loop_carry
--> 1608     return (i + 1, body_fun()(i, x)), None
   1609   return scanned_fun

<ipython-input-2-2e3345899235> in step(timestep, order)
      1 def step(timestep: int, order: int = 100) -> int:
--> 2     order = min(timestep + 1, order)
      3     return order

/usr/local/lib/python3.8/dist-packages/jax/core.py in __bool__(self)
    633   def __nonzero__(self): return self.aval._nonzero(self)
--> 634   def __bool__(self): return self.aval._bool(self)
    635   def __int__(self): return self.aval._int(self)

/usr/local/lib/python3.8/dist-packages/jax/core.py in error(self, arg)
   1266   def error(self, arg):
--> 1267     raise ConcretizationTypeError(arg, fname_context)
   1268   return error

UnfilteredStackTrace: jax._src.errors.ConcretizationTypeError: 遇到抽象追踪值,但预期为具体值:Traced<ShapedArray(bool[], weak_type=True)>with<DynamicJaxprTrace(level=1/0)>
问题出在`bool`函数上。
错误发生在扫描函数scanned_fun的追踪过程中,路径为/usr/local/lib/python3.8/dist-packages/jax/_src/lax/control_flow/loops.py:1606。该具体值在Python中不可用,因为它依赖于参数'loop_carry'的值。

请查看https://jax.readthedocs.io/en/latest/errors.html#jax.errors.ConcretizationTypeError

下面的堆栈跟踪排除了JAX内部帧。
以上是未修改的原始异常。

--------------------

上述异常直接导致以下异常:

ConcretizationTypeError                   Traceback (most recent call last)
<ipython-input-4-9ec280f437cb> in <module>
      1 num_steps = 10
      2 order = 100
----> 3 order = jax.lax.fori_loop(0, num_steps, step, order)

<ipython-input-2-2e3345899235> in step(timestep, order)
      1 def step(timestep: int, order: int = 100) -> int:
--> 2     order = min(timestep + 1, order)
      3     return order

ConcretizationTypeError: 遇到抽象追踪值,但预期为具体值:Traced<ShapedArray(bool[], weak_type=True)>with<DynamicJaxprTrace(level=1/0)>
问题出在`bool`函数上。
错误发生在扫描函数scanned_fun的追踪过程中,路径为/usr/local/lib/python3.8/dist-packages/jax/_src/lax/control_flow/loops.py:1606。该具体值在Python中不可用,因为它依赖于参数'loop_carry'的值。

请查看https://jax.readthedocs.io/en/latest/errors.html#jax.errors.ConcretizationTypeError

使用普通Python循环代码可正常运行,但性能较差,如何修复?

解决方案

错误原因

Python内置min函数执行比较时,会尝试将结果转换为Python原生布尔值,但JAX的fori_loop在追踪过程中处理的是抽象追踪值,无法直接转换为具体的Python布尔值,因此触发ConcretizationTypeError。

修复方法

使用JAX提供的jax.lax.min替代Python内置min,它是专为JAX的追踪和编译流程设计的操作,能在JAX控制流中正常工作,同时保证性能。

修复后的代码:

import jax

def step(timestep: int, order: int = 4) -> int:
    order = jax.lax.min(timestep + 1, order)
    return order

num_steps = 10
order = 100
order = jax.lax.fori_loop(0, num_steps, step, order)

说明

jax.lax.min会生成对应的XLA操作,完全兼容JAX的追踪机制,既避免了具体值依赖的问题,又能利用JAX的编译优化保证运行性能,解决Python循环的性能劣势。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 12:40:11