学习Bert微调

微调介绍

句子分类任务(句子级别)

dense就是全连接层

理论上拿哪个都行,但人家大模型训练是对<cls>"class意思"这个学习的,要用人家训练好的东西微调,那最好就得一样了(用cls是因为我们在预训练bert的时候nsp就是用cls去进行预测的 这样下游任务与预训练尽量保持统一)

说明:

BERT 在出生的时候(也就是它在做预训练、看海量无标注文章的时候),设计了两个核心任务:

  • 任务一(MLM): 完形填空(猜被遮盖的字)。

  • 任务二(NSP,Next Sentence Prediction): 判断第二句话是不是紧跟在第一句话后面。

在做 NSP 任务时,BERT 强制规定:必须用开头 <cls> 位置的向量来预测这两句话是真是假。 因为在整个预训练的漫长过程中,<cls> 已经习惯了背负"总结全句/判断全局"的重任。你在下游做微调(Fine-tuning)时,如果用别人训练好的 BERT,就必须沿用人家的老习惯 ,输入 <cls> 才能完美继承 BERT 预训练时学到的功力。

命名实体识别(词级别)

句子的分类的输出是有权重的,这个权重就可以拥有句子信息

问题回答:qa

<cls>问题<sep>文章<sep>,取出文章所有词的特征,接全连接预测文章中每个词是否是答案的开始/结束/都不是 三分类。 第一个seq句子是问题,第二个seq句子是描述的一段话,输出是当前词元是否为答案的开始/结束/都不是。

先看问题,再做阅读理解 <答案开始> <答案><答案结束>,在给定数据集中找这三个,相当于做阅读理解,最简单的一种qa

总结:可能输出层的类别还有用到的特征不一样

数据集代码:

python 复制代码
import os
import re
import torch
from torch import nn
from d2l import torch as d2l


# 将SNLI数据集添加到d2l的数据下载列表中
# SNLI是一个自然语言推理(Natural Language Inference, NLI)数据集
# 包含:
#   - premise(前提)
#   - hypothesis(假设)
#   - label(关系标签:蕴含、矛盾、中立)
d2l.DATA_HUB['SNLI'] = (
    'https://nlp.stanford.edu/projects/snli/snli_1.0.zip',
    '9fcde07509c7e87ec61c640c1b2753d9041758e4')


# 下载并解压SNLI数据集
# 如果本地已经存在,则直接使用
data_dir = d2l.download_extract('SNLI')


# 定义读取SNLI数据集的函数
# 参数:
#   data_dir: 数据集所在目录
#   is_train: 是否读取训练集
#            True -> 读取训练数据
#            False -> 读取测试数据
def read_snli(data_dir, is_train):
    """
    将SNLI数据集解析为:
    前提(premise)、假设(hypothesis)和标签(label)
    """


    # 定义文本预处理函数
    def extract_text(s):
        # SNLI中的句子包含括号,例如:
        # "( A person is walking )"
        # 删除左括号
        s = re.sub('\\(', '', s)

        # 删除右括号
        s = re.sub('\\)', '', s)


        # 将连续两个或多个空格替换成一个空格
        # 例如:
        # "hello     world"
        # 变成:
        # "hello world"
        s = re.sub('\\s{2,}', ' ', s)


        # 删除字符串两端多余空格
        return s.strip()


    # 定义标签映射关系
    # SNLI原始标签是英文字符串
    # 转换成数字方便神经网络训练
    #
    # entailment:
    #   前提可以推出假设 -> 0
    #
    # contradiction:
    #   前提与假设矛盾 -> 1
    #
    # neutral:
    #   前提与假设没有确定关系 -> 2
    label_set = {
        'entailment': 0,
        'contradiction': 1,
        'neutral': 2
    }


    # 根据训练/测试选择不同文件
    #
    # 训练集:
    # snli_1.0_train.txt
    #
    # 测试集:
    # snli_1.0_test.txt
    file_name = os.path.join(
        data_dir,
        'snli_1.0_train.txt' if is_train
        else 'snli_1.0_test.txt'
    )


    # 打开数据文件
    # 每一行使用tab分割
    #
    # 第一行是表头,所以跳过
    with open(file_name, 'r') as f:
        rows = [
            row.split('\t')
            for row in f.readlines()[1:]
        ]


    # 提取前提句(premise)
    #
    # row[0] 是标签
    # row[1] 是前提
    #
    # 只保留有效标签的数据
    premises = [
        extract_text(row[1])
        for row in rows
        if row[0] in label_set
    ]


    # 提取假设句(hypothesis)
    #
    # row[2] 是假设
    hypotheses = [
        extract_text(row[2])
        for row in rows
        if row[0] in label_set
    ]


    # 提取标签,并转换成数字
    #
    # 例如:
    # entailment -> 0
    # contradiction -> 1
    # neutral -> 2
    labels = [
        label_set[row[0]]
        for row in rows
        if row[0] in label_set
    ]


    # 返回三个列表
    #
    # premises:
    #   前提句列表
    #
    # hypotheses:
    #   假设句列表
    #
    # labels:
    #   标签列表
    return premises, hypotheses, labels



