【发布时间】:2021-06-27 18:11:29
【问题描述】:
假设您有一个损失函数,并且您想在训练时跟踪损失的各个子组件。最“jax”的方式是什么?
def loss(params, x, y):
...
loss_1 = ...
loss_2 = ...
loss = loss_1 + 0.1 * loss_2
return dict(loss=loss, loss_1=loss_1, loss_2=loss_2)
@jax.jit
def update(params, tau, y):
f_value, grads = jax.value_and_grad(loss)(params, tau, y)
# something like this
您是否想只在函数上使用grad 来提取损失,然后再次重新计算值?有没有办法用value_and_grad 更有效地做到这一点?
【问题讨论】:
-
我不完全清楚你的问题是什么。你能举一个你想计算的例子吗?
-
@jakevdp 我只是想尽量减少标准 jax 循环中的损失。我通常在 tensorflow 2.0 中使用的一种模式是从丢失中返回一堆调试信息......我认为这样做的“jax”方式可能是一个数组而不是一个字典。当我在调试迭代中重新评估函数时,它会减慢很多
-
更具体地说...在您的代码 sn-p 中,您希望
f_value和grads包含什么?如所写,它会导致错误,因为grad为返回标量的函数实现反向模式自动差异,而您的函数不返回标量。 -
@jakevdp 是的,这是我的问题,如何避免错误。在 tf 中,我认为您可以通过不调用尝试为您做所有事情的助手来解决此问题。您调用函数(获取所有值)但计算梯度 w.r.t。你想要的一件事。在获得梯度后重新调用函数的唯一方法是这样做吗?有什么方法可以在训练时获取辅助信息吗?必须有一个 jax 模式。
-
@jakevdp 例如,我需要计算 loss_1 和 loss_2 以及它们每次的总和来计算损失。我需要总和的梯度来优化。每 100 个(或其他)步骤,我可能想保存 loss_1、loss_2 等的值以进行调试。也许再次调用损失函数并不是最糟糕的,但它似乎很慢。