迁移学习 - 中文分类案例 - 数据预处理
1 课程概览
本课讲解迁移学习中文分类案例的数据预处理。先定义处理一批数据的函数,再通过循环处理多批,最后获取 DataLoader。使用 BERT 分词器将文本转成模型输入格式。使用 batch_encode_plus 批量编码文本。
2 核心概念与定义
- collate_fn:处理批次数据的函数。
- batch_encode_plus:批量编码文本。
- truncation:启用文本截断。
- max_length:最大序列长度。
- padding='max_length':启用填充。
3 模型与算法详解
数据预处理流程
1. 定义处理一批数据的函数(collate_fn)
2. 提取文本和标签(X 和 Y)
3. 批量编码文本(batch_encode_plus)
4. 转成模型输入格式
5. 通过循环处理多批
6. 封装成 DataLoader
collate_fn 函数
处理一批数据(如 8 条),统一格式,转成模型输入。
def collate_fn(data):
# data 是 8 条数据
# 提取文本和标签
# 批量编码文本
# 返回模型输入
提取文本和标签
从每条数据中提取 text 和 label。
sentences = [item['text'] for item in data]
labels = [item['label'] for item in data]
batch_encode_plus 参数
| 参数 | 说明 |
|---|---|
sentences | 输入文本列表(待处理的文本列表) |
truncation=True | 启用文本截断 |
max_length=300 | 最大序列长度 |
padding='max_length' | 启用填充 |
return_tensors='pt' | 返回张量 |
encode_plus vs batch_encode_plus
| 函数 | 说明 |
|---|---|
encode_plus | 只能处理一条 |
batch_encode_plus | 批量编码,处理多条 |
max_length 的选择
正常应该先分析句子长度分布,再确定 max_length。
- 应该先分析句子长度分布
- 看大多数句子集中在多少
- 这里直接写 300,能囊括绝大多数
4 数学原理与推导
数据预处理
$$\text{inputs} = \text{tokenizer}.\text{batch_encode_plus}(\text{sentences}, \text{truncation}, \text{max_length}, \text{padding})$$
其中:
- $\text{sentences}$ 是文本列表
- $\text{truncation}$ 是截断
- $\text{max_length}$ 是最大长度
- $\text{padding}$ 是填充
批次处理
$$\text{batch} = {(\text{text}_1, \text{label}_1), (\text{text}_2, \text{label}_2), ..., (\text{text}_n, \text{label}_n)}$$
5 代码示例
import torch
from torch.utils.data import DataLoader
from transformers import BertTokenizer
def dm01_collate_fn(data, my_tokenizer, device):
"""数据整理函数:处理批次数据,统一格式,转成模型输入"""
# data 是 8 条数据
# 1. 提取文本和标签(X 和 Y)
sentences = [item['text'] for item in data]
labels = [item['label'] for item in data]
# 2. 批量编码文本(将文本转成模型输入格式)
inputs = my_tokenizer.batch_encode_plus(
sentences, # 输入文本列表
truncation=True, # 启用文本截断
max_length=300, # 最大序列长度
padding='max_length', # 启用填充
return_tensors='pt' # 返回张量
)
# 3. 将标签转成张量
labels_tensor = torch.tensor(labels)
# 4. 转移到设备
inputs = {k: v.to(device) for k, v in inputs.items()}
labels_tensor = labels_tensor.to(device)
return inputs, labels_tensor
def dm02_get_dataloader(dataset, my_tokenizer, device, batch_size=8):
"""获取 DataLoader"""
# 通过循环处理多批,封装成 DataLoader
dataloader = DataLoader(
dataset['train'],
batch_size=batch_size,
shuffle=True,
collate_fn=lambda data: dm01_collate_fn(data, my_tokenizer, device)
)
return dataloader
# 测试
if __name__ == "__main__":
print("=== 迁移学习:中文分类案例 - 数据预处理 ===")
# 假设已加载数据和分词器
model_name = "bert-base-chinese"
my_tokenizer = BertTokenizer.from_pretrained(model_name)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# dataloader = dm02_get_dataloader(dataset, my_tokenizer, device)
print("数据预处理函数已定义")
代码说明
| 代码 | 说明 |
|---|---|
collate_fn(data) | 处理批次数据的函数 |
[item['text'] for item in data] | 提取文本 |
[item['label'] for item in data] | 提取标签 |
batch_encode_plus(...) | 批量编码文本 |
truncation=True | 启用截断 |
max_length=300 | 最大序列长度 |
padding='max_length' | 启用填充 |
DataLoader(..., collate_fn=...) | 封装成 DataLoader |
6 重难点与易错提醒
- ❗重点:先定义处理一批数据的函数,再通过循环处理多批。
- ❗重点:使用
batch_encode_plus批量编码文本。 - ❗重点:
encode_plus只能处理一条,batch_encode_plus可以处理多条。 - ⚠️易错:max_length 应该先分析句子长度分布再确定。
- 💡技巧:之前深度学习时玩过这个数据,可以参考。
7 课堂问答精选
Q1:数据预处理的流程是什么?
A:①定义处理一批数据的函数(collate_fn);②提取文本和标签;③批量编码文本(batch_encode_plus);④转成模型输入格式;⑤通过循环处理多批;⑥封装成 DataLoader。
Q2:encode_plus 和 batch_encode_plus 有什么区别?
A:encode_plus 只能处理一条,batch_encode_plus 可以批量编码,处理多条。
Q3:batch_encode_plus 有哪些参数?
A:①sentences:输入文本列表;②truncation=True:启用截断;③max_length=300:最大序列长度;④padding='max_length':启用填充;⑤return_tensors='pt':返回张量。
Q4:max_length 应该如何确定?
A:应该先分析句子长度分布,看大多数句子集中在多少,再确定 max_length。这里直接写 300,能囊括绝大多数。
8 本课小结
- 数据预处理:先定义处理一批数据的函数,再通过循环处理多批。
- 使用
batch_encode_plus批量编码文本。 encode_plus只能处理一条,batch_encode_plus可以处理多条。- 参数:truncation、max_length、padding、return_tensors。
- max_length 应该先分析句子长度分布再确定。
9 延伸思考
- 如何搭建网络模型?
- 如何进行模型训练和评估?