尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

TensorFlow Lite 设备端模型个性化实战:从生成可个性化 TFLite 模型到 Android 端训练与推理

TensorFlow Lite 设备端模型个性化实战:从生成可个性化 TFLite 模型到 Android 端训练与推理 示例工程【免费下载链接】examplesTensorFlow examples项目地址https://gitcode.com/gh_mirrors/exam/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 GenerationPython CLI定义并生成可个性化模型lite/examples/model_personalization/transfer_learningAndroid 调用库从 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 # 安装依赖当前仓库仅要求 tensorflow2.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 内容为tensorflow2.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 modelTFLite 设备端个性化模型由两部分组成——基座模型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), alpha1.0, include_topFalse, weightsimagenet) # 头模型一个可训练权重矩阵 偏置配合 softmax 激活 self.ws tf.Variable(tf.zeros((self.num_features, self.num_classes)), namews, trainableTrue) self.bs tf.Variable(tf.zeros((1, self.num_classes)), namebs, trainableTrue) # 损失函数与优化器默认 learning_rate0.001 self.loss_fn tf.keras.losses.CategoricalCrossentropy() self.optimizer tf.keras.optimizers.Adam(learning_ratelearning_rate)参数速查表参数默认值说明基座模型MobileNetV2(alpha1.0, include_topFalse, weightsimagenet)图像识别任务的通用特征提取器头模型单个全连接层ws权重 bs偏置 softmax设备端唯一可训练部分损失函数CategoricalCrossentropy()多分类交叉熵优化器Adam(learning_rate0.001)可通过构造参数调整学习率六个具名签名load / train / infer / save / restore / initialize为了在 TFLite 端通过runSignature按名称调用不同功能模型导出了六个tf.function签名见 generate_training_model.py 第 190-197 行签名输入输出用途load图像 batch[None, 224, 224, 3]bottleneck生成瓶颈特征权重冻结不可训练trainbottlenecklabelloss及梯度在设备上执行一步训练更新ws/bsinfer图像 batchoutputsoftmax 概率推理分类savecheckpoint 路径checkpoint_path保存可训练权重restorecheckpoint 路径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 设备上运行模拟器无法完成演示。方式二命令行构建安装Linuxcd 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 32TFLite 相关依赖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摄像头部分基于 CameraXcamera-core/camera-camera2/camera-lifecycle/camera-viewandroidResources { noCompress tflite }模型文件不打压缩便于 mmap 加载界面采用 Jetpack Navigation ViewBinding。模型自动下载android/app/download_models.gradle 会在构建前自动把预训练模型下载到app/src/main/assets/model.tfliteoverwrite 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_KEYloadloadTRAINING_KEYtraintrainINFERENCE_KEYinferinfer采集样本时先调用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 线程训练与推理共用一把locksynchronized(lock)保证同一时刻只有一个线程在训练或推理。图像预处理processInputImage第 255-277 行使用 TFLite Support 的ImageProcessor完成旋转Rot90Op、正方形裁剪ResizeWithCropOrPadOp、双线性缩放至 224×224ResizeOp以及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_rate0.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 验证端侧行为这套模式可平滑迁移到语音、文本等更多任务上。赞分享示例工程【免费下载链接】examplesTensorFlow examples项目地址https://gitcode.com/gh_mirrors/exam/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 的互操作能力使开发者可以在机器学习深度学习上一篇Metallb社区贡献统计贡献者数量与代码提交趋势下一篇SURF完全指南革命性Go HTTP客户端如何实现浏览器指纹与反反爬技术创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表