Zixuan/lstm (#503)

* add fairseq version

* add flask web
This commit is contained in:
Zixuan Chen
2020-05-07 13:52:28 +08:00
committed by GitHub
parent 712778ff0a
commit 39392ed9bc
2 changed files with 211 additions and 16 deletions
@@ -22,8 +22,8 @@ Copyright © Microsoft Corporation. All rights reserved.
* [推荐学习时长](#推荐学习时长)
* [案例详解](#案例详解)
* [程序结构](#程序结构)
* [工具包的选择](#工具包的选择)
* [数据收集](#数据收集)
* [工具包的选择](#工具包的选择)
* [数据预处理](#数据预处理)
* [模型训练](#模型训练)
* [模型推理](#模型推理)
@@ -175,21 +175,6 @@ pip3 install -r train_requirements.txt
由于在该结构中,NLP的核心内容在于上联生成下联,因此我们将会在案例中关注此部分的实现,并搭建一个简单的web应用将模型封装成api。
## 工具包的选择
想要完成一个自动生成对联的小程序,想法十分美好,但想要达到这个目标,光拍拍脑袋想想是不够的,需要训练出一个能完成对联生成的自然语言理解模型。于是乎,就有两个选择:
1. 自己写一套完成对联生成工作的深度学习模型。这个工作量相当之大,可能需要一个NLP专业团队来进行开发,调优。
2. 应用已有的深度学习模型,直接应用。这个选择比较符合客观需要。我们找到了两个工具包:
+ Tensor2Tensor 工具包:Tensor2Tensor(以下简称T2T)是由 Google Brain 团队使用和维护的开源深度学习模型库,支持多种数据集和模型。T2T 在 github 上有完整的介绍和用法,可以访问[这里](https://github.com/tensorflow/tensor2tensor)了解详细信息。
+ Fairseq 工具包:[Fairseq](https://github.com/pytorch/fairseq) 是 Facebook 推出的一个序列建模工具包,这个工具包允许研究和开发人员自定义训练翻译、摘要、语言模型等文本生成任务。这里是它的 PyTorch 实现。
本案例中,我们使用 T2T 工具包进行模型训练。
## 数据收集
有了模型,还需要数据。巧妇难为无米之炊,没有数据,什么都是浮云。数据从哪里来呢?GitHub 上有很多开源贡献者收集和整理了对联数据,可以进行下载使用。
@@ -200,6 +185,24 @@ pip3 install -r train_requirements.txt
3. 微软亚洲研究院提供的10万条对联数据(非公开数据)。
## 工具包的选择
想要完成一个自动生成对联的小程序,想法十分美好,但想要达到这个目标,光拍拍脑袋想想是不够的,需要训练出一个能完成对联生成的自然语言理解模型。于是乎,就有两个选择:
1. 自己写一套完成对联生成工作的深度学习模型。这个工作量相当之大,可能需要一个NLP专业团队来进行开发,调优。
2. 应用已有的深度学习模型,直接应用。这个选择比较符合客观需要。我们找到了两个工具包:Tensor2Tensor和Fairseq。
### Tensor2Tensor
Tensor2Tensor(以下简称T2T)是由 Google Brain 团队使用和维护的开源深度学习模型库,支持多种数据集和模型。T2T 在 github 上有完整的介绍和用法,可以访问[这里](https://github.com/tensorflow/tensor2tensor)了解详细信息。
在本案例中,我们将演示如何使用T2T工具包进行模型训练。
### Fairseq
[Fairseq](https://github.com/pytorch/fairseq) 是 Facebook 推出的一个序列建模工具包,这个工具包允许研究和开发人员自定义训练翻译、摘要、语言模型等文本生成任务。这里是它的 PyTorch 实现。
除了下面的使用T2T训练的版本外,我们也提供了[使用fairseq训练模型](./docs/fairseq.md)的教程。
## 数据预处理
### 生成源数据文件
@@ -0,0 +1,192 @@
# 使用fairseq训练模型
在开始之前,请先确保已成功安装fairseq。
```
pip install fairseq
```
## 数据预处理
开始预处理之前,我们先统一文件路径。
```
HOME_DIR=$(cd `dirname $0`; pwd)
RAW_DATA_DIR=${HOME_DIR}/fairseq-data
PREPROCESSED_DATA_DIR=${HOME_DIR}/data-bin/couplet
MODEL_SAVE_DIR=${HOME_DIR}/output/couplet
```
* `RAW_DATA_DIR`为上述的所有训练数据的存放目录
* `PREPROCESSED_DATA_DIR`是预处理文件的输出目录
* `MODEL_SAVE_DIR`为训练模型结果的保存目录
此后,我们需要将收集的训练数据分为训练集、验证集、测试集三部分。请确保对联的数据分为上联和下联两个文件,用换行符`\n`分隔每条上联或下联数据,每个字以空格隔开。
由于训练数据比较大,我们建议可以使用98:1:1的比例划分训练集、验证集、测试集,并将上联文件分别命名为`train.up``valid.up``test.up`,下联文件命名为`train.down``valid.down``test.down`
完成划分后,将所有文件存放至`RAW_DATA_DIR`目录。
接着执行以下脚本开始预处理数据。
```
fairseq-preprocess \
--source-lang up \
--target-lang down \
--trainpref ${RAW_DATA_DIR}/train \
--validpref ${RAW_DATA_DIR}/valid \
--testpref ${RAW_DATA_DIR}/test \
--destdir ${PREPROCESSED_DATA_DIR}
```
完成预处理后,生成的训练所需的二进制文件及上联和下联的字典文件都会保存在`PREPROCESSED_DATA_DIR`目录下。
## 模型训练
完成数据预处理后,执行以下脚本即可开始训练。
```
fairseq-train ${PREPROCESSED_DATA_DIR} \
--log-interval 100 \
--lr 0.25 \
--clip-norm 0.1 \
--dropout 0.2 \
--criterion label_smoothed_cross_entropy \
--save-dir ${MODEL_SAVE_DIR} \
-a lstm \
--max-tokens 4000 \
--max-epoch 100
```
其中,`-a`参数可以选择训练的模型,此处我们选择lstm进行训练。
更多的参数解释请参考[fairseq文档](https://fairseq.readthedocs.io/en/latest/command_line_tools.html#fairseq-train)。
训练完成后,模型文件会保存在`MODEL_SAVE_DIR`目录下。
## 模型推理
fairseq提供了两种模型推理的命令行工具,分别是fairseq-generate和fairseq-interactive。除此之外,我们还可以使用Python加载模型并完成推理。
### fairseq-generate
fairseq-generate的输入是二进制文件,会自动读取测试集的数据完成推理。
具体命令如下:
```
fairseq-generate ${PREPROCESSED_DATA_DIR} --path ${MODEL_SAVE_DIR}/checkpoint_best.pt --source-lang up --target-lang down
```
### fair-seq-interactive
fairseq-interactive提供了交互式命令行的方式推理,加载模型后,用户输入上联,模型将实时输出下联。
具体命令如下:
```
fairseq-interactive ${PREPROCESSED_DATA_DIR} --path ${MODEL_SAVE_DIR}/checkpoint_best.pt --source-lang up --target-lang down
```
更多参数请参考[fairseq文档](https://fairseq.readthedocs.io/en/latest/command_line_tools.html#fairseq-interactive)。
### 使用Python加载模型
为了搭建后端服务,我们可以使用Python加载模型,并完成推理。
下面以加载LSTM为例。
1. 引入模型
```
from fairseq.models.lstm import LSTMModel
```
2. 读入模型文件
第一个参数为checkpoints所在目录,`checkpoint_file`为需要读入的checkpoint的文件名,`data_name_or_path`为字典文件所在的目录。
```
model = LSTMModel.from_pretrained('./checkpoints',\
checkpoint_file='checkpoint_best.pt',\
data_name_or_path="DICT_PATH")
```
3. 推理
此处需要注意输入的文字之间需用空格隔开。
```
upper = '海内存知己'
down = model.translate(' '.join(list(upper)))
print(down) # 天 涯 若 比 邻
```
## 搭建Flask Web应用
1. 安装flask
```
pip3 install flask
```
2. 搭建服务
这一步我们将使用Python加载模型后,利用flask开启web服务。
```
from flask import Flask
from flask import request
from fairseq.models.lstm import LSTMModel # 引入模型
model = LSTMModel.from_pretrained('./checkpoints',\
checkpoint_file='checkpoint_best.pt',\
data_name_or_path="DICT_PATH") # 读入模型
app = Flask(__name__)
@app.route('/',methods=['GET'])
def get_couplet_down():
couplet_up = request.args.get('upper','')
couplet_down = model.translate(' '.join(list(couplet_up))) # 模型推理
couplet_down = couplet_down.replace(' ','')
return couplet_up + "," + couplet_down
```
3. 启动服务
在测试环境中,我们使用flask自带的web服务即可(注:生产环境应使用uwsgi+nginx部署,有兴趣的同学可以自行查阅资料)。
使用以下两条命令:
在Ubuntu下,
```
export FLASK_APP=app.py
python -m flask run
```
在Windows下,
```
set FLASK_APP=app.py
python -m flask run
```
此时,服务就启动啦。
我们仅需向后端 http://127.0.0.1:5000/ 发起get请求,并带上上联参数upper,即可返回生成的对联到前端。
请求示例: ```http://127.0.0.1:5000/?upper=海内存知己```
返回结果: ```海内存知己,天涯若比邻```
## 模型对比
除了LSTM外,我们还使用了fairseq中内置的几组模型进行训练对比,包括CNN、transformer。具体可用模型可以参考[fairseq Models](https://fairseq.readthedocs.io/en/latest/models.html)。
我们对模型训练了30个epoch。
速度和模型大小对比:
| 模型名称 |推理时间 (100句) | 训练时间 | 模型大小 |
|---|---|---|---|
| LSTM |0.7s| 2h48m| 103.9M |
| CNN | 2.9s|28h42m|679.1M |
| Transformer|1.1s|4h05m|207.9M|
训练效果对比:
| 模型名称 | Valid Loss | BLEU4 | Perplexity |
|---|---|---|---|
| LSTM |3.99| 12.14| 17.00 |
| CNN | 4.39|15.69| 19.39 |
| Transformer|4.24|10.29|10.77|
注:Transformer为2个Encoder和2个Decoder。
可见,LSTM的模型较小,训练和推理速度较快,loss也更小。