# 读取训练集
# 返回:
# train_data[0] -> 所有前提句
# train_data[1] -> 所有假设句
# train_data[2] -> 所有标签
train_data = read_snli(data_dir, is_train=True)



# 查看前三条训练数据
# zip会把三个列表对应位置的数据组合起来
#
# x0:
#   前提
#
# x1:
#   假设
#
# y:
#   标签
for x0, x1, y in zip(
        train_data[0][:3],
        train_data[1][:3],
        train_data[2][:3]):

    print('前提:', x0)
    print('假设:', x1)
    print('标签:', y)

整体流程:

python 复制代码
SNLI文件
   ↓
read_snli(data_dir, True)
   ↓
train_data = (
    前提列表premises,
    假设列表hypotheses,
    标签列表labels
)

train_data
   ├── train_data[0][:3] → 前3个前提
   ├── train_data[1][:3] → 前3个假设
   └── train_data[2][:3] → 前3个标签

        ↓ zip()

(
 前提1, 假设1, 标签1
 前提2, 假设2, 标签2
 前提3, 假设3, 标签3
)

        ↓ for循环

x0 = 前提
x1 = 假设
y  = 标签

        ↓

print输出
python 复制代码
train_data
    |
    | data[2]
    ↓
标签列表
[0,2,1,0,1,2...]

    |
    | count()
    ↓

统计:
0有多少个
1有多少个
2有多少个


输出:
[数量0, 数量1, 数量2]



# 读取测试集
# is_train=False表示读取test文件
test_data = read_snli(data_dir, is_train=False)


# 遍历训练集和测试集
for data in [train_data, test_data]:

    # data[2]表示标签列表
    # 统计标签0、1、2分别有多少个
    print([
        [row for row in data[2]].count(i)
        for i in range(3)   # i依次取0、1、2
    ])

#等价写法

test_data = read_snli(data_dir, is_train=False)  # 读取测试集

for data in [train_data, test_data]:  # 分别统计训练集和测试集

    # data[2]是标签列表,统计每种标签数量
    print([
        data[2].count(i)   # 统计标签i出现次数
        for i in range(3)  # 标签类别:0、1、2
    ])

定义用于加载数据集的类

下面我们来定义一个用于加载SNLI数据集的类。类构造函数中的变量num_steps指定文本序列的长度,使得每个小批量序列将具有相同的形状

换句话说,在较长序列中的前num_steps个标记之后的标记被截断,而特殊标记"<pad>"将被附加到较短的序列后,直到它们的长度变为num_steps。通过实现__getitem__功能,我们可以任意访问带有索引idx的前提、假设和标签。

等长不是为了语言本身,而是为了让不同长度的句子能够组成batch,以固定形状输入神经网络进行并行计算。

因为神经网络训练时通常需要把多个样本组成一个 batch(批次)一起计算,而一个 batch 中的数据必须是相同形状的张量,所以句子需要变成等长。

在NLP里这个操作叫:

padding(填充)和 truncation(截断)

最后 __getitem__()__len__() 是为了让这个类符合 PyTorch 的 Dataset 接口,可以被 DataLoader 按批次读取。

