- 示例工程
【免费下载链接】examples
TensorFlow examples
导读
本文围绕 TensorFlow Lite 官方示例项目 lite/examples/model_personalization 展开,系统讲解如何在完全不上传数据的前提下,于 Android 设备端完成 TFLite 模型的个性化(On-device Model Personalization)。你将掌握完整链路:用 Python 脚本定义并生成"基座模型 + 可训练头模型"结构的 TFLite 文件、通过命令行或 Android Studio 构建示例 App、在端上采集样本并实时训练与推理,以及如何自定义模型结构与超参数。
设备端模型个性化是什么
模型个性化(Model Personalization)解决的是"通用模型 + 个人化数据"的问题:云端或离线训练好的基座模型拥有通用特征提取能力,但无法针对每个用户的专属对象(例如"我的宠物""我的物品")进行分类。传统做法需要把用户数据上传到服务器重新微调,而本示例给出了一种纯端侧方案:
- 不发送任何数据到服务器,隐私友好;
- 复用现有 TFLite 功能,无需引入额外运行时;
- 模型结构为"基座模型(Base Model)+ 头模型(Head Model)"两部分,可适配不同任务与模型。
示例 App 是一个实时摄像头分类器:用户先为每个类别拍摄若干张照片作为训练样本,点击"Train"按钮后模型在设备上完成微调,随后切换至推理模式即可实时预测画面内容。
项目结构总览
示例由三部分组成(对应 lite/examples/model_personalization/README.md 的 Structure 章节):
| 组成部分 | 职责 | 仓库位置 |
|---|---|---|
| 模型生成(Model Generation) | Python CLI,定义并生成可个性化模型 | lite/examples/model_personalization/transfer_learning |
| Android 调用库 | 从 Android App 中使用所生成模型的库能力 | 原版文档将其描述为独立 Gradle 模块android/transfer_api;从当前仓库的 android/settings.gradle 看,示例只包含:app单一模块,该库的核心调用逻辑集成在 App 内的TransferLearningHelper.kt中 |
| Android 分类 App | 演示如何调用模型个性化能力的应用 | android/app |
第一步:准备 TFLite 模型
建立 Python 环境并安装依赖
README 的 Quickstart 要求 Python 3.7+ 及virtualenv。官方流程建议(非强制)创建虚拟环境:
pushd transfer_learning # 创建并激活虚拟环境 python3 -m venv env source env/bin/activate # 安装依赖(当前仓库仅要求 tensorflow>=2.7.*) pip install -r requirements.txt # 生成模型 flatbuffer 文件 model.tflite 到当前目录 python generate_training_model.py popd # 将生成的模型复制到 Android assets 目录 cp transfer_learning/model.tflite android/app/src/main/assets/model/model.tflite依赖文件 lite/examples/model_personalization/transfer_learning/requirements.txt 内容为:
tensorflow>=2.7.*注:除了手动生成并复制模型,android/README.md 还提供了另一种方式——模型文件由 Gradle 脚本在构建时自动下载到 assets(见下文"模型自动下载"一节)。若选择手动流程,可注释掉 android/app/download_models.gradle 的引用。
生成脚本做了什么
执行python generate_training_model.py后,generate_training_model.py 会完成三件事:构建TransferLearningModel实例、以多个具名签名导出 SavedModel、再用TFLiteConverter转换为 TFLite 文件。
关键常量定义在脚本头部(第 24-26 行):
IMG_SIZE = 224 NUM_FEATURES = 7 * 7 * 1280 NUM_CLASSES = 4IMG_SIZE = 224:输入图像尺寸(MobileNetV2 的标准输入);NUM_FEATURES = 7 * 7 * 1280:基座模型输出的瓶颈(bottleneck)特征维度,由 224×224 输入经 MobileNetV2 下采样至 7×7 空间、1280 个通道得到;NUM_CLASSES = 4:示例中的类别数,对应 App 底部四个类别按钮。
深入模型生成原理:双段结构与六个签名
基座模型、头模型与优化器
README 明确指出(Customizing the model):TFLite 设备端个性化模型由两部分组成——
- 基座模型(Base Model):通常为数据丰富任务预训练,负责通用特征提取,其权重在转换时被固定,之后不可修改;
- 头模型(Head Model):将在设备上训练的部分,通常是轻量分类头。
示例的默认组合(见 generate_training_model.py 第 50-57 行):
# 基座模型:ImageNet 预训练的 MobileNetV2,去掉顶层分类器 self.base = tf.keras.applications.MobileNetV2( input_shape=(IMG_SIZE, IMG_SIZE, 3), alpha=1.0, include_top=False, weights='imagenet') # 头模型:一个可训练权重矩阵 + 偏置,配合 softmax 激活 self.ws = tf.Variable(tf.zeros((self.num_features, self.num_classes)), name='ws', trainable=True) self.bs = tf.Variable(tf.zeros((1, self.num_classes)), name='bs', trainable=True) # 损失函数与优化器:默认 learning_rate=0.001 self.loss_fn = tf.keras.losses.CategoricalCrossentropy() self.optimizer = tf.keras.optimizers.Adam(learning_rate=learning_rate)参数速查表:
| 参数 | 默认值 | 说明 |
|---|---|---|
| 基座模型 | MobileNetV2(alpha=1.0, include_top=False, weights='imagenet') | 图像识别任务的通用特征提取器 |
| 头模型 | 单个全连接层(ws权重 +bs偏置)+ softmax | 设备端唯一可训练部分 |
| 损失函数 | CategoricalCrossentropy() | 多分类交叉熵 |
| 优化器 | Adam(learning_rate=0.001) | 可通过构造参数调整学习率 |
六个具名签名:load / train / infer / save / restore / initialize
为了在 TFLite 端通过runSignature按名称调用不同功能,模型导出了六个@tf.function签名(见 generate_training_model.py 第 190-197 行):
| 签名 | 输入 | 输出 | 用途 |
|---|---|---|---|
load | 图像 batch[None, 224, 224, 3] | bottleneck | 生成瓶颈特征(权重冻结,不可训练) |
train | bottleneck+label | loss(及梯度) | 在设备上执行一步训练,更新ws/bs |
infer | 图像 batch | output(softmax 概率) | 推理分类 |
save | checkpoint 路径 | checkpoint_path | 保存可训练权重 |
restore | checkpoint 路径 | ws/bs | 恢复已保存权重 |
initialize | 无 | ws/bs | 随机初始化头模型权重 |
其中train的实现核心(第 91-100 行)是标准梯度下降:前向计算logits = matmul(bottleneck, ws) + bs与 softmax,用GradientTape求出对ws/bs的梯度,再交给 Adam 优化器更新。
转换配置要点
convert_and_save中的转换配置(第 200-205 行)值得注意:
converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir) converter.target_spec.supported_ops = [ tf.lite.OpsSet.TFLITE_BUILTINS, # 启用 TensorFlow Lite 内建算子 tf.lite.OpsSet.SELECT_TF_OPS # 启用 Select TF Ops(加载 TensorFlow 算子) ] converter.experimental_enable_resource_variables = True tflite_model = converter.convert()SELECT_TF_OPS:由于模型内含tf.raw_ops.Save/Restore等 TensorFlow 原生算子,必须启用 Select TF Ops 才能正常转换与运行——这也是 Android 端需要引入tensorflow-lite-select-tf-ops依赖的原因;experimental_enable_resource_variables = True:启用资源变量支持,使ws/bs这些可训练变量在 TFLite 解释器中保持状态、可被反复更新。
第二步:构建并运行 Android 应用
方式一:Android Studio
在 Android Studio 中导入项目(指向顶层build.gradle),连接真机后点击Run。若Run按钮不可用,先为app模块添加 Android Application 运行配置。由于摄像头训练流程需要相机与本地训练,必须在物理 Android 设备上运行,模拟器无法完成演示。
方式二:命令行构建安装(Linux)
cd android gradle wrapper # 若本机未安装 gradle,可参考官方安装文档;仓库已自带 gradlew ./gradlew build adb install ./app/build/outputs/apk/debug/app-debug.apkApp 侧的构建配置(见 android/app/build.gradle)要点:
minSdk 23、targetSdk 32、compileSdk 32;- TFLite 相关依赖:
org.tensorflow:tensorflow-lite:2.9.0、tensorflow-lite-gpu:2.9.0、tensorflow-lite-support:0.4.2、tensorflow-lite-select-tf-ops:2.9.0; - 摄像头部分基于 CameraX(
camera-core/camera-camera2/camera-lifecycle/camera-view); androidResources { noCompress 'tflite' }:模型文件不打压缩,便于 mmap 加载;- 界面采用 Jetpack Navigation + ViewBinding。
模型自动下载
android/app/download_models.gradle 会在构建前自动把预训练模型下载到app/src/main/assets/model.tflite(overwrite false,已存在则跳过):
task downloadModelFile(type: Download) { src 'https://storage.googleapis.com/download.tensorflow.org/models/tflite/task_library/model_personalization/android/model.tflite' dest project.ext.ASSET_DIR + '/model.tflite' overwrite false } preBuild.dependsOn downloadModelFileApp 使用流程:采集样本 → 训练 → 推理
应用启动后,底部四个按钮分别对应模型需要区分的四个类别(示例中编号 1–4)。初始状态下各按钮上的置信度分数要么随机、要么恒定,取决于模型初始化方式。
官方推荐操作流程(README 第 76-93 行):
- 采集样本:至少拍摄一张图片并关联到某个类别——按下对应类别按钮即可拍照。为获得更好的训练效果,每个类别建议至少采集 10 张图片,并尽量覆盖不同背景与物体朝向;
- 训练:当采集样本数 ≥ 1 后,
Train按钮变为可用。按下后等待数秒,观察损失(Loss)下降;训练过程中可通过Pause暂停; - 推理:训练(或暂停)后切换到右上角的Inference(推理)模式,分类器将对摄像头画面进行实时类别预测。
App 内部的状态流转由 MainViewModel.kt 管理,训练状态枚举为PREPARE → TRAINING → PAUSE;采集样本与推理通过captureMode布尔值互斥切换。
Android 端调用细节:TransferLearningHelper 解析
App 的核心逻辑集中在 TransferLearningHelper.kt,它演示了如何在端上串联整个训练闭环。
签名调用与键名约定
代码通过interpreter.runSignature(inputs, outputs, signatureKey)按名称调用模型签名,签名键名必须与 Python 脚本导出的具名签名一致(第 338-353 行):
| 常量 | 值 | 对应 Python 签名 |
|---|---|---|
LOAD_BOTTLENECK_KEY | "load" | load |
TRAINING_KEY | "train" | train |
INFERENCE_KEY | "infer" | infer |
- 采集样本时先调用
load得到瓶颈特征(loadBottleneck),并把(bottleneck, one-hot label)存入trainingSamples列表,避免训练阶段重复跑基座模型; - 训练时把批量瓶颈与标签喂给
train签名,返回的loss通过回调刷新到界面; - 推理时把预处理后的图像喂给
infer签名,用TensorLabel将输出映射为类别与置信度。
训练循环与批处理
EXPECTED_BATCH_SIZE = 20:期望的批大小;当样本不足 20 时,getTrainBatchSize()取min(max(1, 样本数), 20)动态缩小批大小;- 训练在单线程 Executor中持续进行(
while (executor?.isShutdown == false)),每轮先shuffle打乱样本以减少过拟合,再按批送入train签名,最后把平均损失回调到 UI 线程; - 训练与推理共用一把
lock(synchronized(lock)),保证同一时刻只有一个线程在训练或推理。
图像预处理
processInputImage(第 255-277 行)使用 TFLite Support 的ImageProcessor完成旋转(Rot90Op)、正方形裁剪(ResizeWithCropOrPadOp)、双线性缩放至 224×224(ResizeOp)以及NormalizeOp(0f, 255f)归一化——注意输入输出均为FLOAT32。
自定义模型:换基座、调头、改超参
README 明确鼓励"Feel free to create/modify the Transfer Learning model structure and configurations",定制路径非常直接:
- 更换基座模型:将 generate_training_model.py 中
self.base = tf.keras.applications.MobileNetV2(...)换成其他预训练模型,并同步更新NUM_FEATURES(基座输出展平后的维度)与IMG_SIZE(输入尺寸); - 调整头模型:
ws/bs的形状由num_features × num_classes决定,修改NUM_CLASSES即可改变分类数(注意 Android 端 TransferLearningHelper.kt 中推理输出固定为1 × 4、类别映射classes为 4 个,需同步修改); - 修改优化器与学习率:
TransferLearningModel.__init__(learning_rate=0.001)传入不同学习率,或替换Adam为其他tf.keras.optimizers优化器; - 重新生成模型:改完脚本后重新执行
python generate_training_model.py,并把新的model.tflite覆盖到 android/app/src/main/assets/model/model.tflite(或注释掉 download_models.gradle 引用以避免自动下载覆盖)。
自定义时的关键约束
- 基座权重在转换时被固定,之后无法更改——这是"可个性化"与"端侧微调"的边界;
- 模型必须保留
load / train / infer / save / restore / initialize六个签名(至少load / train / infer三个,App 的TransferLearningHelper依赖它们),且签名输入输出张量形状需与 Android 端代码保持一致; - 转换时需保留
SELECT_TF_OPS与experimental_enable_resource_variables = True两项配置,否则设备端变量更新与算子执行可能失败。
小结
从 lite/examples/model_personalization/README.md 出发,本文完整还原了 TensorFlow Lite 设备端模型个性化的落地路径:Python 侧通过双段结构(冻结的 MobileNetV2 基座 + 可训练的线性头)与六个具名签名生成可个性化 TFLite 模型,Android 侧通过Interpreter.runSignature实现瓶颈提取、端上训练与实时推理,全程数据不出设备。若想进一步动手,可参照 generate_training_model.py 调整模型结构,再结合 TransferLearningHelper.kt 验证端侧行为,这套模式可平滑迁移到语音、文本等更多任务上。
- 示例工程
【免费下载链接】examples
TensorFlow examples
相关推荐
Flower Android 端 TFLite 模型生成指南:从 Keras 模型到 `layersSizes` 完整实战
Flower Android 端 TFLite 模型生成指南:从 Keras 模型到 layersSizes 完整实战 本文围绕 Flower 开源联邦学习框架
人工智能联邦学习机器学习深度学习TensorFlow模型性能优化实战:从训练到移动端部署的完整指南
TensorFlow模型性能优化实战:从训练到移动端部署的完整指南 TensorFlow作为业界领先的深度学习框架,其模型性能优化对于移动端部署至关重要。本文将
文档开发工具教程jax2tf 端侧推理实战:用 JAX 训练模型并转换为 TensorFlow Lite 格式
jax2tf 端侧推理实战:用 JAX 训练模型并转换为 TensorFlow Lite 格式 JAX 与 TensorFlow 的互操作能力,使开发者可以在
机器学习深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考