如何移除HuggingFace Transformers中GPT2预训练模型的wte嵌入层?
解决GPT2移除词嵌入层(wte)的问题
原方法无效的原因
你贴的代码是针对BERT模型的编码层删除逻辑,和GPT2的模型结构完全不匹配,而且你的目标是移除wte词嵌入层(不是Transformer的Block层),所以这个方法根本不适用于你的需求。
正确解决方案
wte是GPT2模型中负责将输入ID转为词嵌入的层,属于model.transformer下的直接属性(而非ModuleList中的元素)。要彻底移除它,需要自定义模型类并重写前向传播逻辑,避免原模型依赖该层的计算。
代码实现
import torch import torch.nn as nn from transformers import GPT2LMHeadModel, AutoModelForCausalLM # 自定义移除wte层的GPT2模型 class GPT2WithoutWTE(GPT2LMHeadModel): def __init__(self, config): super().__init__(config) # 删除wte层 del self.transformer.wte def forward(self, inputs_embeds=None, **kwargs): # 移除wte后,无法再通过input_ids生成嵌入,必须直接传入inputs_embeds if inputs_embeds is None: raise ValueError("必须传入inputs_embeds参数,因为wte词嵌入层已被移除") # 调用父类forward,直接使用传入的嵌入向量 return super().forward(inputs_embeds=inputs_embeds, **kwargs) # 加载预训练GPT2模型 original_model = AutoModelForCausalLM.from_pretrained("openai-community/gpt2") # 初始化自定义模型并加载原模型权重(跳过wte相关参数) custom_model = GPT2WithoutWTE(original_model.config) for param_name, param in original_model.named_parameters(): if not param_name.startswith("transformer.wte"): custom_model.state_dict()[param_name].copy_(param) # 查看模型结构,确认wte已被移除 print(custom_model)
注意事项
- 移除
wte后,模型无法再处理input_ids输入,必须直接传入预计算好的嵌入向量inputs_embeds。 - 如果你只是想让
wte层不参与训练(而非彻底移除),可以将其参数设置为不可训练:model.transformer.wte.requires_grad_(False)
内容的提问来源于stack exchange,提问作者dark kk
相关产品推荐
相关产品推荐

