pytorch hook API
发布时间
阅读量:
阅读量
PyTorch中Hook机制的应用
为简化操作,我通常采用单输入单输出的结构进行测试。
PyTorch提供了三种hook机制,其中一种用于观察tensor的梯度信息,另一种用于查看某一层的输入与输出数据,第三种则用于获取某一层输入与输出的梯度信息。
所有hook函数的使用都需要预先定义一个对应的函数作为处理逻辑。
tensor.register_hook(function)是最直观的一种方式,它能够获取该tensor的梯度信息。所定义的函数格式应为hook(grad),其中grad即为该tensor对应的梯度值。
Conv2d.forward_hook(function)(当然其他层也适用,此处以Conv2d为例),其对应的函数格式应为forward_hook(module, input, output)。其中module代表当前使用的Conv2d层,input和output分别表示前向传播过程中该层接收到的输入和输出数据。
Conv2d.backard_hook(function),其对应的函数格式应为backward_hook(module, input_grad, output_grad)。这里的input和output是按照反向传播顺序来定义的(与前向传播方向相反)。在input_grad中包含三个梯度值,依次对应输入张量、层权重以及层偏置的梯度;而
全部评论 (0)
还没有任何评论哟~
