如何在PyTorch中为nn.Transformer编写可区分多输入的前向钩子
nn.Transformer 前向钩子多输入区分方案
PyTorch 前向钩子的第二个入参是所有传入forward方法的位置参数组成的元组,不是单个输入值,你可以直接通过索引或者关键字参数两种方式区分src和tgt:
方法1:按位置索引取参
常规调用transformer_model(src, tgt)时,钩子的入参元组长度为2,直接按顺序取即可:
x[0]对应第一个位置输入srcx[1]对应第二个位置输入tgt
如果调用时还传入了src_mask、tgt_mask等其他位置参数,会按传入顺序依次排在元组的后续索引位。
示例代码:
import torch import torch.nn as nn def transformer_hook(m, x, y): src = x[0] tgt = x[1] # 按需做后续处理,这里打印shape做示例 print(f"捕获输入:src shape = {src.shape}, tgt shape = {tgt.shape}") transformer = nn.Transformer(d_model=512, nhead=8) # 注册前向钩子 hook = transformer.register_forward_hook(transformer_hook) # 测试调用 src = torch.rand(10, 32, 512) tgt = torch.rand(20, 32, 512) out = transformer(src, tgt) # 用完移除钩子,避免内存泄漏 hook.remove()
方法2:按关键字参数取参(更稳定)
如果调用时用了关键字传参,或者不想受参数顺序影响,可以开启钩子的关键字参数捕获能力(PyTorch 1.10及以上版本支持):
注册钩子时加上with_kwargs=True,此时钩子函数定义改为hook_fn(m, args, kwargs, output),可以直接从kwargs字典里按键取对应输入,不用关心参数传入顺序:
def transformer_hook(m, args, kwargs, y): src = kwargs["src"] tgt = kwargs["tgt"] print(f"捕获输入:src shape = {src.shape}, tgt shape = {tgt.shape}") # 注册时开启with_kwargs hook = transformer.register_forward_hook(transformer_hook, with_kwargs=True)
这种方式即使调用时顺序调换(比如transformer(tgt=tgt, src=src)),也能正确取到对应参数,兼容性更好。
内容的提问来源于stack exchange,提问作者Maxwell Albert
相关产品推荐
相关产品推荐

