Jupyter Notebook版本(带详细注释)

保留原标题结构,为每一行代码添加详细注释,确保零基础小白能理解:

# 12.6-GPT模型代码实现

## 1. 数据准备

### 1.1 代码包引入
import torch  # PyTorch核心库,用于张量计算和神经网络构建
import torch.nn as nn  # 神经网络模块,包含各种层和损失函数
import torch.utils.data as Data  # 数据加载工具,用于构建数据集和数据加载器
from torch import optim  # 优化器模块,如Adam、SGD等
import numpy as np  # 数值计算库,用于数组操作
from tqdm import *  # 进度条库,可视化循环进度
import matplotlib.pyplot as plt  # 绘图库,用于可视化损失曲线等
import re  # 正则表达式库,用于文本处理
import string  # 字符串处理库,包含字母、数字、标点等常量
from collections import Counter  # 计数工具,用于统计词频构建词表

### 1.2 数据整理
# 读取对话数据集,编码为utf-8避免中文乱码
# 注意:需确保data文件夹下有dataset.txt文件,文件存储对话文本,每行是多轮对话(\t分隔)
with open('data/dataset.txt','r',encoding='utf-8') as f:
    datas = f.readlines()  # 按行读取所有数据,返回列表
# 查看前10行数据,了解数据格式(每行是多轮对话,用\t分隔不同轮次)
datas[:10]

# 查看数据中的特殊字符,为后续文本清洗做准备
content = ''.join(datas)  # 将所有文本拼接成一个字符串
# 正则替换:把所有中文(\u4e00-\u9fa5是中文Unicode范围)替换为空格,只保留非中文字符
special_char = re.sub(r'[\u4e00-\u9fa5]', ' ', content)  
# 输出特殊字符集合:排除字母、数字后,剩下的就是需要处理的标点/控制字符
print(set(special_char) - set(string.ascii_letters) - set(string.digits))

# 词元化函数:将文本转换为最小语义单位(这里按单个字符拆分,适合中文)
def tokenize(datas):
    # 存储所有文本的词元结果
    tokens = []
    for data in datas:
        # 去除每行首尾的空白字符(包括换行符),并删除换行符
        data=data.strip().replace("\n","")
        # 遍历每个字符:将\t(对话分隔符)替换为<sep>特殊标记,其他字符保留;最后追加<sep>作为行结束标记
        token = [i if i!='\t' else "<sep>" for i in data]+['<sep>']
        tokens.append(token)  # 将当前行的词元加入列表
    return tokens

# 执行词元化
tokens = tokenize(datas)
# 打印前6行的词元结果,验证处理效果
print("tokens:", tokens[:6])

### 1.3 构建词表
# 定义展平函数:将二维的tokens列表(每行一个词元列表)展平为一维列表,方便统计词频
flatten = lambda l: [item for sublist in l for item in sublist]  

# 定义词表类:实现词元<->索引的映射,处理特殊标记
class Vocab:
    def __init__(self, tokens):
        self.tokens = tokens  # 传入的二维词元列表
        # 初始化特殊标记的索引:<pad>填充标记(0)、<unk>未知标记(1)、<seq>(原代码笔误,实际是<sep>)(2)
        self.token2index = {'<pad>': 0, '<unk>': 1, '<seq>': 2}  
        # 统计所有词元的词频,并按词频降序排序
        # Counter(flatten(self.tokens)):展平后统计每个词元出现的次数
        # sorted(..., key=lambda x: x[1], reverse=True):按词频(元组第二个元素)降序排序
        # 为每个词元分配索引(从3开始,因为前3个是特殊标记)
        self.token2index.update({
            token: index + 3
            for index, (token, freq) in enumerate(
                sorted(Counter(flatten(self.tokens)).items(), key=lambda x: x[1], reverse=True))
        })
        # 构建反向映射:索引->词元
        self.index2token = {index: token for token, index in self.token2index.items()}

    def __getitem__(self, query):
        # 实现索引功能:支持传入词元(返回索引)、索引(返回词元)、列表/元组(批量处理)
        # 单一索引:如果是字符串(词元)
        if isinstance(query, str):
            # 存在则返回索引,否则返回<pad>的索引(0)
            return self.token2index.get(query, 0)
        # 单一索引:如果是整数(索引)
        elif isinstance(query, int):
            # 存在则返回词元,否则返回<unk>
            return self.index2token.get(query, '<unk>')
        # 批量索引:如果是列表/元组
        elif isinstance(query, (list, tuple)):
            # 递归处理每个元素
            return [self.__getitem__(item) for item in query]

    def __len__(self):
        # 返回词表大小(索引总数)
        return len(self.index2token)

# 实例化词表
vocab = Vocab(tokens)
# 获取词表大小,后续模型嵌入层会用到
vocab_size = len(vocab)

