1 Bert模型

先看看Bert做了什么:

  1. 静态词嵌入(如word2vec和GloVe)
    • 特点:预训练的向量是固定的,同一个词在不同上下文中使用相同的向量。
    • 局限性:无法处理一词多义或复杂语义,因为它们不考虑词的上下文信息。
  2. 上下文敏感的词表示
    • ELMo
      • 特点:使用双向编码,即考虑词的左侧和右侧上下文。
      • 局限性:需要为每个自然语言处理任务设计特定的架构,这实际上并不容易。
    • GPT(生成预训练)
      • 特点:任务无关,从左到右编码上下文,即只考虑词的左侧上下文。
      • 局限性:不是双向编码,可能无法充分利用右侧的上下文信息。
  3. BERT(双向编码表示模型)
    • 优点:结合了ELMo和GPT的优点,进行双向编码,并且对大量自然语言处理任务需要最小的架构更改。
    • 输入序列的嵌入:由词元嵌入、片段嵌入和位置嵌入组成。
    • 预训练任务
      • 掩蔽语言模型:随机掩盖输入序列中的一些词元,模型需要预测这些被掩盖的词元,从而编码双向上下文来表示单词。
      • 下一句预测:给定两个句子,模型需要预测第二个句子是否是第一个句子的下一句,从而显式地建模文本对之间的逻辑关系。

1.1 编码器

1.1.1 编码器子模块

即下面使用的EncoderBlock,其实就是Transformer的编码器子模块,具体可以去看详情

1.1.2 编码器

class BERTEncoder(nn.Module):
    """BERT编码器"""
    def __init__(self, vocab_size, num_hiddens, norm_shape, ffn_num_input,
                 ffn_num_hiddens, num_heads, num_layers, dropout,
                 max_len=1000, key_size=768, query_size=768, value_size=768,
                 **kwargs):
        super(BERTEncoder, self).__init__(**kwargs)
        self.token_embedding = nn.Embedding(vocab_size, num_hiddens)
        self.segment_embedding = nn.Embedding(2, num_hiddens)
        self.blks = nn.Sequential()
        for i in range(num_layers):
            self.blks.add_module(f"{i}", EncoderBlock(
                key_size, query_size, value_size, num_hiddens, norm_shape,
                ffn_num_input, ffn_num_hiddens, num_heads, dropout, True))
        # 在BERT中,位置嵌入是可学习的,因此我们创建一个足够长的位置嵌入参数
        self.pos_embedding = nn.Parameter(torch.randn(1, max_len, num_hiddens))

    def forward(self, tokens, segments, valid_lens):
        # 在以下代码段中,X的形状保持不变:(批量大小,最大序列长度,num_hiddens)
        X = self.token_embedding(tokens) + self.segment_embedding(segments)
        X = X + self.pos_embedding.data[:, :X.shape[1], :]
        for blk in self.blks:
            X = blk(X, valid_lens)
        return X

注意这里的编码器输入和Transformer不一样,记得在Transformer中,使用的是词嵌入和位置编码,但是Bert使用的是可学习的位置信息(位置嵌入),并且加入了段信息(即词元属于哪一句,0或1)。

image-krms.png

1.2 掩蔽语言模型 (Masked Language Modeling)

class MaskLM(nn.Module):
    """BERT的掩蔽语言模型任务"""
    def __init__(self, vocab_size, num_hiddens, num_inputs=768, **kwargs):
        super(MaskLM, self).__init__(**kwargs)
        self.mlp = nn.Sequential(nn.Linear(num_inputs, num_hiddens),
                                 nn.ReLU(),
                                 nn.LayerNorm(num_hiddens),
                                 nn.Linear(num_hiddens, vocab_size))

    def forward(self, X, pred_positions):
        num_pred_positions = pred_positions.shape[1]
        pred_positions = pred_positions.reshape(-1)
        batch_size = X.shape[0]
        batch_idx = torch.arange(0, batch_size)
        # 假设batch_size=2,num_pred_positions=3
        # 那么batch_idx是np.array([0,0,0,1,1,1])
        batch_idx = torch.repeat_interleave(batch_idx, num_pred_positions)
        masked_X = X[batch_idx, pred_positions]
        masked_X = masked_X.reshape((batch_size, num_pred_positions, -1))
        mlm_Y_hat = self.mlp(masked_X)
        return mlm_Y_hat

