BERT 预训练与微调:参考答案

目录

\[ \begin{align}\begin{aligned}\newcommand{\ba}{\boldsymbol{a}} \newcommand{\bb}{\boldsymbol{b}} \newcommand{\be}{\boldsymbol{e}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bw}{\boldsymbol{w}} \newcommand{\bx}{\boldsymbol{x}} \newcommand{\by}{\boldsymbol{y}} \newcommand{\bz}{\boldsymbol{z}} \newcommand{\bd}{\boldsymbol{d}} \newcommand{\bv}{\boldsymbol{v}} \newcommand{\bs}{\boldsymbol{s}}\\\newcommand{\btheta}{\boldsymbol{\theta}} \newcommand{\bbeta}{\boldsymbol{\beta}} \newcommand{\bgamma}{\boldsymbol{\gamma}} \newcommand{\bsigma}{\boldsymbol{\sigma}} \newcommand{\md}{\mbox{d}} \newcommand{\bmu}{\boldsymbol{\mu}} \newcommand{\bone}{\boldsymbol{1}} \newcommand{\bzero}{\boldsymbol{0}} \newcommand{\bepsilon}{\boldsymbol{\epsilon}} \newcommand{\bphi}{\boldsymbol{\phi}} \newcommand{\bh}{\boldsymbol{h}} \newcommand{\bc}{\boldsymbol{c}} \newcommand{\br}{\boldsymbol{r}} \newcommand{\bQ}{\boldsymbol{Q}} \newcommand{\bK}{\boldsymbol{K}} \newcommand{\bV}{\boldsymbol{V}} \newcommand{\bSigma}{\boldsymbol{\Sigma}} \newcommand{\bg}{\boldsymbol{g}} \newcommand{\bxi}{\boldsymbol{\xi}} \newcommand{\bvarepsilon}{\boldsymbol{\varepsilon}} \newcommand{\bdelta}{\boldsymbol{\delta}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bJ}{\boldsymbol{J}} \newcommand{\bp}{\boldsymbol{p}} \newcommand{\bi}{\boldsymbol{i}} \newcommand{\bo}{\boldsymbol{o}} \newcommand{\bE}{\boldsymbol{E}} \newcommand{\bH}{\boldsymbol{H}} \newcommand{\bL}{\boldsymbol{L}} \newcommand{\bu}{\boldsymbol{u}} \newcommand{\bLambda}{\boldsymbol{\Lambda}} \newcommand{\trans}{^{\rm\scriptsize T}} \newcommand{\var}{\mathrm{var}}\\\newcommand{\bA}{\boldsymbol{A}} \newcommand{\bB}{\boldsymbol{B}} \newcommand{\bC}{\boldsymbol{C}} \newcommand{\bD}{\boldsymbol{D}} \newcommand{\bG}{\boldsymbol{G}} \newcommand{\bI}{\boldsymbol{I}} \newcommand{\bM}{\boldsymbol{M}} \newcommand{\bP}{\boldsymbol{P}} \newcommand{\bS}{\boldsymbol{S}} \newcommand{\bU}{\boldsymbol{U}} \newcommand{\bW}{\boldsymbol{W}} \newcommand{\bX}{\boldsymbol{X}} \newcommand{\bY}{\boldsymbol{Y}} \newcommand{\bZ}{\boldsymbol{Z}} \newcommand{\cotp}{\textcolor[RGB]{48,209,88}{TP}} \newcommand{\cotn}{\textcolor[RGB]{100,210,255}{TN}} \newcommand{\cofp}{\textcolor[RGB]{94,92,230}{FP}} \newcommand{\cofn}{\textcolor[RGB]{191,90,242}{FN}}\\\newcommand{\numcotp}{\textcolor[RGB]{48,209,88}{50}} \newcommand{\numcotn}{\textcolor[RGB]{100,210,255}{30}} \newcommand{\numcofp}{\textcolor[RGB]{94,92,230}{10}} \newcommand{\numcofn}{\textcolor[RGB]{191,90,242}{10}} \DeclareMathOperator*{\argmin}{arg\,min}\end{aligned}\end{align} \]

BERT 预训练与微调:参考答案#

返回正文练习 · 返回答案索引

说明#

