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

将Saved Model转换为TFLite时出现tfl.assign_variable错误求助

Saved Model转TFLite格式报错解决建议

问题概述

将训练好的Saved Model转换为TFLite格式时,出现类型不兼容错误,报错提示tfl.assign_variable和tfl.read_variable操作不支持tensor<2xui32>类型。

错误详情

Error: 'tfl.assign_variable' op operand #1 must be tensor of 32-bit float or 64-bit float or 1-bit signless integer or 8-bit unsigned integer or 8-bit signless integer or QI8 type or QUI8 type or 32-bit signless integer or 64-bit signless integer or QI16 type or complex type with 32-bit float elements or complex type with 64-bit float elements values, but got 'tensor<2xui32>'
error: 'tfl.assign_variable' op operand #1 must be tensor of 32-bit float or 64-bit float or 1-bit signless integer or 8-bit unsigned integer or 8-bit signless integer or QI8 type or QUI8 type or 32-bit signless integer or 64-bit signless integer or QI16 type or complex type with 32-bit float elements or complex type with 64-bit float elements values, but got 'tensor<2xui32>'
loc(callsite(callsite(fused["ReadVariableOp:", "sequential_1/random_flip_1/ReadVariableOp@__inference___call___7760"] at fused["StatefulPartitionedCall:", "StatefulPartitionedCall@__inference_signature_wrapper___call___7851"]) at fused["StatefulPartitionedCall:", "StatefulPartitionedCall_1"])): error: 'tfl.read_variable' op result #0 must be tensor of 32-bit float or 64-bit float or 1-bit signless integer or 8-bit unsigned integer or 8-bit signless integer or QI8 type or QUI8 type or 32-bit signless integer or 64-bit signless integer or QI16 type or complex type with 32-bit float elements or complex type with 64-bit float elements values, but got 'tensor<2xui32>'
loc(callsite(callsite(fused["AssignVariableOp:", "sequential_1/random_flip_1/AssignVariableOp@__inference___call___7760"] at fused["StatefulPartitionedCall:", "StatefulPartitionedCall@__inference_signature_wrapper___call___7851"]) at fused["StatefulPartitionedCall:", "StatefulPartitionedCall_1"])): error: 'tfl.assign_variable' op operand #1 must be tensor of 32-bit float or 64-bit float or 1-bit signless integer or 8-bit unsigned integer or 8-bit signless integer or QI8 type or QUI8 type or 32-bit signless integer or 64-bit signless integer or QI16 type or complex type with 32-bit float elements or complex type with 64-bit float elements values, but got 'tensor<2xui32>'
loc(callsite(callsite(fused["ReadVariableOp:", "sequential_1/random_flip_1/ReadVariableOp_2@__inference___call___7760"] at fused["StatefulPartitionedCall:", "StatefulPartitionedCall@__inference_signature_wrapper___call___7851"]) at fused["StatefulPartitionedCall:", "StatefulPartitionedCall_1"])): error: 'tfl.read_variable' op result #0 must be tensor of 32-bit float or 64-bit float or 1-bit signless integer or 8-bit unsigned integer or 8-bit signless integer or QI8 type or QUI8 type or 32-bit signless integer or 64-bit signless integer or QI16 type or complex type with 32-bit float elements or complex type with 64-bit float elements values, but got 'tensor<2xui32>'
loc(callsite(callsite(fused["AssignVariableOp:", "sequential_1/random_flip_1/AssignVariableOp_1@__inference___call___7760"] at fused["StatefulPartitionedCall:", "StatefulPartitionedCall@__inference_signature_wrapper___call___7851"]) at fused["StatefulPartitionedCall:", "StatefulPartitionedCall_1"])): error: 'tfl.assign_variable' op operand #1 must be tensor of 32-bit float or 64-bit float or 1-bit signless integer or 8-bit unsigned integer or 8-bit signless integer or QI8 type or QUI8 type or 32-bit signless integer or 64-bit signless integer or QI16 type or complex type with 32-bit float elements or complex type with 64-bit float elements values, but got 'tensor<2xui32>'
loc(callsite(callsite(fused["ReadVariableOp:", "sequential_1/random_rotation_1/ReadVariableOp@__inference___call___7760"] at fused["StatefulPartitionedCall:", "StatefulPartitionedCall@__inference_signature_wrapper___call___7851"]) at fused["StatefulPartitionedCall:", "StatefulPartitionedCall_1"])): error: 'tfl.read_variable' op result #0 must be tensor of 32-bit float or 64-bit float or 1-bit signless integer or 8-bit unsigned integer or 8-bit signless integer or QI8 type or QUI8 type or 32-bit signless integer or 64-bit signless integer or QI16 type or complex type with 32-bit float elements or complex type with 64-bit float elements values, but got 'tensor<2xui32>'
loc(callsite(callsite(fused["AssignVariableOp:", "sequential_1/random_rotation_1/AssignVariableOp@__inference___call___7760"] at fused["StatefulPartitionedCall:", "StatefulPartitionedCall@__inference_signature_wrapper___call___7851"]) at fused["StatefulPartitionedCall:", "StatefulPartitionedCall_1"])): error: 'tfl.assign_variable' op operand #1 must be tensor of 32-bit float or 64-bit float or 1-bit signless integer or 8-bit unsigned integer or 8-bit signless integer or QI8 type or QUI8 type or 32-bit signless integer or 64-bit signless integer or QI16 type or complex type with 32-bit float elements or complex type with 64-bit float elements values, but got 'tensor<2xui32>'

