更新于 2026年7月25日

经过前面第1、2节内容的介绍,相信各位大家对于BERT的原理以及实现过程已经有了比较清晰的理解。同时,我们都知道BERT是一个强大的预训练模型,它可以基于谷歌发布的预训练参数在各个下游任务中进行微调。因此,在本节内容中我们将会介绍第一个下游微调场景,即如何在文本分类场景中基于BERT预训练模型进行微调。

4.1 任务构造原理#

总的来说,基于BERT的文本分类(准确的是单文本,也就是输入只包含一个句子)模型就是在原始的BERT模型后再加上一个分类层即可,类似的结构在(「基于Transformer 的分类模型」)中也介绍过,大家可以去看一下。同时,对于分类层的输入(也就是原始BERT的输出),默认情况下取BERT输出结果中[CLS]位置对于的向量即可,当然也可以修改为其它方式,例如所有位置向量的均值等(见「第2.4.3节内容」,将配置文件config.json中的pooler_type 字段设置为"all_token_average"即可)。 因此,对于基于BERT的文本分类模型来说其输入就是BERT的输入,输出则是每个类别对应的logits值。接下来,首先就来介绍如何构造文本分类的数据集。

4.2 数据预处理#

数据预处理完整代码见仓库 BertWithPretrained/utils/data_helper.py 文件。

4.2.1 输入介绍#

在构建数据集之前,我们首先需要知道的是模型到底应该接收什么样的输入,然后才能构建出正确的数据形式。在上面我们说到,基于BERT的文本分类模型的输入就等价于BERT模型的输入,同时根据第2节内容的介绍可以知道BERT模型的输入如图4-1所示。

图 4-1. BERT输入图
图 4-1. BERT输入图

由于对于文本分类这个场景来说其输入只有一个序列,所以在构建数据集的时候并不需要构造Segment Embedding的输入,直接默认使用全为0即可;同时,对于Position Embedding来说在任何场景下都不需要对其指定输入,因为我们在代码实现时已经做了相应默认时的处理。 关于这两部分内容的介绍可以参见「第2.2.4节 Transformer 位置编码教程:图解正弦 Positional Encoding 与编解码过程」内容。

因此,对于文本分类这个场景来说,只需要构造原始文本对应的 Token 序列,并在首尾分别再加上一个 [CLS] 符和 [SEP] 符作为输入即可。

4.2.2 语料介绍#

在这里,我们使用到的数据集是今日头条开放的一个新闻分类数据集[13],一共包含有382688条数据,15个类别。同时我们已近将其进行了格式化处理,以7:2:1的比例划分成了训练集、验证集和测试集3个部分。如下所示便是部分示例数据:

1 千万不要乱申请网贷否则后果很严重_!_4
2 10年前的今天纪念5.12汶川大地震10周年_!_11
3 怎么看待杨毅在一NBA直播比赛中说詹姆斯的球场统治力已经超过乔丹伯德和科比_!_3
4 戴安娜王妃的车祸有什么谜团_!_2

其中_!_左边为新闻标题,即后续需要用到的分类文本,右边为类别标签。

4.2.3 数据集预览#

同样,在正式介绍如何构建数据集之前我们先通过一张图来了解一下整个构建的流程,以便做到心中有数,不会迷路。假如我们现在有两个样本构成了一个batch,那么其整个数据的处理过程则如图4-2所示。

图4-2. 文本分类数据集构建流程图
图4-2. 文本分类数据集构建流程图

如图4-2所示,第1步需要将原始的数据样本进行分字(tokenize)处理;第2步再根据 tokenize 后的结果构造一个字典,不过在使用BERT预训练时并不需要我们自己来构造这个字典,直接载入谷歌开源的 vocab.txt 文件构造字典即可,因为只有 vocab.txt 中每个字的索引顺序才与开源模型中每个字的 Embedding 向量一一对应的。第3步则是根据字典将 tokenize 后的文本序列转换为 Token 序列,同时在 Token 序列的首尾分别加上 [CLS] 和 [SEP] 符号,并进行Padding。第4步则是根据第3步处理后的结果生成对应的Padding Mask向量。

最后,在模型训练时只需要将第3步和第4步处理后的结果一起喂给模型即可。

4.2.4 数据集构建#

第1步:定义tokenize

第1步需要完成的就是将输入进来的文本序列 tokenize 到字符级别。对于中文语料来说就是将每个字和标点符号都给切分开。在这里,我们可以借用 transformers 包中的 BertTokenizer 方法来完成,如下所示:

1 if __name__ == '__main__':
2     model_config = ModelConfig()
3     tokenizer = BertTokenizer.from_pretrained(model_config. pretrained_model_dir).tokenize
4     print(tokenizer("青山不改,绿水长流,我们月来客栈见!"))
5     print(tokenizer("10年前的今天,纪念5.12汶川大地震10周年"))
6 
7 # ['青', '山', '不', '改', ',', '绿', '水', '长', '流', ',', '我', '们', '月', '来', '客', '栈', '见', '!']
8 # ['10', '年', '前', '的', '今', '天', ',', '纪', '念', '5', '.', '12', '汶', '川', '大',        '地', '震', '10', '周', '年']

在上述代码中,第2~3行就是根据指定的路径(BERT预训练模型的路径)来载入一个分字模型;第7~8行便是tokenize后的结果。 第2步:建立词表 由于BERT预训练模型中已经有了一个给定的词表(vocab.txt),因此我们并不需要根据自己的语料来建立一个词表。当然,也不能够根据自己的语料来建立词表,因为相同的字在我们自己构建的词表中和vocab.txt中的索引顺序肯定会不一样,而这就会导致后面根据token id 取出来的向量是错误的。 进一步,我们只需要将vocab.txt中的内容读取进来形成一个词表即可,代码如下:

1 class Vocab:
 2     UNK = '[UNK]'
 3     def __init__(self, vocab_path):
 4         self.stoi = {}
 5         self.itos = []
 6         with open(vocab_path, 'r', encoding='utf-8') as f:
 7             for i, word in enumerate(f):
 8                 w = word.strip('\n')
 9                 self.stoi[w] = i
10                 self.itos.append(w)
11 
12     def __getitem__(self, token):
13         return self.stoi.get(token, self.stoi.get(Vocab.UNK))
14 
15     def __len__(self):
16         return len(self.itos)

接着便可以定义一个方法来实例化一个词表:

1 def build_vocab(vocab_path):
2     return Vocab(vocab_path)
3 
4 if __name__ == '__main__':
5     vocab = build_vocab()

在经过上述代码处理后,我们便能够通过vocab.itos得到一个列表,返回词表中的每一个词;通过vocab.itos[2]返回得到词表中对应索引位置上的词;通过vocab.stoi得到一个字典,返回词表中每个词的索引;通过vocab.stoi[‘月’]返回得到词表中对应词的索引;通过len(vocab)来返回词表的长度。如下便是建立后的词表:

1 {'[PAD]': 0, '[unused1]': 1, '[unused2]': 2, '[unused3]': 3, '[unused4]': 4, 
2 '[unused5]': 5,'[unused6]': 6, '[unused7]': 7, '[unused8]': 8,'[unused9]': 9, 
3 '[unused10]': 10, '[unused11]': 11, '[unused12]': 12, '[unused13]': 13,  ...
4 '[unused42]': 42, ....'乐': 727, '乒': 728, '乓': 729, '乔': 730, '乖': 731,
5  '乗': 732, '乘': 733, '乙': 734, '乜': 735, '九': 736, '乞': 737, '也': 738, 
6  '习':739,'乡': 740,'书': 741,'乩': 742,'买': 743, '乱': 744,'乳': 745, ....}

此时,我们就需要定义一个类,并在类的初始化过程中根据训练语料完成字典的构建等工作,代码如下:

 1 class LoadSingleSentenceClassificationDataset:
 2     def __init__(self,
 3                  vocab_path='./vocab.txt',  #
 4                  tokenizer=None,
 5                  batch_size=32,
 6                  max_sen_len=None,
 7                  split_sep='\n',
 8                  max_position_embeddings=512,
 9                  pad_index=0,
10                  is_sample_shuffle=True):
11 
12         self.tokenizer = tokenizer
13         self.vocab = build_vocab(vocab_path)
14         self.PAD_IDX = pad_index
15         self.SEP_IDX = self.vocab['[SEP]']
16         self.CLS_IDX = self.vocab['[CLS]'] 
17         self.batch_size = batch_size
28         self.split_sep = split_sep
19         self.max_position_embeddings = max_position_embeddings
20         if isinstance(max_sen_len, int) and max_sen_len > max_position_embeddings:
23             max_sen_len = max_position_embeddings
24         self.max_sen_len = max_sen_len
25         self.is_sample_shuffle = is_sample_shuffle

在上述代码中,第3行 vocab_path 表示本地词表的路径。第6行 max_sen_len 表示最大样本长度,当 max_sen_len = None 时,即以每个batch中最长样本长度为标准,对其它进行padding;当 max_sen_len = 'same' 时,以整个数据集中最长样本为标准,对其它进行padding;当 max_sen_len = 50, 表示以某个固定长度符样本进行padding,多余的截掉。第7行split_sep表示样本与标签之间的分隔符。is_sample_shuffle 表示是否打乱数据集。第14~16行为建立词表并取对应特殊字符的索引。第19行中 max_position_embeddings 为最大样本长度,最大为512。第20~23行则是用来判断传入的最大样本长度。

