You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

为何tf.keras推理速度比TFLite慢75倍?

TFLite vs Keras Predict: 75x Speedup for Audio CNN Inference

测试结论

我最近针对基于简单CNN的音频数据预测任务,对比了两种推理方式的速度表现:

  • 使用tf.keras.Model.predict时,平均执行时间约为0.17秒
  • 使用tf.lite.Interpreter时,平均仅需0.002秒,速度提升约75倍!

我分别在两个环境下完成测试,两者的速度差异幅度相近:

  • 桌面端:Ubuntu 18.04,TensorFlow 2.1
  • 树莓派3B+:Raspbian Buster,复用完全相同的测试代码

更新:即使将tf.keras.Model.predict的batch_size手动设置为1后,它的速度仍然比TFLite慢65倍。

测试代码

下面是完整的测试代码test_tflite.py:

import os
import pathlib
import tensorflow as tf
from tensorflow.keras.models import model_from_json
import numpy as np
import time

# disable GPU
tf.config.set_visible_devices([], 'GPU')

parent = pathlib.Path(__file__).parent.absolute()

# path to Tensorflow model and weights
MODEL_PATH = os.path.join(parent, 'models/vd_model.json')
WEIGHTS_PATH = os.path.join(parent, 'models/model.30-0.97.h5')
INPUT_SHAPE = (1, 43, 40, 1)
NUM_RUN = 100

def predict_tflite(interpreter, input_details, output_details, data):
    interpreter.set_tensor(input_details[0]['index'], data)
    interpreter.invoke()
    output_data = interpreter.get_tensor(output_details[0]['index'])
    return output_data

def run():
    # Load Tensorflow model
    with open(MODEL_PATH, 'r') as f:
        model = model_from_json(f.read())
    model.load_weights(WEIGHTS_PATH)

    # Show model
    model.summary()

    # Convert to TFLite
    converter = tf.lite.TFLiteConverter.from_keras_model(model)
    tflite_model = converter.convert()

    interpreter = tf.lite.Interpreter(model_content=tflite_model)
    interpreter.allocate_tensors()

    input_details = interpreter.get_input_details()
    output_details = interpreter.get_output_details()

    predictions = []
    for i in range(NUM_RUN):
        # fake input data
        data = np.random.rand(*INPUT_SHAPE).astype(np.float32)

        # Tensorflow
        start_time = time.time()
        prediction = model.predict(data, batch_size=1)
        elapsed = time.time() - start_time

        # Tensoflow Lite
        start_time = time.time()
        prediction_tflite = predict_tflite(interpreter, input_details, output_details, data)
        elapsed_tflite = time.time() - start_time

        predictions.append(((elapsed, prediction), (elapsed_tflite, prediction_tflite)))

    # Make sure predictions are close
    for pred_tf, pred_tflite in predictions:
        if not np.all(np.isclose(pred_tf[1], pred_tflite[1])):
            print('Predictions are not close')

    # Compute average execution times
    tf_avg = np.mean([p[0] for p, _ in predictions])
    tflite_avg = np.mean([p[0] for _, p in predictions])

    print(f'TF: {tf_avg:.6f}')
    print(f'TFLite: {tflite_avg:.6f}')

if __name__ == "__main__":
    run()

树莓派执行结果

以下是树莓派3B+上的终端输出:

pi@raspberrypi:~/src/audio_monitoring/audio_monitoring/tests $ python3 test_tflite.py
Model: "sequential"
_________________________________________________________________
Layer (type)                 Output Shape              Param #   
=================================================================
conv2d (Conv2D)              (None, 43, 40, 16)        160       
_________________________________________________________________
batch_normalization (BatchNo (None, 43, 40, 16)        64        
_________________________________________________________________
activation (Activation)      (None, 43, 40, 16)        0         
_________________________________________________________________
max_pooling2d (MaxPooling2D) (None, 22, 20, 16)        0         
_________________________________________________________________
conv2d_1 (Conv2D)            (None, 22, 20, 32)        4640      
_________________________________________________________________
batch_normalization_1 (Batch (None, 22, 20, 32)        128       
_________________________________________________________________
activation_1 (Activation)    (None, 22, 20, 32)        0         
_________________________________________________________________
max_pooling2d_1 (MaxPooling2 (None, 1, 1, 32)          0         
_________________________________________________________________
dropout (Dropout)            (None, 1, 1, 32)          0         
_________________________________________________________________
flatten (Flatten)            (None, 32)                0         
_________________________________________________________________
dense (Dense)                (None, 4)                 132       
=================================================================
Total params: 5,124
Trainable params: 5,028
Non-trainable params: 96
_________________________________________________________________
TF: 0.168310
TFLite: 0.002269

内容的提问来源于stack exchange,提问作者jul

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.07 18:43:00