File size: 6,357 Bytes
dee7f43 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 | 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<Int>, val values: FloatArray)
/** Nested primitive arrays carry their current shape to runSignature.
* Image models accept x as Array<Array<Array<FloatArray>>> in NCHW order.
* Pixels must already use the normalization from android_model_config.json.
*/
@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
}
/** 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 }
}
}
}
|