使用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
相关产品推荐
相关产品推荐

