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

如何移除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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 07:51:02