迁移学习 - 中文填空案例 - 模型评估
1 课程概览
本课讲解迁移学习中文填空案例的模型评估。与训练类似,需要修改三个地方:①过滤长度大于 32 的样本;②模型路径;③预测结果解码。准确率约 70%。需要将预测的标签解码成真实的字。
2 核心概念与定义
- 模型评估:使用测试集评估模型性能。
- 过滤长度大于 32 的样本:与训练一致。
- 解码:将预测的标签(id)解码成真实的字。
- 准确率:约 70%。
3 模型与算法详解
评估流程
1. 加载测试集
2. 过滤长度大于 32 的样本
3. 创建数据加载器(shuffle=False)
4. 修改模型路径(fill_mask)
5. 加载训练好的模型
6. 设置模型为评估模式
7. 迭代预测
8. 计算准确率
9. 解码预测结果(将 id 转成字)
10. 打印结果
与训练的区别
| 修改项 | 训练 | 评估 |
|---|---|---|
| ①过滤 | 过滤长度大于 32 | 过滤长度大于 32 |
| ②模型路径 | 保存模型 | 加载模型(fill_mask) |
| ③预测结果 | 不需要 | 解码成真实字 |
解码
将预测的标签(id)解码成真实的字。
- 预测的标签是 id(如 833、7231)
- 需要去词汇表里对字
- 使用
convert_ids_to_tokens()解码
准确率
约 70%。
4 数学原理与推导
准确率
$$\text{accuracy} = \frac{\text{correct}}{\text{total}}$$
解码
$$\text{word} = \text{vocab}[\text{id}]$$
其中:
- $\text{id}$ 是预测的标签
- $\text{vocab}$ 是词汇表
- $\text{word}$ 是解码后的字
5 代码示例
import torch
from tqdm import tqdm
def evaluate_model_fill_mask(test_dataset, my_tokenizer, my_bert_model, device,
batch_size=8, model_path='./model_fill_mask.pth'):
"""模型评估:中文填空"""
# 1. 优化一:过滤长度大于 32 的样本
test_dataset = test_dataset.filter(lambda x: len(x['text']) > 32)
# 2. 创建数据加载器(shuffle=False)
test_dataloader = get_dataloader(test_dataset, my_tokenizer, device, batch_size, shuffle=False)
# 3. 优化二:修改模型路径(fill_mask)
vocab_size = my_tokenizer.vocab_size
my_model = AiModel(my_bert_model, vocab_size).to(device)
my_model.load_state_dict(torch.load(model_path))
# 4. 初始化评估参数
correct = 0
total = 0
# 5. 设置模型为评估模式
my_model.eval()
# 6. 迭代预测
with torch.no_grad():
for i, (inputs, mask_positions) in enumerate(tqdm(test_dataloader, start=1)):
# 将参数移动到 GPU
inputs = {k: v.to(device) for k, v in inputs.items()}
# 前向传播
outputs = my_model(**inputs, mask_position=mask_positions[0])
# 获取预测结果
_, predicted = torch.max(outputs, 1)
# 统计预测正确的样本数
correct += (predicted == labels).sum().item()
total += labels.size(0)
# 优化三:解码预测结果(将 id 转成字)
if i == 1: # 只打印第一批
print(f"\n第一批预测结果:")
for j in range(min(8, len(predicted))):
# 原始文本
original_text = my_tokenizer.decode(inputs['input_ids'][j])
# 预测的标签
predicted_id = predicted[j].item()
predicted_word = my_tokenizer.convert_ids_to_tokens([predicted_id])[0]
# 真实的标签
true_id = labels[j].item()
true_word = my_tokenizer.convert_ids_to_tokens([true_id])[0]
print(f" 样本 {j+1}:")
print(f" 原始文本: {original_text}")
print(f" 预测字: {predicted_word} (id: {predicted_id})")
print(f" 真实字: {true_word} (id: {true_id})")
# 7. 计算准确率
accuracy = correct / total
# 8. 打印结果
print(f'\n测试集大小: {total}')
print(f'预测正确数: {correct}')
print(f'准确率: {accuracy:.4f} ({accuracy*100:.2f}%)')
return accuracy
# 测试
if __name__ == "__main__":
print("=== 迁移学习:中文填空案例 - 模型评估 ===")
print("评估函数已定义")
代码说明
| 代码 | 说明 |
|---|---|
test_dataset.filter(lambda x: len(x['text']) > 32) | 过滤长度大于 32 |
torch.load(model_path) | 加载模型(fill_mask) |
convert_ids_to_tokens([predicted_id]) | 将 id 转成字 |
my_tokenizer.decode(inputs['input_ids'][j]) | 解码原始文本 |
6 重难点与易错提醒
- ❗重点:评估时需要修改三个地方:过滤、模型路径、解码。
- ❗重点:准确率约 70%。
- ⚠️易错:预测的标签是 id,需要解码成字。
- 💡技巧:使用
convert_ids_to_tokens()解码。
7 课堂问答精选
Q1:评估时需要修改哪些地方?
A:修改三个地方:①过滤长度大于 32 的样本;②模型路径(fill_mask);③预测结果解码。
Q2:为什么要解码预测结果?
A:预测的标签是 id(如 833、7231),需要去词汇表里对字,解码成真实的字。
Q3:准确率是多少?
A:约 70%。
Q4:如何解码?
A:使用 convert_ids_to_tokens() 将 id 转成字。
8 本课小结
- 评估时修改三个地方:过滤、模型路径、解码。
- 准确率约 70%。
- 预测的标签是 id,需要解码成字。
- 使用
convert_ids_to_tokens()解码。
9 延伸思考
- NSP 任务是什么?
- 如何构建句子对?