python 复制代码
# 定义SNLI数据集类,用于加载和处理数据
class SNLIDataset(torch.utils.data.Dataset):

    def __init__(self, dataset, num_steps, vocab=None):
        # 保存每个句子的固定长度
        self.num_steps = num_steps

        # 对前提和假设进行分词
        all_premise_tokens = d2l.tokenize(dataset[0])
        all_hypothesis_tokens = d2l.tokenize(dataset[1])

        # 如果没有传入词表,则创建词表
        if vocab is None:
            self.vocab = d2l.Vocab(
                all_premise_tokens + all_hypothesis_tokens,
                min_freq=5,                 # 低频词过滤
                reserved_tokens=['<pad>']   # 添加填充符
            )
        else:
            # 使用已有词表
            self.vocab = vocab

        # 将文本转换成定长数字序列
        self.premises = self.pad(all_premise_tokens)
        self.hypotheses = self.pad(all_hypothesis_tokens)

        # 标签转换为Tensor
        self.labels = torch.tensor(dataset[2])

        # 输出读取样本数量
        print('read ' + str(len(self.premises)) + ' examples')


    # 对句子进行截断和填充
    def pad(self, lines):
        return torch.tensor([
            d2l.truncate_pad(
                self.vocab[line],      # 单词转索引
                self.num_steps,        # 固定长度
                self.vocab['<pad>']    # 不足补pad
            )
            for line in lines
        ])


    # 根据索引获取一条数据
    def __getitem__(self, idx):
        return (
            self.premises[idx],      # 前提
            self.hypotheses[idx],    # 假设
            self.labels[idx]         # 标签
        )


    # 返回数据集大小
    def __len__(self):
        return len(self.premises)

整合代码:

python 复制代码
# 加载SNLI数据集,返回训练集、测试集迭代器和词表
def load_data_snli(batch_size, num_steps=50):

    # 获取DataLoader线程数
    num_workers = d2l.get_dataloader_workers()

    # 下载并读取SNLI数据
    data_dir = d2l.download_extract('SNLI')
    train_data = read_snli(data_dir, True)   # 训练集
    test_data = read_snli(data_dir, False)   # 测试集

    # 创建Dataset,完成分词、数字化、padding
    train_set = SNLIDataset(train_data, num_steps)

    # 测试集使用训练集词表,避免引入新词
    test_set = SNLIDataset(
        test_data,
        num_steps,
        train_set.vocab
    )

    # 创建DataLoader,按batch读取数据
    train_iter = torch.utils.data.DataLoader(
        train_set,
        batch_size,
        shuffle=True,          # 训练集随机打乱
        num_workers=num_workers
    )

    test_iter = torch.utils.data.DataLoader(
        test_set,
        batch_size,
        shuffle=False,         # 测试集不用打乱
        num_workers=num_workers
    )

    # 返回训练迭代器、测试迭代器、词表
    return train_iter, test_iter, train_set.vocab



# batch_size=128,每个句子长度固定为50
train_iter, test_iter, vocab = load_data_snli(128, 50)

# 查看词表大小
len(vocab)

read 549367 examples
read 9824 examples
18678


# 查看一个batch的数据形状
for X, Y in train_iter:
    print(X[0].shape)   # 前提句 shape
    print(X[1].shape)   # 假设句 shape
    print(Y.shape)      # 标签 shape
    break

torch.Size([128, 50])
torch.Size([128, 50])
torch.Size([128])
python 复制代码
SNLI文件
   ↓
read_snli()
   ↓
(premise, hypothesis, label)
   ↓
SNLIDataset
   ↓
分词 → 词表 → 数字索引 → padding
   ↓
DataLoader
   ↓
batch数据

X[0]  前提   [128,50]
X[1]  假设   [128,50]
Y     标签   [128]

Bert微调代码:

python 复制代码
import json
import multiprocessing
import os
import torch
from torch import nn
from d2l import torch as d2l

"""
加载bert数据集:

我们已经在 :numref:sec_bert-dataset和 :numref:sec_bert-pretrainingWikiText-2数据集上预训练BERT(请注意,原始的BERT模型是在更大的语料库上预训练的)。
正如在 :numref:sec_bert-pretraining中所讨论的,原始的BERT模型有数以亿计的参数。
在下面,我们提供了两个版本的预训练的BERT:"bert.base"与原始的BERT基础模型一样大,需要大量的计算资源才能进行微调,而"bert.small"是一个小版本,以便于演示。
"""
d2l.DATA_HUB['bert.base'] = (d2l.DATA_URL + 'bert.base.torch.zip',
                             '225d66f04cae318b841a13d32af3acc165f253ac')
d2l.DATA_HUB['bert.small'] = (d2l.DATA_URL + 'bert.small.torch.zip',
                              'c72329e68a732bef0452e4b96a1c341c8910f81f')
