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

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
        }
    }
}

问题诊断

  1. 格式不兼容:原代码用tf.raw_ops.Save保存的是TensorFlow checkpoint格式,却被命名为.tflite后缀,导致Android端试图用TFLite解析非TFLite格式文件,直接崩溃。
  2. 输出参数未正确初始化:Android端调用runSignature时,outputs传入空HashMap,但save函数有返回值,未分配对应的输出容器,导致Interpreter无法处理返回结果。
  3. 权重保存逻辑适配问题: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存储文件时,无需申请外部存储权限。
  • 若需要恢复训练,调用restore signature时传入checkpoint文件路径即可。

内容的提问来源于stack exchange,提问作者Arcane

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 19:10:32