迁移学习 - 中文分类案例 - 模型评估
1 课程概览
本课讲解迁移学习中文分类案例的模型评估。流程:加载数据 → 创建加载器 → 加载模型 → 批量预测 → 计算准确率 → 打印结果。评估时不需要打乱数据(shuffle=False)。使用 model.eval() 设置评估模式。计算准确率:correct / total。
2 核心概念与定义
- 模型评估:使用测试集评估模型性能。
- model.eval():设置模型为评估模式。
- 准确率(Accuracy):correct / total。
- shuffle=False:评估时不打乱数据。
3 模型与算法详解
评估流程
1. 加载测试集
2. 创建数据加载器(shuffle=False)
3. 加载训练好的模型
4. 初始化评估参数(correct=0, total=0)
5. 设置模型为评估模式(model.eval())
6. 迭代预测
- 将参数移动到 GPU
- 前向传播
- 统计预测正确的样本数
7. 计算准确率
8. 打印结果
评估 vs 训练
| 步骤 | 训练 | 评估 |
|---|---|---|
| 模式 | model.train() | model.eval() |
| shuffle | True | False |
| 梯度 | 计算 | 不计算 |
| 损失 | 计算 | 不计算 |
加载模型参数
使用
load_state_dict加载训练好的模型参数。
my_model.load_state_dict(torch.load(path))
准确率计算
correct / total
accuracy = correct / total
4 数学原理与推导
准确率
$$\text{accuracy} = \frac{\text{correct}}{\text{total}}$$
其中:
- $\text{correct}$ 是预测正确的样本数
- $\text{total}$ 是总样本数
预测
$$\hat{y} = \arg\max \text{model}(x)$$
其中:
- $x$ 是输入
- $\hat{y}$ 是预测标签
5 代码示例
import torch
from torch.utils.data import DataLoader
from tqdm import tqdm
def evaluate_model(test_dataset, my_tokenizer, my_bert_model, device,
batch_size=8, model_path='./model_classification.pth'):
"""模型评估"""
# 1. 加载测试集(创建数据加载器,shuffle=False)
test_dataloader = get_dataloader(
test_dataset, my_tokenizer, device, batch_size, shuffle=False
)
# 2. 加载训练好的模型
my_model = AiModel(my_bert_model).to(device)
my_model.load_state_dict(torch.load(model_path))
# 3. 初始化评估参数
correct = 0 # 预测正确的样本数
total = 0 # 总样本数
# 4. 设置模型为评估模式
my_model.eval()
# 5. 迭代预测
with torch.no_grad():
for i, (inputs, labels) in enumerate(tqdm(test_dataloader, start=1)):
# 5.1 将参数移动到 GPU
inputs = {k: v.to(device) for k, v in inputs.items()}
labels = labels.to(device)
# 5.2 前向传播
outputs = my_model(**inputs)
# 5.3 获取预测结果
_, predicted = torch.max(outputs, 1)
# 5.4 统计预测正确的样本数
correct += (predicted == labels).sum().item()
total += labels.size(0)
# 6. 计算准确率
accuracy = correct / total
# 7. 打印结果
print(f'\n测试集大小: {total}')
print(f'预测正确数: {correct}')
print(f'准确率: {accuracy:.4f} ({accuracy*100:.2f}%)')
return accuracy
# 测试
if __name__ == "__main__":
print("=== 迁移学习:中文分类案例 - 模型评估 ===")
print("评估函数已定义")
代码说明
| 代码 | 说明 |
|---|---|
shuffle=False | 评估时不打乱数据 |
my_model.load_state_dict(...) | 加载模型参数 |
torch.load(path) | 加载保存的模型 |
my_model.eval() | 评估模式 |
with torch.no_grad(): | 不计算梯度 |
torch.max(outputs, 1) | 取最大值索引 |
correct / total | 计算准确率 |
6 重难点与易错提醒
- ❗重点:评估时
shuffle=False,不打乱数据。 - ❗重点:使用
model.eval()设置评估模式。 - ❗重点:使用
with torch.no_grad():不计算梯度。 - ⚠️易错:加载模型参数用
load_state_dict。 - 💡技巧:使用 tqdm 显示进度条。
7 课堂问答精选
Q1:模型评估的流程是什么?
A:①加载测试集;②创建数据加载器(shuffle=False);③加载训练好的模型;④初始化评估参数;⑤设置模型为评估模式;⑥迭代预测;⑦计算准确率;⑧打印结果。
Q2:评估时为什么要 shuffle=False?
A:评估时不需要打乱数据,正常按顺序预测即可。
Q3:如何加载训练好的模型?
A:使用 my_model.load_state_dict(torch.load(path)) 加载模型参数。
Q4:如何计算准确率?
A:准确率 = 预测正确的样本数 / 总样本数,即 correct / total。
8 本课小结
- 评估流程:加载数据 → 创建加载器 → 加载模型 → 批量预测 → 计算准确率。
- 评估时
shuffle=False。 - 使用
model.eval()设置评估模式。 - 使用
with torch.no_grad():不计算梯度。 - 准确率 = correct / total。
9 延伸思考
- 中文填空案例如何进行?
- 中文填空和中文分类有什么区别?