【问题标题】:How to return a dictionary of values from loss function in Jax?如何从 Jax 中的损失函数返回值字典?
【发布时间】: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_valuegrads 包含什么?如所写,它会导致错误,因为grad 为返回标量的函数实现反向模式自动差异,而您的函数不返回标量。
  • @jakevdp 是的,这是我的问题,如何避免错误。在 tf 中,我认为您可以通过不调用尝试为您做所有事情的助手来解决此问题。您调用函数(获取所有值)但计算梯度 w.r.t。你想要的一件事。在获得梯度后重新调用函数的唯一方法是这样做吗?有什么方法可以在训练时获取辅助信息吗?必须有一个 jax 模式。
  • @jakevdp 例如,我需要计算 loss_1 和 loss_2 以及它们每次的总和来计算损失。我需要总和的梯度来优化。每 100 个(或其他)步骤,我可能想保存 loss_1、loss_2 等的值以进行调试。也许再次调用损失函数并不是最糟糕的,但它似乎很慢。

标签: python logging jax


【解决方案1】:

感谢@jakevdp 本人促使我考虑一些替代的谷歌查询,结果证明,截至https://github.com/google/jax/pull/484,grad 函数有一个 aux 选项。我认为这对于迁移到 jax 的 tensorflow 2 用户来说并不是很明显,因为您明确使用 GradientTape 的方式。

类似于以下示例的内容显示了返回的辅助信息。它甚至似乎可以处理一个 dict,这对于在更新循环中定期记录很有用。

import jax
import jax.numpy as jnp

key = jax.random.PRNGKey(0)
theta = jax.random.normal(key, (10, 1))
y = np.random.randn(10, 1)
alpha = 0.01

def loss(theta, y):
    loss_reg = jnp.sum(theta ** 2)
    loss_data = jnp.sum((y - theta) ** 2)
    loss = loss_data + alpha * loss_reg
    return loss, dict(loss_reg=loss_reg, loss_data=loss_data)

grad, aux = jax.grad(loss, has_aux=True)(theta, y)

display(grad)
display(aux)

try:
    jax.grad(loss)(theta, y)
except TypeError as e:
    print(f'yes got error {e}')

输出:

DeviceArray([[-1.4899637 ],
             [-0.71481365],
             [-0.6030376 ],
             [-0.8263864 ],
             [-1.8103108 ],
             [ 0.69435316],
             [-1.5611547 ],
             [-1.6380725 ],
             [ 0.9838154 ],
             [ 0.21186407]], dtype=float32)
{'loss_data': DeviceArray(3.3714797, dtype=float32),
 'loss_reg': DeviceArray(2.658556, dtype=float32)}
yes got error Gradient only defined for scalar-output functions. Output was (DeviceArray(3.3980653, dtype=float32), {'loss_data': DeviceArray(3.3714797, dtype=float32), 'loss_reg': DeviceArray(2.658556, dtype=float32)}).

【讨论】:

  • 这个答案对我很有用。谢谢。
猜你喜欢
  • 1970-01-01
  • 1970-01-01
  • 1970-01-01
  • 2020-08-06
  • 1970-01-01
  • 2018-07-04
  • 2018-05-15
  • 2019-11-22
  • 2011-05-01
相关资源
最近更新 更多