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

如何保存状态形状依赖输入的有状态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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 07:05:13