第3步:转换为Token序列

在得到构建的字典后,便可以通过如下方法来分别将训练集、验证集和测试集转换成Token序列,其作用是将每一句话中的每一个词根据字典转换成索引的形式,同时返回所有样本中最长样本的长度。实现代码如下:

 1     def data_process(self, filepath):
 2         raw_iter = open(filepath, encoding="utf8").readlines()
 3         data = []
 4         max_len = 0
 5         for raw in tqdm(raw_iter, ncols=80):
 6             line = raw.rstrip("\n").split(self.split_sep)
 7             s, l = line[0], line[1]
 8             tmp = [self.CLS_IDX] + [self.vocab[token] for token in self.tokenizer(s)]
 9             if len(tmp) > self.max_position_embeddings - 1:
10                 tmp = tmp[:self.max_position_embeddings - 1]  
11             tmp += [self.SEP_IDX]
12             tensor_ = torch.tensor(tmp, dtype=torch.long)
13             l = torch.tensor(int(l), dtype=torch.long)
14             max_len = max(max_len, tensor_.size(0))
15             data.append((tensor_, l))
16         return data, max_len

在上述代码中,第6~7行便是用来取得文本和标签;第8行则是首先对序列进行tokenize,然后转换成Token序列并在最前面加上分类标志位[CLS]。第9~11行则是用来对Token序列进行截取,最长为 max_position_embeddings 个字符,默认为512,并同时在末尾加上[SEP]符号。不过我们认为其实末尾不加[SEP]应该也不会有影响,因为这本来是单个序列的分类。第14行则是用来保存最长序列的长度。

在处理完成后,4.2.2 节中的4个样本将会被转换成类似如下形式:

1 tensor([[101, 1283,  674,  679, 6206,  744, 4509, 6435, 5381,..,  0,    0],
2         [101, 8108, 2399, 1184, 4638,  791, 2399, 8024, 5279,..,  0,    0],
3         [101, 2582,  720, 4692, 2521, 3342, 3675, 1762,  671,..,8043, 102],
4         [101, 2785, 2128, 2025, 4374, 1964, 4638, 6756, 1730,..,  0,    0]])
5 torch.Size([39, 4])

从上面的输出结果可以看出,101就是[CLS]在词表中的索引位置,102则是[SEP]在词表中的索引;其它非0值就是tokenize后的文本序列转换成的Token序列。同时可以看出,这里的结果是以第3个样本的长度39对其它样本进行padding的,并且padding的Token ID为0。因此,下面我们就来介绍样本的padding处理。

第4步:padding处理与mask

从第3步的输出结果看出,在对原始文本序列tokenize转换为Token ID后还需要对其进行padding处理。对于这一处理过程可以通过如下代码来完成:

 1 def pad_sequence(sequences,batch_first=False,max_len=None, padding_value=0):
 2     if max_len is None:
 3         max_len = max([s.size(0) for s in sequences])
 4     out_tensors = []
 5     for tensor in sequences:
 6         if tensor.size(0) < max_len:
 7             tensor = torch.cat([tensor, torch.tensor([padding_value]  (max_len - tensor.size(0)))], dim=0)
 8         else:
 9             tensor = tensor[:max_len]
10         out_tensors.append(tensor)
11     out_tensors = torch.stack(out_tensors, dim=1)
12     if batch_first:
13         return out_tensors.transpose(0, 1)
14     return out_tensors

在上述代码中,第1行sequences为待padding的序列所构成的列表,其中的每一个元素为一个样本的Token序列;batch_first表示是否将batch_size这个维度放在第1个;max_len表示指定最大序列长度,当max_len = 50时,表示以某个固定长度对样本进行padding多余的截掉,当max_len=None时表示以当前batch中最长样本的长度对其它进行padding。第2~3行用来获取padding的长度;第5~11行则是遍历每一个Token序列,根据max_len来进行padding。第12~13行是将batch_size这个维度放到最前面。

如下为使用示例:

1 if __name__ == '__main__':
2     a = torch.tensor([1, 2, 3])
3     b = torch.tensor([4, 5, 6, 7, 8])
4     c = torch.tensor([9, 10])
5     d = pad_sequence([a, b, c], max_len=None).size()
6         # torch.Size([5, 3])

进一步,我们需要定义一个方法来对每个batch的Token序列进行padding处理:

 1     def generate_batch(self, data_batch):
 2         batch_sentence, batch_label = [], []
 3         for (sen, label) in data_batch: 
 4             batch_sentence.append(sen)
 5             batch_label.append(label)
 6         batch_sentence = pad_sequence(batch_sentence,
 7                                       padding_value=self.PAD_IDX,
 8                                       batch_first=False,
 9                                       max_len=self.max_sen_len)
