diff --git a/app/build.gradle.kts b/app/build.gradle.kts index 8f28ce0..7e5bb0a 100644 --- a/app/build.gradle.kts +++ b/app/build.gradle.kts @@ -51,9 +51,12 @@ dependencies { // ML Kit GenAI — Gemini Nano via AICore (recommended for Pixel 10) implementation("com.google.mlkit:genai-prompt:1.0.0-beta1") - // MediaPipe LLM Inference — for custom models (Gemma, etc.) + // MediaPipe LLM Inference — legacy fallback for .task/.bin models implementation("com.google.mediapipe:tasks-genai:0.10.24") + // LiteRT-LM — primary backend for Gemma 3n .litertlm models + implementation("com.google.ai.edge.litertlm:litertlm-android:0.9.0-alpha05") + // Embedded HTTP server implementation("org.nanohttpd:nanohttpd:2.3.1") diff --git a/app/src/main/java/com/pixel10/ai/inference/LiteRTModel.kt b/app/src/main/java/com/pixel10/ai/inference/LiteRTModel.kt new file mode 100644 index 0000000..a87b5fa --- /dev/null +++ b/app/src/main/java/com/pixel10/ai/inference/LiteRTModel.kt @@ -0,0 +1,137 @@ +package com.pixel10.ai.inference + +import android.content.Context +import android.util.Log +import com.google.ai.edge.litertlm.Backend +import com.google.ai.edge.litertlm.Engine +import com.google.ai.edge.litertlm.EngineConfig +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.flow.catch +import kotlinx.coroutines.withContext +import java.io.File + +/** + * LiteRT-LM backend for Gemma 3n models (.litertlm format). + * + * This replaces MediaPipe for the newer Gemma 3n E4B/E2B models which use + * the LiteRT-LM runtime. Runs fully on-device using the Tensor G5 GPU. + * + * Model files must be placed in the app's files directory (see [ModelDownloader]). + */ +class LiteRTModel private constructor( + private val engine: Engine, + private val modelName: String +) : OnDeviceModel { + + override val backendName = "LiteRT-LM ($modelName)" + + @Volatile + override var isReady: Boolean = true + private set + + override suspend fun generate( + prompt: String, + maxTokens: Int, + temperature: Float + ): String = withContext(Dispatchers.Default) { + val conversation = engine.createConversation() + try { + conversation.sendMessage(prompt).toString() + } catch (e: Exception) { + Log.e(TAG, "LiteRT inference error", e) + throw OnDeviceModel.InferenceException("Generation failed: ${e.message}", e) + } finally { + conversation.close() + } + } + + override suspend fun generateStreaming( + prompt: String, + onToken: (String) -> Unit + ): String = withContext(Dispatchers.Default) { + val conversation = engine.createConversation() + val sb = StringBuilder() + try { + conversation.sendMessageAsync(prompt) + .catch { e -> + throw OnDeviceModel.InferenceException("Streaming failed: ${e.message}", e) + } + .collect { message -> + val token = message.toString() + sb.append(token) + onToken(token) + } + } finally { + conversation.close() + } + sb.toString() + } + + override fun close() { + isReady = false + engine.close() + } + + companion object { + private const val TAG = "LiteRTModel" + + private val MODEL_EXTENSIONS = listOf("litertlm") + + suspend fun create(context: Context): LiteRTModel = withContext(Dispatchers.IO) { + val modelPath = findModelPath(context) + ?: throw OnDeviceModel.InferenceException( + "No LiteRT-LM model file found.\n" + + "Download a .litertlm model via the app or place one in:\n" + + " ${context.filesDir.absolutePath}/" + ) + + val modelName = File(modelPath).name + Log.i(TAG, "Loading LiteRT-LM model: $modelPath") + + try { + val config = EngineConfig( + modelPath = modelPath, + backend = Backend.GPU + ) + val engine = Engine(config) + withContext(Dispatchers.Default) { + engine.initialize() + } + Log.i(TAG, "LiteRT-LM model loaded: $modelName") + LiteRTModel(engine, modelName) + } catch (gpuError: Exception) { + Log.w(TAG, "GPU backend failed, trying CPU: ${gpuError.message}") + try { + val config = EngineConfig( + modelPath = modelPath, + backend = Backend.CPU + ) + val engine = Engine(config) + withContext(Dispatchers.Default) { + engine.initialize() + } + Log.i(TAG, "LiteRT-LM model loaded on CPU: $modelName") + LiteRTModel(engine, modelName) + } catch (e: Exception) { + throw OnDeviceModel.InferenceException( + "Failed to load LiteRT-LM model from $modelPath: ${e.message}", e + ) + } + } + } + + private fun findModelPath(context: Context): String? { + val searchDirs = listOfNotNull( + context.filesDir, + File(context.filesDir, "models"), + context.getExternalFilesDir(null) + ) + for (dir in searchDirs) { + if (!dir.exists()) continue + dir.listFiles()?.firstOrNull { it.extension in MODEL_EXTENSIONS } + ?.let { return it.absolutePath } + } + return null + } + } +} diff --git a/app/src/main/java/com/pixel10/ai/inference/ModelDownloader.kt b/app/src/main/java/com/pixel10/ai/inference/ModelDownloader.kt index 9105ac7..9d493d4 100644 --- a/app/src/main/java/com/pixel10/ai/inference/ModelDownloader.kt +++ b/app/src/main/java/com/pixel10/ai/inference/ModelDownloader.kt @@ -23,7 +23,7 @@ object ModelDownloader { private const val TAG = "ModelDownloader" private const val HF_BASE = "https://huggingface.co" - /** Available model specs downloadable from HuggingFace. */ + /** Available model specs downloadable from HuggingFace (requires token + license acceptance). */ enum class ModelSpec( val displayName: String, val filename: String, @@ -31,21 +31,27 @@ object ModelDownloader { val sizeMb: Int, val description: String ) { - /** Recommended: best size/quality trade-off, runs fast on Tensor G5. */ - GEMMA_3_1B_Q4( - displayName = "Gemma 3 1B IT (Q4)", - filename = "gemma3-1b-it-int4.task", - repo = "litert-community/Gemma3-1B-IT", - sizeMb = 555, - description = "Best balance — fast & capable (~555 MB)" + /** + * Gemma 3n E4B INT4 — best quality, Tensor G5 optimised, background-safe. + * Accept license at: https://huggingface.co/google/gemma-3n-E4B-it-litert-lm + */ + GEMMA_3N_E4B( + displayName = "Gemma 3n E4B", + filename = "gemma-3n-E4B-it-int4.litertlm", + repo = "google/gemma-3n-E4B-it-litert-lm", + sizeMb = 4920, + description = "Best quality — Tensor G5 optimised (~4.9 GB)" ), - /** Higher quality, slower. Good for complex reasoning. */ - GEMMA_3_1B_Q8( - displayName = "Gemma 3 1B IT (Q8)", - filename = "gemma3-1b-it-int8-web.task", - repo = "litert-community/Gemma3-1B-IT", - sizeMb = 1010, - description = "Higher quality, slower (~1 GB)" + /** + * Gemma 3n E4B Web INT4 — smaller variant, slightly lower quality. + * Same license as above. + */ + GEMMA_3N_E4B_WEB( + displayName = "Gemma 3n E4B (Web)", + filename = "gemma-3n-E4B-it-int4-Web.litertlm", + repo = "google/gemma-3n-E4B-it-litert-lm", + sizeMb = 4280, + description = "Slightly smaller variant (~4.3 GB)" ) } @@ -66,10 +72,10 @@ object ModelDownloader { fun modelFile(context: Context, spec: ModelSpec): File = File(context.filesDir, spec.filename) - /** Legacy compat — returns the file of the installed model, or Q4 path as default. */ + /** Returns the file of the installed model, or E4B path as default. */ fun modelFile(context: Context): File = installedSpec(context)?.let { modelFile(context, it) } - ?: modelFile(context, ModelSpec.GEMMA_3_1B_Q4) + ?: modelFile(context, ModelSpec.GEMMA_3N_E4B) /** * Download [spec] from HuggingFace, using [hfToken] for authentication. @@ -79,7 +85,7 @@ object ModelDownloader { */ suspend fun download( context: Context, - spec: ModelSpec = ModelSpec.GEMMA_3_1B_Q4, + spec: ModelSpec = ModelSpec.GEMMA_3N_E4B, hfToken: String, onProgress: (Progress) -> Unit ) = withContext(Dispatchers.IO) { diff --git a/app/src/main/java/com/pixel10/ai/inference/OnDeviceModel.kt b/app/src/main/java/com/pixel10/ai/inference/OnDeviceModel.kt index 8c94e02..fa221d5 100644 --- a/app/src/main/java/com/pixel10/ai/inference/OnDeviceModel.kt +++ b/app/src/main/java/com/pixel10/ai/inference/OnDeviceModel.kt @@ -112,9 +112,19 @@ interface OnDeviceModel { * Tap "Download Model" in the app UI to get the MediaPipe model automatically. */ suspend fun create(context: Context): OnDeviceModel = withContext(Dispatchers.IO) { - // MediaPipe first — background-safe, GPU-accelerated via Tensor G5 + // LiteRT-LM first — Gemma 3n .litertlm format, GPU-accelerated, background-safe try { - Log.i(TAG, "Attempting MediaPipe LLM with local model...") + Log.i(TAG, "Attempting LiteRT-LM with local .litertlm model...") + val litert = LiteRTModel.create(context) + Log.i(TAG, "LiteRT-LM model ready: ${litert.backendName}") + return@withContext litert + } catch (e: Exception) { + Log.w(TAG, "LiteRT-LM not available: ${e.message}") + } + + // MediaPipe fallback — .task/.bin format, background-safe + try { + Log.i(TAG, "Attempting MediaPipe LLM with local .task model...") val mediapipe = MediaPipeModel.create(context) Log.i(TAG, "MediaPipe model ready: ${mediapipe.backendName}") return@withContext mediapipe @@ -122,7 +132,7 @@ interface OnDeviceModel { Log.w(TAG, "MediaPipe not available: ${e.message}") } - // Gemini Nano fallback — only works when app is in foreground + // Gemini Nano last resort — foreground only try { Log.i(TAG, "Attempting Gemini Nano via ML Kit (foreground only)...") val nano = GeminiNanoModel.create(context) @@ -134,11 +144,10 @@ interface OnDeviceModel { throw InferenceException( "No model loaded yet.\n\n" + - "Tap 'Download Model' in the app to download Gemma 2B (~1.3 GB).\n" + + "Tap 'Download Model' in the app to download Gemma 3n E4B.\n" + "Once downloaded the server works fully in the background.\n\n" + - "Or place a compatible model file in:\n" + - " ${context.filesDir.absolutePath}/\n" + - " Supported: gemma-2b-it-gpu-int4.bin, gemma-3n-E2B.task, etc." + "Or place a .litertlm file in:\n" + + " ${context.filesDir.absolutePath}/" ) } } diff --git a/app/src/main/java/com/pixel10/ai/ui/MainActivity.kt b/app/src/main/java/com/pixel10/ai/ui/MainActivity.kt index 27e3a82..8794d1a 100644 --- a/app/src/main/java/com/pixel10/ai/ui/MainActivity.kt +++ b/app/src/main/java/com/pixel10/ai/ui/MainActivity.kt @@ -82,11 +82,11 @@ class MainActivity : AppCompatActivity() { binding.btnDownloadModel.setOnClickListener { saveHfToken() - startModelDownload(ModelSpec.GEMMA_3_1B_Q4) + startModelDownload(ModelSpec.GEMMA_3N_E4B) } binding.btnDownloadGemma3Q8.setOnClickListener { saveHfToken() - startModelDownload(ModelSpec.GEMMA_3_1B_Q8) + startModelDownload(ModelSpec.GEMMA_3N_E4B_WEB) } updateModelCard() diff --git a/app/src/main/res/values/strings.xml b/app/src/main/res/values/strings.xml index d1e6b69..e9d098b 100644 --- a/app/src/main/res/values/strings.xml +++ b/app/src/main/res/values/strings.xml @@ -19,8 +19,8 @@ HuggingFace token (huggingface.co/settings/tokens) - ⭐ Gemma 3 1B IT Q4 — Fast (~555 MB) - Gemma 3 1B IT Q8 — Higher quality (~1 GB) + ⭐ Gemma 3n E4B — Best quality (~4.9 GB) + Gemma 3n E4B Web — Smaller (~4.3 GB) ✓ %s ready — background inference enabled No local model. Enter HuggingFace token and download. Start Server