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

Android集成TensorFlow实现图像描述时模型加载失败求助

Android TensorFlow图像描述功能模型加载失败问题解决

报错信息

java.lang.RuntimeException: Failed to load model from 'file:///android_asset/merged_frozen_graph.pb'

问题说明:需要将merged_frozen_graph.pb放入应用的assets目录,但无法找到该文件,误以为它包含在implementation 'org.tensorflow:tensorflow-android:1.11.0'依赖库中。

相关代码

package com.example.vijay.image_captionanddetection_tensorflow;

import android.content.Context;
import android.graphics.Bitmap;

import org.tensorflow.contrib.android.TensorFlowInferenceInterface;

import java.io.BufferedReader;
import java.io.IOException;
import java.io.InputStream;
import java.io.InputStreamReader;

public class CaptionGenerator {

    private static final String MODEL_FILE = "file:///android_asset/merged_frozen_graph.pb";
    private static final String INPUT1 = "encoder/import/InputImage:0";
    private static final String OUTPUT_NODES = "DecoderOutputs.txt";
    private static final int NUM_TIMESTEPS = 22;
    private static final int IMAGE_SIZE = 299;
    private static final int IMAGE_CHANNELS = 3;
    private static final int[] DIM_IMAGE=new int[]{1, IMAGE_SIZE, IMAGE_SIZE, IMAGE_CHANNELS};
    private TensorFlowInferenceInterface inferenceInterface;
    private String[] OutputNodes = null;
    private String[] WORD_MAP = null;

    Context context;

    CaptionGenerator(Context context){
        this.context=context;
        inferenceInterface = InitSession();

    }

    String[] LoadFile(String fileName){
        InputStream is = null;
        try {
            is = context.getAssets().open(fileName);
        } catch (IOException e) {
            e.printStackTrace();
        }
        BufferedReader r = new BufferedReader(new InputStreamReader(is));
        StringBuilder total = new StringBuilder();
        String line;
        try {
            while ((line = r.readLine()) != null) {
                total.append(line).append('\n');
            }
        } catch (IOException e) {
            e.printStackTrace();
        }
        return total.toString().split("\n");
    }

    TensorFlowInferenceInterface InitSession(){
        inferenceInterface = new TensorFlowInferenceInterface(context.getAssets(),MODEL_FILE);
       // inferenceInterface.initializeTensorFlow(context.getAssets(),MODEL_FILE);
        OutputNodes = LoadFile(OUTPUT_NODES);
        WORD_MAP = LoadFile("idmap");
        return inferenceInterface;
    }

    String runModel(Bitmap imBitmap){
        return  GenerateCaptions(Preprocess(imBitmap));
    }

    float[] Preprocess(Bitmap imBitmap){
        imBitmap = Bitmap.createScaledBitmap(imBitmap, IMAGE_SIZE, IMAGE_SIZE, true);
        int[] intValues = new int[IMAGE_SIZE * IMAGE_SIZE];
        float[] floatValues = new float[IMAGE_SIZE * IMAGE_SIZE * 3];

        imBitmap.getPixels(intValues, 0, IMAGE_SIZE, 0, 0, IMAGE_SIZE, IMAGE_SIZE);

        for (int i = 0; i < intValues.length; ++i) {
            final int val = intValues[i];
            floatValues[i * 3] = ((float)((val >> 16) & 0xFF))/255;//R
            floatValues[i * 3 + 1] = ((float)((val >> 8) & 0xFF))/255;//G
            floatValues[i * 3 + 2] = ((float)((val & 0xFF)))/255;//B
        }
        return floatValues;
    }

    String GenerateCaptions(float[] imRGBMatrix){
//        inferenceInterface.fillNodeFloat(INPUT1, DIM_IMAGE, imRGBMatrix);
//        inferenceInterface.runInference(OutputNodes);
        inferenceInterface.feed(INPUT1, imRGBMatrix, DIM_IMAGE[0], DIM_IMAGE[1], DIM_IMAGE[2], DIM_IMAGE[3]);
        inferenceInterface.run(OutputNodes);
        String result = "";
        int temp[][]= new int[NUM_TIMESTEPS][1];
        for(int i = 0; i<NUM_TIMESTEPS; ++i) {
           // inferenceInterface.readNodeInt(OutputNodes[i], temp[i]);
            inferenceInterface.fetch(OutputNodes[i], temp[i]);
            if(temp[i][0] == 2/*</S>*/){
                return result;
            }
            result += WORD_MAP[temp[i][0]]+" ";
        }
        return null;
    }

}

解决方案

  • 明确模型来源:tensorflow-android库仅提供TensorFlow在Android上的推理接口,不包含图像描述任务的预训练模型。merged_frozen_graph.pb是图像描述专用的冻结模型,需要自行获取或训练。
  • 获取模型的两种方式:
    • 回溯教程资源:如果是跟着特定教程实现,回到教程寻找模型下载链接,多数图像描述教程会提供基于Show and Tell、Neural Image Caption等架构的预训练冻结模型。
    • 自行训练模型:基于TensorFlow 1.x版本构建图像描述模型(比如用InceptionV3做图像编码器,LSTM做文本解码器),训练完成后使用TensorFlow官方的freeze_graph.py工具将训练好的模型冻结为.pb格式文件。
  • 模型文件放置步骤:
    1. 在Android项目的src/main目录下创建assets文件夹(若不存在)。
    2. 将merged_frozen_graph.pb、DecoderOutputs.txt、idmap三个文件全部放入assets目录。
  • 代码修正:
    • 修改模型路径:使用TensorFlowInferenceInterface的构造方法时,传入assets管理器后,模型路径不需要file:///android_asset/前缀,直接写文件名即可:
      private static final String MODEL_FILE = "merged_frozen_graph.pb";
      
    • 确保DecoderOutputs.txt和idmap文件路径正确,否则LoadFile方法会抛出IO异常。
  • 版本兼容检查:使用的tensorflow-android:1.11.0是TensorFlow 1.x版本,确保获取或训练的模型也是基于TensorFlow 1.x版本构建的,避免版本不兼容导致加载失败。

内容的提问来源于stack exchange,提问作者shahzaib ali

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 12:48:21