tflite_flutter如何将图像转换为(48,48,1)单通道灰度TensorImage
Flutter TFLite 单通道灰度图输入适配方案
问题描述
在Flutter中集成TFLite模型时,模型要求输入为48×48尺寸的灰度图像,对应输入张量形状为(48, 48, 1)。实际处理时即便调用Image.grayscale方法对图像做灰度化处理,加载到TensorImage后得到的张量形状仍然为(48, 48, 3),无法匹配模型输入要求。
原有实现代码如下:
final interpreter = await Interpreter.fromAsset('model.tflite'); final inputShape = interpreter.getInputTensor(0).shape; final inputType = interpreter.getInputTensor(0).type; final outputShape = interpreter.getOutputTensor(0).shape; final outputType = interpreter.getOutputTensor(0).type; ByteData data = await rootBundle.load('assets/1'); img.Image dataImg = img.decodeImage(data.buffer.asUint8List())!; TensorImage tensorImage = TensorImage(inputType); tensorImage.loadImage(dataImg);
问题原因
tflite_flutter包提供的TensorImage.loadImage()方法默认会将所有传入的图像自动转换为3通道RGB格式存储,即便传入的是经过灰度化处理的单通道图像,也会被自动补全为3通道,因此无法直接得到单通道的张量数据。
解决步骤
- 先使用
image包完成图像的尺寸缩放、灰度化预处理 - 跳过
TensorImage自动加载逻辑,手动提取灰度图的单通道像素值,构造与模型输入形状完全匹配的张量数据 - 根据模型要求的输入数据类型(uint8/float32)处理像素值,float32类型需要做0-1归一化
- 将构造好的张量reshape为模型要求的形状后传入解释器运行推理
可直接复用的实现代码
import 'dart:typed_data'; import 'package:image/image.dart' as img; import 'package:tflite_flutter/tflite_flutter.dart'; // 省略解释器初始化、资源加载的原有代码 // 1. 缩放图像到模型要求的48*48尺寸,再转灰度 img.Image resizedImg = img.copyResize(dataImg, width: 48, height: 48); img.Image grayscaleImg = img.grayscale(resizedImg); // 2. 根据模型输入类型构造对应格式的存储数组 // 输入为uint8类型时使用Uint8List,输入为float32类型时替换为Float32List final inputSize = 48 * 48 * 1; dynamic inputData; if (inputType == TfLiteType.uint8) { inputData = Uint8List(inputSize); } else if (inputType == TfLiteType.float32) { inputData = Float32List(inputSize); } int index = 0; for (int y = 0; y < 48; y++) { for (int x = 0; x < 48; x++) { final pixel = grayscaleImg.getPixel(x, y); // 灰度图R/G/B通道数值完全一致,取任意通道值即可 final pixelValue = pixel.r.toInt(); if (inputType == TfLiteType.uint8) { inputData[index] = pixelValue; } else { // float32输入需要归一化到0-1区间 inputData[index] = pixelValue / 255.0; } index++; } } // 3. 按照模型输入形状reshape,直接用getInputTensor获取的形状做reshape即可避免维度错误 final input = inputData.reshape(inputShape); // 4. 构造输出数组,运行推理 final outputBuffer = List.filled( outputShape.reduce((curr, next) => curr * next), inputType == TfLiteType.uint8 ? 0 : 0.0 ).reshape(outputShape); interpreter.run(input, outputBuffer);
注意事项
不要依赖
TensorImage自动处理单通道图像,该类默认的图像加载逻辑固定输出3通道数据,单通道输入必须手动构造像素数组。
reshape操作前无需手动修改维度,直接用getInputTensor(0).shape返回的形状做适配即可,自动兼容带batch维度、不带batch维度的不同模型输入要求。
内容的提问来源于stack exchange,提问作者Azzam Alsharafi
相关产品推荐
相关产品推荐

