custom_vjp 和 nondiff_argnums 更新指南

custom_vjpnondiff_argnums 更新指南#

mattjj@ 2020 年 10 月 14 日

本文档假设读者熟悉 jax.custom_vjp,详见 Custom derivative rules for JAX-transformable Python functions (JAX 可转换 Python 函数的自定义导数规则) 笔记本。

需要更新的内容#

在 JAX PR #4008 之后,传递给 custom_vjp 函数的 nondiff_argnums 的参数不能是 Tracer(或包含 Tracer 的容器)。这意味着为了支持任意可转换的代码,nondiff_argnums 不应用于数组值参数。相反,nondiff_argnums 应仅用于非数组值,例如 Python 可调用对象、形状元组或字符串。

在任何曾经使用 nondiff_argnums 处理数组值的地方,我们现在应该直接将它们作为普通参数传递。在 bwd 规则中,我们需要为它们生成对应的值,但我们可以直接产生 None 值,以表明没有对应的梯度值。

例如,以下是 clip_gradient写法,当 hi 和/或 lo 是来自某些 JAX 变换的 Tracer 时,这种写法将无法工作。

from functools import partial
import jax

@partial(jax.custom_vjp, nondiff_argnums=(0, 1))
def clip_gradient(lo, hi, x):
  return x  # identity function

def clip_gradient_fwd(lo, hi, x):
  return x, None  # no residual values to save

def clip_gradient_bwd(lo, hi, _, g):
  return (jnp.clip(g, lo, hi),)

clip_gradient.defvjp(clip_gradient_fwd, clip_gradient_bwd)

以下是支持任意变换的的、更棒的写法

import jax

@jax.custom_vjp  # no nondiff_argnums!
def clip_gradient(lo, hi, x):
  return x  # identity function

def clip_gradient_fwd(lo, hi, x):
  return x, (lo, hi)  # save lo and hi values as residuals

def clip_gradient_bwd(res, g):
  lo, hi = res
  return (None, None, jnp.clip(g, lo, hi))  # return None for lo and hi

clip_gradient.defvjp(clip_gradient_fwd, clip_gradient_bwd)

如果您使用旧写法而不是新写法,在任何可能出错的情况下(即当 Tracer 被传入 nondiff_argnums 参数时),您都会收到明显的报错。

这里是一个确实需要结合 custom_vjp 使用 nondiff_argnums 的情况

from functools import partial
import jax

@partial(jax.custom_vjp, nondiff_argnums=(0,))
def skip_app(f, x):
  return f(x)

def skip_app_fwd(f, x):
  return skip_app(f, x), None

def skip_app_bwd(f, _, g):
  return (g,)

skip_app.defvjp(skip_app_fwd, skip_app_bwd)

解释#

Tracer 传入 nondiff_argnums 参数一直是有问题的。虽然某些情况下可以正常工作,但在其他情况下会导致复杂且令人困惑的错误信息。

该 Bug 的本质在于 nondiff_argnums 的实现方式非常类似于词法闭包 (lexical closure)。但在当时,针对 Tracer 的词法闭包并不打算与 custom_jvp/custom_vjp 协同工作。以那种方式实现 nondiff_argnums 是一个错误!

PR #4008 修复了 custom_jvpcustom_vjp 中所有关于词法闭包的问题。 太棒了!也就是说,现在的 custom_jvpcustom_vjp 函数及规则可以随意闭合 Tracer。对于所有非自动微分的变换,一切都会“自然生效”。对于自动微分变换,我们将获得清晰的错误消息,说明为什么我们不能对 custom_jvpcustom_vjp 所闭合的值进行微分。

检测到对 custom_jvp 函数中闭合的值进行了微分。这是不支持的,因为自定义 JVP 规则仅指定了如何相对于显式输入参数来微分 custom_jvp 函数。

尝试将闭合的值作为参数传入 custom_jvp 函数,并调整自定义 JVP 规则。

在以这种方式加强和稳健化 custom_jvpcustom_vjp 的过程中,我们发现允许 custom_vjp 在其 nondiff_argnums 中接受 Tracer 将需要大量的额外开销:我们需要重写用户的 fwd 函数以将这些值作为残差返回,并重写用户的 bwd 函数以将它们作为常规残差接受(而不是像 nondiff_argnums 那样作为特殊的首个参数接受)。这看起来或许尚可管理,直到您思考如何处理任意的 pytree!此外,这种复杂性是不必要的:如果用户代码将类数组的不可微分参数像普通参数和残差一样对待,一切已经可以正常工作。(在 #4039 之前,JAX 可能会抱怨涉及整数值的输入和输出的自动微分,但在 #4039 之后,这些将直接生效!)

custom_vjp 不同,让 custom_jvp 支持作为 Tracernondiff_argnums 参数很容易。因此,这些更新仅需在 custom_vjp 中进行。