"""
两个预训练好的BERT模型都包含一个定义词表的"vocab.json"文件和一个预训练参数的"pretrained.params"文件。我们实现了以下load_pretrained_model函数来[加载预先训练好的BERT参数]。
"""
def load_pretrained_model(pretrained_model, num_hiddens, ffn_num_hiddens,
                          num_heads, num_layers, dropout, max_len, devices):
    data_dir = d2l.download_extract(pretrained_model)
    # 定义空词表以加载预定义词表
    vocab = d2l.Vocab()
    vocab.idx_to_token = json.load(open(os.path.join(data_dir,
        'vocab.json')))
    vocab.token_to_idx = {token: idx for idx, token in enumerate(
        vocab.idx_to_token)}
    bert = d2l.BERTModel(len(vocab), num_hiddens, norm_shape=[256],
                         ffn_num_input=256, ffn_num_hiddens=ffn_num_hiddens,
                         num_heads=4, num_layers=2, dropout=0.2,
                         max_len=max_len, key_size=256, query_size=256,
                         value_size=256, hid_in_features=256,
                         mlm_in_features=256, nsp_in_features=256)
    # 加载预训练BERT参数
    bert.load_state_dict(torch.load(os.path.join(data_dir,
                                                 'pretrained.params')))
    return bert, vocab

"""
为了便于在大多数机器上演示,我们将在本节中加载和微调经过预训练BERT的小版本("bert.small")。
在练习中,我们将展示如何微调大得多的"bert.base"以显著提高测试精度。
"""
devices = d2l.try_all_gpus()
bert, vocab = load_pretrained_model(
    'bert.small', num_hiddens=256, ffn_num_hiddens=512, num_heads=4,
    num_layers=2, dropout=0.1, max_len=512, devices=devices)


"""
[微调BERT的数据集]:

对于SNLI数据集的下游任务自然语言推断,我们定义了一个定制的数据集类SNLIBERTDataset。
在每个样本中,前提和假设形成一对文本序列,并被打包成一个BERT输入序列,如 :numref:fig_bert-two-seqs所示。回想 :numref:subsec_bert_input_rep,片段索引用于区分BERT输入序列中的前提和假设。

利用预定义的BERT输入序列的最大长度(max_len),持续移除输入文本对中较长文本的最后一个标记,直到满足max_len。
为了加速生成用于微调BERT的SNLI数据集,我们使用4个工作进程并行生成训练或测试样本。
"""
class SNLIBERTDataset(torch.utils.data.Dataset):
    def __init__(self, dataset, max_len, vocab=None):
        all_premise_hypothesis_tokens = [[
            p_tokens, h_tokens] for p_tokens, h_tokens in zip(
            *[d2l.tokenize([s.lower() for s in sentences])
              for sentences in dataset[:2]])]

        self.labels = torch.tensor(dataset[2])
        self.vocab = vocab
        self.max_len = max_len
        (self.all_token_ids, self.all_segments,
         self.valid_lens) = self._preprocess(all_premise_hypothesis_tokens)
        print('read ' + str(len(self.all_token_ids)) + ' examples')

    def _preprocess(self, all_premise_hypothesis_tokens):
        pool = multiprocessing.Pool(4)  # 使用4个进程
        out = pool.map(self._mp_worker, all_premise_hypothesis_tokens)
        all_token_ids = [
            token_ids for token_ids, segments, valid_len in out]
        all_segments = [segments for token_ids, segments, valid_len in out]
        valid_lens = [valid_len for token_ids, segments, valid_len in out]
        return (torch.tensor(all_token_ids, dtype=torch.long),
                torch.tensor(all_segments, dtype=torch.long),
                torch.tensor(valid_lens))

    def _mp_worker(self, premise_hypothesis_tokens):
        p_tokens, h_tokens = premise_hypothesis_tokens
        self._truncate_pair_of_tokens(p_tokens, h_tokens)
        tokens, segments = d2l.get_tokens_and_segments(p_tokens, h_tokens)
        token_ids = self.vocab[tokens] + [self.vocab['<pad>']] \
                             * (self.max_len - len(tokens))
        segments = segments + [0] * (self.max_len - len(segments))
        valid_len = len(tokens)
        return token_ids, segments, valid_len

    def _truncate_pair_of_tokens(self, p_tokens, h_tokens):
        # 为BERT输入中的'<CLS>'、'<SEP>'和'<SEP>'词元保留位置
        while len(p_tokens) + len(h_tokens) > self.max_len - 3:
            if len(p_tokens) > len(h_tokens):
                p_tokens.pop()
            else:
                h_tokens.pop()

    def __getitem__(self, idx):
        return (self.all_token_ids[idx], self.all_segments[idx],
                self.valid_lens[idx]), self.labels[idx]

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

