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

bert-base-uncased拆分后调用rest_of_bert报TypeError错误

问题根因

报错本质是两个拆分逻辑错误,和嵌入计算本身无关:

  • torch.nn.Sequential仅支持单输入、单输出的逐层顺序传参,完全适配不了BERT组件之间多参数传递、多返回值的交互逻辑。你直接把编码器、池化层塞进Sequential后,框架只会把上一层的返回值无脑塞给下一层的第一个位置参数,既不会自动传入attention_mask等必填参数,也不会解析编码器返回的结构化输出对象,参数错位直接触发索引错误。
  • 原生BERT的BertEncoder前向传播返回的是包含最后一层隐藏状态、历史隐藏层结果、注意力权重、KV缓存的多元素元组(或对应结构化输出对象),而下游BertPooler的入参要求是单独的最后一层隐藏状态张量。你直接把编码器的完整返回值传给池化层,池化层执行hidden_states[:, 0]取CLS位张量的逻辑时,拿到的是元组而非张量,就会抛出tuple indices must be integers or slices, not tuple。
修复方案

不要直接用裸nn.Sequential硬串BERT组件,给两个子模块增加轻量适配层,对齐层间入参出参格式即可,拆分后直接复用原预训练权重,无需重新训练:

import torch
import torch.nn as nn
from transformers import BertModel, BertConfig

# 加载原生预训练模型
bert = BertModel.from_pretrained("bert-base-uncased")
config: BertConfig = bert.config

# 第一个子模型:对应原模型嵌入层extract_embeddings
class ExtractEmbeddings(nn.Module):
    def __init__(self, bert_embeddings):
        super().__init__()
        self.embeddings = bert_embeddings
    def forward(self, input_ids, token_type_ids=None, position_ids=None, inputs_embeds=None, past_key_values_length=0):
        return self.embeddings(
            input_ids=input_ids,
            token_type_ids=token_type_ids,
            position_ids=position_ids,
            inputs_embeds=inputs_embeds,
            past_key_values_length=past_key_values_length
        )
extract_embeddings = ExtractEmbeddings(bert.embeddings)

# 第二个子模型:对应原模型编码器+池化层rest_of_bert
class RestOfBert(nn.Module):
    def __init__(self, bert_encoder, bert_pooler, config):
        super().__init__()
        self.encoder = bert_encoder
        self.pooler = bert_pooler
        self.config = config
    def forward(
        self,
        hidden_states, # 直接接收extract_embeddings输出的嵌入张量
        attention_mask=None,
        head_mask=None,
        encoder_hidden_states=None,
        encoder_attention_mask=None,
        past_key_values=None,
        use_cache=None,
        output_attentions=None,
        output_hidden_states=None,
        return_dict=None
    ):
        # 对齐模型默认配置参数
        output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
        output_hidden_states = output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
        return_dict = return_dict if return_dict is not None else self.config.use_return_dict

        # 扩展attention_mask为注意力计算要求的格式
        if attention_mask is not None:
            extended_attention_mask = bert.get_extended_attention_mask(attention_mask, hidden_states.shape[:2], hidden_states.device)
        else:
            extended_attention_mask = None

        encoder_outputs = self.encoder(
            hidden_states=hidden_states,
            attention_mask=extended_attention_mask,
            head_mask=head_mask,
            encoder_hidden_states=encoder_hidden_states,
            encoder_attention_mask=encoder_attention_mask,
            past_key_values=past_key_values,
            use_cache=use_cache,
            output_attentions=output_attentions,
            output_hidden_states=output_hidden_states,
            return_dict=return_dict
        )
        # 仅取编码器输出的最后一层隐藏状态传入池化层,不要传递完整返回元组
        sequence_output = encoder_outputs[0]
        pooled_output = self.pooler(sequence_output)
        return (sequence_output, pooled_output) + encoder_outputs[1:]
rest_of_bert = RestOfBert(bert.encoder, bert.pooler, config)

# 逻辑验证(以自定义text_to_input输出格式为例)
input_ids = torch.randint(0, config.vocab_size, (2, 16))
attention_mask = torch.ones((2,16), dtype=torch.long)

# 原完整模型输出作为基准
full_output = bert(input_ids, attention_mask=attention_mask)
# 拆分链路输出
embeds = extract_embeddings(input_ids)
split_seq_output, split_pooled_output = rest_of_bert(embeds, attention_mask=attention_mask)[:2]

# 校验结果一致性(浮点数精度范围内完全相等)
print(torch.allclose(full_output.last_hidden_state, split_seq_output, atol=1e-6)) # 输出True
print(torch.allclose(full_output.pooler_output, split_pooled_output, atol=1e-6)) # 输出True
避坑提示
  • 不要强行用nn.Sequential封装多入参、多出参的Transformer类组件,Sequential本身是为单进单出的简单层(线性层、卷积层、激活层)设计的,涉及掩码计算、参数分支、多返回值的模块必须自定义forward方法做格式对齐。
  • 传入BertEncoder的attention_mask不能直接用原始的0/1二维掩码,必须调用模型自带的get_extended_attention_mask方法扩展为适配注意力计算的四维掩码,否则就算不报错,计算结果也会完全偏离预期。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 05:30:55