perf(ai): cache feature status, downscale bitmaps, lazy-init clients

- Cache the AICore FeatureStatus check per service instance — was
  re-running an RPC on every image attach.
- Two-pass decode with inSampleSize so 12 MP camera shots become
  ~1024 px before we hand them to the describer (avoids 40+ MB
  ARGB_8888 allocations and the GC churn that follows).
- Switch ML Kit clients to var + lazy-on-first-use so close() no
  longer triggers init for clients we never invoked.
- Wrap the composable's suggestAltText call in try/finally so a
  cancellation mid-inference resets the spinner state.
This commit is contained in:
Claude
2026-04-27 13:54:18 +00:00
parent 952e0a2192
commit f2c58adcf1
2 changed files with 65 additions and 21 deletions
@@ -127,8 +127,12 @@ fun ImageVideoDescription(
LaunchedEffect(firstImageUri) { LaunchedEffect(firstImageUri) {
if (firstImageUri == null || message.isNotEmpty()) return@LaunchedEffect if (firstImageUri == null || message.isNotEmpty()) return@LaunchedEffect
isLabeling = true isLabeling = true
val suggestion = labelService.suggestAltText(firstImageUri) val suggestion =
isLabeling = false try {
labelService.suggestAltText(firstImageUri)
} finally {
isLabeling = false
}
if (suggestion != null && message.isEmpty()) { if (suggestion != null && message.isEmpty()) {
message = suggestion message = suggestion
aiSuggested = true aiSuggested = true
@@ -25,10 +25,12 @@ import android.graphics.Bitmap
import android.graphics.BitmapFactory import android.graphics.BitmapFactory
import android.net.Uri import android.net.Uri
import com.google.mlkit.genai.common.FeatureStatus import com.google.mlkit.genai.common.FeatureStatus
import com.google.mlkit.genai.imagedescription.ImageDescriber
import com.google.mlkit.genai.imagedescription.ImageDescriberOptions import com.google.mlkit.genai.imagedescription.ImageDescriberOptions
import com.google.mlkit.genai.imagedescription.ImageDescription import com.google.mlkit.genai.imagedescription.ImageDescription
import com.google.mlkit.genai.imagedescription.ImageDescriptionRequest import com.google.mlkit.genai.imagedescription.ImageDescriptionRequest
import com.google.mlkit.vision.common.InputImage import com.google.mlkit.vision.common.InputImage
import com.google.mlkit.vision.label.ImageLabeler
import com.google.mlkit.vision.label.ImageLabeling import com.google.mlkit.vision.label.ImageLabeling
import com.google.mlkit.vision.label.defaults.ImageLabelerOptions import com.google.mlkit.vision.label.defaults.ImageLabelerOptions
import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.Dispatchers
@@ -45,24 +47,33 @@ import kotlin.coroutines.resume
class MLKitImageLabelService( class MLKitImageLabelService(
private val context: Context, private val context: Context,
) { ) {
private val labeler by lazy { private var labeler: ImageLabeler? = null
ImageLabeling.getClient(ImageLabelerOptions.DEFAULT_OPTIONS) private var describer: ImageDescriber? = null
}
private val describer by lazy { // FeatureStatus is an Int enum. Cached per-instance — describer availability does not flip
runCatching { // mid-session in practice, and one composer mount only needs to ask AICore once.
ImageDescription.getClient( @Volatile private var cachedGenAiStatus: Int? = null
ImageDescriberOptions.builder(context).build(),
) private fun ensureLabeler(): ImageLabeler =
}.getOrNull() labeler ?: ImageLabeling.getClient(ImageLabelerOptions.DEFAULT_OPTIONS).also { labeler = it }
}
private fun ensureDescriber(): ImageDescriber? =
describer
?: try {
ImageDescription
.getClient(ImageDescriberOptions.builder(context).build())
.also { describer = it }
} catch (_: Exception) {
null
}
suspend fun labelImage(uri: Uri): List<Pair<String, Float>> = suspend fun labelImage(uri: Uri): List<Pair<String, Float>> =
withContext(Dispatchers.IO) { withContext(Dispatchers.IO) {
try { try {
val image = InputImage.fromFilePath(context, uri) val image = InputImage.fromFilePath(context, uri)
val client = ensureLabeler()
suspendCancellableCoroutine { cont -> suspendCancellableCoroutine { cont ->
labeler client
.process(image) .process(image)
.addOnSuccessListener { labels -> .addOnSuccessListener { labels ->
cont.resume(labels.map { it.text to it.confidence }) cont.resume(labels.map { it.text to it.confidence })
@@ -79,14 +90,19 @@ class MLKitImageLabelService(
private suspend fun describeWithGenAi(uri: Uri): String? = private suspend fun describeWithGenAi(uri: Uri): String? =
withContext(Dispatchers.IO) { withContext(Dispatchers.IO) {
val client = describer ?: return@withContext null val client = ensureDescriber() ?: return@withContext null
try { try {
val status = client.checkFeatureStatus().get() val status =
cachedGenAiStatus ?: client.checkFeatureStatus().get().also { cachedGenAiStatus = it }
if (status != FeatureStatus.AVAILABLE) return@withContext null if (status != FeatureStatus.AVAILABLE) return@withContext null
val bitmap = loadBitmap(uri) ?: return@withContext null val bitmap = loadDownscaledBitmap(uri) ?: return@withContext null
val request = ImageDescriptionRequest.builder(bitmap).build() val request = ImageDescriptionRequest.builder(bitmap).build()
val description = client.runInference(request).get().description client
description?.trim()?.takeIf { it.isNotEmpty() } .runInference(request)
.get()
.description
?.trim()
?.takeIf { it.isNotEmpty() }
} catch (_: Exception) { } catch (_: Exception) {
null null
} }
@@ -99,20 +115,44 @@ class MLKitImageLabelService(
return confident.take(MAX_LABELS).joinToString(", ") return confident.take(MAX_LABELS).joinToString(", ")
} }
private fun loadBitmap(uri: Uri): Bitmap? = // Two-pass decode keeps a 12 MP camera shot from blowing past 40 MB of ARGB_8888 — the
// on-device describer downscales internally anyway, so a ~1024 px input is plenty.
private fun loadDownscaledBitmap(uri: Uri): Bitmap? =
try { try {
context.contentResolver.openInputStream(uri)?.use { BitmapFactory.decodeStream(it) } val bounds = BitmapFactory.Options().apply { inJustDecodeBounds = true }
context.contentResolver.openInputStream(uri)?.use { BitmapFactory.decodeStream(it, null, bounds) }
val opts =
BitmapFactory.Options().apply {
inSampleSize = sampleSizeFor(bounds.outWidth, bounds.outHeight, TARGET_DIM_PX)
}
context.contentResolver.openInputStream(uri)?.use { BitmapFactory.decodeStream(it, null, opts) }
} catch (_: Exception) { } catch (_: Exception) {
null null
} }
private fun sampleSizeFor(
width: Int,
height: Int,
target: Int,
): Int {
if (width <= 0 || height <= 0) return 1
var sample = 1
var maxDim = maxOf(width, height)
while (maxDim / sample > target) sample *= 2
return sample
}
fun close() { fun close() {
labeler.close() labeler?.close()
labeler = null
describer?.close() describer?.close()
describer = null
cachedGenAiStatus = null
} }
companion object { companion object {
private const val MIN_CONFIDENCE = 0.6f private const val MIN_CONFIDENCE = 0.6f
private const val MAX_LABELS = 5 private const val MAX_LABELS = 5
private const val TARGET_DIM_PX = 1024
} }
} }