模型代码

import datetime
import time
import numpy as np
from pathlib import Path
import tensorflow as tf
from tensorflow import keras

num_classes = 5

model=keras.Sequential([
  keras.layers.RandomFlip("horizontal_and_vertical"),
  keras.layers.RandomRotation(0.2),
  keras.layers.Rescaling(1./255),
  keras.layers.Conv2D(16, 3, activation='relu'),
  keras.layers.MaxPooling2D(),
  keras.layers.Conv2D(64, 3, activation='relu'),
  keras.layers.MaxPooling2D(),
  keras.layers.Conv2D(128, 3, activation='relu'),
  keras.layers.MaxPooling2D(),
  keras.layers.Conv2D(256, 3, activation='relu'),
  keras.layers.MaxPooling2D(),
  keras.layers.Conv2D(512, 3, activation='relu'),
  keras.layers.MaxPooling2D(),
  keras.layers.Conv2D(1024, 3, activation='relu'),
  keras.layers.MaxPooling2D(),
  keras.layers.Flatten(),
  keras.layers.Dense(256, activation='relu'),
  keras.layers.Dense(64, activation='relu'),
  keras.layers.Dense(num_classes)
])


model.compile(optimizer='adam',
              loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
              metrics=['accuracy'])

file_datetime = datetime.datetime.now().strftime('%Y_%m_%d_%H_%S')
filepath = f'./model'
print(f'Best model will be saved to: {Path.cwd() / filepath}')

# Train the model
NUM_EPOCHS = 1
MIN_VAL_LOSS = np.inf
loss_history = [[],[]]
acc_history = [[],[]]
total_train_time = 0.0
# Training loop
# NOTE: This works better than callbacks
for i in range(0, NUM_EPOCHS):
    print(f'EPOCH: {i+1}/{NUM_EPOCHS}')
    start_epoch = time.time()
    history = model.fit(train_images,
                    epochs=1, batch_size=batch_size,
                    validation_data=val_images)
    end_epoch = time.time()
    delta_epoch = end_epoch - start_epoch
    total_train_time += delta_epoch
    print(f'LOG --> This epoch took {delta_epoch}...')
    if history.history['val_loss'][0] < MIN_VAL_LOSS:
        print(f'LOG --> val_loss improved from {MIN_VAL_LOSS} to {history.history["val_loss"][0]}...')
        print(f'LOG --> saving model as {filepath}')
        MIN_VAL_LOSS = history.history['val_loss'][0]
        model.export(filepath)
    else:
        print(f'LOG --> val_loss did not improve...')
    # Keep track of the training history
    loss_history[0].append(history.history['loss'][0])
    loss_history[1].append(history.history['val_loss'][0])
    acc_history[0].append(history.history['accuracy'][0])
    acc_history[1].append(history.history['val_accuracy'][0])
print(f'LOG--> Total training time: {total_train_time}')

解决建议

原因分析

报错根源在于RandomFlip和RandomRotation这两个训练数据增强层:它们在训练过程中会维护状态变量(比如随机种子相关的变量),这些变量的类型是无符号32位整数(ui32),而TFLite的assign_variable和read_variable操作不支持该类型,导致转换失败。

解决方案

1. 构建不含训练增强层的推理模型(推荐)

推理阶段不需要随机数据增强,直接构建一个只保留核心推理层的模型,复制原模型权重后导出:

# 构建推理模型,跳过训练用的随机增强层
inference_model = keras.Sequential([
    keras.layers.Rescaling(1./255),
    *model.layers[2:]  # 从Rescaling层开始,取原模型后续所有层
])

# 复制原模型的训练权重
inference_model.set_weights(model.get_weights())

# 导出推理模型用于TFLite转换
inference_model.export(filepath)

2. 转换时强制模型进入推理模式

如果需要保留层结构(不推荐,因为推理时不需要增强),可以在转换前让模型处于推理模式,关闭随机增强的状态变量:

# 生成和输入同维度的 dummy 数据,触发模型状态切换
dummy_input = tf.random.normal([1, *train_images.shape[1:]])
_ = model(dummy_input, training=False)

# 重新导出模型
model.export(filepath)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 19:32:03