迁移学习 - 中文分类案例 - 模型训练
1 课程概览
本课讲解迁移学习中文分类案例的模型训练。流程:加载数据 → 初始化模型 → 配置训练参数 → 迭代训练 → 保存模型。冻结 BERT 参数(requires_grad=False)。使用 CrossEntropyLoss 损失函数和 Adam 优化器。训练三轮。
2 核心概念与定义
- requires_grad=False:不计算梯度,冻结参数。
- CrossEntropyLoss:交叉熵损失函数。
- Adam:优化器。
- model.train():设置模型为训练模式。
3 模型与算法详解
训练流程
1. 加载训练集
2. 创建模型并移动到设备
3. 冻结 BERT 预训练模型参数
4. 创建损失函数对象
5. 创建优化器对象
6. 设置模型为训练模式
7. 开始训练(3 轮)
- 获取本轮开始时间
- 分批次训练
- 计算损失
- 梯度清零
- 反向传播
- 梯度更新
8. 保存模型
冻结 BERT 参数
不计算 BERT 的梯度。
for param in my_bert_model.parameters():
param.requires_grad = False
损失函数
使用 CrossEntropyLoss,计算均值。
criterion = nn.CrossEntropyLoss(reduction='mean')
优化器
使用 Adam 优化器。
optimizer = torch.optim.Adam(my_model.parameters(), lr=learning_rate)
训练模式
设置模型为训练模式。
my_model.train()
4 数学原理与推导
损失函数
$$\text{loss} = -\frac{1}{N} \sum_{i=1}^{N} \sum_{c=1}^{C} y_{i,c} \log(\hat{y}_{i,c})$$
其中:
- $N$ 是样本数
- $C$ 是类别数(2)
- $y_{i,c}$ 是真实标签
- $\hat{y}_{i,c}$ 是预测概率
梯度更新
$$\theta = \theta - \eta \cdot \nabla_\theta \text{loss}$$
其中:
- $\theta$ 是模型参数
- $\eta$ 是学习率
- $\nabla_\theta \text{loss}$ 是梯度
5 代码示例
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
def train_model(train_dataset, my_tokenizer, my_bert_model, device,
batch_size=8, num_epochs=3, learning_rate=1e-5):
"""训练模型"""
# 1. 加载训练集(创建数据加载器)
train_dataloader = get_dataloader(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')
# 5. 创建优化器对象
optimizer = optim.Adam(my_model.parameters(), lr=learning_rate)
# 6. 设置模型为训练模式
my_model.train()
# 7. 开始训练(3 轮)
for epoch in range(num_epochs):
# 7.1 获取本轮开始时间
start_time = time.time()
# 7.2 分批次训练
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()
# 打印进度
if (i + 1) % 20 == 0:
print(f'Epoch [{epoch+1}/{num_epochs}], Step [{i+1}/{len(train_dataloader)}], Loss: {loss.item():.4f}')
# 打印本轮结果
elapsed_time = time.time() - start_time
print(f'Epoch {epoch+1} 完成, 平均损失: {total_loss/len(train_dataloader):.4f}, 用时: {elapsed_time:.2f}s')
# 8. 保存模型
torch.save(my_model.state_dict(), './model_classification.pth')
print('模型已保存')
# 测试
if __name__ == "__main__":
print("=== 迁移学习:中文分类案例 - 模型训练 ===")
print("训练函数已定义")
代码说明
| 代码 | 说明 |
|---|---|
param.requires_grad = False | 冻结参数 |
nn.CrossEntropyLoss(reduction='mean') | 交叉熵损失 |
optim.Adam(..., lr=learning_rate) | Adam 优化器 |
my_model.train() | 训练模式 |
optimizer.zero_grad() | 梯度清零 |
loss.backward() | 反向传播 |
optimizer.step() | 梯度更新 |
torch.save(...) | 保存模型 |
6 重难点与易错提醒
- ❗重点:冻结 BERT 参数:
param.requires_grad = False。 - ❗重点:训练三步:梯度清零 → 反向传播 → 梯度更新。
- ❗重点:使用 CrossEntropyLoss,不需要 LogSoftmax。
- ⚠️易错:训练时要设置
my_model.train()。 - 💡技巧:导包时记得
import torch.optim。
7 课堂问答精选
Q1:如何冻结 BERT 参数?
A:遍历 BERT 模型的参数,设置 param.requires_grad = False,不计算梯度。
Q2:训练的三步是什么?
A:①梯度清零(optimizer.zero_grad());②反向传播(loss.backward());③梯度更新(optimizer.step())。
Q3:使用什么损失函数?
A:使用 CrossEntropyLoss,计算均值(reduction='mean')。不需要 LogSoftmax。
Q4:训练几轮?
A:训练 3 轮。
8 本课小结
- 训练流程:加载数据 → 初始化模型 → 配置训练参数 → 迭代训练 → 保存模型。
- 冻结 BERT 参数:
requires_grad=False。 - 损失函数:CrossEntropyLoss。
- 优化器:Adam。
- 训练三步:梯度清零 → 反向传播 → 梯度更新。
9 延伸思考
- 如何进行模型评估?
- 评估流程是什么?