package org.fireviewer.litert.training import android.util.AtomicFile import org.json.JSONObject import org.tensorflow.lite.Interpreter import java.io.File import java.nio.ByteBuffer import java.nio.ByteOrder import java.security.MessageDigest import java.util.UUID /** CPU training and durable optimizer resume; call from a background dispatcher. * One session is serialized. Only reviewed annotations belong in targets. * The model's contract states which parameters are trainable. */ class OnDeviceLearning( private val model: File, private val labelSchemaSha256: String, threads: Int = 2 ) : AutoCloseable { private val modelSha256 = sha256(model) private val engine = Interpreter(model, Interpreter.Options().setNumThreads(threads).setUseXNNPACK(false)) init { require(threads in 1..8) require(labelSchemaSha256.matches(Regex("[0-9a-f]{64}"))) require(engine.signatureKeys.toSet().containsAll(setOf("train", "infer", "save", "restore"))) } data class Output(val shape: List, val values: FloatArray) /** Nested primitive arrays carry their current shape to runSignature. * Image models accept x as Array>> in NCHW order. * Pixels must already use the normalization from android_model_config.json. */ @Synchronized fun infer(inputs: Map): Map { engine.runSignature(inputs, mutableMapOf(), "infer") return engine.getSignatureOutputs("infer").associateWith { name -> val tensor = engine.getOutputTensorFromSignature(name, "infer") require(tensor.dataType().name == "FLOAT32") require(tensor.numBytes() <= 128 * 1024 * 1024) val values = FloatArray(tensor.numElements()) tensor.asReadOnlyBuffer().order(ByteOrder.nativeOrder()).asFloatBuffer().get(values) require(values.all { it.isFinite() }) Output(tensor.shape().toList(), values) } } @Synchronized fun train(inputs: Map, target: FloatArray, learningRate: Float): Float { require(learningRate.isFinite() && learningRate in .000001f..1f) val tensor = engine.getInputTensorFromSignature("y", "train") require(tensor.dataType().name == "FLOAT32" && target.size == tensor.numElements()) require(target.all { it.isFinite() }) val call = inputs.toMutableMap() call["y"] = floats(target) call["learning_rate"] = floats(floatArrayOf(learningRate)) engine.runSignature(call, mutableMapOf(), "train") val loss = engine.getOutputTensorFromSignature("loss", "train") .asReadOnlyBuffer().order(ByteOrder.nativeOrder()).getFloat(0) require(loss.isFinite()) { "Non-finite loss; discard this candidate and restore its last checkpoint" } return loss } /** The generation becomes current only after all checkpoint files are hashed. * Previous generations stay available for rollback; this method never deletes them. */ @Synchronized fun save(root: File, datasetRevision: String, reviewedSamplesSeen: Long): File { require(reviewedSamplesSeen >= 0) root.mkdirs() val generation = File(root, UUID.randomUUID().toString()).apply { check(mkdir()) } val prefix = File(generation, "weights") engine.runSignature(mapOf("checkpoint_path" to prefix.absolutePath), mutableMapOf(), "save") val files = generation.walkTopDown().filter { it.isFile }.toList() require(files.isNotEmpty() && files.all { it.length() > 0 }) val hashes = JSONObject() files.forEach { file -> hashes.put(file.relativeTo(generation).invariantSeparatorsPath, sha256(file)) } val receipt = JSONObject().put("schema", 1).put("modelSha256", modelSha256) .put("labelSchemaSha256", labelSchemaSha256).put("datasetRevision", datasetRevision) .put("reviewedSamplesSeen", reviewedSamplesSeen).put("checkpointPrefix", "weights") .put("files", hashes) writeAtomic(File(generation, "receipt.json"), receipt.toString(2)) writeAtomic(File(root, "CURRENT"), generation.name) return generation } @Synchronized fun restore(generation: File) { val receipt = JSONObject(AtomicFile(File(generation, "receipt.json")).openRead().bufferedReader().use { it.readText() }) require(receipt.getInt("schema") == 1) require(receipt.getString("modelSha256") == modelSha256) { "Checkpoint belongs to another model revision" } require(receipt.getString("labelSchemaSha256") == labelSchemaSha256) { "Checkpoint label schema differs" } val hashes = receipt.getJSONObject("files") require(hashes.length() > 0) hashes.keys().forEach { name -> val file = File(generation, name) require(file.canonicalPath.startsWith(generation.canonicalPath + File.separator)) require(file.isFile && sha256(file) == hashes.getString(name)) { "Checkpoint missing or changed" } } val prefix = File(generation, receipt.getString("checkpointPrefix")) require(prefix.canonicalPath.startsWith(generation.canonicalPath + File.separator)) engine.runSignature(mapOf("checkpoint_path" to prefix.absolutePath), mutableMapOf(), "restore") } @Synchronized override fun close() = engine.close() companion object { private fun floats(values: FloatArray) = ByteBuffer.allocateDirect(values.size * 4) .order(ByteOrder.nativeOrder()).apply { asFloatBuffer().put(values) } private fun sha256(file: File): String { val digest = MessageDigest.getInstance("SHA-256") file.inputStream().use { stream -> val buffer = ByteArray(1024 * 1024) while (true) { val count = stream.read(buffer); if (count < 0) break; digest.update(buffer, 0, count) } } return digest.digest().joinToString("") { "%02x".format(it) } } private fun writeAtomic(file: File, text: String) { val atomic = AtomicFile(file); val stream = atomic.startWrite() try { stream.write(text.toByteArray(Charsets.UTF_8)); atomic.finishWrite(stream) } catch (error: Exception) { atomic.failWrite(stream); throw error } } } }