在进行边缘端微控制器(MCU)开发时,我遇到了一个棘手的问题:手里已有的 INT8 TFLite 模型文件,在加载时无法通过 TFLite Micro 的 Schema 校验。

本文记录了将两个小型语音模型(KWS 语音唤醒、SV 声纹识别)重新转换为完全兼容 Schema 的纯 INT8 TFLite 模型的全过程。涉及环境排雷、工具链选择、维度自动适配以及算子级优化。

一、 问题背景与根本原因

目标设备为资源受限的 MCU,要求模型必须满足:

  1. 纯 INT8 算子,输入输出均为 INT8。
  2. 不能包含 Flex 算子(SELECT_TF_OPS)。
  3. 不能包含 dequantize/quantize 节点(部分纯 INT8 加速后端不支持中间浮点计算)。

我最初拥有的 INT8 TFLite 文件不兼容 Schema,分析后发现根本原因在于转换工具链的选择不当:原始的 ONNX 模型已经被量化过(包含 ConvInteger 等算子),而旧的转换工具无法正确处理这些算子,导致生成的 TFLite 文件存在底层结构错误。

二、 走过的弯路:从体积膨胀到环境死锁

在这一阶段,我走过两次明显的弯路。

1. 第一次弯路:工具链选择与体积膨胀

最初,我尝试使用传统的 onnx-tf 方案,思路是 ONNX -> TF SavedModel -> TFLite
但原始输入的 ONNX 模型已经包含了 ONNX 原生量化算子。onnx-tf 无法直接映射这些量化算子,而是强行将其拆解为底层的浮点运算组合。这导致在后续的 TFLite 量化阶段,转换器无法将这些离散的浮点算子重新融合为高效的 INT8 算子,只能降级引入 Flex 算子以兼容运行。结果就是模型体积从 12KB 暴增至 48KB,且完全无法在 MCU 上运行。

2. 第二次弯路:Python 环境的死锁

为了获取纯 INT8 算子,我决定弃用老旧的 onnx-tf,改用目前针对边缘端优化较好的 onnx2tf 工具。但这一步直接引发了环境灾难:
我直接在原有环境中执行了 pip install onnx2tf。该工具的依赖极其严格,安装时强制将环境中的 TensorFlow 从 2.13 升级到了 2.19。这一升级直接导致原有的 onnx-tf 及其他相关依赖包因版本不兼容而彻底报废。Python 3.12环境直接瘫痪。

3. 环境重建与依赖说明

为了解决这一问题,我直接回退到 Python 3.10 ,并锁定了以下核心依赖版本。如果你要复现本流程,请严格按照以下命令安装:

# 核心深度学习框架
pip install tensorflow==2.19.0

# ONNX 相关工具链
pip install onnx==1.16.1 onnxsim==0.4.33 onnxruntime==1.18.1

# 转换工具
pip install onnx2tf==1.23.0

# 辅助库
pip install flatbuffers==24.3.25 numpy

或者直接一键安装:

pip install tensorflow==2.19.0 onnx==1.16.1 onnxsim==0.4.33 onnxruntime==1.18.1 onnx2tf==1.23.0 flatbuffers==24.3.25 numpy

三、 正确的转换路径:浮点 ONNX + onnx2tf

解决环境问题后,我重新审视了 onnx2tf 的设计逻辑:它期望接收标准的浮点 ONNX 模型,在转换过程中由自身完成量化操作,因此它不支持解析已经包含量化算子的 ONNX 模型。

明确了这一点后,正确的方案浮出水面:

  1. 放弃已量化的 ONNX 模型,寻找并使用这两个模型的原始浮点版本(仅包含 ConvReluGemm 等标准算子)。
  2. 将浮点 ONNX 喂给 onnx2tf 转换为 TensorFlow SavedModel。
  3. 使用 TFLiteConverter 配合校准数据集执行 INT8 量化。

四、 维度错位陷阱:NCHW 与 NHWC

在执行 TFLite 量化校准时,系统报错:维度溢出或元素个数不匹配。

原因分析:
onnx2tf 在转换时会自动调整数据排布以符合 TensorFlow 规范。例如:

  • KWS 模型的输入从 [1,1,100,13] 转换为 [1,100,13,1](NCHW 转 NHWC)。
  • SV 模型将开头的 Transpose 算子直接折叠进输入维度,变成了 [1,13,40]
    而此时提供给 TFLite 的校准数据仍是原始维度,导致校准器初始化崩溃。