"""
下载完SNLI数据集后,我们通过实例化SNLIBERTDataset类来[生成训练和测试样本]。
这些样本将在自然语言推断的训练和测试期间进行小批量读取。
"""
# 如果出现显存不足错误,请减少"batch_size"。在原始的BERT模型中,max_len=512
batch_size, max_len, num_workers = 512, 128, d2l.get_dataloader_workers()
data_dir = d2l.download_extract('SNLI')
train_set = SNLIBERTDataset(d2l.read_snli(data_dir, True), max_len, vocab)
test_set = SNLIBERTDataset(d2l.read_snli(data_dir, False), max_len, vocab)
train_iter = torch.utils.data.DataLoader(train_set, batch_size, shuffle=True,
                                   num_workers=num_workers)
test_iter = torch.utils.data.DataLoader(test_set, batch_size,
                                  num_workers=num_workers)

"""
微调BERT:
如 :numref:fig_bert-two-seqs所示,用于自然语言推断的微调BERT只需要一个额外的多层感知机,该多层感知机由两个全连接层组成(请参见下面BERTClassifier类中的self.hidden和self.output)。

[这个多层感知机将特殊的"<cls>"词元]的BERT表示进行了转换,该词元同时编码前提和假设的信息(为自然语言推断的三个输出):蕴涵、矛盾和中性。
"""
class BERTClassifier(nn.Module):
    def __init__(self, bert):
        super(BERTClassifier, self).__init__()
        self.encoder = bert.encoder
        self.hidden = bert.hidden
        self.output = nn.Linear(256, 3)

    def forward(self, inputs):
        tokens_X, segments_X, valid_lens_x = inputs
        encoded_X = self.encoder(tokens_X, segments_X, valid_lens_x)
        return self.output(self.hidden(encoded_X[:, 0, :]))

#因为自注意力机制的特性,经过了编码器之后,第一个<cls>就已经包含了整个序列的信息了。同时在与训练时,也针对了这个标签进行了分类训练,所以可以直接用来做finetune
"""
在下文中,预训练的BERT模型bert被送到用于下游应用的BERTClassifier实例net中。
在BERT微调的常见实现中,只有额外的多层感知机(net.output)的输出层的参数将从零开始学习。
预训练BERT编码器(net.encoder)和额外的多层感知机的隐藏层(net.hidden)的所有参数都将进行微调。
"""
net = BERTClassifier(bert)


"""
回想一下,在 :numref:sec_bert中,MaskLM类和NextSentencePred类在其使用的多层感知机中都有一些参数。
这些参数是预训练BERT模型bert中参数的一部分,因此是net中的参数的一部分。
然而,这些参数仅用于计算预训练过程中的遮蔽语言模型损失和下一句预测损失。
这两个损失函数与微调下游应用无关,因此当BERT微调时,MaskLM和NextSentencePred中采用的多层感知机的参数不会更新(陈旧的,staled)。

为了允许具有陈旧梯度的参数,标志ignore_stale_grad=True在step函数d2l.train_batch_ch13中被设置。我们通过该函数使用SNLI的训练集(train_iter)和测试集(test_iter)对net模型进行训练和评估。由于计算资源有限,[训练]和测试精度可以进一步提高:我们把对它的讨论留在练习中。
"""
lr, num_epochs = 1e-4, 5
trainer = torch.optim.Adam(net.parameters(), lr=lr)
loss = nn.CrossEntropyLoss(reduction='none')
d2l.train_ch13(net, train_iter, test_iter, loss, trainer, num_epochs,
    devices)

注释版:

python 复制代码
# BERT 自然语言推断:中文注释版
# 目标:输入"前提 + 假设"两句话,判断它们是蕴含、矛盾还是中立。
# 主线:加载预训练模型 -> 整理句子对 -> 接分类器 -> 用 SNLI 数据微调。
# 本文件只补充/修正讲解,保留原代码的执行逻辑,没有运行下载或训练。
# 注意:原代码使用多进程,在 Windows 上直接作为脚本运行前,
# 应把模型加载、数据集/DataLoader 构建和训练等入口放入
# if __name__ == '__main__': 保护中。这里没有代你重构运行入口。

import json                 # 读取保存词表的 JSON 文件
import multiprocessing      # 多进程处理数据
import os                   # 拼接文件路径
import torch
from torch import nn        # 全连接层、损失函数等神经网络组件
from d2l import torch as d2l # 《动手学深度学习》的 PyTorch 教学工具


