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 59cf178..3105fc5 100644
--- a/app/src/main/java/com/pixel10/ai/inference/ModelDownloader.kt
+++ b/app/src/main/java/com/pixel10/ai/inference/ModelDownloader.kt
@@ -10,23 +10,60 @@ import java.net.HttpURLConnection
import java.net.URL
/**
- * Downloads a MediaPipe-compatible Gemma model for background-safe inference.
+ * Downloads a MediaPipe-compatible model for background-safe inference.
*
* Gemini Nano (ML Kit) blocks inference when the app is backgrounded (ErrorCode 30).
* MediaPipe with a local model file has no such restriction — it runs entirely in
* the app process using the Tensor G5 GPU via OpenCL/Vulkan.
*
- * The downloaded model is stored in the app's private files directory and
- * survives app restarts. Only needs to be downloaded once (~1.3 GB).
+ * Three model options (all from Google's MediaPipe CDN):
+ * - [ModelSpec.GEMMA_3N_E4B_CODING] — best coding/reasoning, ~2.5 GB (recommended)
+ * - [ModelSpec.GEMMA_3N_E2B_CODING] — good balance, ~1.5 GB
+ * - [ModelSpec.GEMMA_2B_GENERAL] — lightest, ~1.3 GB
+ *
+ * Custom models (DeepSeek Coder, Qwen2.5-Coder, etc.) can be placed manually in
+ * the app's files directory after converting with ai-edge-torch.
*/
object ModelDownloader {
private const val TAG = "ModelDownloader"
- const val MODEL_FILENAME = "gemma-2b-it-gpu-int4.bin"
- private const val MODEL_URL =
- "https://storage.googleapis.com/mediapipe-models/llm_inference/" +
- "gemma-2b-it-gpu-int4/float16/1/gemma-2b-it-gpu-int4.bin"
+ /** Available model specs that can be downloaded from Google's MediaPipe CDN. */
+ enum class ModelSpec(
+ val displayName: String,
+ val filename: String,
+ val url: String,
+ val sizeMb: Int,
+ val description: String
+ ) {
+ /** Recommended: best coding & reasoning quality via MoE architecture. */
+ GEMMA_3N_E4B_CODING(
+ displayName = "Gemma 3n E4B",
+ filename = "gemma-3n-E4B-it-int4.task",
+ url = "https://storage.googleapis.com/mediapipe-models/llm_inference/" +
+ "gemma-3n-E4B-it-int4/float16/1/gemma-3n-E4B-it-int4.task",
+ sizeMb = 2500,
+ description = "Best coding & reasoning (~2.5 GB)"
+ ),
+ /** Good balance between quality and speed. */
+ GEMMA_3N_E2B_CODING(
+ displayName = "Gemma 3n E2B",
+ filename = "gemma-3n-E2B-it-int4.task",
+ url = "https://storage.googleapis.com/mediapipe-models/llm_inference/" +
+ "gemma-3n-E2B-it-int4/float16/1/gemma-3n-E2B-it-int4.task",
+ sizeMb = 1500,
+ description = "Good balance, faster (~1.5 GB)"
+ ),
+ /** Lightest option — general-purpose, not optimised for code. */
+ GEMMA_2B_GENERAL(
+ displayName = "Gemma 2B",
+ filename = "gemma-2b-it-gpu-int4.bin",
+ url = "https://storage.googleapis.com/mediapipe-models/llm_inference/" +
+ "gemma-2b-it-gpu-int4/float16/1/gemma-2b-it-gpu-int4.bin",
+ sizeMb = 1300,
+ description = "Lightest, general-purpose (~1.3 GB)"
+ )
+ }
data class Progress(
val downloadedBytes: Long,
@@ -34,25 +71,37 @@ object ModelDownloader {
val percent: Int = if (totalBytes > 0) (downloadedBytes * 100 / totalBytes).toInt() else 0
)
+ /** Returns true if any supported model is present in the app's files directory. */
fun isModelPresent(context: Context): Boolean =
- modelFile(context).let { it.exists() && it.length() > 1_000_000L }
+ ModelSpec.values().any { modelFile(context, it).let { f -> f.exists() && f.length() > 1_000_000L } }
- fun modelFile(context: Context): File = File(context.filesDir, MODEL_FILENAME)
+ /** Returns the installed [ModelSpec], or null if no model is present. */
+ fun installedSpec(context: Context): ModelSpec? =
+ ModelSpec.values().firstOrNull { modelFile(context, it).let { f -> f.exists() && f.length() > 1_000_000L } }
+
+ fun modelFile(context: Context, spec: ModelSpec): File =
+ File(context.filesDir, spec.filename)
+
+ /** Legacy compat — returns the file of the installed model, or Gemma 3n E4B path as default. */
+ fun modelFile(context: Context): File =
+ installedSpec(context)?.let { modelFile(context, it) }
+ ?: modelFile(context, ModelSpec.GEMMA_3N_E4B_CODING)
/**
- * Download the model, reporting progress via [onProgress].
+ * Download [spec], reporting progress via [onProgress].
* Supports resume — if a partial file exists, continues from where it left off.
*/
suspend fun download(
context: Context,
+ spec: ModelSpec = ModelSpec.GEMMA_3N_E4B_CODING,
onProgress: (Progress) -> Unit
) = withContext(Dispatchers.IO) {
- val dest = modelFile(context)
+ val dest = modelFile(context, spec)
val alreadyDownloaded = if (dest.exists()) dest.length() else 0L
- Log.i(TAG, "Download starting (already have $alreadyDownloaded bytes)")
+ Log.i(TAG, "Download starting ${spec.displayName} (already have $alreadyDownloaded bytes)")
- val conn = URL(MODEL_URL).openConnection() as HttpURLConnection
+ val conn = URL(spec.url).openConnection() as HttpURLConnection
try {
conn.connectTimeout = 30_000
conn.readTimeout = 60_000
@@ -90,7 +139,7 @@ object ModelDownloader {
}
fun deleteModel(context: Context) {
- modelFile(context).delete()
- Log.i(TAG, "Model deleted")
+ ModelSpec.values().forEach { modelFile(context, it).delete() }
+ Log.i(TAG, "All models deleted")
}
}
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 ab5e782..1c99a1a 100644
--- a/app/src/main/java/com/pixel10/ai/ui/MainActivity.kt
+++ b/app/src/main/java/com/pixel10/ai/ui/MainActivity.kt
@@ -19,6 +19,7 @@ import androidx.lifecycle.lifecycleScope
import com.pixel10.ai.R
import com.pixel10.ai.databinding.ActivityMainBinding
import com.pixel10.ai.inference.ModelDownloader
+import com.pixel10.ai.inference.ModelDownloader.ModelSpec
import com.pixel10.ai.server.ApiServerService
import kotlinx.coroutines.launch
import java.text.SimpleDateFormat
@@ -73,8 +74,14 @@ class MainActivity : AppCompatActivity() {
if (service?.isRunning == true) stopServer() else startServer()
}
+ binding.btnDownloadGemma3nE4b.setOnClickListener {
+ startModelDownload(ModelSpec.GEMMA_3N_E4B_CODING)
+ }
+ binding.btnDownloadGemma3nE2b.setOnClickListener {
+ startModelDownload(ModelSpec.GEMMA_3N_E2B_CODING)
+ }
binding.btnDownloadModel.setOnClickListener {
- startModelDownload()
+ startModelDownload(ModelSpec.GEMMA_2B_GENERAL)
}
updateModelCard()
@@ -121,34 +128,33 @@ class MainActivity : AppCompatActivity() {
}
}
- private fun startModelDownload() {
+ private fun startModelDownload(spec: ModelSpec) {
if (downloading) return
downloading = true
- binding.btnDownloadModel.isEnabled = false
- binding.btnDownloadModel.text = "Downloading…"
+ setDownloadButtonsEnabled(false)
binding.progressDownload.visibility = View.VISIBLE
- binding.tvModelDownloadStatus.text = "Starting download…"
+ binding.tvModelDownloadStatus.text = "Starting download: ${spec.displayName}…"
lifecycleScope.launch {
try {
- ModelDownloader.download(this@MainActivity) { progress ->
+ ModelDownloader.download(this@MainActivity, spec) { progress ->
runOnUiThread {
binding.progressDownload.progress = progress.percent
val mb = progress.downloadedBytes / 1_048_576
val total = progress.totalBytes / 1_048_576
- binding.tvModelDownloadStatus.text = "Downloading… ${mb}MB / ${total}MB (${progress.percent}%)"
+ binding.tvModelDownloadStatus.text =
+ "${spec.displayName}: ${mb}MB / ${total}MB (${progress.percent}%)"
}
}
runOnUiThread {
downloading = false
updateModelCard()
- appendLog("Model downloaded — background inference enabled")
+ appendLog("${spec.displayName} downloaded — background inference enabled")
}
} catch (e: Exception) {
runOnUiThread {
downloading = false
- binding.btnDownloadModel.isEnabled = true
- binding.btnDownloadModel.text = getString(R.string.btn_download_model)
+ setDownloadButtonsEnabled(true)
binding.progressDownload.visibility = View.GONE
binding.tvModelDownloadStatus.text = "Download failed: ${e.message}"
appendLog("Download error: ${e.message}")
@@ -157,17 +163,26 @@ class MainActivity : AppCompatActivity() {
}
}
+ private fun setDownloadButtonsEnabled(enabled: Boolean) {
+ binding.btnDownloadGemma3nE4b.isEnabled = enabled
+ binding.btnDownloadGemma3nE2b.isEnabled = enabled
+ binding.btnDownloadModel.isEnabled = enabled
+ }
+
private fun updateModelCard() {
- val present = ModelDownloader.isModelPresent(this)
- if (present) {
- binding.tvModelDownloadStatus.text = getString(R.string.model_downloaded)
+ val spec = ModelDownloader.installedSpec(this)
+ if (spec != null) {
+ binding.tvModelDownloadStatus.text = getString(R.string.model_downloaded, spec.displayName)
+ binding.btnDownloadGemma3nE4b.visibility = View.GONE
+ binding.btnDownloadGemma3nE2b.visibility = View.GONE
binding.btnDownloadModel.visibility = View.GONE
binding.progressDownload.visibility = View.GONE
} else {
binding.tvModelDownloadStatus.text = getString(R.string.model_not_downloaded)
+ binding.btnDownloadGemma3nE4b.visibility = View.VISIBLE
+ binding.btnDownloadGemma3nE2b.visibility = View.VISIBLE
binding.btnDownloadModel.visibility = View.VISIBLE
- binding.btnDownloadModel.isEnabled = true
- binding.btnDownloadModel.text = getString(R.string.btn_download_model)
+ setDownloadButtonsEnabled(true)
binding.progressDownload.visibility = View.GONE
}
}
diff --git a/app/src/main/res/layout/activity_main.xml b/app/src/main/res/layout/activity_main.xml
index cf00e09..ad9d928 100644
--- a/app/src/main/res/layout/activity_main.xml
+++ b/app/src/main/res/layout/activity_main.xml
@@ -142,11 +142,31 @@
android:visibility="gone" />
+
+
+
+
diff --git a/app/src/main/res/values/strings.xml b/app/src/main/res/values/strings.xml
index eeb9478..062f435 100644
--- a/app/src/main/res/values/strings.xml
+++ b/app/src/main/res/values/strings.xml
@@ -18,9 +18,11 @@
Model: not loaded
- Download Model (~1.3 GB)
- ✓ Gemma 2B ready — background inference enabled
- No local model. Download to enable background inference.
+ ⭐ Gemma 3n E4B — Best coding (~2.5 GB)
+ Gemma 3n E2B — Faster (~1.5 GB)
+ Gemma 2B — Lightest (~1.3 GB)
+ ✓ %s ready — background inference enabled
+ No local model. Download one to enable background inference.
Start Server
Stop Server
Port: