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

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)

关键修改说明

  1. 拆分纯Python计算与JAX包装:把原来的_b_spline_eval拆分为纯Python实现`_b_spline_eval
相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 05:50:24