以下答案与正文10道题逐题对应。前四题给出主要推导与计算;第 5--7 题给出可以运行的程序和检查方法;第 8--10 题说明比较时需要保持哪些条件相同,以及实验结果能够说明什么、不能说明什么。

  1. 初始表示按元素相加:

    \[\bx_i=(1,0)\trans+(0.5,0.5)\trans+(-0.5,1)\trans=(1,1.5)\trans.\]

    若三种嵌入都属于 \(\mathbb R^d\)⁠,每个坐标才有对应的加法;维度不同必须先经过明确的映射,不能依赖广播改变语义。由总和通常不能唯一恢复三项,例如任意向量 \(\boldsymbol{a}\) 都可同时加到第一项、从第二项减去而保持总和不变。模型学习的是合成表示,不是在每层显式分离三部分。

  2. 两个被选位置的平均损失为

    \[\mathcal J_{\mathrm{MLM}} =-\frac{\log0.5+\log0.25}{2} \approx1.03972.\]

    批量损失的分母是整个批次中有效被选位置总数,而不是批量大小,也不是存储位置总数。另一条序列没有目标时不增加分子或分母;若整个批次都没有目标,分母为 0,应重新采样 masking 或明确报错,不能返回 NaN 或把损失设为 0。

  3. 若输入位置 \(i\) 仍含原编号 \(x_i\)⁠,编码器可令该位置表示保留一个区分词元的编码,输出层再近似实现

    \[\begin{split}p(x_i=k\mid\widetilde{\bx})\approx \begin{cases} 1,&k=\widetilde x_i,\\ 0,&k\ne\widetilde x_i. \end{cases}\end{split}\]

    这个规则只复制当前位置,不使用左右上下文,却能让训练损失接近 0。因此大多数目标必须被遮盖或替换,迫使模型利用上下文。原始方案只让被选位置中的约 \(10\%\) 保持不变,而且模型事先不知道属于哪种处理;其余扰动仍提供主要学习信号,并减弱预训练与下游输入之间的分布差异。该经验配比不是防止所有捷径的数学保证。

  4. 张量维度为

    \[\begin{split}\begin{array}{c|c} \text{张量}&\text{维度}\\ \hline \texttt{input\_ids},\ \texttt{attention\_mask}&16\times128\\ \text{编码输出}&16\times128\times768\\ \text{单头查询}&16\times128\times64\\ \text{分类输出}&16\times2 \end{array},\end{split}\]

    其中每头 \(d_k=768/12=64\)⁠。注意力得分元素数为 \(16(12)(128)(128)=3\,145\,728\)⁠,单个 float32 张量占 \(12\,582\,912\) 字节,即 12 MiB。线性分类输出层参数量为 \(768(2)+2=1538\)⁠。真实训练峰值还包含投影、激活、梯度和优化器状态,远大于单个得分张量。

  5. 以下骨架把标签和扰动输入分开保存,避免用修改后的编号作为目标:

    import torch
    
    def mask_batch(input_ids, attention_mask, special_mask,
                   mask_id, ordinary_token_ids, generator, probability=0.15,
                   ignore_index=-100):
        if input_ids.shape != attention_mask.shape or input_ids.shape != special_mask.shape:
            raise ValueError("输入、注意力掩码和特殊词元掩码维度不同")
        candidate = attention_mask.bool() & ~special_mask.bool()
        if not candidate.any():
            raise ValueError("批次中没有可用于 MLM 的普通词元")
        ordinary_token_ids = torch.as_tensor(
            ordinary_token_ids, dtype=input_ids.dtype,
            device=input_ids.device
        )
        if ordinary_token_ids.ndim != 1 or ordinary_token_ids.numel() == 0:
            raise ValueError("ordinary_token_ids 必须是一维非空普通词元编号表")
        draw = torch.rand(input_ids.shape, generator=generator,
                          device=input_ids.device)
        selected = candidate & (draw < probability)
        # 保证每个有候选位置的样本至少选中一个位置。
        for row in range(input_ids.shape[0]):
            if candidate[row].any() and not selected[row].any():
                index = torch.nonzero(candidate[row], as_tuple=False)[0, 0]
                selected[row, index] = True
        labels = torch.full_like(input_ids, ignore_index)
        labels[selected] = input_ids[selected]
        corrupted = input_ids.clone()
        action = torch.rand(input_ids.shape, generator=generator,
                            device=input_ids.device)
        use_mask = selected & (action < 0.8)
        use_random = selected & (action >= 0.8) & (action < 0.9)
        corrupted[use_mask] = mask_id
        random_index = torch.randint(
            0, ordinary_token_ids.numel(), input_ids.shape,
            generator=generator, device=input_ids.device
        )
        random_ids = ordinary_token_ids[random_index]
        corrupted[use_random] = random_ids[use_random]
        return corrupted, labels, attention_mask, selected
    

    调用者应从词表中预先构造不含 [PAD][MASK][CLS][SEP] 等特殊词元的 ordinary_token_ids逐样本强制一个目标会在很短序列上改变精确的 15% 比例,应在实现说明中记录。

  6. 核心流程是先设置训练范围,再创建优化器:

    def configure_model(model, freeze_encoder):
        for parameter in model.base_model.parameters():
            parameter.requires_grad = not freeze_encoder
        trainable = [p for p in model.parameters() if p.requires_grad]
        optimizer = torch.optim.AdamW(trainable, lr=learning_rate)
        return optimizer, sum(p.numel() for p in trainable)
    
    @torch.no_grad()
    def predict_batch(model, batch):
        model.eval()
        output = model(input_ids=batch["input_ids"],
                       attention_mask=batch["attention_mask"])
        return output.logits.argmax(dim=-1)
    

    训练只访问训练集,验证集用于早停和超参数选择,测试集在方案确定后评价一次。完整微调通常需要较小学习率;冻结方案若缓存特征,应将编码器设为 eval 并记录一次性缓存成本。代码还应记录每轮损失、验证指标、学习率和可训练参数量。

  7. 最小测试集合包括:断言 labels[~selected] 和 padding 标签全部等于 ignore_index将 padding 编号换成另一个合法编号后,在正确注意力掩码下有效样本预测不变;全 padding 或无候选位置明确报错;用两个同种子的 torch.Generator 得到相同结果;把有效 MLM 位置展平后,手算损失与 cross_entropy(ignore_index=-100) 一致。

    经验比例测试应累计许多候选位置,再检查被选比例接近 15%,且被选位置中的三类处理接近 80%/10%/10%;不能对单个很短批次要求精确比例。还应验证随机替换不会抽到被禁止的特殊词元。

  8. 三种方法使用相同的原始数据划分、随机种子和最终测试程序。词袋词表只根据训练集建立;BERT 的两种方案使用同一个检查点和相同的词元化结果。规定神经模型最多更新多少次,或者规定编码器最多运行多少次,并在同一个验证集上按照相同标准选择学习率;逻辑回归使用预先规定的最大迭代次数,同时报告实际迭代次数。

    汇总准确率、F1、实际训练时间、固定 256 个样本的预测时间、峰值内存、总参数量、可训练参数量和编码器运行次数。作为简单对照的词袋模型计算成本较低,结果也容易解释;冻结编码器的方法成本居中,完整微调通常更灵活但计算成本更高。三者的精度次序取决于应用领域、样本量和检查点,不能把从预训练语料中学到的知识误称为测试标签泄漏。

  9. 三种方案使用相同的检查点、词元化结果、数据划分、随机种子、最大更新次数和每次更新的样本数,并在同一个验证集上按照相同标准选择模型。三种方案都尝试同一组学习率,也可以事先规定分层的学习率取值;参数高效方法必须记录模块位置、秩等额外设置。报告测试指标、训练时间、固定批量预测时间、峰值内存、可训练参数量、总参数量,以及编码器前向传播和后向传播的运行次数。

    冻结参数少不代表前向计算少;完整微调和参数高效方法推断时通常仍运行整个编码器,所以预测时间可能接近。峰值显存的差异更多来自梯度和优化器状态。若缓存冻结特征,应把缓存时间和存储量单独列出,否则训练成本比较会失真。

  10. 所有最大长度和填充策略共享检查点、划分、种子集合、更新次数、每次更新原始样本数和评价代码。动态填充只补到当前批次最长序列;全局填充补到规定最大长度。固定一批原始文本和批量大小计时,并记录有效词元数、存储词元数、被截断样本比例与被截断词元比例。

    报告测试指标、训练时间、固定批量预测时间、峰值内存、填充比例和编码器运行次数,并按原始长度分组报告误差。较短的最大长度通常计算更快,但可能丢失关键信息;动态填充通常可以减少浪费,但实际收益取决于批次中的长度分布。如果排序或分桶会改变批次组成,应保持样本顺序相同;也可以另外进行一次比较,在其他设置相同的情况下只改变排序或分桶方法,单独分析它带来的影响。