### 1.4 构造数据集
# 自定义数据集类:继承PyTorch的Dataset,实现数据加载和批处理padding
class MyDataSet(Data.Dataset):
    def __init__(self,datas):
        self.datas = datas  # 传入的是词元转换后的索引列表

    def __getitem__(self, item):
        # 获取单个样本:将序列拆分为解码器输入和输出(自回归任务,输入比输出早一个位置)
        data = self.datas[item]  # 获取第item个样本的索引序列
        decoder_input = data[:-1]  # 解码器输入:去掉最后一个元素(因为输出是下一个token)
        decoder_output = data[1:]  # 解码器输出:去掉第一个元素(输入的下一个token)

        # 记录输入和输出的长度,用于后续padding
        decoder_input_len = len(decoder_input)
        decoder_output_len = len(decoder_output)

        # 返回字典格式的样本
        return {"decoder_input":decoder_input,"decoder_input_len":decoder_input_len,
                "decoder_output":decoder_output,"decoder_output_len":decoder_output_len}

    def __len__(self):
        # 返回数据集总长度
        return len(self.datas)

    def padding_batch(self,batch):
        # 自定义批处理函数:对一个批次的样本进行padding,使长度统一
        # 1. 获取当前批次所有样本的输入/输出长度
        decoder_input_lens = [d["decoder_input_len"] for d in batch]
        decoder_output_lens = [d["decoder_output_len"] for d in batch]

        # 2. 找到当前批次的最大长度(作为padding后的统一长度)
        decoder_input_maxlen = max(decoder_input_lens)
        decoder_output_maxlen = max(decoder_output_lens)

        # 3. 对每个样本进行padding:不足最大长度的部分填充<pad>的索引(0)
        for d in batch:
            # 计算需要填充的长度,补充<pad>索引
            d["decoder_input"].extend([vocab["<pad>"]]*(decoder_input_maxlen-d["decoder_input_len"]))
            d["decoder_output"].extend([vocab["<pad>"]]*(decoder_output_maxlen-d["decoder_output_len"]))
        
        # 4. 转换为PyTorch张量(long类型,因为是索引)
        decoder_inputs = torch.tensor([d["decoder_input"] for d in batch], dtype=torch.long)
        decoder_outputs = torch.tensor([d["decoder_output"] for d in batch], dtype=torch.long)
        return decoder_inputs,decoder_outputs

# 定义批次大小
batch_size = 64
# 将词元列表转换为索引列表:每个词元替换为对应的索引
tokens_num = [[vocab[word] for word in line] for line in tokens] 
# 实例化自定义数据集
dataset = MyDataSet(tokens_num)
# 构建数据加载器:collate_fn指定自定义的批处理函数(padding)
data_loader = Data.DataLoader(dataset, batch_size=batch_size, collate_fn=dataset.padding_batch)

## 2. 建立模型

### 2.1 掩码操作
# 定义设备:全局变量,后续模型/张量会用到(需在模型定义前初始化,原代码此处顺序调整,否则get_attn_subsequence_mask会报错)
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")

# 掩码1:Pad掩码 - 屏蔽padding部分(索引0),避免模型关注无意义的填充字符
def get_attn_pad_mask(seq_q, seq_k):                      
    # seq_q: 查询序列 [batch_size, seq_len_q]
    # seq_k: 键序列 [batch_size, seq_len_k]
    batch_size, len_q = seq_q.size()  # 获取批次大小和查询序列长度
    batch_size, len_k = seq_k.size()  # 获取批次大小和键序列长度
    # 1. 找到seq_k中等于0(<pad>)的位置,标记为True
    # eq(0):等于0返回True,否则False;unsqueeze(1):增加维度,变为[batch_size, 1, len_k]
    pad_attn_mask = seq_k.data.eq(0).unsqueeze(1)          
    # 2. 扩展维度,匹配注意力分数的维度 [batch_size, len_q, len_k]
    return pad_attn_mask.expand(batch_size, len_q, len_k)

# 掩码2:未来信息掩码 - 屏蔽未来的token,符合自回归特性(只能看到当前及之前的token)
def get_attn_subsequence_mask(seq):                              
    # seq: [batch_size, tgt_len] 目标序列
    # 1. 构建注意力掩码的形状 [batch_size, tgt_len, tgt_len]
    attn_shape = [seq.size(0), seq.size(1), seq.size(1)]
    # 2. 生成上三角矩阵:k=1表示对角线以上的元素为1(未来token),对角线及以下为0(当前/过去token)
    subsequence_mask = np.triu(np.ones(attn_shape), k=1)          
    # 3. 转换为PyTorch张量(byte类型,用于masked_fill_),并移到指定设备
    subsequence_mask = torch.from_numpy(subsequence_mask).byte()  
    subsequence_mask = subsequence_mask.to(device)
    return subsequence_mask

### 2.2 注意力计算函数
# 定义模型超参数(需在注意力类定义前初始化,否则会报错)
d_model = 768  # Embedding维度
d_ff = 2048  # 前馈层维度
d_k = d_v = 64  # Q/K/V的维度(多头注意力中每个头的维度)
n_layers = 6  # 解码器层数
n_heads = 8  # 多头注意力的头数

