Android中TensorFlow Lite加载ESRGAN模型输出图像异常问题
问题描述
我有一段Android代码,原本用于将TensorFlow Lite模型应用在480x270分辨率的输入图像上,处理后显示结果。使用evsrnet_x4.tflite模型时逻辑正常,能正常显示输出图像。项目里还有另一个esrgan.tflite模型,这个模型应该接收50x50的输入图像,生成200x200的输出图像。但修改代码适配这个尺寸后,输出图像出现损坏:

请问问题出在哪里?还需要修改哪些内容才能让esrgan模型正常工作?
package com.example.mobedsr; import android.content.res.AssetFileDescriptor; import android.content.res.AssetManager; import android.graphics.Bitmap; import org.tensorflow.lite.DataType; import org.tensorflow.lite.Interpreter; import org.tensorflow.lite.gpu.CompatibilityList; import org.tensorflow.lite.gpu.GpuDelegate; import org.tensorflow.lite.support.common.ops.NormalizeOp; 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.FileInputStream; import java.io.IOException; import java.nio.ByteBuffer; import java.nio.channels.FileChannel; /** @brief Super Resolution Model class * @date 23/01/27 */ public class SRModel { private boolean useGpu; public Interpreter interpreter; private Interpreter.Options options; private GpuDelegate gpuDelegate; private AssetManager assetManager; private final String MODEL_NAME = "evsrnet_x4.tflite"; //I want to change to esrgan.tflite SRModel(AssetManager assetManager, boolean useGpu) throws IOException { interpreter = null; gpuDelegate = null; this.assetManager = assetManager; this.useGpu = useGpu; // Initialize the TF Lite interpreter init(); } private void init() throws IOException { options = new Interpreter.Options(); // Set gpu delegate if (useGpu) { CompatibilityList compatList = new CompatibilityList(); GpuDelegate.Options delegateOptions = compatList.getBestOptionsForThisDevice(); gpuDelegate = new GpuDelegate(delegateOptions); options.addDelegate(gpuDelegate); } // Set TF Lite interpreter interpreter = new Interpreter(loadModelFile(), options); } /** @brief Load .tflite model file to ByteBuffer * @date 23/01/25 */ private ByteBuffer loadModelFile() throws IOException { AssetFileDescriptor assetFileDescriptor = assetManager.openFd(MODEL_NAME); FileInputStream fileInputStream = new FileInputStream(assetFileDescriptor.getFileDescriptor()); FileChannel fileChannel = fileInputStream.getChannel(); long startOffset = assetFileDescriptor.getStartOffset(); long declaredLength = assetFileDescriptor.getDeclaredLength(); return fileChannel.map(FileChannel.MapMode.READ_ONLY, startOffset, declaredLength); } public void run(Object a, Object b) { interpreter.run(a, b); } /** @brief Prepare the input tensor from low resolution image * @date 23/01/25 */ public TensorImage prepareInputTensor(Bitmap bitmap_lr) { TensorImage inputImage = TensorImage.fromBitmap(bitmap_lr); int height = bitmap_lr.getHeight(); int width = bitmap_lr.getWidth(); ImageProcessor imageProcessor = new ImageProcessor.Builder() .add(new ResizeOp(height, width, ResizeOp.ResizeMethod.NEAREST_NEIGHBOR)) .add(new NormalizeOp(0.0f, 255.0f)) .build(); inputImage = imageProcessor.process(inputImage); return inputImage; } /** @brief Prepare the output tensor for super resolution * @date 23/01/25 */ public TensorImage prepareOutputTensor() { TensorImage srImage = new TensorImage(DataType.FLOAT32); // int[] srShape = new int[]{1080, 1920, 3}; int[] srShape = new int[]{1920, 1080, 3}; srImage.load(TensorBuffer.createFixedSize(srShape, DataType.FLOAT32)); return srImage; } /** @brief Convert tensor to bitmap image * @date 23/01/25 * @param outputTensor super resolutioned image */ public Bitmap tensorToImage(TensorImage outputTensor) { ByteBuffer srOut = outputTensor.getBuffer(); srOut.rewind(); int height = outputTensor.getHeight(); int width = outputTensor.getWidth(); Bitmap bmpImage = Bitmap.createBitmap(width, height, Bitmap.Config.ARGB_8888); int[] pixels = new int[width * height]; for (int i = 0; i < width * height; i++) { int a = 0xFF; float r = srOut.getFloat() * 255.0f; float g = srOut.getFloat() * 255.0f; float b = srOut.getFloat() * 255.0f; pixels[i] = a << 24 | ((int) r << 16) | ((int) g << 8) | ((int) b); } bmpImage.setPixels(pixels, 0, width, 0, 0, width, height); return bmpImage; } }
更新1
按照建议修改后有进展,但输出图像呈现像素化且偏紫的状态:

解决方案
初始图像损坏问题的修复
修正输入尺寸匹配模型要求
原代码的prepareInputTensor方法中,ResizeOp是按输入Bitmap原尺寸缩放,而esrgan.tflite要求输入为50x50,需修改ResizeOp参数:ImageProcessor imageProcessor = new ImageProcessor.Builder() .add(new ResizeOp(50, 50, ResizeOp.ResizeMethod.NEAREST_NEIGHBOR)) .add(new NormalizeOp(0.0f, 255.0f)) .build();修正输出张量尺寸
prepareOutputTensor中的输出尺寸仍为1920x1080,需改为模型对应的200x200:int[] srShape = new int[]{200, 200, 3};切换目标模型
将MODEL_NAME改为"esrgan.tflite"。
更新后像素化偏紫问题的修复
调整颜色通道顺序
部分ESRGAN模型输出为BGR通道顺序,而Android Bitmap采用RGB顺序,需修改通道读取顺序:for (int i = 0; i < width * height; i++) { int a = 0xFF; float b = srOut.getFloat() * 255.0f; float g = srOut.getFloat() * 255.0f; float r = srOut.getFloat() * 255.0f; pixels[i] = a << 24 | ((int) r << 16) | ((int) g << 8) | ((int) b); }截断超出范围的像素值
模型输出的浮点值可能超出0-1范围,需做截断避免颜色失真:float r = Math.min(Math.max(srOut.getFloat() * 255.0f, 0), 255); float g = Math.min(Math.max(srOut.getFloat() * 255.0f, 0), 255); float b = Math.min(Math.max(srOut.getFloat() * 255.0f, 0), 255);匹配模型的归一化参数
若你的ESRGAN模型训练时采用的是-1~1的归一化范围,需同步修改输入输出的处理逻辑:// 输入处理修改 .add(new NormalizeOp(-1.0f, 1.0f)) // 输出转换修改 float r = Math.min(Math.max((srOut.getFloat() + 1) * 127.5f, 0), 255); float g = Math.min(Math.max((srOut.getFloat() + 1) * 127.5f, 0), 255); float b = Math.min(Math.max((srOut.getFloat() + 1) * 127.5f, 0), 255);
内容的提问来源于stack exchange,提问作者angel_30
相关产品推荐
相关产品推荐

