在Android简易相机应用中嵌入TensorFlow Lite人像分割模型
问题描述
我开发了一个仅显示相机预览的Android简易相机应用,刚接触TensorFlow Lite,希望将人像分割.tflite模型嵌入应用,对相机帧进行处理以得到输出掩码并显示在预览界面。请问如何修改以下代码实现该功能?
import androidx.annotation.NonNull; import androidx.appcompat.app.AppCompatActivity; import androidx.core.app.ActivityCompat; import androidx.core.content.ContextCompat; import android.Manifest; import android.content.pm.PackageManager; import android.hardware.Camera; import android.os.Bundle; import android.util.Log; import android.view.SurfaceHolder; import android.view.SurfaceView; import android.widget.Toast; import java.io.IOException; public class MainActivity extends AppCompatActivity implements SurfaceHolder.Callback { private static final int CAMERA_PERMISSION_REQUEST_CODE = 1001; private static final String TAG = "MainActivity"; private SurfaceView mSurfaceView; private Camera mCamera; private SurfaceHolder mSurfaceHolder; @Override protected void onCreate(Bundle savedInstanceState) { super.onCreate(savedInstanceState); setContentView(R.layout.activity_main); mSurfaceView = findViewById(R.id.surfaceView); mSurfaceHolder = mSurfaceView.getHolder(); mSurfaceHolder.addCallback(this); } @Override protected void onResume() { super.onResume(); if (ContextCompat.checkSelfPermission(this, Manifest.permission.CAMERA) != PackageManager.PERMISSION_GRANTED) { ActivityCompat.requestPermissions(this, new String[]{Manifest.permission.CAMERA}, CAMERA_PERMISSION_REQUEST_CODE); } else { startCameraPreview(); } } @Override protected void onPause() { super.onPause(); releaseCamera(); } private void startCameraPreview() { try { int frontCameraId = -1; int numberOfCameras = Camera.getNumberOfCameras(); for (int i = 0; i < numberOfCameras; i++) { Camera.CameraInfo cameraInfo = new Camera.CameraInfo(); Camera.getCameraInfo(i, cameraInfo); if (cameraInfo.facing == Camera.CameraInfo.CAMERA_FACING_FRONT) { frontCameraId = i; break; } } if (frontCameraId == -1) { Log.e(TAG, "Front camera not found"); return; } mCamera = Camera.open(frontCameraId); mCamera.setPreviewDisplay(mSurfaceHolder); mCamera.setDisplayOrientation(90); // Set orientation for vertical view mCamera.startPreview(); } catch (IOException e) { Log.e(TAG, "Error setting camera preview: " + e.getMessage()); } } private void releaseCamera() { if (mCamera != null) { mCamera.stopPreview(); mCamera.release(); mCamera = null; } } @Override public void surfaceCreated(@NonNull SurfaceHolder holder) { // No implementation needed } @Override public void surfaceChanged(@NonNull SurfaceHolder holder, int format, int width, int height) { if (mSurfaceHolder.getSurface() == null) { return; } try { mCamera.stopPreview(); } catch (Exception e) { Log.e(TAG, "Error stopping camera preview: " + e.getMessage()); } try { mCamera.setPreviewDisplay(mSurfaceHolder); mCamera.startPreview(); } catch (Exception e) { Log.e(TAG, "Error starting camera preview: " + e.getMessage()); } } @Override public void surfaceDestroyed(@NonNull SurfaceHolder holder) { // No implementation needed } @Override public void onRequestPermissionsResult(int requestCode, @NonNull String[] permissions, @NonNull int[] grantResults) { super.onRequestPermissionsResult(requestCode, permissions, grantResults); if (requestCode == CAMERA_PERMISSION_REQUEST_CODE) { if (grantResults.length > 0 && grantResults[0] == PackageManager.PERMISSION_GRANTED) { startCameraPreview(); } else { Toast.makeText(this, "Camera permission denied", Toast.LENGTH_SHORT).show(); finish(); } } } }
实现方案
1. 添加TensorFlow Lite依赖
在app/build.gradle的dependencies块中添加:
implementation 'org.tensorflow:tensorflow-lite:2.15.0' implementation 'org.tensorflow:tensorflow-lite-support:0.4.4'
2. 准备模型文件
将人像分割的tflite模型文件放到app/src/main/assets目录下(无assets文件夹则手动创建)。
3. 修改布局文件
在activity_main.xml中给SurfaceView叠加自定义View,用于绘制分割掩码:
<?xml version="1.0" encoding="utf-8"?> <FrameLayout xmlns:android="http://schemas.android.com/apk/res/android" android:layout_width="match_parent" android:layout_height="match_parent"> <SurfaceView android:id="@+id/surfaceView" android:layout_width="match_parent" android:layout_height="match_parent" /> <com.example.yourpackage.MaskOverlayView android:id="@+id/maskOverlay" android:layout_width="match_parent" android:layout_height="match_parent" /> </FrameLayout>
注意替换com.example.yourpackage为你的实际包名。
4. 创建自定义掩码绘制View
新建MaskOverlayView.java类,负责绘制分割后的掩码:
import android.content.Context; import android.graphics.Bitmap; import android.graphics.Canvas; import android.graphics.Paint; import android.util.AttributeSet; import android.view.View; public class MaskOverlayView extends View { private Bitmap maskBitmap; private Paint paint; public MaskOverlayView(Context context) { super(context); init(); } public MaskOverlayView(Context context, AttributeSet attrs) { super(context, attrs); init(); } private void init() { paint = new Paint(); paint.setAlpha(128); // 设置掩码半透明,避免遮挡相机预览 } public void setMaskBitmap(Bitmap bitmap) { this.maskBitmap = bitmap; invalidate(); } @Override protected void onDraw(Canvas canvas) { super.onDraw(canvas); if (maskBitmap != null) { canvas.drawBitmap(maskBitmap, 0, 0, paint); } } }
5. 修改MainActivity代码
以下是整合模型加载、相机帧处理、推理及掩码绘制的完整代码:
import androidx.annotation.NonNull; import androidx.appcompat.app.AppCompatActivity; import androidx.core.app.ActivityCompat; import androidx.core.content.ContextCompat; import android.Manifest; import android.content.pm.PackageManager; import android.graphics.Bitmap; import android.graphics.BitmapFactory; import android.graphics.ImageFormat; import android.graphics.Matrix; import android.graphics.Rect; import android.graphics.YuvImage; import android.hardware.Camera; import android.os.Bundle; import android.util.Log; import android.view.SurfaceHolder; import android.view.SurfaceView; import android.widget.Toast; import org.tensorflow.lite.DataType; import org.tensorflow.lite.support.image.ImageProcessor; import org.tensorflow.lite.support.image.TensorImage; import org.tensorflow.lite.support.image.ops.ResizeOp; import org.tensorflow.lite.support.tensorbuffer.TensorBuffer; import java.io.ByteArrayOutputStream; import java.io.IOException; import java.nio.MappedByteBuffer; import java.nio.channels.FileChannel; import java.io.FileInputStream; public class MainActivity extends AppCompatActivity implements SurfaceHolder.Callback, Camera.PreviewCallback { private static final int CAMERA_PERMISSION_REQUEST_CODE = 1001; private static final String TAG = "MainActivity"; private static final int MODEL_INPUT_WIDTH = 256; private static final int MODEL_INPUT_HEIGHT = 256; private SurfaceView mSurfaceView; private Camera mCamera; private SurfaceHolder mSurfaceHolder; private MaskOverlayView mMaskOverlay; private MappedByteBuffer tfliteModel; private org.tensorflow.lite.Interpreter tflite; @Override protected void onCreate(Bundle savedInstanceState) { super.onCreate(savedInstanceState); setContentView(R.layout.activity_main); mSurfaceView = findViewById(R.id.surfaceView); mMaskOverlay = findViewById(R.id.maskOverlay); mSurfaceHolder = mSurfaceView.getHolder(); mSurfaceHolder.addCallback(this); // 加载TFLite模型 try { tfliteModel = loadModelFile(); tflite = new org.tensorflow.lite.Interpreter(tfliteModel); } catch (IOException e) { Log.e(TAG, "Failed to load model: " + e.getMessage()); } } @Override protected void onResume() { super.onResume(); if (ContextCompat.checkSelfPermission(this, Manifest.permission.CAMERA) != PackageManager.PERMISSION_GRANTED) { ActivityCompat.requestPermissions(this, new String[]{Manifest.permission.CAMERA}, CAMERA_PERMISSION_REQUEST_CODE); } else { startCameraPreview(); } } @Override protected void onPause() { super.onPause(); releaseCamera(); if (tflite != null) { tflite.close(); } } private void startCameraPreview() { try { int frontCameraId = -1; int numberOfCameras = Camera.getNumberOfCameras(); for (int i = 0; i < numberOfCameras; i++) { Camera.CameraInfo cameraInfo = new Camera.CameraInfo(); Camera.getCameraInfo(i, cameraInfo); if (cameraInfo.facing == Camera.CameraInfo.CAMERA_FACING_FRONT) { frontCameraId = i; break; } } if (frontCameraId == -1) { Log.e(TAG, "Front camera not found"); return; } mCamera = Camera.open(frontCameraId); mCamera.setPreviewDisplay(mSurfaceHolder); mCamera.setDisplayOrientation(90); // 设置预览回调获取相机帧 mCamera.setPreviewCallback(this); // 匹配预览尺寸与模型输入尺寸,减少缩放开销 Camera.Parameters params = mCamera.getParameters(); params.setPreviewSize(MODEL_INPUT_WIDTH, MODEL_INPUT_HEIGHT); mCamera.setParameters(params); mCamera.startPreview(); } catch (IOException e) { Log.e(TAG, "Error setting camera preview: " + e.getMessage()); } } private void releaseCamera() { if (mCamera != null) { mCamera.setPreviewCallback(null); mCamera.stopPreview(); mCamera.release(); mCamera = null; } } @Override public void surfaceCreated(@NonNull SurfaceHolder holder) {} @Override public void surfaceChanged(@NonNull SurfaceHolder holder, int format, int width, int height) { if (mSurfaceHolder.getSurface() == null) { return; } try { mCamera.stopPreview(); mCamera.setPreviewDisplay(mSurfaceHolder); mCamera.startPreview(); } catch (Exception e) { Log.e(TAG, "Error restarting camera preview: " + e.getMessage()); } } @Override public void surfaceDestroyed(@NonNull SurfaceHolder holder) { releaseCamera(); } @Override public void onRequestPermissionsResult(int requestCode, @NonNull String[] permissions, @NonNull int[] grantResults) { super.onRequestPermissionsResult(requestCode, permissions, grantResults); if (requestCode == CAMERA_PERMISSION_REQUEST_CODE) { if (grantResults.length > 0 && grantResults[0] == PackageManager.PERMISSION_GRANTED) { startCameraPreview(); } else { Toast.makeText(this, "相机权限被拒绝", Toast.LENGTH_SHORT).show(); finish(); } } } // 加载assets目录下的模型文件 private MappedByteBuffer loadModelFile() throws IOException { FileInputStream inputStream = new FileInputStream(getAssets().openFd("selfie_segmenter.tflite").getFileDescriptor()); FileChannel fileChannel = inputStream.getChannel(); long startOffset = getAssets().openFd("selfie_segmenter.tflite").getStartOffset(); long declaredLength = getAssets().openFd("selfie_segmenter.tflite").getDeclaredLength(); return fileChannel.map(FileChannel.MapMode.READ_ONLY, startOffset, declaredLength); } @Override public void onPreviewFrame(byte[] data, Camera camera) { if (tflite == null) return; // 将YUV格式相机帧转为Bitmap Camera.Parameters parameters = camera.getParameters(); int width = parameters.getPreviewSize().width; int height = parameters.getPreviewSize().height; YuvImage yuvImage = new YuvImage(data, ImageFormat.NV21, width, height, null); ByteArrayOutputStream out = new ByteArrayOutputStream(); yuvImage.compressToJpeg(new Rect(0, 0, width, height), 100, out); byte[] imageBytes = out.toByteArray(); Bitmap originalBitmap = BitmapFactory.decodeByteArray(imageBytes, 0, imageBytes.length); // 预处理:旋转+镜像翻转,适配前置摄像头与模型输入要求 Matrix matrix = new Matrix(); matrix.postRotate(270); matrix.postScale(-1, 1); Bitmap processedBitmap = Bitmap.createBitmap(originalBitmap, 0, 0, width, height, matrix, true); processedBitmap = Bitmap.createScaledBitmap(processedBitmap, MODEL_INPUT_WIDTH, MODEL_INPUT_HEIGHT, true); // 转换为模型输入格式 ImageProcessor imageProcessor = new ImageProcessor.Builder() .add(new ResizeOp(MODEL_INPUT_HEIGHT, MODEL_INPUT_WIDTH, ResizeOp.ResizeMethod.BILINEAR)) .build(); TensorImage inputImage = TensorImage.fromBitmap(processedBitmap); inputImage = imageProcessor.process(inputImage); // 准备输出缓冲区 TensorBuffer outputBuffer = TensorBuffer.createFixedSize( new int[]{1, MODEL_INPUT_HEIGHT, MODEL_INPUT_WIDTH, 2}, DataType.FLOAT32); // 运行模型推理 tflite.run(inputImage.getBuffer(), outputBuffer.getBuffer()); // 提取人像掩码(第二个通道为人像区域概率值) float[] outputArray = outputBuffer.getFloatArray(); Bitmap maskBitmap = Bitmap.createBitmap(MODEL_INPUT_WIDTH, MODEL_INPUT_HEIGHT, Bitmap.Config.ARGB_8888); for (int y = 0; y < MODEL_INPUT_HEIGHT; y++) { for (int x = 0; x < MODEL_INPUT_WIDTH; x++) { int index = y * MODEL_INPUT_WIDTH * 2 + x * 2 + 1; float maskValue = outputArray[index]; int gray = (int) (maskValue * 255); maskBitmap.setPixel(x, y, 0xFF000000 | (gray << 16) | (gray << 8) | gray); } } // 缩放掩码到预览界面尺寸并显示 Bitmap scaledMask = Bitmap.createScaledBitmap(maskBitmap, mSurfaceView.getWidth(), mSurfaceView.getHeight(), true); mMaskOverlay.setMaskBitmap(scaledMask); // 回收Bitmap避免内存泄漏 originalBitmap.recycle(); processedBitmap.recycle(); maskBitmap.recycle(); } }
关键说明
- 输入适配:前置摄像头帧需做镜像翻转和旋转,才能匹配模型的输入要求
- 输出解析:模型输出两个通道,第二个通道是人像区域的概率值(0-1),转为灰度图作为掩码显示
- 性能优化:预览尺寸与模型输入尺寸保持一致,减少缩放开销;实际项目建议将推理逻辑放到子线程,避免阻塞主线程
内容的提问来源于stack exchange,提问作者angel_30
相关产品推荐
相关产品推荐

