| 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 |
|
|
| |
| |
| |
| |
| 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<Int>, val values: FloatArray) |
|
|
| |
| |
| |
| |
| @Synchronized fun infer(inputs: Map<String, Any>): Map<String, Output> { |
| 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<String, Any>, 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 |
| } |
|
|
| |
| |
| |
| @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 } |
| } |
| } |
| } |
|
|