采集向量优化调试
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user