PyTorch Lightning两种training_step返回值下loss.backward()的差异咨询
关于PyTorch Lightning自动优化中loss.backward()的差异解答
你不用担心场景2里的metric会被反向传播,两种场景下loss.backward()的核心逻辑是一致的,具体说明如下:
- 场景1:直接返回loss张量时,自动优化机制会直接对该张量执行
loss.backward(),完成梯度计算。 - 场景2:返回包含"loss"键的字典时,PyTorch Lightning会自动提取字典中键为
loss的张量,仅对这个张量执行loss.backward(),完全不会处理字典里的metric张量。
额外补充两点:
- 大多数评估指标(比如准确率、F1值)本身的计算过程不会产生需要反向传播的梯度,这类张量的
requires_grad属性默认是False,天然不会参与反向传播。 - 即便你自定义的metric张量意外开启了梯度,PyTorch Lightning的自动优化逻辑也只会聚焦在字典中标记为
loss的张量上,不会对其他键对应的张量执行反向传播操作。
内容的提问来源于stack exchange,提问作者Prithviraj Kanaujia
相关产品推荐
相关产品推荐

