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

如何对cho_solve进行向量化?遇ValueError问题求助

JAX中cho_solve的vectorize使用问题

问题背景

已经实现了在cho_solve上使用vmap的需求,现在想尝试用jnp.vectorize处理,但遇到了问题:cho_solve的第一个参数是元组,而vectorize默认要求所有输入均为数组,运行代码时抛出ValueError: All input arrays must have the same shape。

使用场景是存在数量不确定的batch维度,vectorize的自动广播特性更符合需求,不想通过solve_triangular自行实现cho_solve,同时想了解vectorize是否支持类似vmap的嵌套签名。

报错代码

import jax
import jax.numpy as jnp
import jax.scipy as jsp

key = jax.random.PRNGKey(0)
key, subkey = jax.random.split(key)

k_y = jax.random.normal(subkey, (3, 5, 10, 10))
y = jnp.broadcast_to(jnp.eye(10), k_y.shape)

matmul = jnp.vectorize(jnp.matmul, signature='(a,b),(b,c)->(a,c)')
cholesky = jnp.vectorize(jsp.linalg.cholesky, excluded={1}, signature='(d,d)->(d,d)')
cho_solve = jnp.vectorize(jsp.linalg.cho_solve, signature='(d,d),(d,d)->(d,d)')  # 此处该如何设置?

k_y = matmul(k_y, jnp.moveaxis(k_y, -1, -2))
chol = cholesky(k_y, True)
result = cho_solve((chol, True), y)

解决方案

1. 封装cho_solve拆分元组参数

jnp.vectorize无法直接处理元组输入,我们可以先封装一个函数,把cho_solve的元组参数拆分为独立输入,再对封装后的函数使用vectorize:

import jax
import jax.numpy as jnp
import jax.scipy as jsp

key = jax.random.PRNGKey(0)
key, subkey = jax.random.split(key)

k_y = jax.random.normal(subkey, (3, 5, 10, 10))
y = jnp.broadcast_to(jnp.eye(10), k_y.shape)

matmul = jnp.vectorize(jnp.matmul, signature='(a,b),(b,c)->(a,c)')
cholesky = jnp.vectorize(jsp.linalg.cholesky, excluded={1}, signature='(d,d)->(d,d)')

# 封装cho_solve,拆分元组参数
def cho_solve_wrapper(chol, lower, y):
    return jsp.linalg.cho_solve((chol, lower), y)

# 定义vectorize版本,排除lower参数(标量),指定签名
vectorized_cho_solve = jnp.vectorize(cho_solve_wrapper, excluded={1}, signature='(d,d),(),(d,d)->(d,d)')

k_y = matmul(k_y, jnp.moveaxis(k_y, -1, -2))
chol = cholesky(k_y, True)
result = vectorized_cho_solve(chol, True, y)

2. 灵活使用vmap处理不确定的batch维度

如果batch维度数量不确定,vmap也能通过in_axes参数灵活指定批量轴,实现类似vectorize的广播效果:

# 假设批量轴是前两个维度,in_axes指定元组中chol的批量轴,lower为标量(无批量轴),y的批量轴
vmap_cho_solve = jax.vmap(jsp.linalg.cho_solve, in_axes=((0, 1), None), out_axes=(0, 1))
result = vmap_cho_solve((chol, True), y)

关于嵌套签名的说明

目前jax.numpy.vectorize的signature参数不支持嵌套元组的签名定义,无法直接对接受元组输入的函数使用签名式vectorize,必须先通过封装拆分参数来解决。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 14:13:09