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