Как использовать пользовательскую модель TensorFlow Lite на устройстве Android

Если в вашем приложении используются специальные модели TensorFlow Lite, вы можете развернуть их с помощью Firebase ML. Развертывая модели с помощью Firebase, вы можете уменьшить размер приложения при первой загрузке и обновлять модели машинного обучения без выпуска новой версии приложения. А с помощью Remote Config и A/B Testing вы можете динамически предоставлять разные модели разным группам пользователей.

Модели TensorFlow Lite

Модели TensorFlow Lite – это модели машинного обучения, оптимизированные для работы на мобильных устройствах. Чтобы получить модель TensorFlow Lite:

Подготовка

  1. Если вы ещё этого не сделали, добавьте Firebase в проект для Android.
  2. В файле Gradle модуля (на уровне приложения) (обычно <project>/<app-module>/build.gradle.kts или <project>/<app-module>/build.gradle) добавьте зависимость для библиотеки загрузчика моделей Firebase ML для Android. Мы рекомендуем использовать Firebase Android BoM для управления версиями библиотеки.

    Кроме того, при настройке загрузчика моделей Firebase ML необходимо добавить в приложение TensorFlow Lite SDK.

    dependencies {
        // Import the BoM for the Firebase platform
        implementation(platform("com.google.firebase:firebase-bom:34.19.0"))
    
        // Add the dependency for the Firebase ML model downloader library
        // When using the BoM, you don't specify versions in Firebase library dependencies
        implementation("com.google.firebase:firebase-ml-modeldownloader")
    // Also add the dependency for the TensorFlow Lite library and specify its version implementation("org.tensorflow:tensorflow-lite:2.3.0")
    }

    Благодаря Firebase Android BoM в вашем приложении всегда будут использоваться совместимые версии библиотек Firebase Android.

    (Альтернативный вариант.) Добавьте зависимости библиотеки Firebase без использования BoM.

    Если вы не используете Firebase BoM, вам нужно указать версию каждой библиотеки Firebase в строке зависимости.

    Если в приложении используется несколько библиотек Firebase, мы настоятельно рекомендуем использовать BoM для управления версиями библиотек, чтобы обеспечить их совместимость.

    dependencies {
        // Add the dependency for the Firebase ML model downloader library
        // When NOT using the BoM, you must specify versions in Firebase library dependencies
        implementation("com.google.firebase:firebase-ml-modeldownloader:26.1.1")
    // Also add the dependency for the TensorFlow Lite library and specify its version implementation("org.tensorflow:tensorflow-lite:2.3.0")
    }
  3. В манифесте приложения укажите, что требуется разрешение INTERNET:
    <uses-permission android:name="android.permission.INTERNET" />

1. Как развернуть модель

Развертывать собственные модели TensorFlow можно с помощью консоли Firebase или Firebase Admin SDK для Python и Node.js. Подробнее о том, как развертывать специальные модели и управлять ими…

После того как вы добавите специальную модель в проект Firebase, вы сможете ссылаться на нее в своих приложениях, используя заданное вами имя. Вы можете в любое время развернуть новую модель TensorFlow Lite и скачать ее на устройства пользователей, вызвав функцию getModel() (см. ниже).

2. Скачайте модель на устройство и инициализируйте интерпретатор TensorFlow Lite.

Чтобы использовать модель TensorFlow Lite в приложении, сначала скачайте ее последнюю версию на устройство с помощью SDK Firebase ML. Затем создайте интерпретатор TensorFlow Lite с моделью.

Чтобы начать скачивание модели, вызовите метод getModel() загрузчика модели, указав название, которое вы присвоили модели при загрузке, хотите ли вы всегда скачивать последнюю версию модели и условия, при которых вы хотите разрешить скачивание.

Вы можете выбрать один из трех вариантов:

Тип скачивания Описание
LOCAL_MODEL Получите локальную модель с устройства. Если локальная модель недоступна, это правило работает так же, как LATEST_MODEL. Этот тип скачивания подходит, если вы не хотите проверять наличие обновлений модели. Например, вы используете Remote Config, чтобы получать названия моделей, и всегда загружаете модели под новыми названиями (рекомендуется).
LOCAL_MODEL_UPDATE_IN_BACKGROUND Получите локальную модель с устройства и начните обновлять модель в фоновом режиме. Если локальная модель недоступна, это правило работает так же, как LATEST_MODEL.
LATEST_MODEL Используйте последнюю версию модели. Если локальная модель является последней версией, возвращает локальную модель. В противном случае скачайте последнюю версию модели. В этом случае загрузка последней версии будет блокировать все другие операции до момента ее завершения (не рекомендуется). Используйте это поведение только в тех случаях, когда вам явно нужна последняя версия.

Отключите функции, связанные с моделью, например сделайте неактивными или скройте элементы интерфейса, пока не убедитесь, что модель скачана.

Kotlin

val conditions = CustomModelDownloadConditions.Builder()
        .requireWifi()  // Also possible: .requireCharging() and .requireDeviceIdle()
        .build()
FirebaseModelDownloader.getInstance()
        .getModel("your_model", DownloadType.LOCAL_MODEL_UPDATE_IN_BACKGROUND,
            conditions)
        .addOnSuccessListener { model: CustomModel? ->
            // Download complete. Depending on your app, you could enable the ML
            // feature, or switch from the local model to the remote model, etc.

            // The CustomModel object contains the local path of the model file,
            // which you can use to instantiate a TensorFlow Lite interpreter.
            val modelFile = model?.file
            if (modelFile != null) {
                interpreter = Interpreter(modelFile)
            }
        }

Java

CustomModelDownloadConditions conditions = new CustomModelDownloadConditions.Builder()
    .requireWifi()  // Also possible: .requireCharging() and .requireDeviceIdle()
    .build();
FirebaseModelDownloader.getInstance()
    .getModel("your_model", DownloadType.LOCAL_MODEL_UPDATE_IN_BACKGROUND, conditions)
    .addOnSuccessListener(new OnSuccessListener<CustomModel>() {
      @Override
      public void onSuccess(CustomModel model) {
        // Download complete. Depending on your app, you could enable the ML
        // feature, or switch from the local model to the remote model, etc.

        // The CustomModel object contains the local path of the model file,
        // which you can use to instantiate a TensorFlow Lite interpreter.
        File modelFile = model.getFile();
        if (modelFile != null) {
            interpreter = new Interpreter(modelFile);
        }
      }
    });

Многие приложения начинают скачивание в коде инициализации, но вы можете сделать это в любой момент до того, как модель понадобится.

3. Выполнять логический вывод на основе входных данных.

Как получить информацию о входных и выходных данных модели

Интерпретатор модели TensorFlow Lite принимает на вход и возвращает на выход один или несколько многомерных массивов. Эти массивы содержат значения byte, int, long или float. Прежде чем передавать данные в модель или использовать ее результаты, необходимо узнать количество и размеры массивов, которые она использует.

Если вы создали модель самостоятельно или если формат входных и выходных данных модели задокументирован, у вас уже может быть эта информация. Если вы не знаете форму и тип данных входных и выходных данных модели, вы можете использовать интерпретатор TensorFlow Lite, чтобы проверить модель. Пример:

Python

import tensorflow as tf

interpreter = tf.lite.Interpreter(model_path="your_model.tflite")
interpreter.allocate_tensors()

# Print input shape and type
inputs = interpreter.get_input_details()
print('{} input(s):'.format(len(inputs)))
for i in range(0, len(inputs)):
    print('{} {}'.format(inputs[i]['shape'], inputs[i]['dtype']))

# Print output shape and type
outputs = interpreter.get_output_details()
print('\n{} output(s):'.format(len(outputs)))
for i in range(0, len(outputs)):
    print('{} {}'.format(outputs[i]['shape'], outputs[i]['dtype']))

Пример выходных данных:

1 input(s):
[  1 224 224   3] <class 'numpy.float32'>

1 output(s):
[1 1000] <class 'numpy.float32'>

Как запустить интерпретатор

Определив формат входных и выходных данных модели, получите входные данные и выполните необходимые преобразования, чтобы привести их к нужному формату.

Например, если у вас есть модель классификации изображений с входной формой [1 224 224 3] значений с плавающей запятой, вы можете создать входной объект ByteBuffer из объекта Bitmap, как показано в следующем примере:

Kotlin

