TFLite端侧训练保存签名函数致Android应用崩溃求助
TFLite端侧训练保存模型时Android应用崩溃问题
我在做学校项目,实现TFLite端侧训练时遇到问题:train、infer等signature函数运行正常,但调用save函数时应用直接崩溃。参考TensorFlow官方端侧训练指南,希望训练个性化模型后保存,但每次启动就崩溃,求Android端侧训练后保存模型的正确指引。
附相关代码
原模型代码
class Model(tf.Module): def __init__(self): self.model = tf.keras.Sequential([ tf.keras.layers.Flatten(input_shape=(IMG_SIZE, IMG_SIZE), name='flatten'), tf.keras.layers.Dense(128, activation='relu', name='dense_1'), tf.keras.layers.Dense(10, name='dense_2') ]) self.model.compile( optimizer='sgd', loss=tf.keras.losses.CategoricalCrossentropy(from_logits=True)) @tf.function(input_signature=[ tf.TensorSpec([None, IMG_SIZE, IMG_SIZE], tf.float32), tf.TensorSpec([None, 10], tf.float32), ]) def train(self, x, y): with tf.GradientTape() as tape: prediction = self.model(x) loss = self.model.loss(y, prediction) gradients = tape.gradient(loss, self.model.trainable_variables) self.model.optimizer.apply_gradients( zip(gradients, self.model.trainable_variables)) result = {"loss": loss} return result @tf.function(input_signature=[ tf.TensorSpec([None, IMG_SIZE, IMG_SIZE], tf.float32), ]) def infer(self, x): logits = self.model(x) probabilities = tf.nn.softmax(logits, axis=-1) return { "output": probabilities, "logits": logits } @tf.function(input_signature=[tf.TensorSpec(shape=[], dtype=tf.string)]) def save(self, checkpoint_path): tensor_names = [weight.name for weight in self.model.weights] tensors_to_save = [weight.read_value() for weight in self.model.weights] tf.raw_ops.Save( filename=checkpoint_path, tensor_names=tensor_names, data=tensors_to_save, name='save') return { "checkpoint_path": checkpoint_path } @tf.function(input_signature=[tf.TensorSpec(shape=[], dtype=tf.string)]) def restore(self, checkpoint_path): restored_tensors = {} for var in self.model.weights: restored = tf.raw_ops.Restore( file_pattern=checkpoint_path, tensor_name=var.name, dt=var.dtype, name='restore') var.assign(restored) restored_tensors[var.name] = restored return restored_tensors
原模型转换代码
SAVED_MODEL_DIR = "saved_model" tf.saved_model.save( m, SAVED_MODEL_DIR, signatures={ 'train': m.train.get_concrete_function(), 'infer': m.infer.get_concrete_function(), 'save': m.save.get_concrete_function(), 'restore': m.restore.get_concrete_function(), }) converter = tf.lite.TFLiteConverter.from_saved_model(SAVED_MODEL_DIR) 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('model.tflite', 'wb') as f: f.write(tflite_model)
原Android端代码
class MainActivity : AppCompatActivity() { override fun onCreate(savedInstanceState: Bundle?) { super.onCreate(savedInstanceState) setContentView(R.layout.activity_main) Interpreter(modelPath()).use { interpreter -> val outputFile = File(filesDir, "model.tflite") val inputs: MutableMap<String, Any> = HashMap() inputs["checkpoint_path"] = outputFile.absolutePath val outputs: Map<String, Any> = HashMap() interpreter.runSignature(inputs, outputs, "save") } } private fun modelPath(): File { val file = File(this.filesDir, "model.tflite") if (file.exists() && file.length() > 0) { return file } this.applicationContext.assets.open("model.tflite").use { f -> FileOutputStream(file).use { os -> val buffer = ByteArray(4 * 1024) var read: Int while (f.read(buffer).also { read = it } != -1) { os.write(buffer, 0, read) } os.flush() } return file } } }
问题诊断
- 格式不兼容:原代码用
tf.raw_ops.Save保存的是TensorFlow checkpoint格式,却被命名为.tflite后缀,导致Android端试图用TFLite解析非TFLite格式文件,直接崩溃。 - 输出参数未正确初始化:Android端调用
runSignature时,outputs传入空HashMap,但save函数有返回值,未分配对应的输出容器,导致Interpreter无法处理返回结果。 - 权重保存逻辑适配问题:
tf.raw_ops.Save在TFLite端侧环境下的兼容性不如官方推荐的Checkpoint API,容易引发底层错误。
解决方案
1. 修正模型的save/restore逻辑
改用TensorFlow Checkpoint API管理权重,确保TFLite端侧环境兼容:
class Model(tf.Module): def __init__(self): self.model = tf.keras.Sequential([ tf.keras.layers.Flatten(input_shape=(IMG_SIZE, IMG_SIZE), name='flatten'), tf.keras.layers.Dense(128, activation='relu', name='dense_1'), tf.keras.layers.Dense(10, name='dense_2') ]) self.model.compile( optimizer='sgd', loss=tf.keras.losses.CategoricalCrossentropy(from_logits=True)) # 初始化Checkpoint对象,绑定模型权重 self.checkpoint = tf.train.Checkpoint(model=self.model) # train和infer函数与原代码一致,此处省略 @tf.function(input_signature=[tf.TensorSpec(shape=[], dtype=tf.string)]) def save(self, checkpoint_path): # 保存权重到指定路径 save_path = self.checkpoint.save(checkpoint_path) # 将保存路径转换为Tensor返回 return {"checkpoint_path": tf.convert_to_tensor(save_path)} @tf.function(input_signature=[tf.TensorSpec(shape=[], dtype=tf.string)]) def restore(self, checkpoint_path): # 恢复权重,确保权重匹配 status = self.checkpoint.restore(checkpoint_path) status.assert_consumed() return {"status": tf.convert_to_tensor("success")}
2. 修正Android端调用代码
- 区分checkpoint文件与TFLite模型,使用正确的文件后缀
- 为输出参数分配对应类型的容器
- 启用Interpreter的资源变量支持:
class MainActivity : AppCompatActivity() { private val TAG = "TFLiteOnDeviceTrain" override fun onCreate(savedInstanceState: Bundle?) { super.onCreate(savedInstanceState) setContentView(R.layout.activity_main) // 配置Interpreter,启用资源变量支持 val interpreterOptions = Interpreter.Options().apply { setEnableResourceVariables(true) } Interpreter(modelPath(), interpreterOptions).use { interpreter -> // 保存为checkpoint文件,避免与TFLite模型混淆 val checkpointFile = File(filesDir, "model_weights.ckpt") val inputs = mapOf("checkpoint_path" to checkpointFile.absolutePath) // 为返回的字符串分配ByteBuffer容器 val outputs = mutableMapOf<String, Any>() outputs["checkpoint_path"] = ByteBuffer.allocateDirect(2048).order(ByteOrder.nativeOrder()) // 调用save signature interpreter.runSignature(inputs, outputs, "save") // 读取并打印保存路径 val savedPathBuffer = outputs["checkpoint_path"] as ByteBuffer savedPathBuffer.position(0) val savedPath = String(savedPathBuffer.array()).trim { it == '\u0000' } Log.d(TAG, "权重已保存至: $savedPath") } } private fun modelPath(): File { val file = File(this.filesDir, "model.tflite") if (file.exists() && file.length() > 0) { return file } this.applicationContext.assets.open("model.tflite").use { inputStream -> FileOutputStream(file).use { outputStream -> val buffer = ByteArray(4 * 1024) var readBytes: Int while (inputStream.read(buffer).also { readBytes = it } != -1) { outputStream.write(buffer, 0, readBytes) } outputStream.flush() } return file } } }
3. 完善模型转换配置
确保TFLite转换器支持Checkpoint相关操作:
SAVED_MODEL_DIR = "saved_model" tf.saved_model.save( m, SAVED_MODEL_DIR, signatures={ 'train': m.train.get_concrete_function(), 'infer': m.infer.get_concrete_function(), 'save': m.save.get_concrete_function(), 'restore': m.restore.get_concrete_function(), }) converter = tf.lite.TFLiteConverter.from_saved_model(SAVED_MODEL_DIR) converter.target_spec.supported_ops = [ tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS ] # 必须启用资源变量支持 converter.experimental_enable_resource_variables = True # 允许自定义操作,确保Checkpoint相关操作被正确转换 converter.allow_custom_ops = True tflite_model = converter.convert() with open('model.tflite', 'wb') as f: f.write(tflite_model)
注意事项
- 保存的是权重checkpoint文件,不是完整的TFLite模型。如果需要生成包含新权重的TFLite模型,需在训练完成后将权重导出并重新转换(端侧场景一般保存checkpoint用于后续恢复训练)。
- Android 10及以上版本使用
filesDir存储文件时,无需申请外部存储权限。 - 若需要恢复训练,调用
restoresignature时传入checkpoint文件路径即可。
内容的提问来源于stack exchange,提问作者Arcane
相关产品推荐
相关产品推荐