10         batch_label = torch.tensor(batch_label, dtype=torch.long)
11         return batch_sentence, batch_label

上述代码的作用就是对每个batch的Token序列进行padding处理。

最后,对于每一序列的attention_mask向量,我们只需要判断其是否等于padding_value便可以得到这一结果,可见第5步中的使用示例。

第5步:构造DataLoade与使用示例

经过前面4步的操作,整个数据集的构建就算是已经基本完成了,只需要再构造一个DataLoader迭代器即可,代码如下:

 1     def load_train_val_test_data(self, train_file_path=None,
 2                                  val_file_path=None,
 3                                  test_file_path=None,
 4                                  only_test=False):
 5         test_data, _ = self.data_process(test_file_path)
 6         test_iter = DataLoader(test_data, batch_size=self.batch_size,
 7                                shuffle=False, collate_fn=self.generate_batch)
 8         if only_test:
 9             return test_iter
10         train_data, max_sen_len = self.data_process(train_file_path) 
11         if self.max_sen_len == 'same':
12             self.max_sen_len = max_sen_len
13         val_data, _ = self.data_process(val_file_path)
14         train_iter = DataLoader(train_data, batch_size=self.batch_size,
15                                 shuffle=self.is_sample_shuffle, collate_fn=self.generate_batch)
16         val_iter = DataLoader(val_data, batch_size=self.batch_size,
17                               shuffle=False, collate_fn=self.generate_batch)
18         return train_iter, test_iter, val_iter

在上述代码中,第5~7行用来得到预处理后的数据并构造对应的DataLoader,其中 generate_batch 将作为一个参数传入来对每个batch的样本进行处理;第8~9行则判断是否只返回测试集;同理,第10~18行则是用来构造相应的训练集和验证集。在完成类 LoadSingleSentenceClassificationDataset 所有的编码过程后,便可以通过如下形式进行使用:

 1 from Tasks.TaskForSingleSentenceClassification import ModelConfig
 2 from utils.data_helpers import LoadSingleSentenceClassificationDataset
 3 from transformers import BertTokenizer
 4 
 5 if __name__ == '__main__':
 6     model_config = ModelConfig()
 7     tokenizer= BertTokenizer.from_pretrained(model_config.pretrained_model_dir).tokenize
 8     load_dataset = LoadSingleSentenceClassificationDataset(
 9         vocab_path=model_config.vocab_path,
10         tokenizer=tokenizer,
11         batch_size=model_config.batch_size,
12         max_sen_len=model_config.max_sen_len,
13         split_sep=model_config.split_sep,
14         max_position_embeddings=model_config.max_position_embeddings,
15         pad_index=model_config.pad_token_id,
16         is_sample_shuffle=model_config.is_sample_shuffle)
17 
18     train_iter, test_iter, val_iter = \
19         load_dataset.load_train_val_test_data(model_config.train_file_path,
20                                               model_config.val_file_path,
21                                               model_config.test_file_path)
22     for sample, label in train_iter:
23         print(sample.shape)  # [seq_len,batch_size]
24         print(sample.transpose(0, 1))
25         padding_mask = (sample == load_dataset.PAD_IDX).transpose(0, 1)
26         print(padding_mask)
27         print(label)
28         break

在上述代码中,第6行是载入配置参数。第7行是实例化一个tokenizer。第8~16行是实例化数据集载入对象。第18~21行是返回训练集、测试机和验证集的迭代器;第22~27行是遍历每个batch的样本。

执行完上述代码后便可以得到如下所示的结果:

1 torch.Size([39, 4])
2 tensor([[ 101, 1283,  674,  679, 6206, ... 7028,  102, ...  ,0,    0,   0],
3          ...
4         ])
5 tensor([[False, False, False,False,False, ..., False,... True, True,  True],
6         ...
7         ])
8 tensor([ 4,...])

到此,对于整个数据集构建部分的内容就算是介绍完了,接下来我们再来看如何加载预训练模型进行微调。

4.3 加载预训练模型#

在介绍模型微调之前,我们先来看看当我们拿到一个开源的模型参数后怎么读取以及分析。下面就以huggingface开源的PyTorch训练的bert-base-chinese模型参数[10]为例进行介绍。

4.3.1 查看模型参数#

在第3节内容中,尽管已经大致介绍了如何通过PyTorch来读取和加载模型参数,但是这里仍旧有必要以bert-base-chinese参数为例再进行一次详说明。根据第3.2节内容的介绍可知,我们可以通过如下方式来查看本地模型中的参数情况:

1 import torch
2 loaded_paras = torch.load('./pytorch_model.bin')
3 print(type(loaded_paras))
4 print(len(list(loaded_paras.keys())))
5 print(list(loaded_paras.keys()))

执行完上述代码后,便可以得到如下输出结果:

1 <class 'collections.OrderedDict'>
2 207
3 ['bert.embeddings.word_embeddings.weight', 
4 'bert.embeddings.position_embeddings.weight', 
5 'bert.embeddings.token_type_embeddings.weight', 
6 .....
7 'bert.encoder.layer.11.output.dense.bias', 
8 'bert.encoder.layer.11.output.LayerNorm.gamma', 
9 'bert.encoder.layer.11.output.LayerNorm.beta',.....]

从上面的输出结果可以看到,参数pytorch_model.bin被载入后变成了一个有序的字典OrderedDict,并且其中一共有207个参数,其名字分别就是列表中的各个元素。进一步,我们还可以将各个参数的形状打印出来看一看:

 1 for name in loaded_paras.keys():
 2     print(f"### 参数:{name},形状:{loaded_paras[name].size()}")
 3 
 4 # 参数:bert.embeddings.word_embeddings.weight,形状:torch.Size([21128,768])
 5 # 参数:bert.embeddings.position_embeddings.weight,形状:torch.Size([512, 768])
 6 # 参数:bert.embeddings.token_type_embeddings.weight,形状:torch.Size([2, 768])
 7 # 参数:bert.embeddings.LayerNorm.gamma,形状:torch.Size([768])
 8 ......
 9 # 参数:bert.encoder.layer.11.output.dense.weight,形状:torch.Size([768, 3072])
10 # 参数:bert.encoder.layer.11.output.dense.bias,形状:torch.Size([768])
11 # 参数:bert.encoder.layer.11.output.LayerNorm.gamma,形状:torch.Size([768])
12 # 参数:bert.encoder.layer.11.output.LayerNorm.beta,形状:torch.Size([768])
13 # 参数:bert.pooler.dense.weight,形状:torch.Size([768, 768])
14 # 参数:bert.pooler.dense.bias,形状:torch.Size([768])
15 # 参数:cls.predictions.bias,形状:torch.Size([21128])
16 # 参数:cls.predictions.transform.dense.weight,形状:torch.Size([768, 768])
17 # 参数:cls.predictions.transform.dense.bias,形状:torch.Size([768])
18 # 参数:cls.predictions.transform.LayerNorm.gamma,形状:torch.Size([768])
19 # 参数:cls.predictions.transform.LayerNorm.beta,形状:torch.Size([768])
20 # 参数:cls.predictions.decoder.weight,形状:torch.Size([21128, 768])
21 # 参数:cls.seq_relationship.weight,形状:torch.Size([2, 768])
22 # 参数:cls.seq_relationship.bias,形状:torch.Size([2])

同样,我们还可以直接打印出某个参数具体的值。不过这里并不需要分析所以就不用打印。

到此,对于本地的模型参数的信息就分析完了。不过想要将它迁移到自己所搭建的模型上还要进一步的来分析自己所搭建网络的参数信息。

4.3.2 载入并初始化#

在第2节内容中,我们已经详细地介绍了如何实现整个BERT模型,但是对于如何载入已有参数来初始化网络中的参数还并未介绍。在将本地参数迁移到一个新的模型之前,除了像上面那样分析本地参数之外,我们还需要将网络的参数信息也打印出来看一下,以便将两者一一对应上。

1 json_file = '../bert_base_chinese/config.json'
2 config = BertConfig.from_json_file(json_file)
3 bert_model = BertModel(config)
4 print("\n  =======  BertMolde 参数: ========")
5 print(len(bert_model.state_dict()))
6 for param_tensor in bert_model.state_dict():
7     print(param_tensor, "\t", bert_model.state_dict()[param_tensor].size())

在执行完上述代码后,便可以得到如下输出结果:

 1  =======  BertMolde 参数: ========
 2 200
 3 # bert_embeddings.position_ids  torch.Size([1, 512])
 4 # bert_embeddings.word_embeddings.embedding.weight torch.Size([21128,768])
 5 # bert_embeddings.position_embeddings.embedding.weight torch.Size([512,768])
 6 # bert_embeddings.token_type_embeddings.embedding.weight torch.Size([2,768])
 7 ......
 8 # bert_encoder.bert_layers.11.bert_output.dense.bias   torch.Size([768])
 9 # bert_encoder.bert_layers.11.bert_output.LayerNorm.weight torch.Size([768])
10 # bert_encoder.bert_layers.11.bert_output.LayerNorm.bias   torch.Size([768])
11 # bert_pooler.dense.weight     torch.Size([768, 768])
12 # bert_pooler.dense.bias   torch.Size([768])

从上面的输出结果可以发现,BertMolde一共有200个参数,而bert-base-chinese一共有207个参数。这里需要注意的是BertMolde模型中的position_ids这个参数并不是模型中需要训练的参数,只是一个默认的初始值。最后,经分析(两者一一进行对比)后发现bert-base-chinese中除了最后的8个参数以外,其余的199个参数和BertMolde模型中的199个参数一样且顺序也一样。

