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字符串做长度截断,在运行触发报错的代码前,先执行以下配置修改:
配置生效后重新运行代码,无论是本地环境还是Colab环境,控制台都会输出完整的PyTree结构对比信息,不会再出现末尾省略号截断的问题。import jax # 关闭PyTree结构打印的深度限制 jax.config.update("jax_pytree_def_max_depth", None) - 方案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
相关产品推荐
相关产品推荐

