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

资讯详情

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

TensorFlow Lite 设备端训练:用 TaoToken 统一 Key 打通端侧微调链路

TensorFlow Lite 设备端训练:用 TaoToken 统一 Key 打通端侧微调链路 1. 端侧微调为什么总在“最后一公里”卡住TensorFlow Lite 的设备端训练On-Device Training是这两年被问得越来越多的能力模型不再只做推理而是能在 Android/iOS 本地用用户自己的数据做几轮微调把通用图像分类模型变成“认得出我家那只鸟”的个性化模型。它适合谁适合已经能把.tflite跑起来、但被“训练脚本怎么接、权重怎么存、签名怎么调”卡住的移动端开发者。核心检索词就三个TensorFlow Lite、设备端训练、端侧微调。我自己的踩坑经历很典型模型转换成功、推理正常可一旦把train签名接进脚本就遇到三类问题——转换时没开experimental_enable_resource_variables导致权重变量在端侧不可写训练完的 checkpoint 没落盘下一轮又从零开始以及最隐蔽的训练脚本里调模型服务的 Key 散落在多个文件换环境就要全局搜一遍。前两个是 TFLite 本身的机制问题第三个是工程管理问题而第三个恰恰能用统一 Key 通道一次性解决。这篇就按“先跑通端侧训练闭环再谈统一接入”的顺序写。我会给出config.toml与settings.json的可复制骨架演示怎么通过 TaoToken 的统一 Key/API 通道接入端侧训练脚本最后附上端侧推理验证动作和一份报错排查清单。目标很明确你照着做一次跑通。2. 前置准备TaoToken 统一 Key 与端侧训练环境先说清楚 TaoToken 在这个链路里的位置。端侧训练本身是纯本地计算不依赖网络但训练脚本往往需要拉取预训练权重、调用模型做数据增强标注、或者在 CI 里跑回归验证。这些环节如果各自维护一套 Key维护成本会随环境数量线性上涨。TaoToken 提供的是统一 Key/API 通道把这些调用收敛到一个入口端侧脚本只认一个环境变量。官网入口在这里https://taotoken.net/?utm_sourcetaotoken_aicg_blog_endutm_mediumcsdnutm_campaignrewriteutm_content 。API 基址是 https://taotoken.net/api 这个地址不加 UTM直接用于代码里。你需要准备的东西不多一台能跑 TensorFlow 2.7 的开发机转换和导出用一个 Android 工程示例用 Kotlin以及一个 TaoToken 的 Key。Key 在控制台创建https://taotoken.net/console?utm_sourcetaotoken_aicg_blog_endutm_contentconsoleutm_campaignrewrite 。创建后建议直接写进环境变量别硬编码进脚本。注意端侧训练的低级能力资源变量、权重序列化在 TFLite 里仍属实验特性转换时必须显式开标志否则后面runSignature(train)会直接抛异常。这一步跳过后面全白搭。环境版本建议锁死TensorFlow 2.7 到 2.13 之间对 on-device training 的支持最稳Android 侧org.tensorflow:tensorflow-lite用 2.13.0 附近的版本。版本漂移是端侧训练最常见的“玄学报错”来源。3. 可复制配置config.toml 与 settings.json 骨架统一 Key 的接入点放在配置文件里脚本只读配置不读散落的常量。下面这份config.toml是训练侧Python 导出/验证脚本用的骨架字段名可以直接抄# config.toml —— 端侧训练脚本统一配置 [api] base_url https://taotoken.net/api api_key_env TAOTOKEN_API_KEY # 从环境变量读取不落盘 timeout_sec 30 max_retries 3 [model] name mobilenet_v2_personalize img_size 224 num_classes 10 checkpoint_dir ./checkpoints export_dir ./exported/tflite [tflite] enable_select_tf_ops true enable_resource_variables true # 关键不开则权重不可写 signatures [train, infer, save, restore] [training] epochs 5 batch_size 8 learning_rate 0.001Android 侧用settings.json承接同一套语义放在app/src/main/assets/下{ api: { baseUrl: https://taotoken.net/api, apiKeyEnv: TAOTOKEN_API_KEY, timeoutSec: 30 }, model: { fileName: mobilenet_v2_personalize.tflite, imgSize: 224, numClasses: 10, checkpointName: personalize.ckpt }, training: { epochs: 5, batchSize: 8, numBatches: 4 } }两份配置的api段保持同构是为了让 Python 侧和 Android 侧读同一套 Key 语义。Key 本身永远走环境变量TAOTOKEN_API_KEY配置文件里只存变量名。这样你在 CI、本地、真机调试之间切换时只改环境变量不动代码。读取逻辑用一段 Python 说明Android 侧同理import os, tomllib with open(config.toml, rb) as f: cfg tomllib.load(f) api_key os.environ.get(cfg[api][api_key_env]) if not api_key: raise RuntimeError(TAOTOKEN_API_KEY 未设置) base_url cfg[api][base_url]4. 打通端侧微调链路从转换到 runSignature4.1 构建带 train/infer/save 签名的模型端侧训练要求模型同时暴露训练和推理入口。核心是把train、predict、save三个函数用tf.function固定签名导出成 SavedModel 的多个签名。下面这段是可直接用的最小实现import tensorflow as tf IMG_SIZE 224 NUM_CLASSES 10 class PersonalizeModel(tf.Module): def __init__(self): super().__init__() base tf.keras.applications.MobileNetV2( input_shape(IMG_SIZE, IMG_SIZE, 3), include_topFalse, weightsimagenet) base.trainable True self.model tf.keras.Sequential([ base, tf.keras.layers.GlobalAveragePooling2D(), tf.keras.layers.Dense(NUM_CLASSES) ]) self._LOSS_FN tf.keras.losses.CategoricalCrossentropy(from_logitsTrue) self._OPTIM tf.keras.optimizers.Adam(1e-3) tf.function(input_signature[ tf.TensorSpec([None, IMG_SIZE, IMG_SIZE, 3], tf.float32), tf.TensorSpec([None, NUM_CLASSES], tf.float32)]) def train(self, x, y): with tf.GradientTape() as tape: pred self.model(x, trainingTrue) loss self._LOSS_FN(y, pred) grads tape.gradient(loss, self.model.trainable_variables) self._OPTIM.apply_gradients(zip(grads, self.model.trainable_variables)) return {loss: loss} tf.function(input_signature[ tf.TensorSpec([None, IMG_SIZE, IMG_SIZE, 3], tf.float32)]) def infer(self, x): return {output: self.model(x, trainingFalse)} tf.function(input_signature[tf.TensorSpec(shape[], dtypetf.string)]) def save(self, checkpoint_path): names [w.name for w in self.model.weights] tensors [w.read_value() for w in self.model.weights] tf.raw_ops.Save(filenamecheckpoint_path, tensor_namesnames, datatensors, namesave) return {checkpoint_path: checkpoint_path}4.2 转换时开对标志转换这一步是分水岭。experimental_enable_resource_variables必须为 True否则端侧拿不到可写变量SELECT_TF_OPS也要开因为权重序列化依赖它converter tf.lite.TFLiteConverter.from_saved_model(./exported/saved_model) converter.target_spec.supported_ops [ tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS ] converter.experimental_enable_resource_variables True tflite_model converter.convert() with open(./exported/tflite/personalize.tflite, wb) as f: f.write(tflite_model)4.3 Android 侧调用 train 与 infer 签名模型进 Android 后用Interpreter.runSignature按签名名调用。训练循环和推理调用分开写权重通过 checkpoint 文件在两次调用间传递val interpreter Interpreter(modelBuffer, Interpreter.Options().apply { setNumThreads(4) }) // 训练若干轮 val losses FloatArray(NUM_EPOCHS) for (epoch in 0 until NUM_EPOCHS) { for (batch in 0 until NUM_BATCHES) { val inputs mapOf( x to trainImages[batch], y to trainLabels[batch] ) val lossBuf FloatBuffer.allocate(1) val outputs mapOf(loss to lossBuf) interpreter.runSignature(inputs, outputs, train) if (batch NUM_BATCHES - 1) losses[epoch] lossBuf.get(0) } } // 保存权重 val ckptPath File(context.filesDir, personalize.ckpt).absolutePath interpreter.runSignature( mapOf(checkpoint_path to ckptPath), mutableMapOf(), save )推理侧用infer签名输入输出张量形状要和导出时一致val testInput FloatBuffer.allocateDirect(NUM_TESTS * IMG_SIZE * IMG_SIZE * 3) .order(ByteOrder.nativeOrder()) val output FloatBuffer.allocateDirect(NUM_TESTS * NUM_CLASSES) .order(ByteOrder.nativeOrder()) interpreter.runSignature( mapOf(x to testInput.rewind()), mapOf(output to output), infer ) output.rewind()5. 验证请求与成功结果跑通的标准不是“没报错”而是三个可观测结果同时成立。第一训练循环里loss逐轮下降哪怕只降一点点说明梯度确实在端侧更新了第二save调用后filesDir下出现personalize.ckpt文件且大小非零第三用同一批测试图调infer微调后的输出分布和微调前有可测差异。验证脚本可以这样写把 checkpoint 落盘和推理结果一起打印import numpy as np # 模拟端侧训练后的权重恢复验证 ckpt ./checkpoints/personalize.ckpt import os assert os.path.exists(ckpt), checkpoint 未生成 print(checkpoint size:, os.path.getsize(ckpt)) # 推理输出对比 before np.array([0.1, 0.2, 0.7]) after np.array([0.05, 0.15, 0.8]) print(delta:, np.abs(after - before).sum())如果你在训练脚本里接了 TaoToken 的模型对话能力做数据增强标注验证方式就是发一次请求确认通道通curl -s https://taotoken.net/api/v1/chat/completions \ -H Authorization: Bearer $TAOTOKEN_API_KEY \ -H Content-Type: application/json \ -d {model:gpt-4o-mini,messages:[{role:user,content:ping}]}返回里带choices字段就说明统一 Key 通道正常。模型对话入口在 https://taotoken.net/model-chat?utm_sourcetaotoken_aicg_blog_endutm_contentmodel_chatutm_campaignrewrite 需要长期跑编码和 Agent 任务的可以看 Coding Planhttps://taotoken.net/coding-plan?utm_sourcetaotoken_aicg_blog_endutm_contentcoding_planutm_campaignrewrite 。6. 本篇常见错排查清单报错一Signature train not found。导出 SavedModel 时train函数没被tf.function固定签名或者转换时签名没被保留。检查concrete_functions里是否包含train转换后可以用interpreter.getSignatureRunner(train)确认。报错二Resource variables are not supported。转换时漏了experimental_enable_resource_variables True。这个标志必须在convert()之前设置事后补没用。报错三训练 loss 不下降。大概率是train签名里trainingTrue没传BatchNorm 层没进入训练模式或者学习率过大导致震荡。先把learning_rate降到 1e-4 试。报错四checkpoint 文件为空。save签名里tf.raw_ops.Save的tensor_names和data长度不一致或者权重里有不可序列化的类型。打印len(names)和len(tensors)对齐一下。报错五Android 侧runSignature抛IllegalArgumentException。输入张量的 shape 和导出签名不一致最常见的是 batch 维度没对齐。train签名第一维是None实际传入时按 batch 大小填。报错六Key 读取失败。配置文件里写的是环境变量名但运行时TAOTOKEN_API_KEY没导出。在 shell 里export TAOTOKEN_API_KEYxxxAndroid 侧用BuildConfig或local.properties注入别写进 assets。接入文档在 https://taotoken.net/doc?utm_sourcetaotoken_aicg_blog_endutm_contentdocutm_campaignrewrite API Keys 管理在 https://taotoken.net/api-keys?utm_sourcetaotoken_aicg_blog_endutm_contentapi_keysutm_campaignrewrite 。Claude Code 相关的端侧 Agent 场景可以看 https://taotoken.net/claude-code-anthropic?utm_sourcetaotoken_aicg_blog_endutm_contentclaude_codeutm_campaignrewrite 。最后补一个实用技巧端侧训练的 checkpoint 文件会随训练轮次增长真机上要定期清理旧文件否则filesDir会被撑爆。我一般只保留最近两份用文件名带时间戳的方式滚动覆盖。
返回列表