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
相关产品推荐
相关产品推荐

