android + tflite 分类APP开发-2

APP开发

build.gradle导入库

//implementation 'org.tensorflow:tensorflow-android:+'

implementation 'org.tensorflow:tensorflow-lite:2.4.0'

implementation 'org.tensorflow:tensorflow-lite-support:0.3.1

implementation 'org.tensorflow:tensorflow-lite-metadata:0.3.1'

加载模型

try {

tfLiteClassificationUtil = new TFLiteClassificationUtil(CONST.downPath + "/zjym.tflite");

Toast.makeText(MainActTflite.this, "模型加载成功!", Toast.LENGTH_SHORT).show();

} catch (Exception e) {

Toast.makeText(MainActTflite.this, "模型加载失败!", Toast.LENGTH_SHORT).show();

e.printStackTrace();

finish();

}

模型一般在assets目录下,在编译时会集成到APP中,不利于模型的迭代,这里模型保存在内部存储目录下。

分类预测

try { // 预测图像

FileInputStream fis = new FileInputStream(image_path);

imageView.setImageBitmap(BitmapFactory.decodeStream(fis));

long start = System.currentTimeMillis();

int\[\]\[\] res2Arr = tfLiteClassificationUtil.predictImage(image_path);

long end = System.currentTimeMillis();

String show_text = "预测结果标签:" + (int) res2Arrres2Arr.length-10 +

"\n名称:" + classNames.get((int) res2Arrres2Arr.length-10) +"概率:" + (float) res2Arrres2Arr.length - 11 / 256 +

"\n名称:" + classNames.get((int) res2Arrres2Arr.length-20) +"概率:" + (float) res2Arrres2Arr.length - 21 / 256 +

"\n名称:" + classNames.get((int) res2Arrres2Arr.length-30) +"概率:" + (float) res2Arrres2Arr.length - 31 / 256 +

"\n时间:" + (end - start) + "ms";

textView.setText(show_text);

} catch (Exception e) {

e.printStackTrace();

}

res2Arrres2Arr.length - 11 / 256,两个整数相除显示为0,添加(float)显示字符串

TFLiteClassificationUtil类功能模块

public TFLiteClassificationUtil(String modelPath) throws Exception {

File file = new File(modelPath);

if (!file.exists()) {

throw new Exception("model file is not exists!");

}

try {

Interpreter.Options options = new Interpreter.Options();

options.setNumThreads(NUM_THREADS);// 使用多线程预测

NnApiDelegate delegate = new NnApiDelegate();// 使用Android自带的API或者GPU加速

// GpuDelegate delegate = new GpuDelegate();

options.addDelegate(delegate);

tflite = new Interpreter(file, options);

// 获取输入,shape为{1, height, width, 3}

int\[\] imageShape = tflite.getInputTensor(tflite.getInputIndex("input_1")).shape();

DataType imageDataType = tflite.getInputTensor(tflite.getInputIndex("input_1")).dataType();

inputImageBuffer = new TensorImage(imageDataType);

// 获取输入,shape为{1, NUM_CLASSES}

int\[\] probabilityShape = tflite.getOutputTensor(tflite.getOutputIndex("Identity")).shape();

DataType probabilityDataType = tflite.getOutputTensor(tflite.getOutputIndex("Identity")).dataType();

//outputProbabilityBuffer = TensorBuffer.createFixedSize(probabilityShape, probabilityDataType);

outputProbabilityBuffer = TensorBuffer.createFixedSize(tflite.getOutputTensor(0).shape(), DataType.UINT8);

// 添加图像预处理方式

imageProcessor = new ImageProcessor.Builder()

.add(new ResizeOp(224, 224, ResizeOp.ResizeMethod.NEAREST_NEIGHBOR))

.add(new NormalizeOp(new float\[\] {0.0f}, new float\[\] {255.0f}))

.add(new QuantizeOp(0f, 0.003921569f))

.add(new CastOp(DataType.UINT8))

.build();

TensorProcessor probabilityPostProcessor = new TensorProcessor.Builder()

.add(new DequantizeOp((float) 0, (float) 0.00390625))

.add(new NormalizeOp(new float\[\]{0.0f}, new float\[\]{1.0f}))

.build();

} catch (Exception e) {

e.printStackTrace();

throw new Exception("load model fail!");

}

}

public int\[\]\[\] predictImage(String image_path) throws Exception {

if (!new File(image_path).exists()) {

throw new Exception("image file is not exists!");

}

FileInputStream fis = new FileInputStream(image_path);

Bitmap bitmap = BitmapFactory.decodeStream(fis);

int\[\]\[\] result = predictImage(bitmap);

if (bitmap.isRecycled()) {

bitmap.recycle();

}

return result;

}

// 重载方法,直接使用Bitmap预测

public int\[\]\[\] predictImage(Bitmap bitmap) throws Exception {

return predict(bitmap);

}

private int\[\]\[\] predict(Bitmap bmp) throws Exception {

inputImageBuffer = loadImage(bmp);

try {

tflite.run(inputImageBuffer.getBuffer(), outputProbabilityBuffer.getBuffer().rewind());

} catch (Exception e) {

throw new Exception("predict image fail! log:" + e);

}

int\[\] results = outputProbabilityBuffer.getIntArray();

Log.d("results", Arrays.toString(results));

int\[\]\[\] arr = new intresults.length2;

for (int i=0;i<results.length;i++) {

arri0 = i;

arri1 = resultsi;

}

Arrays.sort(arr, Comparator.comparingInt(e -> e1));

//int l = getMaxResult(results);

return arr;//new float\[\]{l, resultsl};

}

tflite默认保存格式为UINT8,如果不加add(new CastOp(DataType.UINT8))可能显示

Cannot copy to a TensorFlowLite tensor (input_1) with 150528 bytes from a Java Buffer with 602112 bytes

默认的预训练模型是 EfficientNet-Lite0,如果为其他模型,其输入参数等也要修改。可通过下述方法查看。

Android Studio ->File ->open ->other ->tflite,打开tflite模型,build ->Make Project 会自动生成模型接口类,并移动模型到ml目录,查看类中模型参数。

相关推荐
千里马学框架3 天前
一起学 Android 14:ShellTransition 屏幕旋转过程深度剖析
android·智能手机·性能优化·framework·性能·屏幕旋转·rotation
美狐美颜SDK开放平台3 天前
开发直播APP时如何接入视频美颜SDK?开发流程与注意事项
android·人工智能·计算机视觉·音视频·直播美颜sdk
AFinalStone3 天前
Android7 SystemUI源码解析(七)Keyguard锁屏模块深度解析
android·systemui
致远ccc3 天前
Google Play 上架前如何测试 App?多国家 Android 环境测试
android·app测试·googleplay·多国家应用测试
ttyyttemo3 天前
Kotlin 协程中的 Job 结构化并发与取消
android
sun0077003 天前
tbox 4g/5g切换,导致wan ip 改变,导致车机旧网络不可用。需要重启车机才行
android
其实防守也摸鱼3 天前
内网穿透与反向代理:原理、工具与实战指南
android·大数据·运维·安全·网络安全·自动化·渗透
AFinalStone3 天前
Android7 SystemUI 源码解析(四)NavigationBar 导航栏与 SystemBars
android·systemui
JMchen3 天前
属性动画原理与高级动画实现
android·kotlin·canvas
AFinalStone4 天前
Android7 SystemUI 源码解析(二)启动流程深度解析
android·systemui