使用TF2与MobileNetV3-Large做QAT时遇模型嵌套量化不支持错误求助
解决MobileNetV3-Large量化感知训练(QAT)的模型嵌套错误
问题根源
tfmot.quantization.keras.quantize_model不支持量化包含子Keras模型的嵌套结构——你代码中将预训练的MobileNetV3(本身是独立Keras模型)放入Sequential容器,触发了"Quantizing a tf.keras Model inside another tf.keras Model is not supported"错误。
修正方案与完整代码
核心修改点
- 用函数式API替代Sequential构建模型,避免模型嵌套
- 修正QAT训练时的模型调用错误(原代码误用
model.fit而非量化后的模型) - 补全未定义的
class_weight变量(根据数据集类别分布设置) - 确保从QAT模型转换为TFLite,而非原浮点模型
以下是修正后的完整代码:
import numpy as np import tensorflow as tf from tensorflow import keras from tensorflow.keras import layers, regularizers, callbacks from tensorflow.keras.applications import MobileNetV3Large import tensorflow_model_optimization as tfmot print("#### Import GDrive ####") from google.colab import drive drive.mount('/content/drive') # Define model name model_name = "your_model_name" # Declare the training, validation, and testing directories train_dir = r"/content/drive/your_train_dir" val_dir = r"/content/drive/your_val_dir" test_dir = r"/content/drive/your_test_dir" # Load the training, validation, and testing datasets print("#### Dataset Information ####\nTraining Dataset:") train_ds = tf.keras.utils.image_dataset_from_directory( train_dir, label_mode='binary', image_size=(224, 224), batch_size=32) print("Validation Dataset:") val_ds = tf.keras.utils.image_dataset_from_directory( val_dir, label_mode='binary', image_size=(224, 224), batch_size=32) print("Testing Dataset:") test_ds = tf.keras.utils.image_dataset_from_directory( test_dir, label_mode='binary', image_size=(224, 224), batch_size=32) # 计算类别权重(根据你的数据集实际情况调整,示例为二分类) class_weight = {0: 1.0, 1: 1.0} # 类别不平衡时修改对应权重值 # Instantiate the base model print("#### Download Model ####") base_model = MobileNetV3Large(input_shape=(224, 224, 3), alpha=1.0, minimalistic=False, include_top=False, weights='imagenet', include_preprocessing=True) base_model.trainable = False # -------------------------- 关键修改:用函数式API构建模型 -------------------------- inputs = keras.Input(shape=(224, 224, 3)) x = base_model(inputs, training=False) x = layers.GlobalAveragePooling2D()(x) x = layers.Dropout(0.5)(x) outputs = layers.Dense(1, activation='sigmoid', kernel_regularizer=regularizers.l2(0.01))(x) model = keras.Model(inputs, outputs) # --------------------------------------------------------------------------------- # Compile the model model.compile(optimizer=keras.optimizers.Adam(), loss=keras.losses.BinaryCrossentropy(), metrics=[keras.metrics.BinaryAccuracy()]) # Add early stopping early_stopping = callbacks.EarlyStopping(monitor='val_loss', patience=5, restore_best_weights=True) # Train the model print("\n#### Transfer Learning ####") model.fit(train_ds, epochs=25, validation_data=val_ds, callbacks=[early_stopping]) # Save the model model.save(f"{model_name}_initial_raw") print("\nInitial raw model saved.") # Unfreezing the top layers of the base model base_model.trainable = True for layer in base_model.layers[:50]: layer.trainable = False # Implement QAT for fine-tuning with tfmot.quantization.keras.quantize_scope(): # 现在可正常量化,模型为单一实例无嵌套 q_aware_model = tfmot.quantization.keras.quantize_model(model) # Re-compile the QAT model q_aware_model.compile(optimizer=keras.optimizers.Adam(1e-5), loss=keras.losses.BinaryCrossentropy(), metrics=[keras.metrics.BinaryAccuracy()]) # -------------------------- 关键修改:用q_aware_model训练 -------------------------- print("#### Fine Tuning with QAT ####") q_aware_model.fit(train_ds, epochs=25, validation_data=val_ds, callbacks=[early_stopping], class_weight=class_weight) # --------------------------------------------------------------------------------- # Evaluate the fine-tuned QAT model on the test dataset print("\n#### QAT Model Evaluation ####") q_aware_model.evaluate(test_ds) # Save QAT model q_aware_model.save(f"{model_name}_fine_qat") print("\nFine-tuned QAT model saved.") # -------------------------- 关键修改:从QAT模型转换为TFLite -------------------------- converter = tf.lite.TFLiteConverter.from_keras_model(q_aware_model) converter.optimizations = [tf.lite.Optimize.DEFAULT] # 为Coral TPU添加全整数量化配置(需提供校准数据集) def representative_data_gen(): for batch in train_ds.take(10): yield [batch[0]] converter.representative_dataset = representative_data_gen converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8] converter.inference_input_type = tf.int8 converter.inference_output_type = tf.int8 quantized_tflite_model = converter.convert() # 保存量化后的TFLite模型 with open(f"{model_name}_quantized.tflite", "wb") as f: f.write(quantized_tflite_model) print("\nFull integer quantized TFLite model saved.") # ---------------------------------------------------------------------------------
额外说明
- 函数式API构建的模型无嵌套子模型结构,符合tfmot的量化要求
- 转换为TFLite时添加了
representative_dataset和整数量化配置,生成的模型完全兼容Coral TPU - 若数据集类别不平衡,需正确设置
class_weight值,避免训练偏差
内容的提问来源于stack exchange,提问作者Sennsei
相关产品推荐
相关产品推荐

