You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.09 06:02:03