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

如何从可JIT编译函数返回distrax分布对象?

解决JIT编译返回Distrax分布对象的Tracer泄漏问题

问题描述

想要创建可JIT编译的函数,返回Distrax分布对象,示例代码如下:

import distrax
import jax
import jax.numpy as jnp

def f(x):
   dist = distrax.Categorical(logits=jnp.sin(x))
   return dist

jit_f = jax.jit(f)
a = jnp.array([1,2,3])
dist = jit_f(a)

运行后出现UnexpectedTracerError,报错信息如下:

Traceback (most recent call last):
  File "<stdin>", line 1, in <module>
  File "F:\jax_env\lib\site-packages\jax\_src\traceback_util.py", line 162, in reraise_with_filtered_traceback
    return fun(*args, **kwargs)
  File "F:\jax_env\lib\site-packages\jax\_src\api.py", line 628, in cache_miss
    out = tree_unflatten(out_pytree_def, out_flat)
  File "F:\jax_env\lib\site-packages\jax\_src\tree_util.py", line 75, in tree_unflatten
    return treedef.unflatten(leaves)
  File "F:\jax_env\lib\site-packages\distrax\_src\utils\jittable.py", line 40, in tree_unflatten
    obj = cls(*args, **kwargs)
  File "F:\jax_env\lib\site-packages\distrax\_src\distributions\categorical.py", line 60, in __init__
    self._logits = None if logits is None else math.normalize(logits=logits)
  File "F:\jax_env\lib\site-packages\distrax\_src\utils\math.py", line 72, in normalize
    return jax.nn.log_softmax(logits, axis=-1)
  File "F:\jax_env\lib\site-packages\jax\_src\traceback_util.py", line 162, in reraise_with_filtered_traceback
    return fun(*args, **kwargs)
  File "F:\jax_env\lib\site-packages\jax\_src\api.py", line 618, in cache_miss
    keep_unused=keep_unused))
  File "F:\jax_env\lib\site-packages\jax\core.py", line 2031, in call_bind_with_continuation
    top_trace = find_top_trace(args)
  File "F:\jax_env\lib\site-packages\jax\core.py", line 1122, in find_top_trace
    top_tracer._assert_live()
  File "F:\jax_env\lib\site-packages\jax\interpreters\partial_eval.py", line 1486, in _assert_live
    raise core.escaped_tracer_error(self, None)
jax._src.traceback_util.UnfilteredStackTrace: jax._src.errors.UnexpectedTracerError: Encountered an unexpected tracer. A function transformed by JAX had a side effect, allowing for a reference to an intermediate value with type float32[3] wrapped in a DynamicJaxprTracer to escape the scope of the transformation.
JAX transformations require that functions explicitly return their outputs, and disallow saving intermediate values to global state.
The function being traced when the value leaked was f at <stdin>:1 traced for jit.
------------------------------
The leaked intermediate value was created on line <stdin>:2 (f).
------------------------------
When the value was created, the final 5 stack frames (most recent last) excluding JAX-internal frames were:
------------------------------
<stdin>:1 (<module>)
<stdin>:2 (f)
------------------------------

To catch the leak earlier, try setting the environment variable JAX_CHECK_TRACER_LEAKS or using the `jax.checking_leaks` context manager.
See https://jax.readthedocs.io/en/latest/errors.html#jax.errors.UnexpectedTracerError

尝试过jax.block_until_ready(dist),但未解决问题。

问题根源

Distrax的Categorical分布在构造时默认会执行log_softmax归一化操作,而JIT编译函数返回分布对象时,JAX会将其作为pytree序列化;反序列化时会重新调用分布的构造函数,此时归一化操作会尝试使用已脱离JIT追踪上下文的tracer,导致泄漏。

解决方案

方案1:提前完成参数计算,关闭分布内部归一化

在JIT函数内提前执行归一化计算,然后关闭分布构造时的自动归一化逻辑,避免反序列化时触发额外计算:

import distrax
import jax
import jax.numpy as jnp

@jax.jit
def f(x):
    logits = jnp.sin(x)
    # 提前执行归一化计算
    normalized_logits = jax.nn.log_softmax(logits, axis=-1)
    # 关闭分布内部的归一化,直接传入处理后的参数
    return distrax.Categorical(logits=normalized_logits, normalize=False)

a = jnp.array([1,2,3])
dist = f(a)

方案2:返回分布参数,外部构造分布

将函数改为仅返回分布的原始参数,在JIT函数外部再构造Distrax分布,彻底避免JIT处理分布对象的序列化问题:

import distrax
import jax
import jax.numpy as jnp

@jax.jit
def f(x):
    return jnp.sin(x)  # 返回原始logits参数

a = jnp.array([1,2,3])
logits = f(a)
# 外部构造分布
dist = distrax.Categorical(logits=logits)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 08:52:57