因此,最后我们可以通过在BertMolde类(Bert.py文件中)中再加入一个如下所示的方法来用 bert-base-chinese 中的参数初始化BertMolde中的参数:

 1     @classmethod
 2     def from_pretrained(cls, config, pretrained_model_dir=None):
 3         model = cls(config)  
 4         pretrained_model_path = os.path.join(pretrained_model_dir, "pytorch_model.bin")
 5         loaded_paras = torch.load(pretrained_model_path)
 6         state_dict = deepcopy(model.state_dict())
 7         loaded_paras_names = list(loaded_paras.keys())[:-8]
 8         model_paras_names = list(state_dict.keys())[1:]
 9         for i in range(len(loaded_paras_names)):
10             state_dict[model_paras_names[i]] = loaded_paras[loaded_paras_names[i]]
11             logging.info(f"成功将参数{loaded_paras_names[i]}赋值给{model_paras_names[i]}")
12         model.load_state_dict(state_dict)
13         return model

在上述代码中,第3行是初始化模型,cls为未实例化的对象,即一个未实例化的 BertModel 对象;第4~5行用来载入本地的 bert-base-chinese参数;第6行用来拷贝一份BertModel中的网络参数,这是因为我们无法直接修改里面的值;第7~10行则是根据我们上面的分析,将bert-base-chinese中的参数赋值到state_dict中;第12行是用state_dict中的参数来初始化BertModel中的参数。

最后,我们只需要通过如下方式便可以返回一个通过 bert-base-chinese 初始化的 BERT 模型:

1 bert = BertModel.from_pretrained(config, bert_pretrained_model_dir)

当然,如果你需要冻结其中某些层的参数不参与模型训练,那么可以通过类似如下所示的代码来进行设置:

1 for para in bert_model.parameters():
2     if xxxx:
3         para.requires_grad = False

到此,对于整个预训练模型的加载过程就介绍完了,接下来让我们正式进入到基于BERT预训练模型的文本分类场景中。

4.4 文本分类#

4.4.1 前向传播#

在介绍完如何分析和载入本地BERT预训练模型后,接下来我们首先要做的就是实现文本分类的前向传播过程。如图2-4所示,在BertForSentenceClassification.py文件中,我们通过定义如下一个类来完成整个前向传播的过程:

 1 from ..BasicBert.Bert import BertModel
 2 import torch.nn as nn
 3 class BertForSentenceClassification(nn.Module):
 4     def __init__(self, config, pretrained_model_dir=None):
 5         super(BertForSentenceClassification, self).__init__()
 6         self.num_labels = config.num_labels
 7         if bert_pretrained_model_dir is not None:
 8             self.bert=BertModel.from_pretrained(config,pretrained_model_dir)
 9         else:
10             self.bert = BertModel (config)
11         self.dropout = nn.Dropout(config.hidden_dropout_prob)
12         self.classifier = nn.Linear(config.hidden_size, self.num_labels)

在上述代码中,第4行代码分别就是用来指定模型配置和预训练模型的路径;第7~10行代码则是用来定义一个BERT模型,可以看到如果预训练模型的路径存在则会返回一个由bert-base-chinese参数初始化后的BERT模型,否则则会返回一个随机初始化参数的BERT模型;第12行则是定义最后的分类层。

最后,整个前向传播的实现代码如下所示:

 1     def forward(self, input_ids, # [src_len, batch_size]
 2                 attention_mask=None, # [batch_size, src_len]
 3                 token_type_ids=None, # [src_len, batch_size] 
 4                 position_ids=None, # [1,src_len]
 5                 labels=None): # [batch_size,]
 6         pooled_output, _ = self.bert(input_ids=input_ids,
 7                                      attention_mask=attention_mask,
 8                                      token_type_ids=token_type_ids,
 9                                      position_ids=position_ids)  
10         pooled_output = self.dropout(pooled_output)
11         logits = self.classifier(pooled_output)  # [batch_size, num_label]
12         if labels is not None:
13             loss_fct = nn.CrossEntropyLoss()
14             loss = loss_fct(logits.view(-1,self.num_labels),labels.view(-1))
15             return loss, logits
16         else:
17             return logits

在上述代码中,第6~9行返回的就是原始BERT网络的输出,其中pooled_output为BERT第1个位置的向量经过一个全连接层后的结果,第2个参数是BERT中所有位置的向量;第10~11行便是用来进行文本分类的分类层;第12~17行则是用来判断返回损失值还是返回logits值。

4.4.2 模型训练#

如图2-5所示,我们将在Tasks目录下新建一个名为TaskForSingleSentenceClassification.py的模块来完成分类模型的微调训练任务。