解决方案:动态 Shape 读取与自适应
不要在代码中写死输入维度。在量化前,增加一段逻辑动态读取 onnx2tf 生成的 SavedModel 的真实输入 Shape,然后编写自适应函数,根据读取到的真实 Shape 对校准数据进行 Transpose 或 Reshape 操作,确保数据维度严格匹配。

五、 算子级优化:剥离 SV 模型的 L2 归一化

经过上述修改,两个模型都成功生成了无 Flex 算子的 TFLite 文件。但通过解析生成的二进制文件,我发现 SV 模型中仍然存在 3 个 dequantize 和 1 个 quantize 节点。

原因分析:
SV 模型末尾包含 L2 归一化操作(SquareSqrtMaximumDiv)。在 TFLite 中,Square 和 Sqrt 算子不支持 INT8 计算,系统被迫在这些算子前后插入 int8 转 float、float 转 int8 的节点。这在标准 TFLite Micro 环境下可以运行,但破坏了“纯 INT8”的初衷,且在部分极端受限的加速后端上依然会报错。

解决方案:图截断与端侧实现
L2 归一化逻辑简单,没必要在模型内部用复杂的算子组合实现。我使用 Python 的 onnx API 对原始 ONNX 图进行了截断:

  1. 定位到 fc/Gemm 节点。
  2. 移除其后的 ReduceL2ClipDiv 等所有归一化节点。
  3. 将 Gemm 的输出直接设为模型输出。
  4. 重新执行转换流程。

优化后,SV 模型成功变为纯 INT8。被剥离的 L2 归一化,在 MCU 端只需用 5 行 C 代码实现即可:

void l2_normalize(float* output, int len) {
    float norm = 0.0f;
    for (int i = 0; i < len; i++) norm += output[i] * output[i];
    norm = sqrtf(norm);
    if (norm < 1e-12f) norm = 1e-12f; // Maximum(eps)
    for (int i = 0; i < len; i++) output[i] /= norm;
}

六、 最终结果

经过全流程优化,最终生成的模型完全符合 Schema 标准:

  • KWS 模型:10.9 KB,无 Flex 算子,无 dequantize 节点,纯 INT8。
  • SV 模型:15.6 KB,无 Flex 算子,无 dequantize 节点,纯 INT8。

附:完整工作流 Python 脚本

以下脚本整合了环境适配、ONNX 图截断、维度自适应匹配和 TFLite 纯 INT8 量化全流程。修改路径参数后即可直接运行。