# 一、加载已经预训练好的 BERT
# 这里只登记下载地址和校验值,还没有下载模型。
# 下载包包含 vocab.json(词表)和 pretrained.params(模型参数)。
d2l.DATA_HUB['bert.base'] = (d2l.DATA_URL + 'bert.base.torch.zip',
                             '225d66f04cae318b841a13d32af3acc165f253ac')
d2l.DATA_HUB['bert.small'] = (d2l.DATA_URL + 'bert.small.torch.zip',
                              'c72329e68a732bef0452e4b96a1c341c8910f81f')


def load_pretrained_model(pretrained_model, num_hiddens, ffn_num_hiddens,
                          num_heads, num_layers, dropout, max_len, devices):
    # 下载并解压模型;工具会利用已有缓存。
    data_dir = d2l.download_extract(pretrained_model)

    # 模型只能接收数字,词表负责"词元 <-> 数字编号"的转换。
    vocab = d2l.Vocab()
    vocab.idx_to_token = json.load(open(os.path.join(data_dir,
        'vocab.json')))
    vocab.token_to_idx = {token: idx for idx, token in enumerate(
        vocab.idx_to_token)}

    # 先建立模型结构,下一步再把预训练得到的权重填进去。
    # num_hiddens=256:每个位置最终用 256 维向量表示。
    # ffn_num_hiddens=512:编码器内部前馈网络的中间层维度,
    # 与后面"三分类器"的输出类别数 3 不是一回事。
    # num_heads=4:4 个注意力头;num_layers=2:2 层编码器。
    # max_len=512:这里设置位置表示能覆盖的最大序列长度。
    # 注意原代码把不少参数写死了,未使用传入的 num_heads、num_layers、
    # dropout 和 devices。实际 dropout 是下面的 0.2,不是调用处的 0.1。
    # 因而不能只把模型名换成 bert.base 就认为其它配置自动适配了。
    bert = d2l.BERTModel(len(vocab), num_hiddens, norm_shape=[256],
                         ffn_num_input=256, ffn_num_hiddens=ffn_num_hiddens,
                         num_heads=4, num_layers=2, dropout=0.2,
                         max_len=max_len, key_size=256, query_size=256,
                         value_size=256, hid_in_features=256,
                         mlm_in_features=256, nsp_in_features=256)

    # 用已经学好的参数初始化模型,而不是从随机参数开始训练。
    # 模型结构必须与参数文件匹配,且应只加载可信来源的权重文件。
    bert.load_state_dict(torch.load(os.path.join(data_dir,
                                                 'pretrained.params')))
    return bert, vocab


# 查找可用 GPU,没有 GPU 时返回 CPU 设备列表。
# 真正的训练设备分配交给最后的训练函数处理。
devices = d2l.try_all_gpus()
bert, vocab = load_pretrained_model(
    'bert.small', num_hiddens=256, ffn_num_hiddens=512, num_heads=4,
    num_layers=2, dropout=0.1, max_len=512, devices=devices)


