迁移学习 - NSP 案例 - 模型训练和评估
1 课程概览
本课讲解迁移学习 NSP 案例的模型训练和评估。与案例一(中文分类)一模一样,只改一个地方:模型保存名改为 NSP。训练速度快,准确率约 88%。评估时使用方式三的 collate_fn,模型路径改为 NSP,解码时跳过特殊符号。
2 核心概念与定义
- 模型训练:与案例一一模一样,只改模型保存名。
- 模型评估:与案例一类似,改 collate_fn(方式三)和模型路径。
- skip_special_tokens:跳过特殊符号(如 [CLS]、[SEP])。
- clean_up_tokenization_spaces:保留原始分词空格。
3 模型与算法详解
训练流程
与案例一一模一样,只改一个地方。
| 修改项 | 案例一 | NSP |
|---|---|---|
| 模型保存名 | classification | NSP |
训练过程
1. 加载数据
2. 创建模型
3. 冻结 BERT 参数
4. 创建损失函数和优化器
5. 分批次训练(20 批保存一次)
6. 保存模型(model_NSP.pth)
评估流程
与案例一类似,改两个地方。
| 修改项 | 案例一 | NSP |
|---|---|---|
| collate_fn | 方式一 | 方式三 |
| 模型路径 | classification | NSP |
解码参数
| 参数 | 说明 |
|---|---|
skip_special_tokens=True | 跳过特殊符号([CLS]、[SEP]) |
clean_up_tokenization_spaces=True | 保留原始分词空格 |
准确率
约 88%。
4 数学原理与推导
训练
$$\text{loss} = \text{CrossEntropyLoss}(\text{model}(x), y)$$
评估
$$\text{accuracy} = \frac{\text{correct}}{\text{total}}$$
5 代码示例
import torch
import torch.nn as nn
import torch.optim as optim
from tqdm import tqdm
def train_model_nsp(train_dataset, my_tokenizer, my_bert_model, device,
batch_size=8, num_epochs=3, learning_rate=1e-5):
"""训练模型:NSP 任务"""
# 1. 创建数据加载器(方式三)
train_dataloader = get_dataloader_nsp(train_dataset, my_tokenizer, device, batch_size)
# 2. 创建模型并移动到设备
my_model = AiModel(my_bert_model).to(device)
# 3. 冻结 BERT 参数
for param in my_model.bert.parameters():
param.requires_grad = False
# 4. 创建损失函数和优化器
criterion = nn.CrossEntropyLoss(reduction='mean')
optimizer = optim.Adam(my_model.parameters(), lr=learning_rate)
# 5. 设置模型为训练模式
my_model.train()
# 6. 分批次训练
for epoch in range(num_epochs):
total_loss = 0
for i, (inputs, labels) in enumerate(train_dataloader):
# 前向传播
outputs = my_model(**inputs)
# 计算损失
loss = criterion(outputs, labels)
# 梯度清零
optimizer.zero_grad()
# 反向传播
loss.backward()
# 梯度更新
optimizer.step()
total_loss += loss.item()
# 每 20 批打印一次
if (i + 1) % 20 == 0:
print(f'Epoch [{epoch+1}/{num_epochs}], Step [{i+1}], Loss: {loss.item():.4f}')
print(f'Epoch {epoch+1} 完成, 平均损失: {total_loss/len(train_dataloader):.4f}')
# 7. 保存模型(模型名改为 NSP)
torch.save(my_model.state_dict(), './model_NSP.pth')
print('模型已保存')
def evaluate_model_nsp(test_dataset, my_tokenizer, my_bert_model, device,
batch_size=8, model_path='./model_NSP.pth'):
"""模型评估:NSP 任务"""
# 1. 创建数据加载器(方式三,shuffle=False)
test_dataloader = get_dataloader_nsp(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)):
# 前向传播
outputs = my_model(**inputs)
# 获取预测结果
_, predicted = torch.max(outputs, 1)
# 统计预测正确的样本数
correct += (predicted == labels).sum().item()
total += labels.size(0)
# 打印第一批的预测结果
if i == 1:
print(f"\n第一批预测结果:")
for j in range(min(8, len(predicted))):
# 解码原始文本(跳过特殊符号)
original_text = my_tokenizer.decode(
inputs['input_ids'][j],
skip_special_tokens=True,
clean_up_tokenization_spaces=True
)
# 预测结果
predicted_label = predicted[j].item()
# 真实标签
true_label = labels[j].item()
print(f" 样本 {j+1}:")
print(f" 原始文本: {original_text}")
print(f" 预测: {predicted_label} ({'是' if predicted_label == 1 else '不是'})")
print(f" 真实: {true_label} ({'是' if true_label == 1 else '不是'})")
# 6. 计算准确率
accuracy = correct / total
print(f'\n准确率: {accuracy:.4f} ({accuracy*100:.2f}%)')
return accuracy
# 测试
if __name__ == "__main__":
print("=== 迁移学习:NSP 案例 - 模型训练和评估 ===")
print("训练和评估函数已定义")
print("准确率约 88%")
代码说明
| 代码 | 说明 |
|---|---|
torch.save(..., 'model_NSP.pth') | 保存模型(模型名改为 NSP) |
get_dataloader_nsp(...) | 使用方式三的 collate_fn |
skip_special_tokens=True | 跳过特殊符号 |
clean_up_tokenization_spaces=True | 保留原始分词空格 |
6 重难点与易错提醒
- ❗重点:训练时只改一个地方:模型保存名改为 NSP。
- ❗重点:评估时改两个地方:collate_fn(方式三)和模型路径(NSP)。
- ❗重点:准确率约 88%。
- ❗重点:解码时使用
skip_special_tokens=True跳过特殊符号。 - 💡技巧:训练速度快,因为句子长度短(44 vs 300)。
7 课堂问答精选
Q1:训练时需要修改哪些地方?
A:只改一个地方:模型保存名改为 NSP。其他与案例一一模一样。
Q2:评估时需要修改哪些地方?
A:改两个地方:①collate_fn 改为方式三;②模型路径改为 NSP。
Q3:准确率是多少?
A:约 88%。
Q4:skip_special_tokens 是什么意思?
A:跳过特殊符号(如 [CLS]、[SEP]),解码时忽略这些特殊符号。
8 本课小结
- 训练时只改一个地方:模型保存名改为 NSP。
- 评估时改两个地方:collate_fn(方式三)和模型路径(NSP)。
- 准确率约 88%。
- 解码时使用
skip_special_tokens=True跳过特殊符号。
9 延伸思考
- BERT 模型是什么?
- BERT 模型的架构是怎样的?