package com.sw.dualscreen.objbox import android.annotation.SuppressLint import android.content.Context import android.graphics.Bitmap import android.net.Uri import androidx.core.graphics.scale import com.sw.dualscreen.MyApp import com.sw.dualscreen.utils.GsonUtils import com.sw.dualscreen.utils.ImageUtil import com.sw.plate.utils.ToastUtils import io.objectbox.Box import io.objectbox.kotlin.boxFor import kotlinx.coroutines.Dispatchers import kotlinx.coroutines.withContext import org.pytorch.IValue import org.pytorch.Module import org.pytorch.torchvision.TensorImageUtils import timber.log.Timber import java.io.File import java.io.FileOutputStream import java.io.IOException import java.io.InputStream object FoodModule { private const val TAG = "FoodModule" private const val THRESHOLD = 0.8 // private const val THRESHOLD = 0.0 private var module: Module? = null // private lateinit var embeddingsList: List> // private lateinit var labelsList: IntArray // private lateinit var classInfo: FoodClassInfo val NO_MEAN_RGB = floatArrayOf(0.0f, 0.0f, 0.0f) val NO_STD_RGB = floatArrayOf(1.0f, 1.0f, 1.0f) // val DEFAULT_FOOD_INDEX = -1 // 1. 定义你的模型固定输入尺寸 (根据你的tflite模型修改,比如224x224) private const val MODEL_INPUT_WIDTH = 224 private const val MODEL_INPUT_HEIGHT = 224 @SuppressLint("SuspiciousIndentation") suspend fun init(context: Context, block: () -> Unit = {}) { // Thread {}.start() withContext(Dispatchers.IO) { val modelPath = copyAssetToCache(context, "best_embedding_model_mobile.pt") module = Module.load(modelPath) //初始化默认重新拉取数据,先清空本地数据 val list = ObjectBox.getAll() if (list.isNotEmpty()) { ObjectBox.removeAll() } //if (box.all.isEmpty()) { // initDefFoodData(context) //} block() } } // fun uri2FloatArray(uri: Uri): FloatArray? { // return MyApp.instance?.let { context -> // ImageUtil.uriToBitmap(context, uri)?.let { // bitmap2FloatArray(it) // } // } // } fun bitmap2FloatArray(originBitmap: Bitmap, isRecycle: Boolean): FloatArray? { var rgb565Bitmap: Bitmap? = null try { val scaledBitmap = originBitmap.scale( MODEL_INPUT_WIDTH, MODEL_INPUT_HEIGHT ) if (isRecycle) { originBitmap.recycle() } rgb565Bitmap = scaledBitmap.copy(Bitmap.Config.RGB_565, false) scaledBitmap.recycle() val inputTensor = TensorImageUtils.bitmapToFloat32Tensor( rgb565Bitmap, NO_MEAN_RGB, // [0.485, 0.456, 0.406] TORCHVISION_NORM_MEAN_RGB NO_STD_RGB // [0.229, 0.224, 0.225] TORCHVISION_NORM_STD_RGB ) if (module == null) { return null } val outputTensor = module?.forward(IValue.from(inputTensor))?.toTensor() return outputTensor?.dataAsFloatArray } catch (e: OutOfMemoryError) { e.printStackTrace() } finally { if (rgb565Bitmap != null && rgb565Bitmap.isRecycled.not()) { rgb565Bitmap.recycle() } //System.gc() //System.runFinalization() } return null } // fun queryFood(uri: Uri, queryCount: Int = 15): List? { // return uri2FloatArray(uri)?.let { // queryFood(it, queryCount) // } // } /** * 返回识别物品名称列表 */ // fun queryFood(bitmap: Bitmap, queryCount: Int = 15): List { // val floatArray = bitmap2FloatArray(bitmap) // return queryFood(floatArray, queryCount) // } /** * 返回识别物品IdNameScore对象列表 */ // fun queryFoodNameScore(bitmap: Bitmap, queryCount: Int = 15): List { // val floatArray = bitmap2FloatArray(bitmap) // return queryFoodNameScore(floatArray, queryCount) // } suspend fun queryFoodNameScore( floatArray: FloatArray?, queryCount: Int = 50 ): List { if (floatArray == null) return emptyList() val startTime = System.currentTimeMillis() val size = ObjectBox.getAll().filter { it.isDel.not() }.size Timber.tag(TAG).d("queryFoodNameScore-已采集向量总数:${size}") //查询比较分数 // val tempList = query.findWithScores().sortedBy { it.score }.map { "${it.get().name}|${it.get().foodIdx}|${it.score}" } // val idScoreList = query.findIdsWithScores() val objScoreList = ObjectBox.query(floatArray, queryCount) Timber.tag(TAG).d("idScoreList:${GsonUtils.toJson(objScoreList)}") val nameScoreList = mutableListOf() objScoreList.forEach { val food = it.get() nameScoreList.add( IdNameScore( id = food.id, name = food.foodName ?: "", score = it.score ) ) } val nameScoreData = GsonUtils.toJson(nameScoreList) Timber.tag(TAG).d("queryFood,耗时:${System.currentTimeMillis() - startTime},数据:$nameScoreData") return nameScoreList } suspend fun getFoodScoreList(bitmap: Bitmap, queryCount: Int = 50): List { val startTime = System.currentTimeMillis() val floatArray = bitmap2FloatArray(bitmap, false) Timber.tag(TAG).d("bitmap2FloatArray,耗时:${System.currentTimeMillis() - startTime}") val nameScoreList = queryFoodNameScore(floatArray, queryCount) if (nameScoreList.isEmpty()) { return emptyList() } val maxScoreList = nameScoreList // .filter { it.score < 1 - THRESHOLD } .groupBy { it.name } .map { (_, value) -> value.minByOrNull { it.score }!! } .toMutableList() // val map = mutableMapOf() // nameScoreList.forEach { // val key = it.name // val count = map[key] ?: 0 // map[key] = count + 1 // } // val orderList = map.entries.sortedByDescending { it.value }.map { it.key }.toMutableList() // val firstFood = nameScoreList[0].name // orderList.remove(firstFood) // orderList.add(0, firstFood) // // val sortedScoreList = maxScoreList.sortedWith(compareBy { // orderList.indexOf(it.name) // }) val sortedScoreList = maxScoreList.sortedBy { it.score } Timber.tag("FoodModule").d("getFoodScoreList数据:${sortedScoreList.toString()}") return sortedScoreList } // fun queryFood(floatArray: FloatArray, queryCount: Int = 15): List { // val query = box.query() // .equal(Food_.isDel, false) // .and() // .nearestNeighbors(Food_.foodVector, floatArray, queryCount) // .build() //// val query: Query = //// box.query(Food_.foodVector.nearestNeighbors(floatArray, queryCount)).build() // //查询比较分数 //// val tempList = query.findWithScores().sortedBy { it.score }.map { "${it.get().name}|${it.get().foodIdx}|${it.score}" } // val map = mutableMapOf() // query.findIdsWithScores().forEach { // val food = box.get(it.id) // Timber.d("${food.foodName}|${food.foodId}|${food.collectId}|${food.version}|${food.otherField}|${it.score}") // if (1 - it.score >= THRESHOLD) { // //FoodQueryResult(id = it.id, name = food.name, foodIdx = food.foodIdx, score = it.score) // food.foodName?.let { key -> // val count = map[key] ?: 0 // map.put(key, count + 1) // } // } // } // val list = map.entries.sortedByDescending { it.value }.map { it.key } // return list // //// val map = mutableMapOf() //// val nameScoreList = queryFoodNameScore(floatArray, queryCount) //// nameScoreList.filter { it.score < 0.05 }.forEach { //// val count = map[it.name] ?: 0 //// map[it.name] = count + 1 //// } //// val list = map.entries.sortedByDescending { it.value }.map { it.key } //// return list // } data class IdNameScore( val id: Long, var name: String, val score: Double ) // fun initDefFoodData(context: Context, action: () -> Unit = {}) { // //val count = box.all.count { it.foodIdx == DEFAULT_FOOD_INDEX } // //if (count > 0) { // // return // //} // val embeddingsJson = AssetsTool.readAssetsFile(context, "data/embeddings.json") // val labelsJson = AssetsTool.readAssetsFile(context, "data/labels.json") // val classInfoJson = AssetsTool.readAssetsFile(context, "data/class_info.json") // // embeddingsList = // Gson().fromJson(embeddingsJson, object : TypeToken>>() {}.type) // labelsList = Gson().fromJson(labelsJson, IntArray::class.java) // classInfo = // Gson().fromJson(classInfoJson, FoodClassInfo::class.java) // // val foodMap = classInfo.idx_to_class // embeddingsList.forEachIndexed { index, floatList -> // val classIdx = labelsList[index] // val foodName = foodMap["$classIdx"] // val array = floatList.toFloatArray() // box.put( // Food( // foodName = foodName, // foodVector = array, // //foodIdx = DEFAULT_FOOD_INDEX // ) // ) // } // action() // } /** * ,此方法的主要目的是:从assets 拷贝到 app的cache目录 * @param context * @param fileName * @return 例如是这样:/data/user/0/com.frizzle.pluginhookandroid9/cache/plugin-debug.apk * * 不可能反正SD */ // fun copyAssetToCache(context: Context, fileName: String): String? { // // 此app的缓存目录 --> 会默认在 cache目录...,可以自己去看看哦 // val cacheDir = context.cacheDir // if (!cacheDir.exists()) { // cacheDir.mkdirs() // TODO 如果没有缓存目录,就创建 // } // val outPath = File(cacheDir, fileName) // TODO 创建输出的文件位置 // if (outPath.exists()) { // outPath.delete() // TODO 如果该文件已经存在,就删掉 // } // var `is`: InputStream? = null // 读取 // var fos: FileOutputStream? = null // 写入 // try { // // 创建文件,如果创建成功,就返回true // val res = outPath.createNewFile() // if (res) { // `is` = context.assets.open(fileName) // 拿到main/assets目录的输入流,用于读取字节 // fos = FileOutputStream(outPath) // 读取出来的字节最终写到outPath // val buf = ByteArray(`is`.available()) // 缓存区 // var byteCount: Int // // // 开始循环读取 // while ((`is`.read(buf).also { byteCount = it }) != -1) { // fos.write(buf, 0, byteCount) // } // return outPath.absolutePath // } // } catch (e: IOException) { // e.printStackTrace() // } finally { // try { // // TODO 一定要记得关闭资源,为了不去性能的磨损 // fos!!.flush() // `is`!!.close() // fos.close() // } catch (e: IOException) { // e.printStackTrace() // } // } // return null // } fun copyAssetToCache(context: Context, fileName: String): String? { val cacheFile = File(context.cacheDir, fileName) val buffer = ByteArray(8 * 1024) var inputStream: InputStream? = null var outputStream: FileOutputStream? = null try { inputStream = context.assets.open(fileName) outputStream = FileOutputStream(cacheFile) var byteCount: Int while (inputStream.read(buffer).also { byteCount = it } != -1) { outputStream.write(buffer, 0, byteCount) } outputStream.channel.force(true) // 强制物理落盘,比flush更彻底 return cacheFile.absolutePath } catch (e: IOException) { e.printStackTrace() Timber.tag("FoodModule").d("文件拷贝失败 fileName=$fileName, error=${e.message}") // 拷贝失败时删除残缺文件,避免下次读取到损坏文件 if (cacheFile.exists()) { cacheFile.delete() } } finally { outputStream?.close() inputStream?.close() } return null } }