You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PyTorch中为nn.Transformer编写可区分多输入的前向钩子

nn.Transformer 前向钩子多输入区分方案

PyTorch 前向钩子的第二个入参是所有传入forward方法的位置参数组成的元组,不是单个输入值,你可以直接通过索引或者关键字参数两种方式区分src和tgt:

方法1:按位置索引取参

常规调用transformer_model(src, tgt)时,钩子的入参元组长度为2,直接按顺序取即可:

  • x[0] 对应第一个位置输入 src
  • x[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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.10.04 19:09:01