# 缩放点积注意力 - Transformer核心组件
class ScaledDotProductAttention(nn.Module):
    def __init__(self):
        super(ScaledDotProductAttention, self).__init__()
    def forward(self, Q, K, V, attn_mask):
        '''
        前向传播:计算缩放点积注意力
        参数说明:
        Q: 查询矩阵 [batch_size, n_heads, len_q, d_k]
        K: 键矩阵 [batch_size, n_heads, len_k, d_k]
        V: 值矩阵 [batch_size, n_heads, len_v(=len_k), d_v]
        attn_mask: 注意力掩码 [batch_size, n_heads, seq_len, seq_len]
        '''
        # 1. 计算Q和K的点积(注意力分数),并除以根号d_k(缩放,避免分数过大)
        scores = torch.matmul(Q, K.transpose(-1, -2)) / np.sqrt(d_k) 
        # 2. 应用掩码:掩码位置填充为-1e9(Softmax后接近0,相当于屏蔽)
        scores.masked_fill_(attn_mask, -1e9) 
        # 3. Softmax归一化,得到注意力权重
        attn = nn.Softmax(dim=-1)(scores)
        # 4. 注意力权重乘以V,得到上下文向量(加权求和)
        context = torch.matmul(attn, V) 
        return context, attn  # 返回上下文向量和注意力权重

# 多头注意力 - 将缩放点积注意力并行化
class MultiHeadAttention(nn.Module):
    def __init__(self):
        super(MultiHeadAttention, self).__init__()
        # 定义线性层:将d_model维度的输入投影到n_heads*d_k维度(拆分到多个头)
        self.W_Q = nn.Linear(d_model, d_k * n_heads, bias=False)  # Q的投影层
        self.W_K = nn.Linear(d_model, d_k * n_heads, bias=False)  # K的投影层
        self.W_V = nn.Linear(d_model, d_v * n_heads, bias=False)  # V的投影层
        # 拼接所有头的输出后,投影回d_model维度
        self.fc = nn.Linear(n_heads * d_v, d_model, bias=False)
        # 层归一化:稳定训练
        self.layernorm = nn.LayerNorm(d_model)
    def forward(self, input_Q, input_K, input_V, attn_mask):
        '''
        前向传播:多头注意力计算
        参数说明:
        input_Q: 输入Q [batch_size, len_q, d_model]
        input_K: 输入K [batch_size, len_k, d_model]
        input_V: 输入V [batch_size, len_v(=len_k), d_model]
        attn_mask: 注意力掩码 [batch_size, seq_len, seq_len]
        '''
        # 保存残差连接的输入(后续相加)
        residual, batch_size = input_Q, input_Q.size(0)
        # 1. 线性投影 + 拆分多头
        # (B, S, D) -> (B, S, H*D_k) -> (B, S, H, D_k) -> (B, H, S, D_k)
        Q = self.W_Q(input_Q).view(batch_size, -1, n_heads, d_k).transpose(1,2) 
        K = self.W_K(input_K).view(batch_size, -1, n_heads, d_k).transpose(1,2) 
        V = self.W_V(input_V).view(batch_size, -1, n_heads, d_v).transpose(1,2) 
        # 2. 扩展掩码维度:适配多头 [batch_size, n_heads, seq_len, seq_len]
        attn_mask = attn_mask.unsqueeze(1).repeat(1, n_heads, 1, 1) 
        # 3. 计算缩放点积注意力
        context, attn = ScaledDotProductAttention()(Q, K, V, attn_mask)
        # 4. 拼接所有头的输出:(B, H, S, D_v) -> (B, S, H, D_v) -> (B, S, H*D_v)
        context = context.transpose(1, 2).reshape(batch_size, -1, n_heads * d_v)
        # 5. 线性投影回d_model维度
        output = self.fc(context)
        # 6. 残差连接 + 层归一化
        return self.layernorm(output + residual), attn

### 2.3 构建前馈网络
class PoswiseFeedForwardNet(nn.Module):
    def __init__(self):
        super(PoswiseFeedForwardNet, self).__init__()
        # 前馈网络结构:两层线性层 + ReLU激活
        self.fc = nn.Sequential(
            nn.Linear(d_model, d_ff, bias=False),  # 升维到d_ff
            nn.ReLU(),  # 非线性激活
            nn.Linear(d_ff, d_model, bias=False))  # 降维回d_model
        # 层归一化
        self.layernorm = nn.LayerNorm(d_model)

    def forward(self, inputs):  
        '''
        前向传播:前馈网络 + 残差 + 层归一化
        inputs: [batch_size, seq_len, d_model]
        '''
        residual = inputs  # 残差连接
        output = self.fc(inputs)  # 前馈网络计算
        return self.layernorm(output + residual)  # 残差 + 层归一化

