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

Jax实现Transformer的三个问题:JIT失效、softmax调用及NaN异常

JAX实现Transformer训练中的三个问题解决

1. jax.disable_jit()无法移除隐式JIT编译

  • 原因:jax.disable_jit()仅能禁用隐式触发的JIT(比如JAX自动对纯函数的JIT优化),但无法取消显式用jax.jit装饰的函数(比如你的step函数大概率被jax.jit装饰了)。另外,JAX部分内置函数(如jax.nn.softmax)内部也可能包含显式JIT逻辑,不受全局禁用开关影响。
  • 解决方法:
    • 检查训练流程中的step函数,若被jax.jit装饰,暂时移除该装饰器;
    • 若需完全禁用JIT,避免调用任何内部带显式JIT的函数,必要时手动实现替代逻辑(比如手动写softmax)。

2. jax.nn.softmax默认调用_softmax_deprecated

  • 原因:这是JAX的版本兼容机制导致的——要么你使用的JAX版本较旧,旧版本的softmax核心实现就是_softmax_deprecated;要么你调用softmax时传递的参数(如where、initial)匹配了旧版本的接口规范,触发了兼容回退逻辑。
  • 解决方法:
    • 升级JAX到最新稳定版本,新版本的jax.nn.softmax已经替换为更高效的非deprecated实现;
    • 对照JAX官方文档检查softmax的参数传递方式,确保符合当前版本的接口要求,避免传递过时的参数组合。

3. _softmax_deprecated中减法操作出现NaN

  • 原因:你的create_mask用np.NINF填充padding位置,当某条输入序列全为padding时,scaled_dot_prod会被mask填充为全NINF。此时jnp.max(x)得到NINF,执行x - lax.stop_gradient(x_max)就会出现NINF - NINF,结果为NaN,最终导致jnp.exp(NaN)触发浮点错误。
  • 解决方法:
    1. 统一使用JAX常量替换numpy常量,避免跨库类型问题:
      def create_mask(arr):
          return jnp.where(arr == 0, jnp.NINF, 0)  # 用jnp.NINF代替np.NINF
      
    2. 调用softmax时指定where参数,仅计算有效位置的最大值:
      # 在SelfAttention的__call__中修改softmax调用
      valid_positions = mask != jnp.NINF
      return (jax.nn.softmax(scaled_dot_prod, where=valid_positions) @ value)
      
    3. 给jnp.max设置initial参数,避免全NINF时得到无效值:
      若自定义softmax逻辑,将x_max的计算改为:
      x_max = jnp.max(x, axis, where=where, initial=-1e10, keepdims=True)
      
    4. 检查数据加载器,过滤掉全padding的样本,避免这类无效输入进入模型。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 16:33:13