- 将默认查询数量从15增加到50 - 移除不必要的分数阈值过滤条件 - 移除复杂的按名称分组和计数逻辑 - 简化排序算法,直接按分数排序 - 添加调试日志输出结果数据
339 lines
13 KiB
Kotlin
339 lines
13 KiB
Kotlin
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<List<Float>>
|
|
// 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<String>? {
|
|
// return uri2FloatArray(uri)?.let {
|
|
// queryFood(it, queryCount)
|
|
// }
|
|
// }
|
|
|
|
/**
|
|
* 返回识别物品名称列表
|
|
*/
|
|
// fun queryFood(bitmap: Bitmap, queryCount: Int = 15): List<String> {
|
|
// val floatArray = bitmap2FloatArray(bitmap)
|
|
// return queryFood(floatArray, queryCount)
|
|
// }
|
|
|
|
/**
|
|
* 返回识别物品IdNameScore对象列表
|
|
*/
|
|
// fun queryFoodNameScore(bitmap: Bitmap, queryCount: Int = 15): List<IdNameScore> {
|
|
// val floatArray = bitmap2FloatArray(bitmap)
|
|
// return queryFoodNameScore(floatArray, queryCount)
|
|
// }
|
|
|
|
suspend fun queryFoodNameScore(
|
|
floatArray: FloatArray?,
|
|
queryCount: Int = 50
|
|
): List<IdNameScore> {
|
|
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<IdNameScore>()
|
|
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<IdNameScore> {
|
|
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<String, Int>()
|
|
// 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<String> {
|
|
// val query = box.query()
|
|
// .equal(Food_.isDel, false)
|
|
// .and()
|
|
// .nearestNeighbors(Food_.foodVector, floatArray, queryCount)
|
|
// .build()
|
|
//// val query: Query<Food> =
|
|
//// 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<String, Int>()
|
|
// 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<String, Int>()
|
|
//// 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<List<List<Float>>>() {}.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
|
|
}
|
|
|
|
} |