unicorn who dev
Complete FireViewer Android learning SDK, source and verification records
dee7f43 verified Download android/OnDeviceLearning.kt from fireviewer/litert-models: direct link, hf CLI and curl.
- Browser
- Download file 6.36 kB
-
https://huggingface.co/fireviewer/litert-models/resolve/main/android/OnDeviceLearning.kt
- Command line
-
hf download hf://fireviewer/litert-models/android/OnDeviceLearning.kt
-
curl -L -o OnDeviceLearning.kt https://huggingface.co/fireviewer/litert-models/resolve/main/android/OnDeviceLearning.kt
6.36 kB
| 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. | |
| */ | |
| 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) | |
| } | |
| } | |
| 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. | |
| */ | |
| 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 | |
| } | |
| 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") | |
| } | |
| 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 } | |
| } | |
| } | |
| } | |