You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

从零构建并训练基于自定义词汇表的MLM任务BERT模型

从零构建适配自定义词汇表的MLM-BERT模型步骤

1. 自定义词汇表处理

  • 把你的乱码词汇整理成一个每行一个词的文本文件(比如vocab.txt),格式示例:
    xddccfff988
    9900zxzxxx
    zzzTzz
    [PAD]
    [MASK]
    [CLS]
    [SEP]
    
    必须加入BERT核心特殊标记:[PAD](填充)、[MASK](掩码)、[CLS](句首)、[SEP](分隔),这些是MLM任务的必需组件。
  • 给每个词汇分配唯一ID,用字典映射实现:word_to_idx = {word: idx for idx, word in enumerate(vocab_list)},后续编码、解码都依赖该映射。

2. 自定义Tokenizer实现

无需BERT自带tokenizer,针对你的完整乱码词汇实现简单编码逻辑:

  • 核心功能:将输入文本(空格分隔的乱码词序列)转换为ID序列,同时处理特殊标记
  • 示例代码(Python):
    class CustomTokenizer:
        def __init__(self, vocab_path):
            with open(vocab_path, 'r', encoding='utf-8') as f:
                self.vocab = [line.strip() for line in f]
            self.word_to_idx = {w: i for i, w in enumerate(self.vocab)}
            self.idx_to_word = {i: w for w, i in self.word_to_idx.items()}
            self.pad_token = '[PAD]'
            self.mask_token = '[MASK]'
            self.cls_token = '[CLS]'
            self.sep_token = '[SEP]'
    
        def encode(self, text, max_len=None):
            # 输入text格式:"xddccfff988 zzzTzz"(空格分隔的乱码词)
            tokens = text.split()
            tokens = [self.cls_token] + tokens + [self.sep_token]
            # 过滤不在词汇表中的词,转成ID
            ids = [self.word_to_idx[token] for token in tokens if token in self.word_to_idx]
            # 填充或截断到指定长度
            if max_len:
                if len(ids) < max_len:
                    ids += [self.word_to_idx[self.pad_token]] * (max_len - len(ids))
                else:
                    ids = ids[:max_len]
            return ids
    
        def decode(self, ids):
            return ' '.join([self.idx_to_word[idx] for idx in ids if idx != self.word_to_idx[self.pad_token]])
    

3. 构建BERT模型主体

基于PyTorch实现核心Transformer编码器与MLM预测头部:

3.1 基础组件实现

  • 词嵌入层:输入词汇ID,输出对应维度的词嵌入(比如768维,与BERT-base对齐)
  • 位置嵌入层:可学习的位置信息嵌入,长度设为你的最大序列长度
  • Transformer编码器:多层多头注意力+前馈网络结构

3.2 MLM头部与完整模型代码

import torch
import torch.nn as nn
from torch.nn import functional as F

class BertMLMHead(nn.Module):
    def __init__(self, hidden_size, vocab_size):
        super().__init__()
        self.dense = nn.Linear(hidden_size, hidden_size)
        self.layer_norm = nn.LayerNorm(hidden_size)
        self.decoder = nn.Linear(hidden_size, vocab_size)

    def forward(self, hidden_states):
        x = self.dense(hidden_states)
        x = F.gelu(x)
        x = self.layer_norm(x)
        x = self.decoder(x)
        return x

class CustomBERT(nn.Module):
    def __init__(self, vocab_size, hidden_size=768, num_layers=12, num_heads=12, max_seq_len=512):
        super().__init__()
        self.word_embeddings = nn.Embedding(vocab_size, hidden_size)
        self.position_embeddings = nn.Embedding(max_seq_len, hidden_size)
        self.token_type_embeddings = nn.Embedding(2, hidden_size)
        self.embedding_layer_norm = nn.LayerNorm(hidden_size)
        self.embedding_dropout = nn.Dropout(0.1)

        # Transformer编码器层
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=hidden_size, nhead=num_heads, 
            dim_feedforward=hidden_size*4, dropout=0.1
        )
        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)

        # MLM预测头部
        self.mlm_head = BertMLMHead(hidden_size, vocab_size)

    def forward(self, input_ids, token_type_ids=None, attention_mask=None):
        seq_len = input_ids.size(1)
        positions = torch.arange(seq_len, device=input_ids.device).unsqueeze(0).expand_as(input_ids)

        if token_type_ids is None:
            token_type_ids = torch.zeros_like(input_ids)

        # 计算总嵌入:词嵌入+位置嵌入+类型嵌入
        word_embeds = self.word_embeddings(input_ids)
        pos_embeds = self.position_embeddings(positions)
        type_embeds = self.token_type_embeddings(token_type_ids)
        embeddings = word_embeds + pos_embeds + type_embeds
        embeddings = self.embedding_layer_norm(embeddings)
        embeddings = self.embedding_dropout(embeddings)

        # Transformer编码(注意维度转置)
        encoder_output = self.encoder(
            embeddings.transpose(0,1), 
            src_key_padding_mask=~attention_mask
        ).transpose(0,1)

        # MLM预测输出
        mlm_logits = self.mlm_head(encoder_output)
        return mlm_logits, encoder_output

4. MLM训练数据生成

生成符合MLM任务要求的训练样本,核心逻辑是随机掩码词汇:

  • 对每个序列,随机选择15%的token:80%替换为[MASK],10%替换为随机词汇,10%保持原词
  • 示例代码:
    def create_mlm_data(tokenizer, text, max_len):
        input_ids = tokenizer.encode(text, max_len)
        input_ids = torch.tensor(input_ids).unsqueeze(0)
        labels = input_ids.clone()
    
        # 生成掩码位置,避开特殊标记
        mask_prob = 0.15
        mask_indices = torch.rand(input_ids.shape) < mask_prob
        special_token_ids = [
            tokenizer.word_to_idx[token] 
            for token in [tokenizer.cls_token, tokenizer.sep_token, tokenizer.pad_token]
        ]
        for sid in special_token_ids:
            mask_indices = mask_indices & (input_ids != sid)
    
        # 处理掩码逻辑
        for i in range(input_ids.shape[0]):
            masked_pos = mask_indices[i].nonzero().squeeze(1)
            for pos in masked_pos:
                rand = torch.rand(1).item()
                if rand < 0.8:
                    input_ids[i, pos] = tokenizer.word_to_idx[tokenizer.mask_token]
                elif rand < 0.9:
                    # 随机选择词汇表中的词
                    random_id = torch.randint(0, len(tokenizer.vocab), (1,)).item()
                    input_ids[i, pos] = random_id
                # 剩余10%保持原词不变
    
        attention_mask = (input_ids != tokenizer.word_to_idx[tokenizer.pad_token]).float()
        return input_ids, attention_mask, labels
    

5. 训练与获取词嵌入

  • 训练流程:使用交叉熵损失,优化器选择AdamW,训练逻辑与标准BERT一致
  • 获取词嵌入:
    • 静态词嵌入:训练完成后,model.word_embeddings.weight即为所有词汇的静态嵌入向量
    • 上下文相关嵌入:模型前向输出的encoder_output是每个token的上下文相关隐藏层输出,可作为动态词嵌入

内容的提问来源于stack exchange,提问作者Saeed Asadi Bagloee

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.19 04:26:05