如何在JAX中用梯度下降训练含多输出损失函数的模型?
多输出模型的梯度下降训练问题
我尝试用梯度下降训练一个具有两个输出的模型,损失函数会返回两个误差,训练时出现报错,目前还没找到合适的解决办法。以下是复现问题的示例代码:
from jax import jit, random, grad import optax @jit def my_model(forz, params): a, b = params a_vect = a + forz**b b_vect = b + forz**a return a_vect, b_vect*50. @jit def rmse(predictions, targets): rmse = jnp.sqrt(jnp.mean((predictions - targets) ** 2)) return rmse @jit def my_loss(forz, params, true_a, true_b): sim_a, sim_b = my_model(forz, params) loss_a = rmse(sim_a, true_a) loss_b = rmse(sim_b, true_b) return loss_a, loss_b grad_myloss = jit(grad(my_loss, argnums=1)) # synthetic true data key = random.PRNGKey(758493) forz = random.uniform(key, shape=(1000,)) true_params = [8.9, 6.6] true_a, true_b = my_model(forz, true_params) # Train model_params = random.uniform(key, shape=(2,)) optimizer = optax.adabelief(1e-1) opt_state = optimizer.init(model_params) for i in range(1000): grads = grad_myloss(forz, model_params, true_a, true_b) # 此处报错 updates, opt_state = optimizer.update(grads, opt_state) model_params = optax.apply_updates(model_params, updates)
我了解到可以通过归一化将两个误差聚合为单个损失(因为输出向量单位不可比),示例代码如下:
@jit def normalized_rmse(predictions, targets): std_dev_targets = jnp.std(targets) rmse = jnp.sqrt(jnp.mean((predictions - targets) ** 2)) return rmse/std_dev_targets @jit def my_loss_single(forz, params, true_a, true_b): sim_a, sim_b = my_model(forz, params) loss_a = normalized_rmse(sim_a, true_a) loss_b = normalized_rmse(sim_b, true_b) return jnp.sqrt((loss_a ** 2) + (loss_b * 2))
除了这种聚合单个损失的方式,是否应该使用Jacobian矩阵(jacrev)来解决该问题?
内容的提问来源于stack exchange,提问作者Lacococha
相关产品推荐
相关产品推荐