首先,我们需要在其中定义一个ModelConfig类来对分类模型中的超参数进行管理,代码如下所示:

 1 class ModelConfig:
 2     def __init__(self):
 3         self.project_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
 4         self.dataset_dir = os.path.join(self.project_dir, 'data', 'SingleSentenceClassification')
 5         self.pretrained_model_dir = os.path.join(self.project_dir, "bert_base_chinese")
 6         self.vocab_path=os.path.join(self.pretrained_model_dir, 'vocab.txt')
 7         self.device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')
 8         self.train_file_path = os.path.join(self.dataset_dir,'toutiao_train.txt')
 9         self.val_file_path = os.path.join(self.dataset_dir, 'toutiao_val.txt')
10         self.test_file_path = os.path.join(self.dataset_dir, 'toutiao_test.txt')
11         self.model_save_dir = os.path.join(self.project_dir, 'cache')
12         self.logs_save_dir = os.path.join(self.project_dir, 'logs')
13         self.split_sep = '_!_'
14         self.is_sample_shuffle = True
15         self.batch_size = 64
16         self.max_sen_len = None
17         self.num_labels = 15
18         self.epochs = 10
19         self.model_val_per_epoch = 2
20         logger_init(log_file_name='single', log_level=logging.INFO,
21                     log_dir=self.logs_save_dir)
22         if not os.path.exists(self.model_save_dir):
23             os.makedirs(self.model_save_dir)
24 
25         # 把原始bert中的配置参数也导入进来
26         bert_config_path = os.path.join(self.pretrained_model_dir, "config.json")
27         bert_config = BertConfig.from_json_file(bert_config_path)
28         for key, value in bert_config.__dict__.items():
29             self.__dict__[key] = value
30         logging.info(" ### 将当前配置打印到日志文件中 ")
31         for key, value in self.__dict__.items():
32             logging.info(f"###  {key} = {value}")

在上述代码中,第2~23行则是分别用来定义模型中的一些数据集目录、超参数和初始化日志打印类等;第25~29行则是将原始bert_base_chinese配置文件,即config.json中的参数也导入到类ModelConfig中;第31~32行则是将所有的超参数配置情况一同打印到日志文件中方便后续分析,更多关于日志的内容可以参加文章[14]。

最后,我们只需要再定义一个train()函数来完成模型的训练即可,代码如下:

 1 def train(config):
 2     model = BertForSentenceClassification(config,
 3                                           config.pretrained_model_dir)
 4     #......
 5     optimizer = torch.optim.Adam(model.parameters(), lr=5e-5)
 6     model.train()
 7     tokenizer = BertTokenizer.from_pretrained(config.pretrained_model_dir)
 8     data_loader = LoadSingleSentenceClassificationDataset(
 9                     vocab_path=config.vocab_path,
10                     tokenizer=tokenizer.tokenize,
11                     batch_size=config.batch_size,
12                     max_sen_len=config.max_sen_len,
13                     split_sep=config.split_sep,
14                     max_position_embeddings=config.max_position_embeddings,
15                     pad_index=config.pad_token_id)
16     train_iter, test_iter, val_iter = data_loader.load_train_val_test_data(
17         config.train_file_path,config.val_file_path,config.test_file_path)
18     max_acc = 0
19     for epoch in range(config.epochs):
20         losses = 0
21         start_time = time.time()
22         for idx, (sample, label) in enumerate(train_iter):
23             sample = sample.to(config.device)  # [src_len, batch_size]
24             label = label.to(config.device)
25             padding_mask = (sample == data_loader.PAD_IDX).transpose(0, 1)
26             loss, logits = model(input_ids=sample,
27                                  attention_mask=padding_mask,
28                                  token_type_ids=None,
29                                  position_ids=None,
30                                  labels=label)
31                         #.........
32             acc = (logits.argmax(1) == label).float().mean()
33             #.........
34         if (epoch + 1) % config.model_save_per_epoch == 0:
35             acc = evaluate(val_iter, model, config.device)
36             logging.info(f"Accuracy on val {acc:.3f}")
37             #.........

在上述代码中,第2~3行用来初始化一个基于BERT的文本分类模型;第8~17行则是载入相应的数据集;第19~36行则是整个模型的训练过程,完整示例代码可参见[6],我们也在代码中进行了详细的注释。

紧接着,可以通过如下方式来进行模型训练:

1 if __name__ == '__main__':
2     model_config = ModelConfig()
3     train(model_config)

