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

使用HF Transformers时,如何在绑定嵌入中仅冻结部分嵌入索引?

解决BERT自定义嵌入层的权重绑定问题

问题核心

HF Transformers的默认权重绑定逻辑仅针对原生nn.Embedding做了适配,当你使用自定义嵌入层(用于部分索引冻结)时,自动绑定机制失效;直接替换解码器模块也无法正确处理权重转置与共享的逻辑。

解决方案

1. 手动实现权重绑定

自定义嵌入层需暴露可直接访问的权重张量,然后手动将解码器权重设置为嵌入层权重的转置,并确保两者共享梯度逻辑:

import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import BertForMaskedLM

class CustomEmbedding(nn.Module):
    def __init__(self, num_embeddings, embedding_dim, freeze_indices):
        super().__init__()
        self.weight = nn.Parameter(torch.randn(num_embeddings, embedding_dim))
        self.freeze_indices = freeze_indices
        # 初始化时冻结指定索引的权重
        self.weight.data[self.freeze_indices] = self.weight.data[self.freeze_indices].detach()
        self.weight.requires_grad = True
    
    def forward(self, input_ids):
        return F.embedding(input_ids, self.weight)

# 初始化模型并替换嵌入层
model = BertForMaskedLM.from_pretrained('bert-base-uncased')
freeze_indices = [0, 1, 2]  # 可根据需求修改冻结的token索引
model.bert.embeddings.word_embeddings = CustomEmbedding(
    model.config.vocab_size, model.config.hidden_size, freeze_indices
)

# 手动绑定解码器与嵌入层权重
model.cls.predictions.decoder.weight = model.bert.embeddings.word_embeddings.weight.t()
# 禁止解码器权重单独更新,完全依赖嵌入层的梯度控制
model.cls.predictions.decoder.weight.requires_grad = False

2. 完善梯度冻结逻辑

为确保冻结的索引在训练过程中始终不更新,可在自定义嵌入层中添加反向传播钩子,强制清零冻结索引的梯度:

class CustomEmbedding(nn.Module):
    def __init__(self, num_embeddings, embedding_dim, freeze_indices):
        super().__init__()
        self.weight = nn.Parameter(torch.randn(num_embeddings, embedding_dim))
        self.freeze_indices = freeze_indices
        self.weight.data[self.freeze_indices] = self.weight.data[self.freeze_indices].detach()
        self.weight.requires_grad = True
        # 注册梯度钩子
        self.weight.register_hook(self._freeze_grad_hook)
    
    def _freeze_grad_hook(self, grad):
        # 清零冻结索引的梯度
        grad[self.freeze_indices] = 0.0
        return grad
    
    def forward(self, input_ids):
        return F.embedding(input_ids, self.weight)

3. 适配HF的自动绑定逻辑(可选)

如果你希望沿用HF的tie_weights机制,可以修改模型的绑定逻辑,让它识别你的自定义嵌入层:

# 重写模型的tie_weights方法
def custom_tie_weights(self):
    # 保留原有绑定逻辑,新增自定义嵌入层的处理
    if hasattr(self, "cls") and hasattr(self.cls, "predictions"):
        embedding_layer = self.bert.embeddings.word_embeddings
        if isinstance(embedding_layer, CustomEmbedding):
            self.cls.predictions.decoder.weight = embedding_layer.weight.t()
        else:
            self.cls.predictions.decoder.weight = self.bert.embeddings.word_embeddings.weight
    # 执行其他原有绑定逻辑(如位置嵌入、token类型嵌入等)
    self.bert.embeddings.tie_weights()

# 替换模型的tie_weights方法
model.tie_weights = custom_tie_weights.__get__(model, BertForMaskedLM)
# 触发绑定
model.tie_weights()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 14:07:29