### 2.4 解码器模块
# 解码器层 - GPT的核心层(仅包含自注意力+前馈网络)
class DecoderLayer(nn.Module):
    def __init__(self):
        super(DecoderLayer, self).__init__()
        # 解码器自注意力(自回归,输入=键=值)
        self.dec_self_attn = MultiHeadAttention()
        # 前馈网络
        self.pos_ffn = PoswiseFeedForwardNet()

    def forward(self, dec_inputs, dec_self_attn_mask):
        '''
        前向传播:单解码器层计算
        参数说明:
        dec_inputs: 解码器输入 [batch_size, tgt_len, d_model]
        dec_self_attn_mask: 解码器自注意力掩码 [batch_size, tgt_len, tgt_len]
        '''
        # 1. 自注意力计算(残差+层归一化)
        dec_outputs, dec_self_attn = self.dec_self_attn(dec_inputs, dec_inputs, dec_inputs, dec_self_attn_mask)
        # 2. 前馈网络计算(残差+层归一化)
        dec_outputs = self.pos_ffn(dec_outputs)  
        return dec_outputs, dec_self_attn

# 解码器模块 - 多层解码器层堆叠
class Decoder(nn.Module):
    def __init__(self):
        super(Decoder, self).__init__()
        # 词嵌入层:将索引转换为d_model维度的向量
        self.tgt_emb = nn.Embedding(vocab_size, d_model)
        # 位置嵌入层:将位置索引转换为d_model维度的向量(GPT是绝对位置编码)
        self.pos_emb = nn.Embedding(seq_len, d_model)
        # 堆叠n_layers个解码器层
        self.layers = nn.ModuleList([DecoderLayer() for _ in range(n_layers)])

    def forward(self, dec_inputs):
        '''
        前向传播:解码器整体计算
        dec_inputs: [batch_size, tgt_len] 索引序列
        '''
        # 1. 获取序列长度
        seq_len = dec_inputs.size(1)
        # 2. 构建位置索引:[0,1,...,seq_len-1] -> [batch_size, seq_len]
        pos = torch.arange(seq_len, dtype=torch.long, device=device)
        pos = pos.unsqueeze(0).expand_as(dec_inputs)  

        # 3. 计算词嵌入和位置嵌入,并相加
        word_emb = self.tgt_emb(dec_inputs)  # 词嵌入 [batch_size, tgt_len, d_model]
        pos_emb = self.pos_emb(pos)  # 位置嵌入 [batch_size, tgt_len, d_model]
        dec_outputs = word_emb + pos_emb  # 嵌入总和

        # 4. 构建解码器自注意力掩码(Pad掩码 + 未来信息掩码)
        # Pad掩码:屏蔽padding
        dec_self_attn_pad_mask = get_attn_pad_mask(dec_inputs, dec_inputs)  
        # 未来信息掩码:屏蔽未来token
        dec_self_attn_subsequent_mask = get_attn_subsequence_mask(dec_inputs)  
        # 合并掩码:只要有一个掩码为1(需要屏蔽),最终掩码就为1
        dec_self_attn_mask = torch.gt((dec_self_attn_pad_mask + dec_self_attn_subsequent_mask), 0)  

        # 5. 逐层计算解码器
        dec_self_attns = []  # 保存每一层的注意力权重(可选,用于分析)
        for layer in self.layers:
            dec_outputs, dec_self_attn = layer(dec_outputs, dec_self_attn_mask)
            dec_self_attns.append(dec_self_attn)

        return dec_outputs, dec_self_attns

### 2.5 GPT模型
class GPT(nn.Module):
    def __init__(self):
        super(GPT, self).__init__()
        self.decoder = Decoder()  # 解码器模块
        # 输出投影层:将d_model维度映射到词表大小(预测下一个token)
        self.projection = nn.Linear(d_model, vocab_size, bias=False)

    def forward(self, dec_inputs):
        """
        前向传播:GPT整体计算
        dec_inputs: [batch_size, tgt_len] 索引序列
        """
        # 解码器计算:得到输出和注意力权重
        dec_outputs, dec_self_attns = self.decoder(dec_inputs)
        # 投影到词表大小:[batch_size, tgt_len, vocab_size]
        dec_logits = self.projection(dec_outputs)
        # 重塑形状:[batch_size*tgt_len, vocab_size](适配CrossEntropyLoss的输入格式)
        return dec_logits.view(-1, dec_logits.size(-1)), dec_self_attns

    def answer(self, above): 
        '''
        生成回复:自回归生成文本
        above: 输入的查询文本(字符串)
        return: 生成的回复文本
        '''
        # 1. 预处理输入文本:转换为索引序列,并添加<sep>标记
        dec_input = [vocab[word] for word in above]  # 字符转索引
        dec_input.append(vocab['<sep>'])  # 添加分隔符
        # 转换为张量:[1, seq_len](batch_size=1),移到指定设备
        dec_input = torch.tensor(dec_input, dtype=torch.long, device=device).unsqueeze(0)

        # 2. 自回归生成(最多生成100个token,避免无限循环)
        for i in range(100):
            # 解码器前向计算
            dec_outputs, _ = self.decoder(dec_input)
            # 投影到词表大小,预测下一个token
            projected = self.projection(dec_outputs)
            # 取概率最大的token索引:squeeze(0)去掉batch维度,max(dim=-1)取最后一维最大值,[1]取索引
            prob = projected.squeeze(0).max(dim=-1, keepdim=False)[1]
            next_id = prob.data[-1]  # 取最后一个token的预测索引

            # 终止条件:生成<sep>标记则停止
            if next_id == vocab["<sep>"]:
                break

            # 拼接新生成的token到输入序列,继续生成
            dec_input = torch.cat(
                [dec_input.detach(), torch.tensor([[next_id]], dtype=dec_input.dtype, device=device)], -1)

        # 3. 处理生成结果:转换为文本
        output = dec_input.squeeze(0)  # 去掉batch维度
        sequence = [vocab[int(id)] for id in output] # 索引转字符

        # 4. 提取回复内容:取最后一个<sep>之后的部分
        answer = "".join(sequence)
        # 找到最后一个<sep>的位置,截取后面的内容(+5是因为"<sep>"长度为5)
        answer = answer[answer.rindex("<sep>")+5:]  

        return answer