# 二、把 SNLI 的"前提、假设、关系标签"整理成 BERT 能接收的数据
class SNLIBERTDataset(torch.utils.data.Dataset):
    def __init__(self, dataset, max_len, vocab=None):
        # dataset[0]:所有前提;dataset[1]:所有假设;dataset[2]:所有标签。
        # 下面这段较长的写法实际做了三件事:
        # 1. 把每句话转为小写。
        # 2. 用 d2l.tokenize 分词,得到词元列表。
        # 3. 用 zip 把同一个样本的前提和假设配成一对。
        # 注意这里使用教材的分词与词表,不是在调用完整的 Hugging Face tokenizer。
        all_premise_hypothesis_tokens = [[
            p_tokens, h_tokens] for p_tokens, h_tokens in zip(
            *[d2l.tokenize([s.lower() for s in sentences])
              for sentences in dataset[:2]])]

        self.labels = torch.tensor(dataset[2])  # 每个样本的真实关系类别
        self.vocab = vocab
        self.max_len = max_len

        # 提前处理全部样本,得到三个张量:词元编号、句子片段编号、有效长度。
        (self.all_token_ids, self.all_segments,
         self.valid_lens) = self._preprocess(all_premise_hypothesis_tokens)
        print('read ' + str(len(self.all_token_ids)) + ' examples')

    def _preprocess(self, all_premise_hypothesis_tokens):
        pool = multiprocessing.Pool(4)  # 用 4 个进程处理句子对
        # 每个句子对交给 _mp_worker,map 返回的结果顺序与输入顺序一致。
        # 原代码未显式关闭进程池,正式整理时可使用 with Pool(...) 管理资源。
        out = pool.map(self._mp_worker, all_premise_hypothesis_tokens)

        # 每个结果都是 (token_ids, segments, valid_len),这里分别收集。
        all_token_ids = [
            token_ids for token_ids, segments, valid_len in out]
        all_segments = [segments for token_ids, segments, valid_len in out]
        valid_lens = [valid_len for token_ids, segments, valid_len in out]

        # 假设有 N 个样本、统一长度 L:
        # all_token_ids / all_segments 的形状是 [N, L]。
        # valid_lens 的形状是 [N],每个样本对应一个有效长度。
        return (torch.tensor(all_token_ids, dtype=torch.long),
                torch.tensor(all_segments, dtype=torch.long),
                torch.tensor(valid_lens))

    def _mp_worker(self, premise_hypothesis_tokens):
        p_tokens, h_tokens = premise_hypothesis_tokens
        # 如果两句话总共太长,先截短。
        self._truncate_pair_of_tokens(p_tokens, h_tokens)

        # 拼接为:<cls> 前提 <sep> 假设 <sep>
        # segments:<cls>、前提及第一个 <sep> 标 0,假设及最后一个 <sep> 标 1。
        # 它用来区分两句话,不是"蕴含/矛盾/中立"的分类标签。
        tokens, segments = d2l.get_tokens_and_segments(p_tokens, h_tokens)

        # 把词元换成数字编号,不足 max_len 的位置补 <pad>。
        token_ids = self.vocab[tokens] + [self.vocab['<pad>']] \
                             * (self.max_len - len(tokens))
        # 片段编号也补齐长度,填充位置统一写 0。
        segments = segments + [0] * (self.max_len - len(segments))

        # 有效长度包含 <cls>/<sep>,但不包含 <pad>。
        # 编码器据此屏蔽注意力对填充位置的读取。
        valid_len = len(tokens)
        return token_ids, segments, valid_len

    def _truncate_pair_of_tokens(self, p_tokens, h_tokens):
        # 为一个 <cls> 和两个 <sep> 预留 3 个位置。
        # 每次从当前较长的句子末尾删一个词元,直到能放进 max_len。
        while len(p_tokens) + len(h_tokens) > self.max_len - 3:
            if len(p_tokens) > len(h_tokens):
                p_tokens.pop()
            else:
                h_tokens.pop()

    def __getitem__(self, idx):
        # DataLoader 通过这个方法取一个样本。
        # 返回 ((词元编号, 片段编号, 有效长度), 真实关系标签)。
        return (self.all_token_ids[idx], self.all_segments[idx],
                self.valid_lens[idx]), self.labels[idx]

    def __len__(self):
        return len(self.all_token_ids)  # 数据集样本数


# 三、创建训练集、测试集与批量读取器
# batch_size=512:每批最多处理 512 个句子对,显存不足时需要调小。
# max_len=128:本任务把每个句子对补齐/截断到 128 个位置。
# 这不矛盾:模型支持最长 512,本次实际只使用 128。
batch_size, max_len, num_workers = 512, 128, d2l.get_dataloader_workers()
data_dir = d2l.download_extract('SNLI')
train_set = SNLIBERTDataset(d2l.read_snli(data_dir, True), max_len, vocab)
test_set = SNLIBERTDataset(d2l.read_snli(data_dir, False), max_len, vocab)

# 训练时打乱样本顺序;测试时按顺序读取即可。
# 这里的 num_workers 管的是"批量读取",不是前面预处理的 Pool(4)。
train_iter = torch.utils.data.DataLoader(train_set, batch_size, shuffle=True,
                                   num_workers=num_workers)
test_iter = torch.utils.data.DataLoader(test_set, batch_size,
                                  num_workers=num_workers)


