JAX+NumPyro启用JIT时,如何实现B样条输入的自定义微分?
可JIT兼容的JAX B样条求导实现(适配NumPyro)
我在NumPyro中使用JAX,希望通过scipy.interpolate.BSpline实现B样条函数,将依赖模型参数的输入点转换为样条曲线。需求是仅对输入参数x求导,不对节点(knots)或样条阶数(order)求导。
用jax.custom_vjp能实现无JIT场景下的功能,但在NumPyro启用JIT后无法正常运行。考虑使用callback方案解决,但不清楚具体实现方式——JAX文档中反向模式自动微分的TensorFlow示例未启用JIT。
原无JIT兼容代码
from scipy.interpolate import BSpline import numpy as np from numpy import typing as npt from functools import partial import jax doubleArray = npt.NDArray[np.double] # 参考B样条导数公式实现 def _b_spline_deriv_inner(spline: BSpline, deriv_basis: doubleArray) -> doubleArray: out = np.zeros((deriv_basis.shape[0], deriv_basis.shape[1] - 1)) for col_index in range(out.shape[1] - 1): scale = spline.t[col_index + spline.k + 1] - spline.t[col_index + 1] if scale != 0: out[:, col_index] = -deriv_basis[:, col_index + 1] / scale for col_index in range(1, out.shape[1]): scale = spline.t[col_index + spline.k] - spline.t[col_index] if scale != 0: out[:, col_index] += deriv_basis[:, col_index] / scale return float(spline.k) * out def _b_spline_eval(spline: BSpline, x: doubleArray, deriv: int) -> doubleArray: if deriv == 0: return spline.design_matrix(x=x, t=spline.t, k=spline.k).todense() elif spline.k <= 0: return np.zeros((x.shape[0], spline.t.shape[0] - spline.k - 1)) return _b_spline_deriv_inner( spline=spline, deriv_basis=_b_spline_eval( BSpline(t=spline.t, k=spline.k - 1, c=np.zeros(spline.c.shape[0] + 1)), x=x, deriv=deriv - 1 ), ) @partial(jax.custom_vjp, nondiff_argnums=(0, 1, 2)) def b_spline_basis(knots: doubleArray, order: int, deriv: int, x: doubleArray) -> doubleArray: return _b_spline_eval(spline=BSpline(t=knots, k=order, c=np.zeros((order + knots.shape[0] - 1))), x=x, deriv=deriv)[ :, 1: ] def b_spline_basis_fwd(knots: doubleArray, order: int, deriv: int, x: doubleArray) -> tuple[doubleArray, doubleArray]: spline = BSpline(t=knots, k=order, c=np.zeros(order + knots.shape[0] - 1)) return ( _b_spline_eval(spline=spline, x=x, deriv=deriv)[:, 1:], _b_spline_eval(spline=spline, x=x, deriv=deriv + 1)[:, 1:], ) def b_spline_basis_bwd( knots: doubleArray, order: int, deriv: int, partials: doubleArray, grad: doubleArray ) -> tuple[doubleArray]: return (jax.numpy.sum(partials * grad, axis=1),) b_spline_basis.defvjp(b_spline_basis_fwd, b_spline_basis_bwd) if __name__ == "__main__": knots = np.array([0, 0, 0, 0, 0.25, 1, 1, 1, 1]) x = np.array([0.1, 0.5, 0.9]) order = 3 def test_jax(basis: doubleArray, partials: doubleArray, deriv: int) -> None: weights = jax.numpy.arange(1, basis.shape[1] + 1) def test_func(x: doubleArray) -> doubleArray: return jax.numpy.sum(jax.numpy.dot(b_spline_basis(knots=knots, order=order, deriv=deriv, x=x), weights)) assert np.allclose(test_func(x), np.sum(np.dot(basis, weights))) assert np.allclose(jax.grad(test_func)(x), np.dot(partials, weights)) # 预定义的测试基准值 deriv0 = np.transpose( np.array( [ 0.684, 0.166666666666667, 0.00133333333333333, 0.096, 0.444444444444444, 0.0355555555555555, 0.004, 0.351851851851852, 0.312148148148148, 0, 0.037037037037037, 0.650962962962963, ] ).reshape(-1, 3) ) deriv1 = np.transpose( np.array( [ 2.52, -1, -0.04, 1.68, -0.666666666666667, -0.666666666666667, 0.12, 1.22222222222222, -2.29777777777778, 0, 0.444444444444444, 3.00444444444444, ] ).reshape(-1, 3) ) test_jax(deriv0, deriv1, deriv=0) deriv2 = np.transpose( np.array( [ -69.6, 4, 0.8, 9.6, -5.33333333333333, 5.33333333333333, 2.4, -2.22222222222222, -15.3777777777778, 0, 3.55555555555556, 9.24444444444445, ] ).reshape(-1, 3) ) test_jax(deriv1, deriv2, deriv=1) deriv3 = np.transpose( np.array( [ 504, -8, -8, -144, 26.6666666666667, 26.6666666666667, 24, -32.8888888888889, -32.8888888888889, 0, 14.2222222222222, 14.2222222222222, ] ).reshape(-1, 3) ) test_jax(deriv2, deriv3, deriv=2)
适配JIT的解决方案
要让代码兼容JIT,核心是用jax.pure_callback把SciPy的非JAX兼容代码包装起来,同时保留自定义VJP的导数规则。修改后的代码如下:
from scipy.interpolate import BSpline import numpy as np from numpy import typing as npt from functools import partial import jax import jax.numpy as jnp doubleArray = npt.NDArray[np.double] # 定义纯Python的B样条计算函数(供callback调用) def _b_spline_deriv_inner_py(spline: BSpline, deriv_basis: doubleArray) -> doubleArray: out = np.zeros((deriv_basis.shape[0], deriv_basis.shape[1] - 1)) for col_index in range(out.shape[1] - 1): scale = spline.t[col_index + spline.k + 1] - spline.t[col_index + 1] if scale != 0: out[:, col_index] = -deriv_basis[:, col_index + 1] / scale for col_index in range(1, out.shape[1]): scale = spline.t[col_index + spline.k] - spline.t[col_index] if scale != 0: out[:, col_index] += deriv_basis[:, col_index] / scale return float(spline.k) * out def _b_spline_eval_py(knots: doubleArray, order: int, x: doubleArray, deriv: int) -> doubleArray: spline = BSpline(t=knots, k=order, c=np.zeros(order + knots.shape[0] - 1)) if deriv == 0: return spline.design_matrix(x=x, t=spline.t, k=spline.k).todense() elif spline.k <= 0: return np.zeros((x.shape[0], spline.t.shape[0] - spline.k - 1)) return _b_spline_deriv_inner_py( spline=spline, deriv_basis=_b_spline_eval_py(knots, order-1, x, deriv-1) ) # 用jax.pure_callback包装计算函数,指定输出形状和dtype def _b_spline_eval_jax(knots: jnp.ndarray, order: int, x: jnp.ndarray, deriv: int) -> jnp.ndarray: # 定义输出形状:(x.shape[0], len(knots) - order - 1 + (deriv>0)) output_shape = jax.ShapeDtypeStruct( shape=(x.shape[0], len(knots) - order - 1 + (0 if deriv ==0 else 1)), dtype=jnp.float64 ) return jax.pure_callback( _b_spline_eval_py, output_shape, knots, order, x, deriv, vectorized=True # 支持输入x的批量处理 ) @partial(jax.custom_vjp, nondiff_argnums=(0, 1, 2)) def b_spline_basis(knots: jnp.ndarray, order: int, deriv: int, x: jnp.ndarray) -> jnp.ndarray: return _b_spline_eval_jax(knots, order, x, deriv)[:, 1:] def b_spline_basis_fwd(knots: jnp.ndarray, order: int, deriv: int, x: jnp.ndarray) -> tuple[jnp.ndarray, jnp.ndarray]: basis = _b_spline_eval_jax(knots, order, x, deriv)[:, 1:] partials = _b_spline_eval_jax(knots, order, x, deriv+1)[:, 1:] return basis, partials def b_spline_basis_bwd( knots: jnp.ndarray, order: int, deriv: int, partials: jnp.ndarray, grad: jnp.ndarray ) -> tuple[jnp.ndarray]: return (jnp.sum(partials * grad, axis=1),) b_spline_basis.defvjp(b_spline_basis_fwd, b_spline_basis_bwd) if __name__ == "__main__": knots = jnp.array([0, 0, 0, 0, 0.25, 1, 1, 1, 1]) x = jnp.array([0.1, 0.5, 0.9]) order = 3 def test_jax(basis: doubleArray, partials: doubleArray, deriv: int) -> None: weights = jnp.arange(1, basis.shape[1] + 1) def test_func(x: jnp.ndarray) -> jnp.ndarray: return jnp.sum(jnp.dot(b_spline_basis(knots=knots, order=order, deriv=deriv, x=x), weights)) # 测试无JIT场景 assert np.allclose(test_func(x), np.sum(np.dot(basis, weights))) assert np.allclose(jax.grad(test_func)(x), np.dot(partials, weights)) # 测试JIT场景 test_func_jit = jax.jit(test_func) assert np.allclose(test_func_jit(x), np.sum(np.dot(basis, weights))) assert np.allclose(jax.jit(jax.grad(test_func))(x), np.dot(partials, weights)) print("JIT场景测试通过") # 预定义的测试基准值 deriv0 = np.transpose( np.array( [ 0.684, 0.166666666666667, 0.00133333333333333, 0.096, 0.444444444444444, 0.0355555555555555, 0.004, 0.351851851851852, 0.312148148148148, 0, 0.037037037037037, 0.650962962962963, ] ).reshape(-1, 3) ) deriv1 = np.transpose( np.array( [ 2.52, -1, -0.04, 1.68, -0.666666666666667, -0.666666666666667, 0.12, 1.22222222222222, -2.29777777777778, 0, 0.444444444444444, 3.00444444444444, ] ).reshape(-1, 3) ) test_jax(deriv0, deriv1, deriv=0) deriv2 = np.transpose( np.array( [ -69.6, 4, 0.8, 9.6, -5.33333333333333, 5.33333333333333, 2.4, -2.22222222222222, -15.3777777777778, 0, 3.55555555555556, 9.24444444444445, ] ).reshape(-1, 3) ) test_jax(deriv1, deriv2, deriv=1) deriv3 = np.transpose( np.array( [ 504, -8, -8, -144, 26.6666666666667, 26.6666666666667, 24, -32.8888888888889, -32.8888888888889, 0, 14.2222222222222, 14.2222222222222, ] ).reshape(-1, 3) ) test_jax(deriv2, deriv3, deriv=2)
关键修改说明
- 拆分纯Python计算与JAX包装:把原来的
_b_spline_eval拆分为纯Python实现`_b_spline_eval
相关产品推荐
相关产品推荐

