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

如何在PyTorch Transformer中获取编码器中间层的输出

方案1:修改模型前向逻辑(推荐长期需要提取特征的场景)

PyTorch的TransformerEncoder内部的网络层以ModuleList的形式存放在self.encoder.layers属性中,你可以自定义层的执行逻辑,在执行到倒数第二层时直接返回结果,无需执行最后一层:

class LitModel(pl.LightningModule):
    # 其他原有代码保持不变
    def forward(self, indices, mask, return_penultimate=False):
        x = self.embed(indices)
        # 不需要提取倒数第二层时走原有逻辑
        if not return_penultimate:
            x = self.encoder(x, src_key_padding_mask=mask)
            return x
        # 逐层执行,到倒数第二层直接返回
        total_layers = len(self.encoder.layers)
        for idx, layer in enumerate(self.encoder.layers):
            x = layer(x, src_key_padding_mask=mask)
            if idx == total_layers - 2:
                return x

调用时只需传入return_penultimate=True即可直接获取倒数第二层的输出:

# 推理示例
model.eval()
with torch.no_grad():
    penultimate_out = model(your_indices, your_mask, return_penultimate=True)

方案2:使用前向钩子(无需修改模型定义,适合临时提取需求)

如果你不想改动已经调试完成的模型代码,可以通过PyTorch的前向钩子机制,在推理时自动捕捉指定层的输出:

# 加载训练完成的模型
model = LitModel.load_from_checkpoint("your_checkpoint_path.ckpt")
model.eval()

penultimate_output = None
# 定义钩子函数,用于存储指定层的输出
def capture_output(module, input, output):
    global penultimate_output
    penultimate_output = output

# 给倒数第二层注册钩子
target_layer = model.encoder.layers[-2]
hook_handle = target_layer.register_forward_hook(capture_output)

# 正常执行推理
with torch.no_grad():
    model(your_indices, your_mask)
    # 此时penultimate_output已经被赋值为倒数第二层的输出
    print(penultimate_output.shape)

# 钩子使用完成后必须移除,避免内存泄漏
hook_handle.remove()

注意事项

  • 若你的模型共6个TransformerEncoder层,layers[-2]对应的就是倒数第二层,无需手动调整索引
  • 如果需要同时获取多层输出,只需要在方案1的遍历过程中把每层结果存入列表返回即可,或者给多个层分别注册钩子

内容的提问来源于stack exchange,提问作者Evan Zamir

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.07 13:18:01