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

JAX向量化指南:如何规范实现协方差函数向量化以生成协方差矩阵?

Idiomatic JAX Approach to Vectorize Covariance Matrix Calculation

Great question! Your two-step vmap implementation is totally valid, but there are more concise and idiomatic ways to generate the covariance matrix in JAX—whether you want to stick close to your original cov function or leverage JAX's built-in array operations for better efficiency.

Option 1: Leverage Matrix Multiplication (Most Efficient)

First, let's recall the mathematical definition behind your cov function: the covariance (as you've defined it) between two sequences is the dot product of their centered versions. For a data matrix D (shape [n_samples, n_features]), the full covariance matrix is just the Gram matrix of the centered data.

Instead of using vmap twice, you can compute this directly with matrix multiplication (or einsum for explicit clarity):

import jax.numpy as jnp

def cov(x, y):
    return jnp.dot((x - jnp.mean(x)), (y - jnp.mean(y)))

def cov_matrix(D):
    # Center the data: subtract feature-wise mean from each sample
    centered = D - jnp.mean(D, axis=0, keepdims=True)
    # Compute Gram matrix (each entry (i,j) is dot product of centered column i and j)
    return centered.T @ centered

This is the most efficient approach because JAX optimizes matrix multiplication operations heavily, and it avoids the overhead of nested vmap calls. If you prefer explicit indexing, you can also use jnp.einsum:

return jnp.einsum('ni,nj->ij', centered, centered)

Option 2: Concise Nested vmap (Stays True to Your Original Function)

If you want to keep using your original cov function and stick with vmap, you can collapse the two vmap calls into a single line for cleaner, more compact code:

cov_matrix = jax.vmap(jax.vmap(cov, in_axes=(None, 1)), in_axes=(1, None))

This does exactly the same thing as your original code—first vectorizing over the second argument of cov (calculating covariance between a fixed column and all columns), then vectorizing over the first argument (repeating this for every column in D). It’s just a more streamlined way to write the same logic.

Bonus: Scaling to Unbiased Covariance

If you later want the unbiased covariance estimate (dividing by n_samples - 1 instead of n_samples), you can easily adjust either approach:

# For matrix multiplication approach
n_samples = D.shape[0]
return (centered.T @ centered) / (n_samples - 1)

# For vmap approach
def cov_unbiased(x, y):
    return jnp.dot((x - jnp.mean(x)), (y - jnp.mean(y))) / (len(x) - 1)
cov_matrix_unbiased = jax.vmap(jax.vmap(cov_unbiased, in_axes=(None, 1)), in_axes=(1, None))

内容的提问来源于stack exchange,提问作者Igor Rivin

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 23:42:38