这里的输入是编码器的输出X,以及需要预测位置pred_positions。

至于为什么叫掩蔽(Masked),因为在这个预训练任务中,大约15%的词元被随机选择作为掩蔽词元以进行预测。为了防止模型在预训练时直接看到答案而“作弊”,采用了以下策略:如果某个词元被选为掩蔽词元(例如,在“this movie is great”中选择掩蔽和预测“great”),则在输入中将其替换为:

  • 80%的情况下替换为特殊的“<mask>”词元(例如,“this movie is great”变为“this movie is <mask>”;
  • 10%的情况下替换为随机选择的词元(例如,“this movie is great”变为“this movie is drink”);
  • 10%的情况下保持原词元不变(例如,“this movie is great”变为“this movie is great”)。

请注意,在选为掩蔽的词元中,10%的词元被替换为随机词元。这种偶然的噪声鼓励BERT在其双向上下文编码中不那么偏向于掩蔽词元(尤其是当标签词元保持不变时)。

为什么只需要预测masked_X? 因为X是经过编码器的,每个词元特征维度已经有了上下文信息。

1.3 下一句预测 (Next Sentence Prediction)

class NextSentencePred(nn.Module):
    """BERT的下一句预测任务"""
    def __init__(self, num_inputs, **kwargs):
        super(NextSentencePred, self).__init__(**kwargs)
        # 定义一个线性层,将输入映射到两个输出,对应于两个类别:真或假
        self.output = nn.Linear(num_inputs, 2)

    def forward(self, X):
        # X的形状:(batch_size, num_hiddens)
        # X是编码后的"<cls>"词元的表示,包含了两个句子的信息
        # 通过线性层输出两个值,表示第二个句子是否是第一个句子的下一句
        return self.output(X)

尽管掩蔽语言建模能够有效地编码单词的双向上下文信息,但它并不直接建模文本对之间的逻辑关系。为了增强模型对两个文本序列之间关系的理解,BERT在预训练阶段引入了一个二元分类任务——下一句预测。在这个任务中,预训练数据中的一半句子对是真实连续的,标记为“真”;而另一半的第二个句子则是随机从语料库中选取的,与第一个句子不相关,标记为“假”。

NextSentencePred类实现了一个简单的多层感知机(MLP),用于预测第二个句子是否是BERT输入序列中第一个句子的直接后续句子。得益于Transformer编码器的自注意力机制,特殊词元“<cls>”的BERT表示已经融合了两个句子的信息。因此,MLP的输出层(self.output)直接以编码后的“<cls>”词元作为输入,这个输入包含了两个句子的综合信息,用于预测句子对之间的关系。

1.4 整体结构

class BERTModel(nn.Module):
    """BERT模型"""
    def __init__(self, vocab_size, num_hiddens, norm_shape, ffn_num_input,
                 ffn_num_hiddens, num_heads, num_layers, dropout,
                 max_len=1000, key_size=768, query_size=768, value_size=768,
                 hid_in_features=768, mlm_in_features=768,
                 nsp_in_features=768):
        super(BERTModel, self).__init__()
        self.encoder = BERTEncoder(vocab_size, num_hiddens, norm_shape,
                                   ffn_num_input, ffn_num_hiddens, num_heads, num_layers,
                                   dropout, max_len=max_len, key_size=key_size,
                                   query_size=query_size, value_size=value_size)
        self.hidden = nn.Sequential(nn.Linear(hid_in_features, num_hiddens), nn.Tanh())
        self.mlm = MaskLM(vocab_size, num_hiddens, mlm_in_features)
        self.nsp = NextSentencePred(nsp_in_features)

    def forward(self, tokens, segments, valid_lens=None, pred_positions=None):
        encoded_X = self.encoder(tokens, segments, valid_lens)
        if pred_positions is not None:
            mlm_Y_hat = self.mlm(encoded_X, pred_positions)
        else:
            mlm_Y_hat = None
        # 用于下一句预测的多层感知机分类器的隐藏层,0是“<cls>”标记的索引
        nsp_Y_hat = self.nsp(self.hidden(encoded_X[:, 0, :]))
        return encoded_X, mlm_Y_hat, nsp_Y_hat

1.4.1 前向传播

在前向传播方法中,模型按照以下步骤处理输入:

  1. 编码​:首先,输入的词元(tokens)和段信息(segments)通过BERTEncoder进行编码。编码过程中还考虑了有效长度(valid_lens),以确保注意力机制只关注有效的词元。
  2. 掩蔽语言模型预测​:如果提供了预测位置信息(pred_positions),则使用MaskLM组件对被掩蔽的词元进行预测。
  3. 下一句预测​:无论是否进行掩蔽语言模型预测,都会使用NextSentencePred组件对“”词元的表示进行下一句预测。这是通过将“”词元的表示传递给隐藏层,然后输出层的线性层来实现的。
  4. 返回结果​:前向传播方法返回编码后的词元表示(encoded_X)、掩蔽语言模型预测结果(mlm_Y_hat)以及下一句预测结果(nsp_Y_hat)。

1.4.2 训练和推理

在训练过程中,BERT模型同时学习掩蔽语言模型任务和下一句预测任务。这两个任务通过联合训练使得模型能够更好地理解语言中的上下文信息和句子间的关系。

在推理过程中,可以根据需要使用模型的不同部分,具体需要根据不同应用场景进行微调。例如,如果只需要进行句子表示,则可以使用编码器的输出;如果需要预测被掩蔽的词元,则可以使用掩蔽语言模型部分;如果需要判断两个句子是否连续,则可以使用下一句预测部分。

2 预训练数据

2.1 必要的工具函数、类

def count_corpus(tokens):
    """统计词元的频率。 """
    # 这里 `tokens` 可以是一维列表或二维列表
    if len(tokens) == 0 or isinstance(tokens[0], list):
        # 将词元列表的列表展平成一个词元列表
        tokens = [token for line in tokens for token in line]
    return collections.Counter(tokens)
class Vocab:
    """文本的词汇表类。"""
    def __init__(self, tokens=None, min_freq=0, reserved_tokens=None):
        """初始化词汇表。
        tokens -- 词元列表。
        min_freq -- 最小频率,词元出现次数少于这个值将被忽略。
        reserved_tokens -- 保留的词元列表,如特殊符号。
        """
        if tokens is None:
            tokens = []
        if reserved_tokens is None:
            reserved_tokens = []
        # 根据频率排序
        counter = count_corpus(tokens)
        self._token_freqs = sorted(counter.items(), key=lambda x: x[1], reverse=True)
        # 未知词元的索引为 0
        self.idx_to_token = ['<unk>'] + reserved_tokens
        self.token_to_idx = {token: idx
                             for idx, token in enumerate(self.idx_to_token)}
        for token, freq in self._token_freqs:
            if freq < min_freq:
                break
            if token not in self.token_to_idx:
                self.idx_to_token.append(token)
                self.token_to_idx[token] = len(self.idx_to_token) - 1

    def __len__(self):
        """返回词汇表的大小。"""
        return len(self.idx_to_token)

    def __getitem__(self, tokens):
        """根据词元返回其索引。
        如果词元不在词汇表中,返回未知词元的索引。
        """
        if not isinstance(tokens, (list, tuple)):
            return self.token_to_idx.get(tokens, self.unk)
        return [self.__getitem__(token) for token in tokens]

    def to_tokens(self, indices):
        """根据索引列表返回对应的词元列表。"""
        if not isinstance(indices, (list, tuple)):
            return self.idx_to_token[indices]
        return [self.idx_to_token[index] for index in indices]

    @property
    def unk(self):
        """返回未知词元的索引。"""
        return 0

    @property
    def token_freqs(self):
        """返回词元频率列表。"""
        return self._token_freqs```

2.2 下一句预测的数据处理

def _get_next_sentence(sentence, next_sentence, paragraphs):
    if random.random() < 0.5:
        is_next = True
    else:
        # paragraphs是三重列表的嵌套
        next_sentence = random.choice(random.choice(paragraphs))
        is_next = False
    return sentence, next_sentence, is_next
  • 功能:生成下一句预测任务的数据。
  • 输入
    • sentence:当前句子。
    • next_sentence:下一个句子。
    • paragraphs:段落列表,每个段落是一个句子列表。
  • 处理
    • 有50%的概率,next_sentence是实际的下一句,is_next设置为True
    • 另50%的概率,从随机段落中随机选择一个句子作为next_sentenceis_next设置为False
  • 输出:返回当前句子、下一句子和是否为实际下一句的标记。
def get_tokens_and_segments(tokens_a, tokens_b=None):
    """获取输入序列的词元及其片段索引"""
    tokens = ['<cls>'] + tokens_a + ['<sep>']
    # 0和1分别标记片段A和B
    segments = [0] * (len(tokens_a) + 2)
    if tokens_b is not None:
        tokens += tokens_b + ['<sep>']
        segments += [1] * (len(tokens_b) + 1)
    return tokens, segments
  • 功能:将两个句子转换为BERT模型的输入格式,包括词元序列和片段索引。
  • 输入
    • tokens_a:第一个句子的词元列表。
    • tokens_b:第二个句子的词元列表,默认为None
  • 处理
    • 在词元序列的开始添加<cls>标记,在两个句子之间和句子末尾添加<sep>标记。
    • 片段索引用于区分两个句子,第一个句子的片段索引为0,第二个句子的片段索引为1。
  • 输出:返回词元序列和片段索引列表。
def _get_nsp_data_from_paragraph(paragraph, paragraphs, vocab, max_len):
    nsp_data_from_paragraph = []
    for i in range(len(paragraph) - 1):
        tokens_a, tokens_b, is_next = _get_next_sentence(paragraph[i], paragraph[i + 1], paragraphs)
        # 考虑1个'<cls>'词元和2个'<sep>'词元
        if len(tokens_a) + len(tokens_b) + 3 > max_len:
            continue
        tokens, segments = get_tokens_and_segments(tokens_a, tokens_b)
        nsp_data_from_paragraph.append((tokens, segments, is_next))
    return nsp_data_from_paragraph
  • 功能:从段落中生成下一句预测任务的数据。
  • 输入
    • paragraph:当前段落,是一个句子列表。
    • paragraphs:所有段落的列表。
    • vocab:词汇表。
    • max_len:最大序列长度。
  • 处理
    • 遍历段落中的句子对,使用_get_next_sentence函数生成下一句预测数据。
    • 检查生成的序列长度是否超过max_len,如果超过则跳过。
    • 使用get_tokens_and_segments函数将句子对转换为词元序列和片段索引。
    • 将生成的数据添加到列表中。
  • 输出:返回当前段落中所有有效的下一句预测数据。

2.3 遮蔽语言模型的数据处理

def _replace_mlm_tokens(tokens, candidate_pred_positions, num_mlm_preds, vocab):
    # 为遮蔽语言模型的输入创建新的词元副本,其中输入可能包含替换的“<mask>”或随机词元
    mlm_input_tokens = [token for token in tokens]
    pred_positions_and_labels = []
    # 打乱后用于在遮蔽语言模型任务中获取15%的随机词元进行预测
    random.shuffle(candidate_pred_positions)
    for mlm_pred_position in candidate_pred_positions:
        if len(pred_positions_and_labels) >= num_mlm_preds:
            break
        masked_token = None
        # 80%的时间:将词替换为“<mask>”词元
        if random.random() < 0.8:
            masked_token = '<mask>'
        else:
            # 10%的时间:保持词不变
            if random.random() < 0.5:
                masked_token = tokens[mlm_pred_position]
            # 10%的时间:用随机词替换该词
            else:
                masked_token = random.choice(vocab.idx_to_token)
        mlm_input_tokens[mlm_pred_position] = masked_token
        pred_positions_and_labels.append((mlm_pred_position, tokens[mlm_pred_position]))
    return mlm_input_tokens, pred_positions_and_labels
  • 功能:替换遮蔽语言模型任务中的词元,生成输入序列和预测位置及标签。
  • 输入
    • tokens:原始词元序列。
    • candidate_pred_positions:候选预测位置列表。
    • num_mlm_preds:需要预测的词元数量。
    • vocab:词汇表。
  • 处理
    • 创建原始词元序列的副本mlm_input_tokens
    • 打乱候选预测位置列表。
    • 遍历候选预测位置,根据一定概率替换词元:
      • 80%的概率替换为<mask>
      • 10%的概率保持不变。
      • 10%的概率替换为随机词元。
    • 记录替换位置和原始词元作为预测标签。
  • 输出:返回替换后的词元序列和预测位置及标签的列表。
def _get_mlm_data_from_tokens(tokens, vocab):
    candidate_pred_positions = []
    # tokens是一个字符串列表
    for i, token in enumerate(tokens):
        # 在遮蔽语言模型任务中不会预测特殊词元
        if token in ['<cls>', '<sep>']:
            continue
        candidate_pred_positions.append(i)
    # 遮蔽语言模型任务中预测15%的随机词元
    num_mlm_preds = max(1, round(len(tokens) * 0.15))
    mlm_input_tokens, pred_positions_and_labels = _replace_mlm_tokens(tokens, candidate_pred_positions, num_mlm_preds, vocab)
    pred_positions_and_labels = sorted(pred_positions_and_labels, key=lambda x: x[0])
    pred_positions = [v[0] for v in pred_positions_and_labels]
    mlm_pred_labels = [v[1] for v in pred_positions_and_labels]
    return vocab[mlm_input_tokens], pred_positions, vocab[mlm_pred_labels]
  • 功能:从词元序列中生成遮蔽语言模型任务的数据。
  • 输入
    • tokens:词元序列。
    • vocab:词汇表。
  • 处理
    • 初始化候选预测位置列表,排除特殊词元<cls><sep>
    • 计算15%的词元数量作为需要预测的词元数量num_mlm_preds
    • 调用_replace_mlm_tokens函数生成替换后的词元序列和预测位置及标签。
    • 对预测位置及标签列表进行排序。
    • 分离预测位置和预测标签。
    • 将词元序列和预测标签转换为词汇表中的索引。
  • 输出:返回替换后的词元序列索引、预测位置和预测标签索引。

2.5 填充数据

def _pad_bert_inputs(examples, max_len, vocab):
    max_num_mlm_preds = round(max_len * 0.15)
    all_token_ids, all_segments, valid_lens,  = [], [], []
    all_pred_positions, all_mlm_weights, all_mlm_labels = [], [], []
    nsp_labels = []
    for (token_ids, pred_positions, mlm_pred_label_ids, segments, is_next) in examples:
        all_token_ids.append(torch.tensor(token_ids + [vocab['<pad>']] * (max_len - len(token_ids)), dtype=torch.long))
        all_segments.append(torch.tensor(segments + [0] * (max_len - len(segments)), dtype=torch.long))
        # valid_lens不包括'<pad>'的计数
        valid_lens.append(torch.tensor(len(token_ids), dtype=torch.float32))
        all_pred_positions.append(torch.tensor(pred_positions + [0] * (max_num_mlm_preds - len(pred_positions)), dtype=torch.long))
        # 填充词元的预测将通过乘以0权重在损失中过滤掉
        all_mlm_weights.append(torch.tensor([1.0] * len(mlm_pred_label_ids) + [0.0] * (max_num_mlm_preds - len(pred_positions)), dtype=torch.float32))
        all_mlm_labels.append(torch.tensor(mlm_pred_label_ids + [0] * (max_num_mlm_preds - len(mlm_pred_label_ids)), dtype=torch.long))
        nsp_labels.append(torch.tensor(is_next, dtype=torch.long))
    return (all_token_ids, all_segments, valid_lens, all_pred_positions, all_mlm_weights, all_mlm_labels, nsp_labels)
  • 功能​:将BERT模型的输入数据填充到相同的长度,以便于批量处理。同时处理遮蔽语言模型(MLM)和下一句预测(NSP)任务的输入数据。
  • 输入​:
    • examples:一个列表,每个元素是一个元组,包含以下内容:
      • token_ids:词元序列的索引。
      • pred_positions:预测位置的索引。
      • mlm_pred_label_ids:MLM预测标签的索引。
      • segments:片段索引。
      • is_next:NSP任务的标签。
    • max_len:最大序列长度。
    • vocab:词汇表。
  • 处理​:
    1. 初始化列表​:
      • all_token_ids:用于存储所有填充后的词元序列索引。
      • all_segments:用于存储所有填充后的片段索引。
      • valid_lens:用于存储所有有效长度(不包括填充词元)。
      • all_pred_positions:用于存储所有填充后的预测位置索引。
      • all_mlm_weights:用于存储所有MLM预测的权重。
      • all_mlm_labels:用于存储所有MLM预测标签的索引。
      • nsp_labels:用于存储所有NSP任务的标签。
    2. 遍历每个示例​:
      • 填充词元序列索引​:将词元序列索引填充到max_len长度,不足部分用<pad>的索引填充。
      • 填充片段索引​:将片段索引填充到max_len长度,不足部分用0填充。
      • 计算有效长度​:记录原始词元序列的长度(不包括填充词元)。
      • 填充预测位置索引​:将预测位置索引填充到max_num_mlm_preds长度,不足部分用0填充。
      • 设置MLM预测权重​:为真实的MLM预测位置设置权重为1.0,填充位置设置权重为0.0,以便在计算损失时忽略填充位置。
      • 填充MLM预测标签索引​:将MLM预测标签索引填充到max_num_mlm_preds长度,不足部分用0填充。
      • 设置NSP标签​:将NSP任务的标签转换为张量。
    3. 返回填充后的数据​:返回一个元组,包含所有填充后的数据列表。
  • 输出​:
    • all_token_ids:所有填充后的词元序列索引张量列表。
    • all_segments:所有填充后的片段索引张量列表。
    • valid_lens:所有有效长度张量列表。
    • all_pred_positions:所有填充后的预测位置索引张量列表。
    • all_mlm_weights:所有MLM预测权重张量列表。
    • all_mlm_labels:所有MLM预测标签索引张量列表。
    • nsp_labels:所有NSP任务标签张量列表。

2.6 构建数据集(组合上述功能)

class _WikiTextDataset(torch.utils.data.Dataset):
    def __init__(self, paragraphs, max_len):
        # 输入paragraphs[i]是代表段落的句子字符串列表;
        # 而输出paragraphs[i]是代表段落的句子列表,其中每个句子都是词元列表
        paragraphs = [[paragraph.split() for paragraph in paragraph] for paragraph in paragraphs]
        sentences = [sentence
                     for paragraph in paragraphs
                     for sentence in paragraph]
        self.vocab = Vocab(sentences, min_freq=5, reserved_tokens=['<pad>', '<mask>', '<cls>', '<sep>'])
        # 获取下一句子预测任务的数据
        examples = []
        for paragraph in paragraphs:
            examples.extend(_get_nsp_data_from_paragraph(paragraph, paragraphs, self.vocab, max_len))
        # 获取遮蔽语言模型任务的数据
        examples = [(_get_mlm_data_from_tokens(tokens, self.vocab) + (segments, is_next))
                    for tokens, segments, is_next in examples]
        # 填充输入
        (self.all_token_ids, self.all_segments, self.valid_lens,
         self.all_pred_positions, self.all_mlm_weights,
         self.all_mlm_labels, self.nsp_labels) = _pad_bert_inputs(examples, max_len, self.vocab)

    def __getitem__(self, idx):
        return (self.all_token_ids[idx], self.all_segments[idx],
                self.valid_lens[idx], self.all_pred_positions[idx],
                self.all_mlm_weights[idx], self.all_mlm_labels[idx],
                self.nsp_labels[idx])

    def __len__(self):
        return len(self.all_token_ids)
  • 功能:自定义的PyTorch数据集类,用于处理和准备WikiText数据集的输入数据,以供BERT模型训练。结合了下一句预测(NSP)和遮蔽语言模型(MLM)任务的数据预处理。
  • 初始化方法 __init__
    • 输入
      • paragraphs:一个列表,每个元素是一个段落,段落是由句子组成的列表,句子是字符串。
      • max_len:最大序列长度。
    • 处理
      1. 词元化:将每个段落的句子分割成词元列表。
      2. 构建词汇表:使用所有句子的词元列表构建词汇表,设置最小频率为5,并保留特殊词元(<pad>, <mask>, <cls>, <sep>)。
      3. 获取NSP数据:遍历每个段落,使用_get_nsp_data_from_paragraph函数生成NSP任务的数据。
      4. 获取MLM数据:遍历NSP数据,使用_get_mlm_data_from_tokens函数生成MLM任务的数据。
      5. 填充输入:使用_pad_bert_inputs函数将数据填充到相同的长度,以便于批量处理。
    • 输出:初始化后的数据集对象,包含了填充后的输入数据和张量。
  • __getitem__方法
    • 功能:根据索引获取数据集中的一个样本。
    • 输入idx,样本的索引。
    • 输出:一个元组,包含以下内容:
      • all_token_ids[idx]:填充后的词元序列索引张量。
      • all_segments[idx]:填充后的片段索引张量。
      • valid_lens[idx]:有效长度张量。
      • all_pred_positions[idx]:填充后的预测位置索引张量。
      • all_mlm_weights[idx]:MLM预测权重张量。
      • all_mlm_labels[idx]:MLM预测标签索引张量。
      • nsp_labels[idx]:NSP任务标签张量。
  • __len__方法
    • 功能:返回数据集的样本数量。
    • 输出:数据集的长度,即all_token_ids列表的长度。

2.7 读取和处理

def _read_wiki(data_dir):
    file_name = os.path.join(data_dir, 'wiki.train.tokens')
    with open(file_name, 'r', encoding='utf-8') as f:
        lines = f.readlines()
    # 大写字母转换为小写字母
    paragraphs = [line.strip().lower().split(' . ')
                  for line in lines if len(line.split(' . ')) >= 2]
    random.shuffle(paragraphs)
    return paragraphs
def load_data_wiki(batch_size, max_len):
    """加载WikiText-2数据集"""
    num_workers = 0
    data_dir = "path/to/data"
    paragraphs = _read_wiki(data_dir)
    train_set = _WikiTextDataset(paragraphs, max_len)
    train_iter = DataLoader(train_set, batch_size, shuffle=True, num_workers=num_workers)
    return train_iter, train_set.vocab

以batch_size=512、max_len=64为例:

  1. tokens_X​: 词元序列索引张量。
    • 形状​: torch.Size([512, 64])
    • 含义​: 表示一个批次中的512个样本,每个样本由64个词元索引组成。这些索引对应于词汇表中的词元。
  2. segments_X​: 片段索引张量。
    • 形状​: torch.Size([512, 64])
    • 含义​: 表示一个批次中的512个样本,每个样本由64个片段索引组成。片段索引用于区分不同的句子(例如,第一个句子和第二个句子)。
  3. valid_lens_x​: 有效长度张量。
    • 形状​: torch.Size([512])
    • 含义​: 表示一个批次中的512个样本的有效长度。有效长度是指不包括填充词元(<pad>)的实际词元数量。
  4. pred_positions_X​: 预测位置索引张量。
    • 形状​: torch.Size([512, 10])
    • 含义​: 表示一个批次中的512个样本,每个样本有10个预测位置的索引。这些位置是用于遮蔽语言模型(MLM)任务的预测位置。
  5. mlm_weights_X​: MLM预测权重张量。
    • 形状​: torch.Size([512, 10])
    • 含义​: 表示一个批次中的512个样本,每个样本有10个预测位置的权重。权重为1.0的位置是真实的预测位置,权重为0.0的位置是填充位置。
  6. mlm_Y​: MLM预测标签索引张量。
    • 形状​: torch.Size([512, 10])
    • 含义​: 表示一个批次中的512个样本,每个样本有10个预测位置的标签索引。这些标签是用于遮蔽语言模型(MLM)任务的正确词元索引。
  7. nsp_y​: NSP任务标签张量。
    • 形状​: torch.Size([512])
    • 含义​: 表示一个批次中的512个样本的下一句预测(NSP)任务标签。标签为1表示下一句是真实的,标签为0表示下一句是随机选择的。

3 预训练

loss = nn.CrossEntropyLoss()

def _get_batch_loss_bert(net, loss, vocab_size, tokens_X, segments_X, valid_lens_x, pred_positions_X,
                         mlm_weights_X, mlm_Y, nsp_y):
    # 前向传播
    _, mlm_Y_hat, nsp_Y_hat = net(tokens_X, segments_X, valid_lens_x.reshape(-1), pred_positions_X)
    # 计算遮蔽语言模型损失
    mlm_l = loss(mlm_Y_hat.reshape(-1, vocab_size), mlm_Y.reshape(-1)) * mlm_weights_X.reshape(-1, 1)
    mlm_l = mlm_l.sum() / (mlm_weights_X.sum() + 1e-8)
    # 计算下一句子预测任务的损失
    nsp_l = loss(nsp_Y_hat, nsp_y)
    l = mlm_l + nsp_l
    return mlm_l, nsp_l, l

net = nn.DataParallel(net, device_ids=devices).to(devices[0])
trainer = torch.optim.Adam(net.parameters(), lr=0.01)
step, num_steps = 0, 100
num_steps_reached = False
while step < num_steps and not num_steps_reached:
    for tokens_X, segments_X, valid_lens_x, pred_positions_X, mlm_weights_X, mlm_Y, nsp_y in train_iter:
        tokens_X = tokens_X.to(devices[0])
        segments_X = segments_X.to(devices[0])
        valid_lens_x = valid_lens_x.to(devices[0])
        pred_positions_X = pred_positions_X.to(devices[0])
        mlm_weights_X = mlm_weights_X.to(devices[0])    # mlm_weights_X其实是一个mask矩阵,用于屏蔽掉pad的token
        mlm_Y, nsp_y = mlm_Y.to(devices[0]), nsp_y.to(devices[0])
 

        trainer.zero_grad()
        mlm_l, nsp_l, l = _get_batch_loss_bert(net, loss, vocab_size, tokens_X, segments_X, valid_lens_x,
                                               pred_positions_X, mlm_weights_X, mlm_Y, nsp_y)
        l.backward()
        trainer.step()
 
        step += 1
        if step == num_steps:
            num_steps_reached = True
            break