## 3. 模型训练
# 定义超参数(补充seq_len,原代码此处重复定义,统一整理)
seq_len = 300  # 序列最大长度
epochs = 30  # 训练轮数

# 实例化模型并移到指定设备
model = GPT().to(device)
# 定义损失函数:CrossEntropyLoss,ignore_index=0(忽略padding的损失)
criterion = nn.CrossEntropyLoss(ignore_index=0).to(device)
# 定义优化器:Adam优化器,学习率1e-4
optimizer = optim.Adam(model.parameters(), lr=1e-4)

loss_history = [] # 记录每轮的训练损失,用于绘图
# 训练循环
for epoch in range(epochs):
    model.train()  # 切换到训练模式
    epoch_loss = 0  # 累计当前轮的总损失
    # 遍历数据加载器,tqdm显示进度条
    for i, (dec_inputs, dec_outputs) in enumerate(tqdm(data_loader)):
        optimizer.zero_grad()  # 清空梯度
        # 将数据移到指定设备
        dec_inputs, dec_outputs =dec_inputs.to(device), dec_outputs.to(device)
        # 模型前向计算
        outputs, dec_self_attns = model(dec_inputs)
        # 计算损失:outputs是[batch_size*tgt_len, vocab_size],dec_outputs.view(-1)是[batch_size*tgt_len]
        loss = criterion(outputs, dec_outputs.view(-1))
        epoch_loss += loss.item()  # 累计损失
        loss.backward()  # 反向传播计算梯度
        optimizer.step()  # 更新参数
    # 计算当前轮的平均损失
    train_loss = epoch_loss / len(data_loader)
    loss_history.append(train_loss)   # 记录损失
    print(f'\tTrain Loss: {train_loss:.3f}')  # 打印损失
    # 保存模型参数(需确保model文件夹已创建)
    torch.save(model.state_dict(), 'model/gpt_chat.pt')

# 绘制损失曲线:可视化训练过程
plt.plot(loss_history)  # 横轴:轮数,纵轴:训练损失
plt.ylabel('train loss')  # 纵轴标签:训练损失
plt.xlabel('epoch')  # 横轴标签:训练轮数(补充,更清晰)
plt.title('GPT Training Loss Curve')  # 标题(补充,更清晰)
plt.show()  # 显示图像

## 4. 效果测试
# 重新实例化模型(可选,也可直接用训练后的model)
model = GPT().to(device)
# 加载训练好的模型参数
model.load_state_dict(torch.load('model/gpt_chat.pt'))
# 切换到评估模式(禁用Dropout等训练特有的层)
model.eval()

# 测试案例1
ask = "你好啊"
print(f"输入:{ask},输出:{model.answer(ask)}")

# 测试案例2
ask = "你叫什么名字"
print(f"输入:{ask},输出:{model.answer(ask)}")

# 测试案例3
ask = "今天天气不错"
print(f"输入:{ask},输出:{model.answer(ask)}")

PyCharm版本代码(无if name,可直接运行)

适配PyCharm运行,调整路径逻辑,确保依赖和路径正确:

import torch
import torch.nn as nn
import torch.utils.data as Data
from torch import optim
import numpy as np
from tqdm import tqdm
import matplotlib.pyplot as plt
import re
import string
from collections import Counter
import os

# ======================== 1. 环境准备 ========================
# 创建必要文件夹(避免路径不存在报错)
os.makedirs("data", exist_ok=True)
os.makedirs("model", exist_ok=True)

# 注意:需将dataset.txt放入data文件夹,或修改以下路径为实际路径
DATA_PATH = "data/dataset.txt"
MODEL_SAVE_PATH = "model/gpt_chat.pt"

# 设备配置
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")

# 模型超参数(全局定义,避免顺序问题)
seq_len = 300  # 序列最大长度
d_model = 768  # Embedding维度
d_ff = 2048  # 前馈层维度
d_k = d_v = 64  # Q/K/V维度
n_layers = 6  # 解码器层数
n_heads = 8  # 多头注意力头数
batch_size = 64
epochs = 30  # 训练轮数