# 运行前请确保已安装以下依赖:
# pip install tensorflow==2.19.0 onnx==1.16.1 onnxsim==0.4.33 onnxruntime==1.18.1 onnx2tf==1.23.0 flatbuffers==24.3.25 numpy
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
模型优化转换脚本 (以 SV 模型为例)
- 从 ONNX 模型中截断 L2 归一化部分 (ReduceL2/Constant/Clip/Shape/Expand/Div)
- 只保留到 fc/Gemm 输出
- 动态适配 onnx2tf 产生的维度变化
- 执行纯 INT8 量化
"""

import os
import shutil
import subprocess
import numpy as np
import onnx
from onnx import helper, TensorProto
from onnx2tf import onnx2tf
import tensorflow as tf

# ================= 路径配置 =================
BASE = "/mnt/workspace"

SV_ONNX = f"{BASE}/vs/himia/models/model_sv.onnx"
SV_SIM = f"{BASE}/vs/himia/models/model_sv_sim.onnx"
SV_TRIMMED = f"{BASE}/vs/himia/models/model_sv_trimmed.onnx"
SV_DATA = f"{BASE}/vs/himia/features/sv_test.npy"
SV_OUTPUT = f"{BASE}/vs/himia/models/model_sv_pure_int8.tflite"

NUM_CALIBRATION_SAMPLES = 200


# ================= 核心函数 =================
def trim_l2_normalization(onnx_path, output_path):
    """从 SV 模型中移除 L2 归一化部分,截断到 Gemm 输出"""
    print("  [trim] 加载 ONNX 模型...")
    model = onnx.load(onnx_path)
    graph = model.graph

    # 找到 fc/Gemm 的输出节点
    gemm_node = None
    for node in graph.node:
        if node.op_type == "Gemm" and "fc" in node.name:
            gemm_node = node
            break

    if gemm_node is None:
        for node in graph.node:
            if node.op_type == "Gemm":
                gemm_node = node

    if gemm_node is None:
        raise ValueError("找不到 Gemm 节点")

    fc_output_name = gemm_node.output[0]
    print(f"  [trim] fc/Gemm 输出: {fc_output_name}")

    # 收集需要保留的节点 (到 Gemm 为止)
    nodes_to_keep = []
    for node in graph.node:
        nodes_to_keep.append(node)
        if fc_output_name in node.output:
            break

    print(f"  [trim] 保留 {len(nodes_to_keep)} 个节点 (原始 {len(graph.node)} 个)")
    print(f"  [trim] 移除的节点:")
    for node in graph.node:
        if node not in nodes_to_keep:
            print(f"    - {node.op_type}: {node.name}")

    # 获取 fc 输出的 ValueInfo (保持原始输出名,不改名)
    fc_output_vi = None
    for vi in graph.value_info:
        if vi.name == fc_output_name:
            fc_output_vi = vi
            break

    if fc_output_vi is not None:
        new_output = fc_output_vi
        print(f"  [trim] 输出 shape: {[d.dim_value for d in fc_output_vi.type.tensor_type.shape.dim]}")
    else:
        new_output = helper.make_tensor_value_info(
            fc_output_name,
            TensorProto.FLOAT,
            [1, 16]
        )
        print(f"  [trim] 使用默认输出 shape: [1, 16]")

    # 收集保留节点用到的所有 initializer
    keep_names = set()
    for node in nodes_to_keep:
        for inp in node.input:
            keep_names.add(inp)

    initializers_to_keep = []
    for init in graph.initializer:
        if init.name in keep_names:
            initializers_to_keep.append(init)

    print(f"  [trim] 保留 {len(initializers_to_keep)} 个 initializer (原始 {len(graph.initializer)} 个)")

    # 重建 graph,保持原始输出名
    new_graph = helper.make_graph(
        nodes_to_keep,
        graph.name,
        graph.input,
        [new_output],
        initializers_to_keep,
    )

    new_model = helper.make_model(new_graph)
    new_model.ir_version = model.ir_version
    del new_model.opset_import[:]
    for opset in model.opset_import:
        new_model.opset_import.add().CopyFrom(opset)

    try:
        onnx.checker.check_model(new_model)
        print("  [trim] 模型验证通过")
    except Exception as e:
        print(f"  [trim] 模型验证警告: {e}")

    onnx.save(new_model, output_path)
    print(f"  [trim] 保存截断模型: {output_path}")
    return output_path


def get_onnx_input_info(onnx_path):
    model = onnx.load(onnx_path)
    info = []
    for inp in model.graph.input:
        name = inp.name
        shape = []
        for d in inp.type.tensor_type.shape.dim:
            if d.dim_value > 0:
                shape.append(d.dim_value)
            elif d.dim_param:
                shape.append(d.dim_param)
            else:
                shape.append("?")
        info.append((name, shape))
    return info


def simplify_onnx(onnx_path, sim_path, input_shape_str):
    print("  [onnxsim] 简化模型...")
    input_info = get_onnx_input_info(onnx_path)
    input_name = input_info[0][0]

    cmd = [
        "onnxsim", onnx_path, sim_path,
        "--overwrite-input-shape",
        f"{input_name}:{input_shape_str}",
    ]
    result = subprocess.run(cmd, capture_output=True, text=True)
    if result.returncode != 0:
        print(f"  [onnxsim] 简化失败: {result.stderr}")
        return None
    print(f"  [onnxsim] 完成: {sim_path}")
    return sim_path


def get_saved_model_input_shape(saved_model_dir):
    """动态读取 SavedModel 的真实输入维度"""
    loaded = tf.saved_model.load(saved_model_dir)
    concrete_func = loaded.signatures['serving_default']
    input_signature = concrete_func.structured_input_signature[1]
    input_key = list(input_signature.keys())[0]
    input_spec = input_signature[input_key]
    shape = tuple(input_spec.shape.as_list())
    print(f"  SavedModel 输入 shape: {shape}")
    return shape


def make_representative_dataset(calib_data, target_shape):
    """根据 SavedModel 真实 Shape 自适应生成校准数据"""
    target_shape_no_batch = tuple(target_shape[1:])
    total_elements = int(np.prod(target_shape_no_batch))

    def gen():
        for i in range(len(calib_data)):
            sample = calib_data[i].astype(np.float32)
            current_shape = sample.shape

            if len(current_shape) == len(target_shape_no_batch):
                if current_shape == target_shape_no_batch:
                    pass
                elif len(current_shape) == 2:
                    sample = sample.T
                else:
                    from itertools import permutations
                    found = False
                    for perm in permutations(range(len(current_shape))):
                        if tuple(current_shape[p] for p in perm) == target_shape_no_batch:
                            sample = np.transpose(sample, perm)
                            found = True
                            break
                    if not found:
                        sample = sample.reshape(target_shape_no_batch)
            elif sample.size == total_elements:
                sample = sample.reshape(target_shape_no_batch)
            else:
                continue

            sample = np.expand_dims(sample, axis=0).astype(np.float32)
            yield [sample]

    return gen


# ================= 主流程 =================
def convert_sv():
    print(f"\n{'='*60}")
    print(f"SV 模型优化转换 (移除 L2 归一化)")
    print(f"{'='*60}")

    # Step 1: 加载校准数据
    print("\n[Step 1] 加载校准数据...")
    calib_data = np.load(SV_DATA)
    print(f"  数据 shape: {calib_data.shape}")
    calib_data = calib_data[:NUM_CALIBRATION_SAMPLES].astype(np.float32)

    # Step 2: 截断 L2 归一化
    print("\n[Step 2] 截断 L2 归一化...")
    trim_path = trim_l2_normalization(SV_ONNX, SV_TRIMMED)

    # Step 3: onnxsim
    print("\n[Step 3] onnxsim 简化...")
    sim_path = simplify_onnx(SV_TRIMMED, SV_SIM, "1,40,13")
    use_path = sim_path if sim_path else SV_TRIMMED
    print(f"  使用: {use_path}")

    # Step 4: onnx2tf
    print(f"\n[Step 4] onnx2tf 转换...")
    onnx2tf_output = SV_OUTPUT + "_onnx2tf_out"
    if os.path.exists(onnx2tf_output):
        shutil.rmtree(onnx2tf_output)

    onnx2tf.convert(
        input_onnx_file_path=use_path,
        output_folder_path=onnx2tf_output,
        output_signaturedefs=True,
        non_verbose=True,
    )

    saved_model_dir = os.path.join(onnx2tf_output, "saved_model")
    if not os.path.exists(saved_model_dir):
        saved_model_dir = onnx2tf_output

    # Step 5: 检测输入 shape
    print("\n[Step 5] 检测输入 shape...")
    actual_shape = get_saved_model_input_shape(saved_model_dir)

    # Step 6: TFLite INT8 量化
    print("\n[Step 6] INT8 量化...")
    converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir)
    converter.optimizations = [tf.lite.Optimize.DEFAULT]
    converter.representative_dataset = make_representative_dataset(calib_data, actual_shape)
    converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8]
    converter.inference_input_type = tf.int8
    converter.inference_output_type = tf.int8

    try:
        tflite_model = converter.convert()
        print("  纯 INT8 转换成功!")
    except Exception as e:
        print(f"  纯 INT8 失败: {e}")
        print("  尝试允许 SELECT_TF_OPS...")
        converter.target_spec.supported_ops = [
            tf.lite.OpsSet.TFLITE_BUILTINS_INT8,
            tf.lite.OpsSet.SELECT_TF_OPS,
        ]
        tflite_model = converter.convert()

    with open(SV_OUTPUT, 'wb') as f:
        f.write(tflite_model)

    size_kb = os.path.getsize(SV_OUTPUT) / 1024
    print(f"\n[完成] {SV_OUTPUT} ({size_kb:.1f} KB)")

    # 清理中间文件
    shutil.rmtree(onnx2tf_output, ignore_errors=True)
    for p in [SV_TRIMMED, SV_SIM]:
        if os.path.exists(p):
            os.remove(p)

    return SV_OUTPUT


if __name__ == "__main__":
    print("=" * 60)
    print("SV 模型优化: 移除 L2 归一化 → 纯 INT8 TFLite")
    print("=" * 60)
    print()
    print("原理:")
    print("  原始: ...→ Gemm(fc) → ReduceL2 → Clip → Div → output")
    print("  优化: ...→ Gemm(fc) → output")
    print("  L2 归一化在 MCU 端用 C 代码实现 (仅需 5 行)")
    print()

    result = convert_sv()

    if result:
        size = os.path.getsize(result) / 1024
        print(f"\n优化结果: {result} ({size:.1f} KB)")
Logo

免费领 150 小时云算力,进群参与显卡、AI PC 幸运抽奖

更多推荐