# 四、最关键的部分:给 BERT 接上"三分类器"
class BERTClassifier(nn.Module):
    def __init__(self, bert):
        super(BERTClassifier, self).__init__()
        self.encoder = bert.encoder  # 复用预训练 BERT 编码器,负责处理两句话
        self.hidden = bert.hidden    # 复用其隐藏变换层,教材实现为 Linear + Tanh
        self.output = nn.Linear(256, 3)  # 新建输出层:256 维向量 -> 3 个类别分数

    def forward(self, inputs):
        # B 表示当前批次的样本数,L 表示输入长度(这里是 128)。
        # tokens_X:[B, L],每个位置的词元编号。
        # segments_X:[B, L],每个位置属于前提(0)还是假设(1)。
        # valid_lens_x:[B],每个样本不含填充的长度。
        tokens_X, segments_X, valid_lens_x = inputs

        # 编码器给每个位置输出一个 256 维向量:[B, L, 256]。
        encoded_X = self.encoder(tokens_X, segments_X, valid_lens_x)

        # encoded_X[:, 0, :] 的三个索引分别表示:
        # 第一个 : -> 取批次里的所有样本。
        # 中间的 0 -> 只取第 0 个位置,也就是 <cls>。
        # 最后的 : -> 保留该位置完整的 256 个特征。
        # 因此取出的形状为 [B, 256]。
        # <cls> 能通过注意力聚合两句话的信息,训练让它适合做整体分类,
        # 但这不意味着它能无损保存两句话的所有信息。
        # 接着:hidden 变换 [B,256] -> output 得到 [B,3]。
        # 返回值是原始分数 logits,不是概率,这里不需要先做 softmax。
        return self.output(self.hidden(encoded_X[:, 0, :]))


net = BERTClassifier(bert)

# 五、微调:预训练部分继续学习,新输出层从零学习
# 这份代码没有冻结 encoder 或 hidden,因此它们也会更新。
# net.output 是随机初始化的新层,要学习怎样分辨三种关系。
# 纠正原说明中容易混淆的一点:这个 PyTorch 分类器仅引用
# bert.encoder 和 bert.hidden,没有把 bert.mlm / bert.nsp 挂到 net 上,
# 所以这些预训练任务专用头不在 net.parameters() 中,也不参加本次训练。
# 原说明中的 ignore_stale_grad 是 MXNet 相关说明,不适用于这里的 Adam 调用。
lr, num_epochs = 1e-4, 5  # 学习率 0.0001,遍历训练集 5 轮
trainer = torch.optim.Adam(net.parameters(), lr=lr)

# 三分类交叉熵:比较模型的三个分数与真实类别。
# reduction='none':先保留每个样本各自的损失,供训练函数汇总。
loss = nn.CrossEntropyLoss(reduction='none')

# 教材封装的训练循环:读取批次、前向计算、计算损失、反向传播、更新参数,
# 并用 test_iter 评估。训练函数内部完成这些步骤,不是"只调用一下就不用训练"。
d2l.train_ch13(net, train_iter, test_iter, loss, trainer, num_epochs,
    devices)

# 辅助核对的教材来源(原代码由用户提供):
# https://d2l.ai/chapter_natural-language-processing-applications/natural-language-inference-bert.html
# https://d2l.ai/chapter_natural-language-processing-pretraining/bert.html

微调时,BERT 本身也会继续学习,不是只训练最后的分类层。**其中 encoderhidden 复用预训练参数,output 是新建的三分类层

小结:

  • 我们可以针对下游应用对预训练的BERT模型进行微调,例如在SNLI数据集上进行自然语言推断。
  • 在微调过程中,BERT模型成为下游应用模型的一部分。仅与训练前损失相关的参数在微调期间不会更新。
相关推荐
Java后端的Ai之路1 小时前
20、Python - 备忘录模式
开发语言·人工智能·python·外观模式·备忘录模式
飞哥数智坊1 小时前
我对 AI 生图的一点工程化理解
人工智能·aigc
ACP广源盛139246256731 小时前
M6/M5 Pro Mac mini 端侧 AI 新形态@ACP#GSV5800 Serdes 长距离视频传输在 AI 服务中的机会与落地场景
大数据·网络·数据库·人工智能·嵌入式硬件·macos·音视频
xian_wwq1 小时前
【学习笔记】深度认知系列-第14讲 端侧AI崛起——为什么AI正在从云端走向本地
人工智能·笔记·学习
火山引擎开发者社区1 小时前
火山引擎 Milvus Vector Lakebase 正式公测
人工智能
ocean21031 小时前
2025-2026年AI部署与MLOps大厂面试高频问题
人工智能·面试·大模型推理·ai部署
嘿嘿-662 小时前
Windows 一键使用 GPT-6 Astra:Codex CLI 配置教程
java·人工智能·windows·gpt·chatgpt·web
我爱cope2 小时前
【计算机网络 | 传输层3:TCP 协议概述:面向连接、可靠传输到底意味着什么?】
网络·网络协议·学习·tcp/ip·计算机网络·传输层
冬奇Lab2 小时前
Code Agent 解剖(22):从零扩展——用 Markdown 写一个 Skill
人工智能