如下便是网络的部分训练结果:

 1 -- INFO: Epoch: 0, Batch[0/4186], Train loss :2.862, Train acc: 0.125
 2 -- INFO: Epoch: 0, Batch[10/4186], Train loss :2.084, Train acc: 0.562
 3 -- INFO: Epoch: 0, Batch[20/4186], Train loss :1.136, Train acc: 0.812        
 4 -- INFO: Epoch: 0, Batch[30/4186], Train loss :1.000, Train acc: 0.734
 5 ...
 6 -- INFO: Epoch: 0, Batch[4180/4186], Train loss :0.418, Train acc: 0.875
 7 -- INFO: Epoch: 0, Train loss: 0.481, Epoch time = 1123.244s
 8 ...
 9 -- INFO: Epoch: 9, Batch[4180/4186], Train loss :0.102, Train acc: 0.984
10 -- INFO: Epoch: 9, Train loss: 0.100, Epoch time = 1130.071s
11 -- INFO: Accurcay on val 0.884
12 -- INFO: Accurcay on test 0.888

4.4.3 模型推理#

在完成模型的训练过程后,便可以将训练过程中保存好的模型用于任务的推理场景中。这部分代码实现起来也比较容易,只需要按照第3.3节内容中介绍的方式使用即可,具体代码实现如下:

 1 def inference(config):
 2     model = BertForSentenceClassification(config,
 3                                           config.pretrained_model_dir)
 4     model_save_path = os.path.join(config.model_save_dir, 'model.pt')
 5     if os.path.exists(model_save_path):
 6         loaded_paras = torch.load(model_save_path)
 7         model.load_state_dict(loaded_paras)
 8         logging.info("## 成功载入已有模型,进行预测......")
 9     model = model.to(config.device)
10     tokenizer = BertTokenizer.from_pretrained(config.pretrained_model_dir)
11     data_loader = LoadSingleSentenceClassificationDataset(
12                   vocab_path=config.vocab_path,
13                   tokenizer=tokenizer.tokenize,
14                   batch_size=config.batch_size,
15                   max_sen_len=config.max_sen_len,
16                   split_sep=config.split_sep,
17                   max_position_embeddings=config.max_position_embeddings,
18                   pad_index=config.pad_token_id,
19                   is_sample_shuffle=config.is_sample_shuffle)
20     train_iter, test_iter, val_iter = data_loader.load_train_val_test_data(
21           config.train_file_path,config.val_file_path,config.test_file_path)
22     acc = evaluate(test_iter,model,config.device, data_loader.PAD_IDX)
23     logging.info(f"Acc on test:{acc:.3f}")

在上述代码中,第2~3行是用来实例化一个文本分类模型;第4~8行是载入本地模型参数来重新初始化网络中的参数;第11~21行是载入新的测试数据;第22行是返回模型在测试集上的准确率,当然在实际情况中也可以按照需求进行改动。

同时,测试集上准确率的实现过程如下所示:

 1 def evaluate(data_iter, model, device, PAD_IDX):
 2     model.eval()
 3     with torch.no_grad():
 4         acc_sum, n = 0.0, 0
 5         for x, y in data_iter:
 6             x, y = x.to(device), y.to(device)
 7             padding_mask = (x == PAD_IDX).transpose(0, 1)
 8             logits = model(x, attention_mask=padding_mask)
 9             acc_sum += (logits.argmax(1) == y).float().sum().item()
10             n += len(y)
11         model.train()
12         return acc_sum / n

紧接着,可以通过如下方式来进行模型的推理过程:

1 if __name__ == '__main__':
2     model_config = ModelConfig()
3     inference(model_config)

到此,对于简单的单文本分类任务就介绍完了。在下一节内容中,我们将会介绍如何在文本蕴含任务中,即输入两个句子来进行分类的场景下,进行BERT预训练模型的微调。

您当前阅读的系列内容有完整版高清 PDF ,点击右侧了解

近200页高清 PDF + 80页配套PPT、高清无水印网络结构图,打印学习更方便!

查看详情
阅读 --

第1节 多头注意力机制原理

本文是《Attention Is All You Need》精读系列第1讲,从 RNN 串行瓶颈出发讲清 Transformer 为什么必须引入自注意力,再用图解与公式拆解 Scaled Dot-Product …

第1节 BERT原理与预训练任务

本文是 BERT 精读系列第1讲,从 BERT 与 Transformer Encoder 的关系出发,系统讲清 BERT 论文中那张易混淆的结构图背后的真实网络组成 …

第2节 BERT 从零实现过程

本文是 BERT 精读系列第2讲,配套开源仓库 moon-hotel/BertWithPretrained 系统讲解如何用 PyTorch 从零实现 BERT 模型。覆盖工程目录结构(预训练模型、cache、数据集 …

第3节 模型的保存与迁移

本文是 BERT 精读系列第3讲,作为 BERT 下游任务微调前的过渡篇,系统讲清 PyTorch 中模型保存与加载的三大运用场景(模型推理、再训练、迁移学习)。内容覆盖 nn.Module 参数字典 state_dict 的结构与遍历 …