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

给BartForConditionalGeneration添加额外嵌入层,还需修改哪些位置?

在BART编码器中添加额外嵌入层的必要修改位置

除了你提到的BartEncoder中嵌入合并的代码行外,要添加2个带位置嵌入的额外嵌入,还需要修改以下几个核心位置:

1. 模型配置类(BartConfig)

  • 在BartConfig类中新增对应额外嵌入的配置参数,比如控制额外嵌入数量的num_additional_embeddings(设为2)、额外嵌入维度的additional_embedding_dim(若与主嵌入维度不同需指定)、额外嵌入最大位置长度的additional_max_position_embeddings。
  • 示例代码(在__init__方法中添加):
    self.num_additional_embeddings = num_additional_embeddings if num_additional_embeddings is not None else 0
    self.additional_embedding_dim = additional_embedding_dim if additional_embedding_dim is not None else self.hidden_size
    self.additional_max_position_embeddings = additional_max_position_embeddings if additional_max_position_embeddings is not None else self.max_position_embeddings
    

2. BartEmbeddings类

  • 初始化对应数量的额外词嵌入层和位置嵌入层:
    self.additional_embeddings = nn.ModuleList([
        nn.Embedding(self.config.vocab_size, self.config.additional_embedding_dim)
        for _ in range(self.config.num_additional_embeddings)
    ])
    self.additional_position_embeddings = nn.ModuleList([
        nn.Embedding(self.config.additional_max_position_embeddings, self.config.additional_embedding_dim)
        for _ in range(self.config.num_additional_embeddings)
    ])
    
  • 若额外嵌入维度与主隐藏层维度不一致,需新增线性投影层做维度对齐:
    self.additional_embedding_projections = nn.ModuleList([
        nn.Linear(self.config.additional_embedding_dim, self.config.hidden_size)
        for _ in range(self.config.num_additional_embeddings)
    ])
    
  • 可在forward方法中封装额外嵌入的计算逻辑(词嵌入+位置嵌入+可选投影),方便后续调用。

3. BartEncoder的forward方法

  • 生成与输入序列长度匹配的位置ID,用于额外位置嵌入计算:
    seq_length = input_ids.size(1)
    additional_position_ids = torch.arange(seq_length, dtype=torch.long, device=input_ids.device)
    additional_position_ids = additional_position_ids.unsqueeze(0).expand_as(input_ids)
    
  • 遍历每个额外嵌入,完成计算并与主嵌入相加:
    for idx in range(self.config.num_additional_embeddings):
        add_emb = self.embeddings.additional_embeddings[idx](input_ids)
        add_pos_emb = self.embeddings.additional_position_embeddings[idx](additional_position_ids)
        if self.config.additional_embedding_dim != self.config.hidden_size:
            add_emb = self.embeddings.additional_embedding_projections[idx](add_emb + add_pos_emb)
        else:
            add_emb = add_emb + add_pos_emb
        embeddings = embeddings + add_emb
    

4. 权重初始化逻辑

  • 在BartPreTrainedModel的_init_weights方法中,添加对额外嵌入层、位置嵌入层及投影层的初始化,对齐BART原有初始化规则:
    # 对额外嵌入层初始化
    for emb_layer in self.embeddings.additional_embeddings:
        self._init_weights(emb_layer)
    # 对额外位置嵌入层初始化
    for pos_emb_layer in self.embeddings.additional_position_embeddings:
        self._init_weights(pos_emb_layer)
    # 对投影层初始化(若有)
    for proj_layer in self.embeddings.additional_embedding_projections:
        self._init_weights(proj_layer)
    

5. 预训练权重加载逻辑(可选但重要)

  • 若需加载原有BART预训练权重,需手动处理权重字典,移除额外嵌入相关的键(预训练权重中无对应参数),避免加载报错:
    state_dict = torch.load(pretrained_model_path)
    # 删除额外嵌入相关的权重键
    for key in list(state_dict.keys()):
        if any(kw in key for kw in ["additional_embeddings", "additional_position_embeddings", "additional_embedding_projections"]):
            del state_dict[key]
    

6. 输入扩展(若额外嵌入需独立输入)

  • 如果额外嵌入需要基于独立的token序列生成,需修改BartEncoder的forward方法参数,接收additional_input_ids(列表或张量),并对应计算每个额外嵌入的词嵌入。

内容的提问来源于stack exchange,提问作者Animesh Kumar Paul

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 22:13:20