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

使用JAX grad函数计算复数数组梯度报错的解决方案咨询

问题:JAX grad计算复数数组Loss梯度报错

运行JAX的grad函数计算输入为复数数组的loss函数梯度时,出现如下TypeError:

---------------------------------------------------------------------------
TypeError                                 Traceback (most recent call last)
<ipython-input-79-b455310c3caa> in <module>
----> 1 grads = jax.grad(loss)(params, data.T, Pk, Pk, num_kraus)
      2 grads = jnp.conj(grads)
      3 updated_params = stiefel_update(params, grads, 0.00001)

    [... skipping hidden 6 frame]

~/miniconda3/envs/redes-neuronales/lib/python3.7/site-packages/jax/_src/api.py in _check_output_dtype_revderiv(name, holomorphic, x)
   1225                       f"but got {aval.dtype.name}."
   1226   elif dtypes.issubdtype(aval.dtype, np.complexfloating):
-> 1227     raise TypeError(f"{name} requires real-valued outputs (output dtype that is "
   1228                     f"a sub-dtype of np.floating), but got {aval.dtype.name}. "
   1229                     "For holomorphic differentiation, pass holomorphic=True. "

TypeError: grad requires real-valued outputs (output dtype that is a sub-dtype of np.floating), but got complex128. For holomorphic differentiation, pass holomorphic=True. For differentiation of non-holomorphic functions involving complex outputs, use jax.vjp directly.

相关代码如下:

Loss函数:

@partial(jit, static_argnums=4)
def loss(params, data=None, probes=None, measurements=None, num_kraus=None):
    """Loss function for the training assuming a predict function that can 
    generate probabilities for a measurement from the given process representation
    captured in params.

    Args:
        params (array): Parameters to optimize, e.g., Kraus operators.
        data (array): Data representing measured probabilities.
        probes (array): The probe operators.
        measurements (array): The measurement operators as Pauli vectors.
        num_kraus (int): The number of Kraus operators.

    Returns:
        loss (float): A scalar loss
    """
    k_ops = get_unblock(params, num_kraus)
    data_pred = predicta(k_ops, probes, measurements )

    l2 = jnp.sum(((data - data_pred)**2))
    return l2 + 0.001*jnp.linalg.norm(params, 1)

梯度计算代码:

grads = jax.grad(loss)(params, data.T, Pk, Pk, num_kraus)
grads = jnp.conj(grads)
updated_params = stiefel_update(params, grads, 0.00001)

需求:了解如何用grad处理复数数组,或正确使用vjp计算loss梯度的方法。


解决方法

方法一:使用grad并设置holomorphic=True

如果你的loss函数是**全纯(holomorphic)**的(满足柯西-黎曼条件,复数域上可微),直接给jax.grad传入holomorphic=True参数即可:

grads = jax.grad(loss, holomorphic=True)(params, data.T, Pk, Pk, num_kraus)
grads = jnp.conj(grads)
updated_params = stiefel_update(params, grads, 0.00001)

注意:该方法仅适用于全纯函数,若loss涉及取共轭、实部/虚部等非全纯操作,此方法不生效。

方法二:使用jax.vjp手动计算梯度

如果loss函数不是全纯的,用jax.vjp(向量雅可比乘积)手动计算梯度,步骤如下:

# 计算vjp:获取函数输出和vjp函数
loss_val, vjp_fn = jax.vjp(loss, params, data.T, Pk, Pk, num_kraus)
# 标量loss传入cotangent为1.0(需共轭梯度可传jnp.conj(1.0))
grads = vjp_fn(jnp.array(1.0, dtype=loss_val.dtype))[0]
# 后续更新逻辑不变
grads = jnp.conj(grads)
updated_params = stiefel_update(params, grads, 0.00001)

解释:jax.vjp返回两个值,第一个是函数输出loss_val,第二个是vjp_fn——该函数接受cotangent向量(标量loss对应标量值),返回输入参数的梯度。取返回值第一个元素即为params对应的梯度。


内容的提问来源于stack exchange,提问作者Nicolás Legnazzi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 03:05:21