val bitmap = Bitmap.createScaledBitmap(yourInputImage, 224, 224, true)
val input = ByteBuffer.allocateDirect(224*224*3*4).order(ByteOrder.nativeOrder())
for (y in 0 until 224) {
    for (x in 0 until 224) {
        val px = bitmap.getPixel(x, y)

        // Get channel values from the pixel value.
        val r = Color.red(px)
        val g = Color.green(px)
        val b = Color.blue(px)

        // Normalize channel values to [-1.0, 1.0]. This requirement depends on the model.
        // For example, some models might require values to be normalized to the range
        // [0.0, 1.0] instead.
        val rf = (r - 127) / 255f
        val gf = (g - 127) / 255f
        val bf = (b - 127) / 255f

        input.putFloat(rf)
        input.putFloat(gf)
        input.putFloat(bf)
    }
}

Java

Bitmap bitmap = Bitmap.createScaledBitmap(yourInputImage, 224, 224, true);
ByteBuffer input = ByteBuffer.allocateDirect(224 * 224 * 3 * 4).order(ByteOrder.nativeOrder());
for (int y = 0; y < 224; y++) {
    for (int x = 0; x < 224; x++) {
        int px = bitmap.getPixel(x, y);

        // Get channel values from the pixel value.
        int r = Color.red(px);
        int g = Color.green(px);
        int b = Color.blue(px);

        // Normalize channel values to [-1.0, 1.0]. This requirement depends
        // on the model. For example, some models might require values to be
        // normalized to the range [0.0, 1.0] instead.
        float rf = (r - 127) / 255.0f;
        float gf = (g - 127) / 255.0f;
        float bf = (b - 127) / 255.0f;

        input.putFloat(rf);
        input.putFloat(gf);
        input.putFloat(bf);
    }
}

Затем выделите буфер ByteBuffer, достаточно большой для выходных данных модели, и передайте входной и выходной буферы методу run() интерпретатора TensorFlow Lite. Например, для выходной формы [1 1000] с плавающей запятой:

Kotlin

val bufferSize = 1000 * java.lang.Float.SIZE / java.lang.Byte.SIZE
val modelOutput = ByteBuffer.allocateDirect(bufferSize).order(ByteOrder.nativeOrder())
interpreter?.run(input, modelOutput)

Java

int bufferSize = 1000 * java.lang.Float.SIZE / java.lang.Byte.SIZE;
ByteBuffer modelOutput = ByteBuffer.allocateDirect(bufferSize).order(ByteOrder.nativeOrder());
interpreter.run(input, modelOutput);

Способ использования выходных данных зависит от модели.

Например, если вы выполняете классификацию, следующим шагом может быть сопоставление индексов результата с ярлыками, которые они представляют:

Kotlin

modelOutput.rewind()
val probabilities = modelOutput.asFloatBuffer()
try {
    val reader = BufferedReader(
            InputStreamReader(assets.open("custom_labels.txt")))
    for (i in probabilities.capacity()) {
        val label: String = reader.readLine()
        val probability = probabilities.get(i)
        println("$label: $probability")
    }
} catch (e: IOException) {
    // File not found?
}

Java

modelOutput.rewind();
FloatBuffer probabilities = modelOutput.asFloatBuffer();
try {
    BufferedReader reader = new BufferedReader(
            new InputStreamReader(getAssets().open("custom_labels.txt")));
    for (int i = 0; i < probabilities.capacity(); i++) {
        String label = reader.readLine();
        float probability = probabilities.get(i);
        Log.i(TAG, String.format("%s: %1.4f", label, probability));
    }
} catch (IOException e) {
    // File not found?
}

Приложение. Безопасность моделей

Независимо от того, как вы предоставляете модели TensorFlow Lite для Firebase ML, Firebase ML хранит их в стандартном сериализованном формате protobuf в локальном хранилище.

Теоретически это означает, что любой может скопировать вашу модель. Однако на практике большинство моделей настолько специфичны для приложений и запутаны оптимизациями, что риск аналогичен риску того, что конкуренты разберут и повторно используют ваш код. Тем не менее, прежде чем использовать в приложении специальную модель, вам следует знать об этом риске.

На устройствах с Android API уровня 21 (Lollipop) и более поздних версий модель скачивается в каталог, который исключен из автоматического резервного копирования.

На устройствах с Android API уровня 20 и ниже модель скачивается в каталог com.google.firebase.ml.custom.models во внутреннем хранилище, доступном только приложению. Если вы включили резервное копирование файлов с помощью BackupAgent, вы можете исключить этот каталог.