如何保存状态形状依赖输入的有状态TFLite模型?
实现状态形状依赖输入的TensorFlow Lite有状态图
问题背景
我需要在TensorFlow Lite中实现一个有状态功能:让状态变量的形状在模型加载时确定,而非保存模型时。以计算连续输入的差值为例,模型需要返回连续两次调用输入的差值,预期通过以下测试:
func = load_tflite_model_func(tflite_model_file_path) runtime_shape = 60, 80 rng = np.random.RandomState(1234) ims = [rng.randn(*runtime_shape).astype(np.float32) for _ in range(3)] assert np.allclose(func(ims[0]), ims[0]) assert np.allclose(func(ims[1]), ims[1]-ims[0]) assert np.allclose(func(ims[2]), ims[2]-ims[1])
当前实现及问题
我当前的模型创建和保存代码如下:
from dataclasses import dataclass import tempfile import tensorflow as tf from typing import Optional @dataclass class TimeDelta(tf.Module): _last_val: Optional[tf.Tensor] = None def compute_delta(self, arr: tf.Tensor): if self._last_val is None: self._last_val = tf.Variable(tf.zeros(tf.shape(arr))) delta = arr-self._last_val self._last_val.assign(arr) return delta compile_time_shape = 30, 40 # compile_time_shape = None, None # 会触发UnliftableError tflite_model_file_path = tempfile.mktemp() delta = TimeDelta() save_signatures_to_tflite_model( {'delta': tf.function(delta.compute_delta, input_signature=[tf.TensorSpec(shape=compile_time_shape)])}, path=tflite_model_file_path, parent_object=delta )
现在遇到两个核心问题:
- 若编译时设置的形状与运行时输入形状不一致,程序直接崩溃;
- 尝试设置动态形状
compile_time_shape = None, None时,保存模型会触发UnliftableError,因为TensorFlow要求变量必须有具体维度。
解决方案
要实现状态形状依赖输入的有状态TFLite模型,需绕过静态变量维度限制,采用延迟初始化+动态形状适配的思路,具体步骤如下:
1. 修改模块实现,适配动态形状的状态变量
调整TimeDelta类,确保状态变量在第一次收到输入时才初始化,且能适配输入形状的变化:
@dataclass class TimeDelta(tf.Module): _last_val: Optional[tf.Variable] = None def compute_delta(self, arr: tf.Tensor): # 首次调用或输入形状变化时,重新初始化状态变量 if self._last_val is None or not tf.reduce_all(tf.equal(tf.shape(arr), tf.shape(self._last_val))): # 用输入的形状创建变量,允许动态维度 self._last_val = tf.Variable(tf.zeros_like(arr), shape=tf.TensorShape(None)) delta = arr - self._last_val self._last_val.assign(arr) return delta
2. 用动态签名保存模型
保存模型时使用动态输入签名,同时启用TFLite对TensorFlow原生操作的支持:
compile_time_shape = (None, None) tflite_model_file_path = tempfile.mktemp() delta = TimeDelta() # 定义带动态输入签名的推理函数 infer_func = tf.function( delta.compute_delta, input_signature=[tf.TensorSpec(shape=compile_time_shape, dtype=tf.float32)] ) # 转换并保存模型 converter = tf.lite.TFLiteConverter.from_concrete_functions( [infer_func.get_concrete_function()], delta ) # 启用动态形状所需的TensorFlow操作支持 converter.target_spec.supported_ops = [ tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS ] tflite_model = converter.convert() with open(tflite_model_file_path, 'wb') as f: f.write(tflite_model)
3. 加载模型并验证功能
加载模型时,需根据输入形状动态调整张量内存:
def load_tflite_model_func(model_path): interpreter = tf.lite.Interpreter(model_path=model_path) input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() def infer(input_arr): # 根据输入形状调整张量大小并重新分配内存 interpreter.resize_tensor_input(input_details[0]['index'], input_arr.shape) interpreter.allocate_tensors() interpreter.set_tensor(input_details[0]['index'], input_arr) interpreter.invoke() return interpreter.get_tensor(output_details[0]['index']) return infer # 执行测试 func = load_tflite_model_func(tflite_model_file_path) runtime_shape = 60, 80 rng = np.random.RandomState(1234) ims = [rng.randn(*runtime_shape).astype(np.float32) for _ in range(3)] assert np.allclose(func(ims[0]), ims[0]) assert np.allclose(func(ims[1]), ims[1]-ims[0]) assert np.allclose(func(ims[2]), ims[2]-ims[1])
关键说明
- 状态变量初始化时用
tf.zeros_like(arr)保证与输入形状完全匹配,shape=tf.TensorShape(None)允许动态维度; - 启用
SELECT_TF_OPS确保动态形状相关的TensorFlow操作能被TFLite兼容; - 每次输入形状变化后,需调用
resize_tensor_input和allocate_tensors重新分配内存。
内容的提问来源于stack exchange,提问作者Peter
相关产品推荐
相关产品推荐

