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

Google Colab中如何打印被截断的完整Jax错误信息

问题描述

使用JAX的jax.lax.scan函数时会触发TypeError,报错提示scan的carry输出与输入必须具备相同类型结构。报错信息中返回的PyTreeDef结构信息末尾带有省略号,内容被截断,无法查看完整结构详情。
典型截断报错片段如下:

TypeError: scan carry output and input must have same type structure, got PyTreeDef((CustomNode(<class 'brax.experimental.braxlines.training.env.EnvState'>[()], [CustomNode(<class 'brax.envs.env.State'>[()], [CustomNode(<class 'brax.physics.base.QP'>[()], [*, *, *, *]), *, *, *, {'agent_idx': *, 'reward': *, 'reward_contact_cost': *, 'reward_ctrl_cost': *, 'reward_forward': *, 'reward_survive': *, 'score': *}, {'agent_idx': *, 'first_obs': *, 'first_qp': CustomNode(<class 'brax.physics.base.QP'>[()], [*, *, *, *]), 'rng': *, 'static_agent_policy': {'normalizer': (*, *, *), 'policy': [{'params': {'hidden_0': {'bias': *, 'kernel': *}, 'hidden_1': {'bias': *, 'kernel': *}, 'hidden_2': {'bias': *, 'kernel': *}, 'hidden_3': {'bias': *, 'kernel': *}, 'hidden_4': {'bias': *, 'kernel': *}}}, {'params': {'hidden_0': {'bias': *, 'kernel': *}, 'hidden_1': {'bias': *, 'kernel': *}, 'hidden_2': {'bias': *, 'kernel': *}, 'hidden_3': {'bias': *, 'kernel': *}, 'hidden_4': {'bias': *, 'kernel': *}}}]}, 'steps': *, 'truncation': *}]), {'agent_idx': *, 'reward': *, 'reward_contact_cost': *, 'reward_ctrl_cost': *, 'reward_forward': *, 'reward_survive': *, 'score': *}, *]), [CustomNode(<class 'flax.core.frozen_dict.FrozenDict'>[()], [{'params': {'hidden_0': {'bias': *, 'kernel': *}, 'hidden_1': {'bias': *, 'kernel': *}, 'hidden_2': {'bias': *, 'kernel': *}, 'hidden_3': {'bias': *, 'kernel': *}, 'hidden_4': {'bias': *, 'kernel': *}}}])], (*, *, *), [None], *)) and PyTreeDe...

以下代码可复现同类截断报错:

def f(carry, xslice):
  new_carry = carry['this'] * 2
  return new_carry, xslice

jax.lax.scan(f, init={'this': 1}, xs=(), length=2)

需要获取完整无截断的报错信息,验证两种可行路径:一是将完整报错写入txt文件,二是通过配置强制Colab输出完整报错内容。

解决方法

两种路径均可实现,操作步骤如下:

  • 方案1:修改JAX全局配置,关闭PyTree打印截断
    JAX默认对过长的PyTreeDef字符串做长度截断,在运行触发报错的代码前,先执行以下配置修改:
    import jax
    # 关闭PyTree结构打印的深度限制
    jax.config.update("jax_pytree_def_max_depth", None)
    
    配置生效后重新运行代码,无论是本地环境还是Colab环境,控制台都会输出完整的PyTree结构对比信息,不会再出现末尾省略号截断的问题。
  • 方案2:捕获完整异常栈写入txt文件
    通过Python原生异常捕获逻辑拿到完整报错traceback,直接写入本地文件留存,示例代码如下:
    import traceback
    import jax
    
    try:
        # 替换为实际触发报错的业务代码
        def f(carry, xslice):
          new_carry = carry['this'] * 2
          return new_carry, xslice
        jax.lax.scan(f, init={'this': 1}, xs=(), length=2)
    except TypeError:
        full_error_log = traceback.format_exc()
        with open("full_scan_error.txt", "w", encoding="utf-8") as f:
            f.write(full_error_log)
    
    代码运行完成后,当前工作目录下会生成full_scan_error.txt文件,存储无截断的完整报错信息。

这类报错的根因是scan传入的step函数返回的新carry结构,和初始传入的carry结构不匹配。比如上方复现示例中,初始carry是字典结构{'this': 1},但返回的新carry是单个标量值,结构不一致就会触发错误。拿到完整PyTree结构后,逐节点对比输入输出carry的结构差异,即可快速定位问题点。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 09:54:21