自动模型方式 - 文本摘要任务
1 课程概览
本课讲解自动模型方式进行文本摘要。思路与之前类似。使用 AutoModelForSeq2SeqLM 加载模型。纯英文不需要 encode 编码。使用 my_model.generate() 生成摘要。使用 convert_ids_to_tokens() 将 id 转成可读文本。
2 核心概念与定义
- 文本摘要(Summarization):给一段长文本,输出一段概括或简单的文字。
- AutoModelForSeq2SeqLM:专门加载序列到序列语言模型的类。
- generate():生成摘要。
- convert_ids_to_tokens():将 id 转成可读文本。
3 模型与算法详解
模型加载
使用
AutoModelForSeq2SeqLM加载模型。
my_model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
纯英文不需要 encode
处理纯英文时,不需要 encode 编码。
inputs = my_tokenizer(text, return_tensors='pt')
文本摘要流程
0. 定义变量记录模型名
1. 加载 Tokenizer(分词器)
2. 加载模型(AutoModelForSeq2SeqLM)
3. 定义输入文本
4. 将输入文本转成模型可接收的输入格式
5. 喂给模型,获取预测结果
6. 打印预测结果(将 id 转成可读文本)
模型选择
使用 DistilBART 模型。
| 任务 | 模型 |
|---|---|
| 文本摘要 | distilbart |
generate() 参数
只传入 input_ids,不要传入所有内容。
output = my_model.generate(input_ids=inputs['input_ids'])
convert_ids_to_tokens()
将 id 转成可读文本。
text = my_tokenizer.convert_ids_to_tokens(output[0])
4 数学原理与推导
文本摘要
$$\text{summary} = \text{model}.\text{generate}(\text{input_ids})$$
其中:
- $\text{input_ids}$ 是输入文本的 id 序列
- $\text{summary}$ 是生成的摘要
序列到序列
$$\text{summary} = \text{Decoder}(\text{Encoder}(\text{text}))$$
5 代码示例
import torch
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
def dm05_summarization():
"""自动模型方式:文本摘要"""
# 0. 定义变量记录模型名
model_name = "C:/software/softwallg/pretrained_model/distilbart"
# 1. 加载 Tokenizer(分词器)
my_tokenizer = AutoTokenizer.from_pretrained(model_name)
# 2. 加载模型(使用 AutoModelForSeq2SeqLM)
my_model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
# 3. 定义输入文本(一段关于 BERT 模型的英文介绍)
text = """
BERT is a transformer-based machine learning technique for NLP.
It was developed by Google and introduced in 2018.
BERT stands for Bidirectional Encoder Representations from Transformers.
It has achieved state-of-the-art results on a wide array of NLP tasks.
"""
# 4. 将输入文本转成模型可接收的输入格式
# 纯英文不需要 encode 编码
inputs = my_tokenizer(
text,
return_tensors='pt' # 返回二维张量
)
print(f"输入张量: {inputs}")
# 5. 喂给模型,获取预测结果
my_model.eval()
with torch.no_grad():
# 只传入 input_ids
output = my_model.generate(input_ids=inputs['input_ids'])
print(f"输出 id: {output}")
# 6. 打印预测结果(将 id 转成可读文本)
summary = my_tokenizer.convert_ids_to_tokens(output[0])
print(f"\n摘要文本: {' '.join(summary)}")
return my_model
# 测试
if __name__ == "__main__":
print("=== 自动模型方式:文本摘要 ===")
model = dm05_summarization()
代码说明
| 代码 | 说明 |
|---|---|
AutoModelForSeq2SeqLM.from_pretrained() | 加载序列到序列模型 |
my_tokenizer(text, return_tensors='pt') | 纯英文不需要 encode |
my_model.generate(input_ids=...) | 生成摘要 |
convert_ids_to_tokens(output[0]) | 将 id 转成可读文本 |
6 重难点与易错提醒
- ❗重点:使用
AutoModelForSeq2SeqLM加载模型。 - ❗重点:纯英文不需要 encode 编码。
- ❗重点:generate() 只传入 input_ids,不要传入所有内容。
- ⚠️易错:使用
convert_ids_to_tokens()将 id 转成可读文本。 - 💡技巧:记不住 API 时,可在官网查看测试用例。
7 课堂问答精选
Q1:文本摘要使用什么类加载模型?
A:使用 AutoModelForSeq2SeqLM 加载模型。
Q2:纯英文需要 encode 编码吗?
A:不需要。纯英文可以直接使用 my_tokenizer(text, return_tensors='pt')。
Q3:generate() 如何传参?
A:只传入 input_ids,如 my_model.generate(input_ids=inputs['input_ids']),不要传入所有内容。
Q4:如何将 id 转成可读文本?
A:使用 convert_ids_to_tokens(),如 my_tokenizer.convert_ids_to_tokens(output[0])。
8 本课小结
- 文本摘要:给一段长文本,输出概括。
- 使用
AutoModelForSeq2SeqLM加载模型。 - 纯英文不需要 encode 编码。
- 使用
generate()生成摘要,只传入 input_ids。 - 使用
convert_ids_to_tokens()将 id 转成可读文本。
9 延伸思考
- 自动模型方式如何进行 NER 任务?
- NER 任务需要加载什么配置?