为什么在完整结构中包含None时,JAX会把前缀广播为None值,而不是广播前缀的值?
我对以下 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的结构 指的是什么。
默认情况下,jax 将 None 视为pytree节点,而不是叶子节点。具体来说,None 在 full_tree 中默认被解释为节点。因此,为了再现 full_tree的结构,还必须再现其非叶节点的 None。这解释了这种行为。
从更实际的角度来看,这使得像 jax.tree.map 这样的函数的默认行为能够在 out_tree 和 full_tree 上无缝工作,默认情况下忽略两者中相同的 None 节点。
站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。