netrans/examples/tensorflow/README.md

231 lines
6.0 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# 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
```