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格式文件。
- 模型文件放置步骤:
- 在Android项目的
src/main目录下创建assets文件夹(若不存在)。 - 将
merged_frozen_graph.pb、DecoderOutputs.txt、idmap三个文件全部放入assets目录。
- 在Android项目的
- 代码修正:
- 修改模型路径:使用
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
相关产品推荐
相关产品推荐

