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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 06:00:57