迁移学习 - NSP 案例 - 数据预处理
1 课程概览
本课讲解迁移学习 NSP 案例的数据预处理。需要创建数据加载器,使用 collate_fn(方式三)处理批次数据。提取句子对(sentence1, sentence2)和标签(label)。使用 batch_encode_plus 批量编码,最大长度为 50。
2 核心概念与定义
- collate_fn(方式三):NSP 任务的数据整理函数。
- 句子对:(sentence1, sentence2, label)。
- batch_encode_plus:批量编码句子对。
- max_length=50:最大长度。
3 模型与算法详解
数据格式
data 的维度是 [(sentence1, sentence2, label), ...]。
data = [(sentence1, sentence2, label), ...]
collate_fn(方式三)
处理批次数据,提取句子对和标签。
1. 提取句子对(sentence1, sentence2)
2. 提取标签(label)
3. 批量编码文本
4. 转成模型输入格式
提取句子对
使用
item[:2]提取前两个数据(sentence1, sentence2)。
sentences = [item[:2] for item in data]
提取标签
使用
item[2]提取标签。
labels = [item[2] for item in data]
batch_encode_plus 参数
| 参数 | 说明 |
|---|---|
sentences | 句子对列表 |
truncation=True | 启用截断 |
max_length=50 | 最大长度 |
padding='max_length' | 启用填充 |
return_tensors='pt' | 返回张量 |
最大长度 50
两个句子各 22 个字符,加上特殊字符,最大长度设为 50。
- 句子一:22 个字符
- 句子二:22 个字符
- 特殊字符:[CLS]、[SEP] 等
- 最大长度:50
4 数学原理与推导
数据预处理
$$\text{inputs} = \text{tokenizer}.\text{batch_encode_plus}(\text{sentences}, \text{max_length}=50)$$
其中:
- $\text{sentences}$ 是句子对列表
- $\text{max_length}=50$ 是最大长度
句子对
$$\text{sentences} = [(\text{s1}_1, \text{s2}_1), (\text{s1}_2, \text{s2}_2), ..., (\text{s1}_n, \text{s2}_n)]$$
5 代码示例
import torch
from transformers import BertTokenizer
def collate_fn_nsp(data, my_tokenizer, device, max_length=50):
"""数据整理函数:NSP 任务(方式三)"""
# data 的维度:[(sentence1, sentence2, label), ...]
# 1. 提取句子对(sentence1, sentence2)
sentences = [item[:2] for item in data] # 取前两个数据
# 2. 提取标签(label)
labels = [item[2] for item in data] # 取第三个数据
# 3. 批量编码文本
inputs = my_tokenizer.batch_encode_plus(
sentences, # 句子对列表
truncation=True, # 启用截断
max_length=max_length, # 最大长度 50
padding='max_length', # 启用填充
return_tensors='pt' # 返回张量
)
# 4. 将标签转成张量
labels_tensor = torch.tensor(labels)
# 5. 转移到设备
inputs = {k: v.to(device) for k, v in inputs.items()}
labels_tensor = labels_tensor.to(device)
return inputs, labels_tensor
def get_dataloader_nsp(dataset, my_tokenizer, device, batch_size=8, shuffle=True):
"""获取 DataLoader"""
from torch.utils.data import DataLoader
dataloader = DataLoader(
dataset,
batch_size=batch_size,
shuffle=shuffle,
collate_fn=lambda data: collate_fn_nsp(data, my_tokenizer, device)
)
return dataloader
# 测试
if __name__ == "__main__":
print("=== 迁移学习:NSP 案例 - 数据预处理 ===")
model_name = "bert-base-chinese"
my_tokenizer = BertTokenizer.from_pretrained(model_name)
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 模拟数据
data = [
('床前明月光疑是地上霜', '举头望明月低头思故乡', 1),
('春眠不觉晓处处闻啼鸟', '床前明月光疑是地上霜', 0),
]
# 测试
inputs, labels = collate_fn_nsp(data, my_tokenizer, device)
print(f"输入张量: {inputs['input_ids'].shape}")
print(f"标签: {labels}")
代码说明
| 代码 | 说明 |
|---|---|
item[:2] | 提取前两个数据(sentence1, sentence2) |
item[2] | 提取标签(label) |
batch_encode_plus(sentences, ...) | 批量编码句子对 |
max_length=50 | 最大长度 |
collate_fn=lambda data: ... | 自定义整理函数 |
6 重难点与易错提醒
- ❗重点:data 的维度是 [(sentence1, sentence2, label), ...]。
- ❗重点:使用
item[:2]提取句子对,item[2]提取标签。 - ❗重点:最大长度设为 50。
- ⚠️易错:collate_fn 使用方式三,与前两个案例不同。
7 课堂问答精选
Q1:data 的维度是什么?
A:data 的维度是 [(sentence1, sentence2, label), ...],每个元素是一个三元组。
Q2:如何提取句子对和标签?
A:使用 item[:2] 提取句子对(sentence1, sentence2),使用 item[2] 提取标签。
Q3:最大长度设为多少?
A:最大长度设为 50。两个句子各 22 个字符,加上特殊字符。
Q4:collate_fn 使用哪种方式?
A:使用方式三(collate_fn_nsp),与前两个案例不同。
8 本课小结
- data 的维度:[(sentence1, sentence2, label), ...]。
- 使用
item[:2]提取句子对,item[2]提取标签。 - 使用
batch_encode_plus批量编码。 - 最大长度设为 50。
- collate_fn 使用方式三。
9 延伸思考
- 如何搭建 NSP 模型?
- NSP 模型和中文分类模型有什么区别?