将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
相关产品推荐
相关产品推荐

