为什么在完整结构中包含None时,JAX会把前缀广播为None值,而不是广播前缀的值?

编程语言 2026-07-11

我对以下 jax 行为有点困惑(版本0.6.2)。

通常,如果我将前缀树广播成结构树,前缀的值会广播到结构中:

out_tree = jax.tree.broadcast(True, (False, False), is_leaf=lambda x: x is None) # Outputs (True, True)

然而,如果结构树包含 None(标记为叶节点),那么广播会携带 None 的值作为代替:

out_tree = jax.tree.broadcast(True, (None, False), is_leaf=lambda x: x is None) # Outputs (None, True)

就广播而言,我原本期望结构树只是一个结构,其叶节点的值不应影响广播。这种行为的目的是什么?

解决方案

文档说:

返回值:一个pytree,与full_tree的结构相匹配,其中prefix_tree的叶子节点已被广播到每个相应子树的叶子节点。

这里的细微差别在于所说的 full_tree的结构 指的是什么。

默认情况下,jaxNone 视为pytree节点,而不是叶子节点。具体来说,Nonefull_tree 中默认被解释为节点。因此,为了再现 full_tree的结构,还必须再现其非叶节点的 None。这解释了这种行为。

从更实际的角度来看,这使得像 jax.tree.map 这样的函数的默认行为能够在 out_treefull_tree 上无缝工作,默认情况下忽略两者中相同的 None 节点。

站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。

相关文章