# ======================== 2. 数据预处理 ========================
# 2.1 读取数据
def load_data(path):
    """读取数据集"""
    if not os.path.exists(path):
        raise FileNotFoundError(f"数据集文件不存在:{path},请检查路径")
    with open(path, 'r', encoding='utf-8') as f:
        datas = f.readlines()
    return datas

# 2.2 文本预处理
def preprocess_data(datas):
    """文本清洗+词元化"""
    # 查看特殊字符(可选,用于分析)
    content = ''.join(datas)
    special_char = re.sub(r'[\u4e00-\u9fa5]', ' ', content)
    print("数据中的特殊字符:", set(special_char) - set(string.ascii_letters) - set(string.digits))
    
    # 词元化
    def tokenize(datas):
        tokens = []
        for data in datas:
            data = data.strip().replace("\n", "")
            token = [i if i != '\t' else "<sep>" for i in data] + ['<sep>']
            tokens.append(token)
        return tokens
    
    tokens = tokenize(datas)
    print("词元化示例(前6行):", tokens[:6])
    return tokens

# 2.3 构建词表
class Vocab:
    def __init__(self, tokens):
        self.tokens = tokens
        self.token2index = {'<pad>': 0, '<unk>': 1, '<seq>': 2}
        flatten = lambda l: [item for sublist in l for item in sublist]
        self.token2index.update({
            token: index + 3
            for index, (token, freq) in enumerate(
                sorted(Counter(flatten(self.tokens)).items(), key=lambda x: x[1], reverse=True))
        })
        self.index2token = {index: token for token, index in self.token2index.items()}

    def __getitem__(self, query):
        if isinstance(query, str):
            return self.token2index.get(query, 0)
        elif isinstance(query, int):
            return self.index2token.get(query, '<unk>')
        elif isinstance(query, (list, tuple)):
            return [self.__getitem__(item) for item in query]

    def __len__(self):
        return len(self.index2token)

# 2.4 自定义数据集
class MyDataSet(Data.Dataset):
    def __init__(self, tokens_num):
        self.datas = tokens_num

    def __getitem__(self, item):
        data = self.datas[item]
        decoder_input = data[:-1]
        decoder_output = data[1:]
        decoder_input_len = len(decoder_input)
        decoder_output_len = len(decoder_output)
        return {
            "decoder_input": decoder_input,
            "decoder_input_len": decoder_input_len,
            "decoder_output": decoder_output,
            "decoder_output_len": decoder_output_len
        }

    def __len__(self):
        return len(self.datas)

    def padding_batch(self, batch):
        decoder_input_lens = [d["decoder_input_len"] for d in batch]
        decoder_output_lens = [d["decoder_output_len"] for d in batch]
        decoder_input_maxlen = max(decoder_input_lens)
        decoder_output_maxlen = max(decoder_output_lens)

        for d in batch:
            d["decoder_input"].extend([0] * (decoder_input_maxlen - d["decoder_input_len"]))
            d["decoder_output"].extend([0] * (decoder_output_maxlen - d["decoder_output_len"]))
        
        decoder_inputs = torch.tensor([d["decoder_input"] for d in batch], dtype=torch.long)
        decoder_outputs = torch.tensor([d["decoder_output"] for d in batch], dtype=torch.long)
        return decoder_inputs, decoder_outputs

# 2.5 构建数据加载器
def build_dataloader(tokens, vocab):
    tokens_num = [[vocab[word] for word in line] for line in tokens]
    dataset = MyDataSet(tokens_num)
    data_loader = Data.DataLoader(
        dataset, 
        batch_size=batch_size, 
        collate_fn=dataset.padding_batch,
        shuffle=True  # 训练时打乱数据(补充,提升训练效果)
    )
    return data_loader

# ======================== 3. 模型定义 ========================
# 3.1 掩码函数
def get_attn_pad_mask(seq_q, seq_k):
    batch_size, len_q = seq_q.size()
    batch_size, len_k = seq_k.size()
    pad_attn_mask = seq_k.data.eq(0).unsqueeze(1)
    return pad_attn_mask.expand(batch_size, len_q, len_k)

def get_attn_subsequence_mask(seq):
    attn_shape = [seq.size(0), seq.size(1), seq.size(1)]
    subsequence_mask = np.triu(np.ones(attn_shape), k=1)
    subsequence_mask = torch.from_numpy(subsequence_mask).byte().to(device)
    return subsequence_mask

# 3.2 缩放点积注意力
class ScaledDotProductAttention(nn.Module):
    def __init__(self):
        super(ScaledDotProductAttention, self).__init__()

    def forward(self, Q, K, V, attn_mask):
        scores = torch.matmul(Q, K.transpose(-1, -2)) / np.sqrt(d_k)
        scores.masked_fill_(attn_mask, -1e9)
        attn = nn.Softmax(dim=-1)(scores)
        context = torch.matmul(attn, V)
        return context, attn

