迁移学习 - 中文填空案例 - 模型训练
1 课程概览
本课讲解迁移学习中文填空案例的模型训练。与中文分类案例类似,但需要修改三个地方:①过滤长度大于 32 的样本;②模型名;③存储时的模型名。过滤长度大于 32 的样本是为了避免填充后预测无意义。
2 核心概念与定义
- 过滤长度大于 32 的样本:避免填充后预测无意义。
- filter 函数:用于过滤数据。
- 20 批:不是 20 轮,是 20 批(20P)。
3 模型与算法详解
与中文分类案例的区别
修改三个地方。
| 修改项 | 说明 |
|---|---|
| ①过滤长度大于 32 的样本 | 避免填充后预测无意义 |
| ②模型名 | 使用填空模型 |
| ③存储时的模型名 | 保存为 fill_mask |
为什么要过滤长度大于 32 的样本
如果文本长度小于 32,填充后预测第 16 个位置无意义。
- 最大长度设置为 32
- 如果文本只有 3 个字,填充成 32 个字
- 预测第 16 个位置时,本来就是填充的
- 预测意义不大
- 所以要真实长度大于 32
filter 函数
使用 filter 函数过滤数据。
train_dataset = train_dataset.filter(lambda x: len(x['text']) > 32)
训练流程
1. 加载训练集
2. 过滤长度大于 32 的样本
3. 创建模型并移动到设备
4. 冻结 BERT 参数
5. 创建损失函数和优化器
6. 设置模型为训练模式
7. 分批次训练(20 批保存一次)
8. 保存模型
4 数学原理与推导
过滤
$$\text{filtered_data} = \text{filter}(\text{data}, \text{len}(x) > 32)$$
训练
$$\text{loss} = \text{CrossEntropyLoss}(\text{model}(x), y)$$
5 代码示例
import torch
import torch.nn as nn
import torch.optim as optim
def train_model_fill_mask(train_dataset, my_tokenizer, my_bert_model, device,
batch_size=8, num_epochs=3, learning_rate=1e-5):
"""训练模型:中文填空"""
# 1. 优化一:过滤长度大于 32 的样本
train_dataset = train_dataset.filter(lambda x: len(x['text']) > 32)
# 2. 创建数据加载器
train_dataloader = get_dataloader(train_dataset, my_tokenizer, device, batch_size)
# 3. 创建模型并移动到设备
vocab_size = my_tokenizer.vocab_size
my_model = AiModel(my_bert_model, vocab_size).to(device)
# 4. 冻结 BERT 参数
for param in my_model.bert.parameters():
param.requires_grad = False
# 5. 创建损失函数和优化器
criterion = nn.CrossEntropyLoss(reduction='mean')
optimizer = optim.Adam(my_model.parameters(), lr=learning_rate)
# 6. 设置模型为训练模式
my_model.train()
# 7. 分批次训练
for epoch in range(num_epochs):
for i, (inputs, mask_positions) in enumerate(train_dataloader):
# 前向传播
outputs = my_model(**inputs, mask_position=mask_positions[0])
# 计算损失
loss = criterion(outputs, labels)
# 梯度清零
optimizer.zero_grad()
# 反向传播
loss.backward()
# 梯度更新
optimizer.step()
# 每 20 批打印一次
if (i + 1) % 20 == 0:
print(f'Epoch [{epoch+1}/{num_epochs}], Step [{i+1}], Loss: {loss.item():.4f}')
# 8. 保存模型(模型名改为 fill_mask)
torch.save(my_model.state_dict(), './model_fill_mask.pth')
print('模型已保存')
# 测试
if __name__ == "__main__":
print("=== 迁移学习:中文填空案例 - 模型训练 ===")
print("训练函数已定义")
代码说明
| 代码 | 说明 |
|---|---|
train_dataset.filter(lambda x: len(x['text']) > 32) | 过滤长度大于 32 的样本 |
my_tokenizer.vocab_size | 词汇表大小 |
optim.Adam(...) | Adam 优化器 |
torch.save(..., 'model_fill_mask.pth') | 保存模型(模型名改为 fill_mask) |
6 重难点与易错提醒
- ❗重点:过滤长度大于 32 的样本,避免填充后预测无意义。
- ❗重点:修改三个地方:过滤、模型名、存储模型名。
- ⚠️易错:20 批不是 20 轮,是 20P。
- 💡技巧:使用 filter 函数过滤数据。
7 课堂问答精选
Q1:为什么要过滤长度大于 32 的样本?
A:如果文本长度小于 32,填充成 32 个字后,预测第 16 个位置时本来就是填充的,预测意义不大。所以要真实长度大于 32。
Q2:与中文分类案例相比,需要修改哪些地方?
A:修改三个地方:①过滤长度大于 32 的样本;②模型名;③存储时的模型名。
Q3:如何过滤数据?
A:使用 filter 函数:train_dataset.filter(lambda x: len(x['text']) > 32)。
Q4:20 批和 20 轮有什么区别?
A:20 批是 20P(20 个 batch),不是 20 轮(20 个 epoch)。
8 本课小结
- 与中文分类案例类似,修改三个地方。
- 过滤长度大于 32 的样本,避免填充后预测无意义。
- 使用 filter 函数过滤数据。
- 20 批不是 20 轮。
9 延伸思考
- 如何进行模型评估?
- 评估时需要修改哪些地方?