采集向量优化调试

This commit is contained in:
2026-01-08 15:02:44 +08:00
parent 378e6088e7
commit 2a55ddb617
35 changed files with 663 additions and 1112 deletions
@@ -1,18 +1,14 @@
package com.sw.dualscreen.objbox
import android.annotation.SuppressLint
import android.content.Context
import android.graphics.Bitmap
import android.net.Uri
import com.google.gson.Gson
import com.google.gson.reflect.TypeToken
import com.sw.dualscreen.MyApp
import com.sw.dualscreen.utils.AssetsTool
import com.sw.dualscreen.utils.GsonUtils
import com.sw.dualscreen.utils.ImageUtil
import com.sw.plate.App
import io.objectbox.Box
import io.objectbox.kotlin.boxFor
import io.objectbox.query.Query
import org.pytorch.IValue
import org.pytorch.Module
import org.pytorch.torchvision.TensorImageUtils
@@ -23,22 +19,27 @@ import java.io.IOException
import java.io.InputStream
object FoodModule {
private const val TAG = "FoodModule"
private const val THRESHOLD = 0.8
private var module_mobile: Module? = null
// private const val THRESHOLD = 0.0
private lateinit var module: Module
private lateinit var box: Box<Food>
private lateinit var embeddingsList: List<List<Float>>
private lateinit var labelsList: IntArray
private lateinit var classInfo: FoodClassInfo
// 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
// val DEFAULT_FOOD_INDEX = -1
@SuppressLint("SuspiciousIndentation")
fun init(context: Context, block: () -> Unit = {}) {
Thread {
module_mobile = Module.load(copyAssetToCache(context, "best_embedding_model_mobile.pt"))
module = Module.load(copyAssetToCache(context, "best_embedding_model_mobile.pt"))
box = ObjectBox.boxStore.boxFor(Food::class)
//初始化默认重新拉取数据,先清空本地数据
// ObjectBox.boxStore.runInTx {
if (box.all.isNotEmpty()) {
box.removeAll()
}
@@ -58,68 +59,79 @@ object FoodModule {
}
fun bitmap2FloatArray(bitmap: Bitmap): FloatArray {
var tensorStartTime = System.currentTimeMillis()
val inputTensor = TensorImageUtils.bitmapToFloat32Tensor(
bitmap,
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_mobile == null) {
init(App.getContext())
}
val outputTensor = module_mobile!!.forward(IValue.from(inputTensor)).toTensor()
Timber.tag(TAG).d("bitmap2FloatArray-inputTensor耗时:${System.currentTimeMillis() - tensorStartTime}")
// if (module_mobile == null) {
// init(App.getContext())
// }
tensorStartTime = System.currentTimeMillis()
val outputTensor = module.forward(IValue.from(inputTensor)).toTensor()
Timber.tag(TAG).d("bitmap2FloatArray-toTensor耗时:${System.currentTimeMillis() - tensorStartTime}")
return outputTensor.dataAsFloatArray
}
fun queryFood(uri: Uri, queryCount: Int = 15): List<String>? {
return uri2FloatArray(uri)?.let {
queryFood(it, queryCount)
}
}
// 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)
}
// 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)
}
// fun queryFoodNameScore(bitmap: Bitmap, queryCount: Int = 15): List<IdNameScore> {
// val floatArray = bitmap2FloatArray(bitmap)
// return queryFoodNameScore(floatArray, queryCount)
// }
fun queryFoodNameScore(floatArray: FloatArray, queryCount: Int = 15): List<IdNameScore> {
val startTime = System.currentTimeMillis()
val query: Query<Food> =
box.query(Food_.foodVector.nearestNeighbors(floatArray, queryCount)).build()
Timber.tag(TAG).d("queryFoodNameScore-已采集向量总数:${box.all.filter { it.isDel.not() }.size}")
box.store.runInTx { }
val query = box.query()
.equal(Food_.isDel, false)
.and()
.nearestNeighbors(Food_.foodVector, floatArray, queryCount)
.build()
//查询比较分数
// val tempList = query.findWithScores().sortedBy { it.score }.map { "${it.get().name}|${it.get().foodIdx}|${it.score}" }
val idScoreList = query.findIdsWithScores()
// val idScoreList = query.findIdsWithScores()
val objScoreList = query.findWithScores()
Timber.tag(TAG).d("idScoreList:${GsonUtils.toJson(objScoreList)}")
val nameScoreList = mutableListOf<IdNameScore>()
idScoreList.forEach {
objScoreList.forEach {
val food = it.get()
nameScoreList.add(
IdNameScore(
id = it.id,
name = box.get(it.id).foodName ?: "",
id = food.id,
name = food.foodName ?: "",
score = it.score
)
)
}
val nameScoreData = GsonUtils.toJson(nameScoreList)
Timber.tag("FoodModule")
.d("registerDataChange,queryFood,耗时:${System.currentTimeMillis() - startTime},数据:$nameScoreData")
Timber.tag(TAG).d("queryFood,耗时:${System.currentTimeMillis() - startTime},数据:$nameScoreData")
query.close()
return nameScoreList
}
fun getFoodScoreList(bitmap: Bitmap, queryCount: Int = 15): List<IdNameScore> {
val startTime = System.currentTimeMillis()
val floatArray = bitmap2FloatArray(bitmap)
Timber.tag("FoodModule")
.d("registerDataChange,bitmap2FloatArray,耗时:${System.currentTimeMillis() - startTime}")
Timber.tag(TAG).d("bitmap2FloatArray,耗时:${System.currentTimeMillis() - startTime}")
val nameScoreList = queryFoodNameScore(floatArray, queryCount)
if (nameScoreList.isEmpty()) {
return emptyList()
@@ -147,35 +159,40 @@ object FoodModule {
return sortedScoreList
}
fun queryFood(floatArray: FloatArray, queryCount: Int = 15): List<String> {
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
// 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>()
// val nameScoreList = queryFoodNameScore(floatArray, queryCount)
// nameScoreList.filter { it.score < 0.05 }.forEach {
// val count = map[it.name] ?: 0
// map[it.name] = count + 1
// 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,
@@ -183,36 +200,36 @@ object FoodModule {
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()
}
// 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()
// }
/**
@@ -225,7 +242,7 @@ object FoodModule {
*/
fun copyAssetToCache(context: Context, fileName: String): String? {
// 此app的缓存目录 --> 会默认在 cache目录...,可以自己去看看哦
val cacheDir = context.getCacheDir()
val cacheDir = context.cacheDir
if (!cacheDir.exists()) {
cacheDir.mkdirs() // TODO 如果没有缓存目录,就创建
}
@@ -239,7 +256,7 @@ object FoodModule {
// 创建文件,如果创建成功,就返回true
val res = outPath.createNewFile()
if (res) {
`is` = context.getAssets().open(fileName) // 拿到main/assets目录的输入流,用于读取字节
`is` = context.assets.open(fileName) // 拿到main/assets目录的输入流,用于读取字节
fos = FileOutputStream(outPath) // 读取出来的字节最终写到outPath
val buf = ByteArray(`is`.available()) // 缓存区
var byteCount: Int
@@ -248,7 +265,7 @@ object FoodModule {
while ((`is`.read(buf).also { byteCount = it }) != -1) {
fos.write(buf, 0, byteCount)
}
return outPath.getAbsolutePath()
return outPath.absolutePath
}
} catch (e: IOException) {
e.printStackTrace()