aiflow-math-ink-06-intermediate / android /aiflow-math-ink-runtime /src /test /java /ai /aiflow /mathink /AIFlowMathInkTest.kt
| package ai.aiflow.mathink | |
| import org.junit.Assert.assertEquals | |
| import org.junit.Assert.assertFalse | |
| import org.junit.Assert.assertTrue | |
| import org.junit.Test | |
| class AIFlowMathInkTest { | |
| /** 필요 변수: 결정적 logit. 작동 원리: native runtime 없이 API와 privacy 계약을 검증한다. */ | |
| private class FakeSession : LiteRtSession { | |
| override val modelVersion = "test-0.6" | |
| override val exactLabelCount = 3 | |
| var onlineCalls = 0 | |
| var rasterCalls = 0 | |
| override fun runOnline(features: FloatArray): FloatArray { | |
| assertEquals(128 * 19, features.size) | |
| onlineCalls += 1 | |
| return floatArrayOf(0.0f, 3.0f, 1.0f) | |
| } | |
| override fun runRaster(rasterInk: FloatArray): FloatArray { | |
| assertEquals(128 * 128, rasterInk.size) | |
| rasterCalls += 1 | |
| return floatArrayOf(2.0f, 0.0f, 1.0f) | |
| } | |
| override fun runRasterDebug(rasterInk: FloatArray): RasterInference { | |
| assertEquals(128 * 128, rasterInk.size) | |
| rasterCalls += 1 | |
| return RasterInference( | |
| exactLogits = floatArrayOf(2.0f, 0.0f, 1.0f), | |
| virtualHypotheses = List(4) { hypothesis -> | |
| VirtualStrokeHypothesis( | |
| points = FloatArray(128 * 2) { hypothesis.toFloat() }, | |
| stateLogits = FloatArray(128 * 3), | |
| progress = FloatArray(128) { it / 127.0f }, | |
| logProbability = -hypothesis.toFloat(), | |
| ) | |
| }, | |
| ) | |
| } | |
| override fun close() = Unit | |
| } | |
| private fun stroke(withTime: Boolean = true) = InkStroke( | |
| strokeId = 7, | |
| order = 0, | |
| points = listOf( | |
| InkPoint(10.0f, 10.0f, if (withTime) 0L else null), | |
| InkPoint(40.0f, 40.0f, if (withTime) 200L else null), | |
| InkPoint(80.0f, 20.0f, if (withTime) 400L else null), | |
| ), | |
| ) | |
| fun canonicalEncoderPreservesRawAndObservedAnchors() { | |
| val ink = CanonicalTapEncoder().encodeOnline( | |
| listOf(stroke()), | |
| InkCanvas(100.0f, 100.0f), | |
| ) | |
| assertEquals(TimestampMode.OBSERVED, ink.timestampMode) | |
| assertEquals(7, ink.rawStrokes.single().strokeId) | |
| assertTrue(ink.canonicalTaps.first().strokeStart) | |
| assertTrue(ink.canonicalTaps.last().strokeEnd) | |
| assertTrue(ink.canonicalTaps.last().penUp) | |
| assertEquals(128 * 19, ink.features.size) | |
| assertTrue((0 until 128).all { ink.features[it * 19 + 17] == 0.0f }) | |
| } | |
| fun canonicalFeaturesMatchPythonTrainingPreprocessorWithinFloatTolerance() { | |
| val ink = CanonicalTapEncoder().encodeOnline( | |
| listOf(stroke()), | |
| InkCanvas(100.0f, 100.0f), | |
| ) | |
| val expected = mapOf( | |
| 0 to floatArrayOf( | |
| 0.0f, 0.0f, 0.0625f, 0.3125f, 0.0f, 0.0f, 0.0f, 1.0f, 0.0f, | |
| 2.3333333f, 0.1f, 0.4f, 0.3f, 0.25f, 1.0f, 0.0f, 0.0f, 0.0f, 0.0f, | |
| ), | |
| 1 to floatArrayOf( | |
| 0.006602364f, 0.0185185f, 0.06827707f, 0.31827706f, 0.33582273f, | |
| 0.94192517f, 0.0f, 0.0f, 0.007874016f, 2.3333333f, 0.1f, 0.4f, | |
| 0.3f, 0.25f, 1.0f, 0.0030811024f, 1.3258247f, 0.0f, 0.0f, | |
| ), | |
| 63 to floatArrayOf( | |
| 0.4375688f, 0.9423623f, 0.4453727f, 0.60648084f, 0.804563f, | |
| -0.5938673f, -0.000012785f, 0.0f, 0.496063f, 2.3333333f, 0.1f, | |
| 0.4f, 0.3f, 0.25f, 1.0f, 0.0033267364f, 1.2279294f, 0.0f, 0.0f, | |
| ), | |
| 127 to floatArrayOf( | |
| 1.0f, 0.40069035f, 0.9375f, 0.4375f, 0.580574f, -0.81420743f, | |
| 0.0f, 0.0f, 1.0f, 2.3333333f, 0.1f, 0.4f, 0.3f, 0.25f, 1.0f, | |
| 0.002923004f, 1.3975347f, 0.0f, 0.0f, | |
| ), | |
| ) | |
| expected.forEach { (row, values) -> | |
| values.forEachIndexed { channel, value -> | |
| assertEquals( | |
| "row=$row channel=$channel", | |
| value, | |
| ink.features[row * 19 + channel], | |
| 1e-4f, | |
| ) | |
| } | |
| } | |
| } | |
| fun missingTimestampUsesCanonicalModeAndMask() { | |
| val ink = CanonicalTapEncoder().encodeOnline( | |
| listOf(stroke(withTime = false)), | |
| InkCanvas(100.0f, 100.0f), | |
| ) | |
| assertEquals(TimestampMode.CANONICAL, ink.timestampMode) | |
| assertTrue((0 until 128).all { ink.features[it * 19 + 17] == 1.0f }) | |
| } | |
| fun recognizeReturnsOnlyPublicFields() { | |
| val session = FakeSession() | |
| val runtime = AIFlowMathInk(session, listOf("0", "x", "+")) | |
| val result = runtime.recognizeOnline( | |
| listOf(stroke()), | |
| InkCanvas(100.0f, 100.0f), | |
| ) | |
| assertEquals("x", result.candidates.first().token) | |
| assertEquals(result.candidates.first().probability, result.confidence) | |
| assertEquals(1, session.onlineCalls) | |
| val payload = result.toServerPayload() | |
| assertEquals(setOf("candidates", "confidence", "modelVersion", "latencyMs"), payload.keys) | |
| assertFalse(payload.containsKey("rawStrokes")) | |
| assertFalse(payload.containsKey("canonicalTaps")) | |
| assertFalse(payload.containsKey("virtualHypotheses")) | |
| } | |
| fun rasterUsesLocalSessionAndReturnsTopK() { | |
| val session = FakeSession() | |
| val runtime = AIFlowMathInk(session, listOf("0", "x", "+"), topK = 2) | |
| val result = runtime.recognizeRaster( | |
| RasterInput(128, 128, FloatArray(128 * 128)), | |
| ) | |
| assertEquals(listOf("0", "+"), result.candidates.map { it.token }) | |
| assertEquals(1, session.rasterCalls) | |
| } | |
| fun rasterDebugReturnsFourLocalHypothesesWithoutServerPayload() { | |
| val session = FakeSession() | |
| val runtime = AIFlowMathInk(session, listOf("0", "x", "+"), topK = 2) | |
| val debug = runtime.recognizeRasterDebug( | |
| RasterInput(128, 128, FloatArray(128 * 128)), | |
| ) | |
| assertEquals("0", debug.symbol.candidates.first().token) | |
| assertEquals(4, debug.virtualHypotheses.size) | |
| assertEquals(128 * 2, debug.virtualHypotheses.first().points.size) | |
| assertEquals(1, session.rasterCalls) | |
| assertFalse(debug.symbol.toServerPayload().containsKey("virtualHypotheses")) | |
| } | |
| fun androidBenchmarkUsesNearestRankAndAllReleaseChecks() { | |
| class Probe : AndroidBenchmarkProbe { | |
| var nanos = 0L | |
| var pss = 20L * 1024L * 1024L | |
| override fun nowNanos() = nanos | |
| override fun totalPssBytes() = pss | |
| override fun batteryChargeMicroAh() = 1000 | |
| } | |
| val probe = Probe() | |
| val session = object : LiteRtSession { | |
| override val modelVersion = "benchmark-0.6" | |
| override val exactLabelCount = 3 | |
| override fun runOnline(features: FloatArray): FloatArray { | |
| probe.nanos += 10_000_000L | |
| return floatArrayOf(1.0f, 0.0f, 0.0f) | |
| } | |
| override fun runRaster(rasterInk: FloatArray): FloatArray { | |
| probe.nanos += 100_000_000L | |
| return floatArrayOf(1.0f, 0.0f, 0.0f) | |
| } | |
| override fun close() = Unit | |
| } | |
| val report = AndroidBenchmarkRunner(session, probe).run( | |
| deviceTier = AndroidDeviceTier.LOW, | |
| onlineModelSha256 = "a".repeat(64), | |
| rasterModelSha256 = "b".repeat(64), | |
| onlineInputs = listOf(FloatArray(128 * 19)), | |
| rasterInputs = listOf(FloatArray(128 * 128)), | |
| config = AndroidBenchmarkConfig(warmupRuns = 1, measuredRuns = 5), | |
| deviceManufacturer = "fixture", | |
| deviceModel = "fixture-low", | |
| sdkInt = 24, | |
| ) | |
| assertEquals(10.0, report.online.p95Ms, 0.0) | |
| assertEquals(100.0, report.raster.p95Ms, 0.0) | |
| assertEquals(20L * 1024L * 1024L, report.peakPssBytes) | |
| assertTrue(report.gatePassed) | |
| assertEquals(false, report.productValidation) | |
| assertEquals("low", report.toMap()["device_tier"]) | |
| assertEquals( | |
| modelBundleSha256("a".repeat(64), "b".repeat(64)), | |
| report.modelBundleSha256, | |
| ) | |
| } | |
| fun androidBenchmarkFailsAnyExceededThreshold() { | |
| val summary = summarizeLatency(listOf(1.0, 2.0, 3.0, 100.0)) | |
| assertEquals(100.0, summary.p95Ms, 0.0) | |
| assertEquals(2.0, summary.p50Ms, 0.0) | |
| } | |
| } | |