# 3.3 多头注意力
class MultiHeadAttention(nn.Module):
    def __init__(self):
        super(MultiHeadAttention, self).__init__()
        self.W_Q = nn.Linear(d_model, d_k * n_heads, bias=False)
        self.W_K = nn.Linear(d_model, d_k * n_heads, bias=False)
        self.W_V = nn.Linear(d_model, d_v * n_heads, bias=False)
        self.fc = nn.Linear(n_heads * d_v, d_model, bias=False)
        self.layernorm = nn.LayerNorm(d_model)

    def forward(self, input_Q, input_K, input_V, attn_mask):
        residual, batch_size = input_Q, input_Q.size(0)
        Q = self.W_Q(input_Q).view(batch_size, -1, n_heads, d_k).transpose(1, 2)
        K = self.W_K(input_K).view(batch_size, -1, n_heads, d_k).transpose(1, 2)
        V = self.W_V(input_V).view(batch_size, -1, n_heads, d_v).transpose(1, 2)
        attn_mask = attn_mask.unsqueeze(1).repeat(1, n_heads, 1, 1)
        context, attn = ScaledDotProductAttention()(Q, K, V, attn_mask)
        context = context.transpose(1, 2).reshape(batch_size, -1, n_heads * d_v)
        output = self.fc(context)
        return self.layernorm(output + residual), attn

# 3.4 前馈网络
class PoswiseFeedForwardNet(nn.Module):
    def __init__(self):
        super(PoswiseFeedForwardNet, self).__init__()
        self.fc = nn.Sequential(
            nn.Linear(d_model, d_ff, bias=False),
            nn.ReLU(),
            nn.Linear(d_ff, d_model, bias=False)
        )
        self.layernorm = nn.LayerNorm(d_model)

    def forward(self, inputs):
        residual = inputs
        output = self.fc(inputs)
        return self.layernorm(output + residual)

# 3.5 解码器层
class DecoderLayer(nn.Module):
    def __init__(self):
        super(DecoderLayer, self).__init__()
        self.dec_self_attn = MultiHeadAttention()
        self.pos_ffn = PoswiseFeedForwardNet()

    def forward(self, dec_inputs, dec_self_attn_mask):
        dec_outputs, dec_self_attn = self.dec_self_attn(dec_inputs, dec_inputs, dec_inputs, dec_self_attn_mask)
        dec_outputs = self.pos_ffn(dec_outputs)
        return dec_outputs, dec_self_attn

# 3.6 解码器
class Decoder(nn.Module):
    def __init__(self, vocab_size):
        super(Decoder, self).__init__()
        self.tgt_emb = nn.Embedding(vocab_size, d_model)
        self.pos_emb = nn.Embedding(seq_len, d_model)
        self.layers = nn.ModuleList([DecoderLayer() for _ in range(n_layers)])

    def forward(self, dec_inputs):
        seq_len = dec_inputs.size(1)
        pos = torch.arange(seq_len, dtype=torch.long, device=device).unsqueeze(0).expand_as(dec_inputs)
        word_emb = self.tgt_emb(dec_inputs)
        pos_emb = self.pos_emb(pos)
        dec_outputs = word_emb + pos_emb

        dec_self_attn_pad_mask = get_attn_pad_mask(dec_inputs, dec_inputs)
        dec_self_attn_subsequent_mask = get_attn_subsequence_mask(dec_inputs)
        dec_self_attn_mask = torch.gt((dec_self_attn_pad_mask + dec_self_attn_subsequent_mask), 0)

        dec_self_attns = []
        for layer in self.layers:
            dec_outputs, dec_self_attn = layer(dec_outputs, dec_self_attn_mask)
            dec_self_attns.append(dec_self_attn)
        return dec_outputs, dec_self_attns

# 3.7 GPT模型
class GPT(nn.Module):
    def __init__(self, vocab_size):
        super(GPT, self).__init__()
        self.decoder = Decoder(vocab_size)
        self.projection = nn.Linear(d_model, vocab_size, bias=False)

    def forward(self, dec_inputs):
        dec_outputs, dec_self_attns = self.decoder(dec_inputs)
        dec_logits = self.projection(dec_outputs)
        return dec_logits.view(-1, dec_logits.size(-1)), dec_self_attns

    def answer(self, above, vocab):
        dec_input = [vocab[word] for word in above]
        dec_input.append(vocab['<sep>'])
        dec_input = torch.tensor(dec_input, dtype=torch.long, device=device).unsqueeze(0)

        for i in range(100):
            dec_outputs, _ = self.decoder(dec_input)
            projected = self.projection(dec_outputs)
            prob = projected.squeeze(0).max(dim=-1, keepdim=False)[1]
            next_id = prob.data[-1]

            if next_id == vocab["<sep>"]:
                break

            dec_input = torch.cat(
                [dec_input.detach(), torch.tensor([[next_id]], dtype=dec_input.dtype, device=device)], -1)

        output = dec_input.squeeze(0)
        sequence = [vocab[int(id)] for id in output]
        answer = "".join(sequence)
        if "<sep>" in answer:
            answer = answer[answer.rindex("<sep>")+5:]
        return answer

