使用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
相关产品推荐
相关产品推荐

