231 lines
6.0 KiB
Markdown
231 lines
6.0 KiB
Markdown
# TensorFlow模型转换示例
|
||
|
||
本文档以 lenet 为例,介绍如何使用 Netrans 对 Tensorflow 模型进行转换。
|
||
|
||
Netrans 支持 TensorFlow 版本1.4.x, 2.0.x, 2.3.x, 2.6.x, 2.8.x, 2.10.x, 2.12.x 以tf.io.write_graph()保存的模型。
|
||
|
||
## 安装 Netrans
|
||
|
||
创建虚拟环境。
|
||
|
||
```bash
|
||
# 下载 mamba 安装脚本
|
||
wget "https://mirrors.tuna.tsinghua.edu.cn/github-release/conda-forge/miniforge/LatestRelease//Miniforge3-$(uname)-$(uname -m).sh"
|
||
# 创建 mamba 的安装目录
|
||
mkdir -p ~/app
|
||
# 安装 mamba 到 ~/app/
|
||
bash Miniforge3-Linux-x86_64.sh -b -p ${HOME}/app/miniforge3
|
||
# 添加 mamba 的初始化脚本到环境配置文件
|
||
echo "source " ${HOME}/app/miniforge3/etc/profile.d/mamba.sh"" >> ${HOME}/.bashrc
|
||
# 重新加载 ~/.bashrc 文件,使 mamba 初始化生效
|
||
source ${HOME}/.bashrc
|
||
# 创建一个名为 netrans 的虚拟环境,并安装 Python 3.8
|
||
mamba create -n netrans python=3.10 -y
|
||
# 激活 netrans 虚拟环境
|
||
mamba activate netrans
|
||
```
|
||
|
||
下载 Netrans
|
||
|
||
```bash
|
||
cd ~/app
|
||
git clone https://gitlink.org.cn/nudt_dsp/netrans.git
|
||
```
|
||
|
||
Netrans_cli 是基于 Netrans_api 封装的命令行工具,执行 `setup.sh` 可安装 Netrans_cli。
|
||
|
||
```bash
|
||
cd ~/app/netrans
|
||
# 执行 setup.sh
|
||
bash setup.sh
|
||
# setup.sh 会修改系统环境变量,需要 source 重新生效
|
||
source ~/.bashrc
|
||
# 重新激活 netrans 环境
|
||
mamba activate netrans
|
||
```
|
||
|
||
## 数据准备
|
||
|
||
示例使用 ONNX 格式的 yolov8s 模型,已经完成数据准备,可以使用下面命令进入目录执行。
|
||
|
||
```bash
|
||
cd netrans/
|
||
cd examples/infer_with_pre_post_process
|
||
# 激活 netrans 环境
|
||
mamba activate netrans
|
||
```
|
||
|
||
|
||
## 数据准备
|
||
转换 TensorFlow 模型时,模型工程目录应包含以下文件:
|
||
- .pb 文件:冻结图模型文件
|
||
- inputs_outputs.txt:输入输出节点定义文件
|
||
- dataset.txt:数据路径配置文件
|
||
|
||
我们的示例 已经完成数据准备,可以使用下面命令进入目录执行。
|
||
|
||
```bash
|
||
cd ~/app/netrans/examples/tensorflow
|
||
# 激活 netrans 环境
|
||
mamba activate netrans
|
||
```
|
||
|
||
此时目录如下:
|
||
```bash
|
||
lenet/
|
||
├── 0.jpg # 校准数据
|
||
├── dataset.txt # 指定数据地址的文件
|
||
├── inputs_outputs.txt # 输入输出节点定义文件
|
||
└── lenet.pb # 冻结图模型文件
|
||
```
|
||
|
||
|
||
## 使用 nertans_cli 命令行工具
|
||
|
||
### 模型导入
|
||
```bash
|
||
load lenet
|
||
```
|
||
|
||
该命令会在工程目录下生成包含模型信息的 .json 和 .data 数据文件。
|
||
此时 lenet 的目录结构如下:
|
||
|
||
```bash
|
||
lenet/
|
||
├── 0.jpg
|
||
├── dataset.txt
|
||
├── inputs_outputs.txt
|
||
├── lenet.data
|
||
├── lenet_inputmeta.yml
|
||
├── lenet.json
|
||
├── lenet.pb
|
||
└── lenet_postprocess_file.yml
|
||
```
|
||
|
||
### 模型量化
|
||
|
||
量化处理可优化模型的推理效率,加快模型的推理速度,我们使用以下命令对模型进行量化处理。量化模型需要两个参数:目录(模型)名字和量化类型。支持的量化类型包括:
|
||
symi8: 对称量化算法,使用 int8 类型
|
||
asymu8: 非对称量化算法,使用 uint8 类型
|
||
symi16: 对称量化算法,使用 int16 类型
|
||
|
||
```bash
|
||
quantize lenet asymu8
|
||
```
|
||
|
||
此时 lenet 的目录结构如下:
|
||
|
||
```bash
|
||
lenet/
|
||
├── 0.jpg
|
||
├── dataset.txt
|
||
├── inputs_outputs.txt
|
||
├── lenet_asymu8.quantize
|
||
├── lenet.data
|
||
├── lenet_inputmeta.yml
|
||
├── lenet.json
|
||
├── lenet.pb
|
||
└── lenet_postprocess_file.yml
|
||
```
|
||
|
||
### 模型导出
|
||
|
||
使用 `export` 将模型导出为 `nbg` 格式并生成应用程序工程。
|
||
|
||
```bash
|
||
export lenet asymu8
|
||
```
|
||
|
||
此时 lenet 的目录结构如下:
|
||
|
||
```bash
|
||
lenet/
|
||
├── 0.jpg
|
||
├── dataset.txt
|
||
├── inputs_outputs.txt
|
||
├── lenet_asymu8.quantize
|
||
├── lenet.data
|
||
├── lenet_inputmeta.yml
|
||
├── lenet.json
|
||
├── lenet.pb
|
||
├── lenet_postprocess_file.yml
|
||
└── wksp
|
||
├── lenet_asymu8
|
||
│ ├── analysis.json
|
||
│ ├── BUILD
|
||
│ ├── dump_core_graph.json
|
||
│ ├── graph.json
|
||
│ ├── lenetasymu8.2012.vcxproj
|
||
│ ├── lenet_asymu8.export.data
|
||
│ ├── lenetasymu8.vcxproj
|
||
│ ├── main.c
|
||
│ ├── makefile.linux
|
||
│ ├── vnn_global.h
|
||
│ ├── vnn_lenetasymu8.c
|
||
│ ├── vnn_lenetasymu8.h
|
||
│ ├── vnn_lenetasymu8_tensor.c
|
||
│ ├── vnn_post_process.c
|
||
│ ├── vnn_post_process.h
|
||
│ ├── vnn_pre_process.c
|
||
│ └── vnn_pre_process.h
|
||
└── lenet_asymu8_nbg_unify
|
||
├── BUILD
|
||
├── cmd.sh
|
||
├── lenetasymu8.2012.vcxproj
|
||
├── lenetasymu8.vcxproj
|
||
├── main.c
|
||
├── makefile.linux
|
||
├── nbg_meta.json
|
||
├── network_binary.nb
|
||
├── vnn_global.h
|
||
├── vnn_lenetasymu8.c
|
||
├── vnn_lenetasymu8.h
|
||
├── vnn_lenetasymu8_tensor.c
|
||
├── vnn_post_process.c
|
||
├── vnn_post_process.h
|
||
├── vnn_pre_process.c
|
||
└── vnn_pre_process.h
|
||
```
|
||
|
||
## 使用 Netrans_py Python API
|
||
|
||
### 示例代码
|
||
|
||
```python
|
||
# example.py
|
||
from netrans import Netrans
|
||
|
||
def main(model_path: str, quantize_type: str):
|
||
# 初始化 Netrans
|
||
net = Netrans()
|
||
|
||
# 导入模型并配置预处理参数
|
||
net.load(model_path, mean=[128, 128, 128], scale=[1, 1, 1])
|
||
|
||
# 模型量化
|
||
net.quantize(quantize_type)
|
||
|
||
# 配置前后处理加入推理计算图
|
||
net.add_pre_post(quantize_type, pre=True, post=True)
|
||
|
||
# 模型导出
|
||
net.export(quantize_type)
|
||
|
||
if __name__ == "__main__":
|
||
import argparse
|
||
|
||
parser = argparse.ArgumentParser(description="Netrans Model Conversion")
|
||
parser.add_argument("model_path", type=str, help="Path to the model directory")
|
||
parser.add_argument("-q", "--quantize", type=str, default="asymu8", help="Quantization type (default: asymu8)")
|
||
args = parser.parse_args()
|
||
|
||
main(args.model_path, args.quantize)
|
||
```
|
||
|
||
### 运行示例
|
||
|
||
```bash
|
||
python example.py lenet -q asymu8
|
||
```
|
||
|