Files
Inbound/app/src/main/java/com/sw/inbound/objbox/FoodModule.kt
T

226 lines
8.3 KiB
Kotlin

package com.sw.inbound.objbox
import android.content.Context
import android.graphics.Bitmap
import android.net.Uri
import android.util.SparseLongArray
import com.google.gson.Gson
import com.google.gson.reflect.TypeToken
import com.sw.inbound.MyApp
import com.sw.inbound.utils.AssetsTool
import com.sw.inbound.utils.ImageUtil
import com.sw.inbound.utils.ext.toJsonString
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
import timber.log.Timber
import java.io.File
import java.io.FileOutputStream
import java.io.IOException
import java.io.InputStream
object FoodModule {
private lateinit var module_mobile: Module
private lateinit var box: Box<Food>
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
fun init(context: Context) {
Thread {
module_mobile = Module.load(copyAssetToCache(context, "best_embedding_model_mobile.pt"))
box = ObjectBox.boxStore.boxFor(Food::class)
//if (box.all.isNotEmpty()) {
// box.removeAll()
//}
//if (box.all.isEmpty()) {
// initDefFoodData(context)
//}
}.start()
}
fun uri2FloatArray(uri: Uri): FloatArray? {
return MyApp.instance?.let { context ->
ImageUtil.uriToBitmap(context, uri)?.let {
bitmap2FloatArray(it)
}
}
}
fun bitmap2FloatArray(bitmap: Bitmap): FloatArray {
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
)
val outputTensor = module_mobile.forward(IValue.from(inputTensor)).toTensor()
val floatArray = outputTensor.dataAsFloatArray
// try {
// if (bitmap.isRecycled.not()) {
// bitmap.recycle()
// }
// } catch (e: Exception) {
// e.printStackTrace()
// }
return floatArray
}
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)
}
fun queryFoodNameScore(floatArray: FloatArray, queryCount: Int = 15): List<IdNameScore> {
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 idScoreList = query.findIdsWithScores()
val nameScoreList = mutableListOf<IdNameScore>()
idScoreList.forEach {
nameScoreList.add(IdNameScore(id = it.id, name = box.get(it.id).name?:"", score = it.score))
}
Timber.tag("FoodModule").d("queryFood数据:${nameScoreList.toJsonString()}")
return nameScoreList
}
fun getFoodScoreList(bitmap: Bitmap, queryCount: Int = 15): List<IdNameScore> {
val floatArray = bitmap2FloatArray(bitmap)
val nameScoreList = queryFoodNameScore(floatArray, queryCount)
val maxScoreList = nameScoreList.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)
})
Timber.tag("FoodModule").d("getFoodScoreList数据:${sortedScoreList.toJsonString()}")
return sortedScoreList
}
fun queryFood(floatArray: FloatArray, queryCount: Int = 15): List<String> {
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,
val 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(name = 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.getCacheDir()
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.getAssets().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.getAbsolutePath()
}
} catch (e: IOException) {
e.printStackTrace()
} finally {
try {
// TODO 一定要记得关闭资源,为了不去性能的磨损
fos!!.flush()
`is`!!.close()
fos.close()
} catch (e: IOException) {
e.printStackTrace()
}
}
return null
}
}