Jax实现Hutchinson-Skilling估计器:简单场景精度不足疑问
首先明确,你的测试案例中,函数f(x)=x²在输入[1.0,1.0]处的真实散度是4(雅可比矩阵的迹:21 + 21=4)。
你遇到的“需要大量高斯样本才能收敛”的核心原因是高斯随机向量的单个估计样本方差过大,导致均值收敛速度慢:
1. 高斯样本的方差计算
对于你的案例,单个高斯样本v对应的估计值是:v·Jf(x)v = 2x₁v₁² + 2x₂v₂²
其中v₁,v₂是独立标准正态变量,E[v_i²]=1,但Var(v_i²)=E[v_i⁴]-E[v_i²]²=3-1=2。
单个估计值的方差为:Var(2x₁v₁² + 2x₂v₂²) = (2x₁)²*2 + (2x₂)²*2 = 4*2 +4*2=16
n个样本均值的方差为16/n,根据中心极限定理,均值的标准差是4/√n。要让估计值与真实值4的误差控制在1%(即±0.04)以内,需要3*(4/√n) ≤0.04,解得n≥90000——这和你观察到的“约10k样本才能勉强控制误差”的现象一致(实际中随机性会让部分情况提前达标,但整体收敛慢)。
2. 优化方案:改用Rademacher随机变量
如果用拉德马赫(Rademacher)随机变量(每个元素独立取±1,概率各0.5),单个样本的估计值方差会骤降为0,因为v_i²=1,此时:v·Jf(x)v =2x₁*1 +2x₂*1=4
单个样本就能得到完全准确的结果,根本不需要大量采样。
修改后的代码示例:
import jax import jax.numpy as jnp from jax import random def f(x): return x**2 def get_div_fn(f, step_rng, num_samples=1): key = step_rng def hutchinson_skilling_divergence(x): def single_sample_divergence(x, key): # 生成Rademacher随机向量:先取正态,再取符号 v = jnp.sign(jax.random.normal(key, shape=x.shape)) jacobian_vector_product = jax.jvp(f, (x,), (v,))[1] return jnp.dot(v, jacobian_vector_product) keys = jax.random.split(key, num_samples) divergence_estimates = jax.vmap(single_sample_divergence, in_axes=(None, 0))(x, keys) return jnp.mean(divergence_estimates) return hutchinson_skilling_divergence input = jnp.array([1.0, 1.0]) key = random.PRNGKey(42) f_div = get_div_fn(f, key) result = f_div(input) print(result) # 输出应为4.0
3. 原代码的正确性说明
你的原代码实现本身是正确的,Hutchinson-Skilling估计器对于高斯向量的数学期望确实等于散度。只是高斯样本的方差特性导致收敛速度远慢于Rademacher这类离散各向同性向量。
内容的提问来源于stack exchange,提问作者Antoine M