# ======================== 4. 模型训练 ========================
def train_model(model, data_loader, vocab, epochs, device):
    """训练模型"""
    criterion = nn.CrossEntropyLoss(ignore_index=0).to(device)
    optimizer = optim.Adam(model.parameters(), lr=1e-4)
    loss_history = []

    for epoch in range(epochs):
        model.train()
        epoch_loss = 0
        pbar = tqdm(data_loader, desc=f"Epoch {epoch+1}/{epochs}")
        for i, (dec_inputs, dec_outputs) in enumerate(pbar):
            optimizer.zero_grad()
            dec_inputs, dec_outputs = dec_inputs.to(device), dec_outputs.to(device)
            outputs, _ = model(dec_inputs)
            loss = criterion(outputs, dec_outputs.view(-1))
            epoch_loss += loss.item()
            loss.backward()
            optimizer.step()
            pbar.set_postfix({"batch_loss": loss.item()})
        
        train_loss = epoch_loss / len(data_loader)
        loss_history.append(train_loss)
        print(f'Epoch {epoch+1}/{epochs}, Train Loss: {train_loss:.3f}')
        torch.save(model.state_dict(), MODEL_SAVE_PATH)
    
    # 绘制损失曲线
    plt.plot(loss_history)
    plt.ylabel('Train Loss')
    plt.xlabel('Epoch')
    plt.title('GPT Training Loss Curve')
    plt.savefig("train_loss.png")  # 保存图片(PyCharm中更方便)
    plt.show()

# ======================== 5. 模型测试 ========================
def test_model(model, vocab, device):
    """测试模型生成效果"""
    model.load_state_dict(torch.load(MODEL_SAVE_PATH, map_location=device))
    model.eval()
    
    test_cases = ["你好啊", "你叫什么名字", "今天天气不错"]
    for ask in test_cases:
        answer = model.answer(ask, vocab)
        print(f"输入:{ask} → 输出:{answer}")

# ======================== 6. 主流程执行 ========================
# 加载数据
datas = load_data(DATA_PATH)
# 预处理
tokens = preprocess_data(datas)
# 构建词表
vocab = Vocab(tokens)
vocab_size = len(vocab)
print(f"词表大小:{vocab_size}")
# 构建数据加载器
data_loader = build_dataloader(tokens, vocab)
# 实例化模型
model = GPT(vocab_size).to(device)
# 训练模型
train_model(model, data_loader, vocab, epochs, device)
# 测试模型
test_model(model, vocab, device)

核心知识点系统梳理

知识点分类 核心概念 详细解释 代码对应模块
数据预处理 词元化 将文本拆分为最小语义单位(中文按字符拆分),用<sep>标记对话分隔符 tokenize函数
词表构建 统计词频,建立词元<->索引映射,包含<pad>(填充)、<unk>(未知)等特殊标记 Vocab
自定义Dataset 继承PyTorch Dataset,实现样本读取和批处理padding(统一序列长度) MyDataSet
GPT核心机制 自回归掩码 未来信息掩码(屏蔽未来token)+ Pad掩码(屏蔽填充),保证自回归特性 get_attn_pad_mask/get_attn_subsequence_mask
缩放点积注意力 Q·K^T / √d_k 缩放,避免梯度消失,Softmax后加权V得到上下文向量 ScaledDotProductAttention
多头注意力 将Q/K/V拆分为多个头并行计算注意力,拼接后投影回原维度,提升模型表达能力 MultiHeadAttention
残差连接+层归一化 残差连接缓解梯度消失,层归一化稳定训练过程,是Transformer的核心优化 多头注意力/前馈网络中均有实现
前馈网络 两层线性层+ReLU激活,对每个位置的特征独立变换,提升模型非线性表达能力 PoswiseFeedForwardNet
模型架构 GPT解码器 仅包含解码器(无编码器),多层解码器层堆叠,每层层包含自注意力+前馈网络 Decoder/DecoderLayer
位置编码 绝对位置编码,将位置索引转换为向量,与词嵌入相加(GPT用绝对位置,区别于BERT) Decoder类中pos_emb
模型训练 损失函数 CrossEntropyLoss,ignore_index=0忽略padding部分的损失 训练部分criterion定义
自回归训练 输入序列为x1,x2,...,xn-1,输出序列为x2,x3,...,xn,预测下一个token MyDataSetdecoder_input/output拆分
文本生成 自回归生成 输入初始文本,循环预测下一个token,直到生成<sep>或达到最大长度 GPT类中answer方法

总结

  1. 数据预处理核心:中文GPT实现需先对文本进行字符级词元化,构建包含特殊标记的词表,通过padding统一批次内序列长度,适配PyTorch张量计算。
  2. GPT模型核心:仅由解码器构成,核心是自注意力机制(含Pad掩码和未来信息掩码),通过残差连接+层归一化稳定训练,最终通过投影层预测下一个token。
  3. 文本生成逻辑:采用自回归方式,输入初始文本后循环生成token,直到触发终止标记(<sep>),实现对话回复生成。
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