【论文解读】GPT Understands, Too

一.论文

1.1 P-tuning

区别于之前的工作,这篇工作认为promote可以在句子中的任意位置起到作用,可以将它们插入上下文或目标中

上图中,左图是不使用任何操作,右图是选择在居首和目标前插入promote的embedding,插入promote的过程可以表示为

其中x代表一系列离散的输入令牌,y代表目标(可以理解为希望模型想要给你的回答),e()表示对应的embedding,其实就是将其参数化映射成为伪tokens,即

通过最小化这些参数

1.2 promote生成

嵌入的promote实际上可以理解为不一定离散不相互关联的,而实际上的promote其实应该是高度离散的且具有关联性的,因此作者选择使用双向长短期记忆网络(LSTM),激活函数和MLP来建模这种关系

在推理中,我们只需要输出嵌入h,并且可以丢弃LSTM头

二.代码

本质上是使用一个PromptEncoder来生成伪的embedding添加到原先的embedding中

2.1 训练

训练过程只更新promote_encoder中的参数

2.1.1 PromptEncoder

PTuneForLAMA中实例化了PromptEncoder

PromptEncoder本质上是一个(嵌入 + LSTM + MLP)

python 复制代码
import torch
import torch.nn as nn


class PromptEncoder(torch.nn.Module):
    def __init__(self, template, hidden_size, tokenizer, device, args):
        super().__init__()
        self.device = device
        self.spell_length = sum(template)
        self.hidden_size = hidden_size
        self.tokenizer = tokenizer
        self.args = args
        # ent embedding
        self.cloze_length = template
        self.cloze_mask = [
            [1] * self.cloze_length[0]  # first cloze
            + [1] * self.cloze_length[1]  # second cloze
            + [1] * self.cloze_length[2]  # third cloze
        ]
        self.cloze_mask = torch.LongTensor(self.cloze_mask).bool().to(self.device)

        self.seq_indices = torch.LongTensor(list(range(len(self.cloze_mask[0])))).to(self.device)
        # embedding
        self.embedding = torch.nn.Embedding(len(self.cloze_mask[0]), self.hidden_size).to(self.device)
        # LSTM
        self.lstm_head = torch.nn.LSTM(input_size=self.hidden_size,
                                       hidden_size=self.hidden_size // 2,
                                       num_layers=2,
                                       dropout=self.args.lstm_dropout,
                                       bidirectional=True,
                                       batch_first=True)
        self.mlp_head = nn.Sequential(nn.Linear(self.hidden_size, self.hidden_size),
                                      nn.ReLU(),
                                      nn.Linear(self.hidden_size, self.hidden_size))
        print("init prompt encoder...")

    def forward(self):
        input_embeds = self.embedding(self.seq_indices).unsqueeze(0)
        output_embeds = self.mlp_head(self.lstm_head(input_embeds)[0]).squeeze()
        return output_embeds

2.1.2 调用

在PTuneForLAMA的forward函数中调用了embed_input来实现

相关推荐
Highcharts.js10 分钟前
五大痛点拖慢数据分析平台决策效率:Highchart可视化与AI分析解决方案解析
人工智能·信息可视化·数据分析·数据可视化·highcharts·ai可视化分析
义嘉泰17 分钟前
手腕上的小手电:智能手表怎么把“绿光”变成心率
人工智能·智能手表
Canace24 分钟前
AI 生成到 90% 突然断了:你的解决方案是?
前端·人工智能
湘美书院--湘美谈教育25 分钟前
AI时代的奥德赛:算法星空,寻找精神归航
大数据·人工智能·深度学习·机器学习·生活
瓦学妹40 分钟前
2026亚马逊本土店怎么注册?入驻流程、资料准备与常见问题
大数据·运维·人工智能
搞科研的小刘选手42 分钟前
【昌吉学院主办】第三届大数据、神经网络与深度学习研讨会(BDNNDL 2026)
大数据·深度学习·神经网络·学术会议·会议推荐
YucongCai44 分钟前
Opengovernment(智慧城市) v0.0.3 攻防一体的智能体主动网络安全平台:ActiveCybersecurity v0.2
人工智能·web安全·智慧城市
腾讯云大数据1 小时前
从多模态数据处理到模型训练:腾讯云EMR-Ray打通Data+AI全流程
人工智能·云计算·腾讯云·mapreduce·腾讯云大数据
xiakq1 小时前
5 分钟接入 GPT-5.6 完整指南
gpt·openai·claude·gemini·anthropic
louyu6668881 小时前
国产楼宇自控系统哪家靠谱?
大数据·人工智能