Flutter中超分辨率模型推理输出图像异常排查求助
超分辨率模型Flutter推理输出异常排查
在Flutter中使用超分辨率模型进行图像推理时,输出图像不符合预期,疑似问题出在图像归一化或输出转图像环节。
原图像/输出图像情况
- 原图像:正常待超分图像
- 输出图像:色彩/维度异常,不符合超分预期
模型形状
Model input shape: ['batch_size', 3, 'width', 'height']
Model output shape: ['batch_size', 3, 'width', 'height']
运行日志
flutter: Is normalized: true flutter: Image normalized successfully. flutter: Input tensor created successfully. flutter: Width: 1428, Height: 804, Channel: 3
问题代码
Future<void> inference() async { if (selectedImage == null) { debugPrint('No image selected'); return; } if (selectedImage != null) { Float32List? floatData; try { img.Image normalizedImage = img.normalize(selectedImage!, min: 0, max: 255); Uint8List imageData = normalizedImage.getBytes(order: img.ChannelOrder.rgb); floatData = Float32List.fromList( imageData.map((byte) => byte / 255.0).toList()); debugPrint("Is normalized: ${isNormalized(floatData)}"); } catch (e) { debugPrint("Error during normalization: $e"); } final shape = [1, 3, selectedImage!.width, selectedImage!.height]; debugPrint('Image normalized successfully.'); final inputOrt = OrtValueTensor.createTensorWithDataList(floatData!, shape); final inputs = {'input': inputOrt}; debugPrint('Input tensor created successfully.'); final runOptions = OrtRunOptions(); final outputs = await ortSession.runAsync(runOptions, inputs); inputOrt.release(); runOptions.release(); outputs?.forEach((element) { final outputValue = element?.value; if (outputValue is List<List<List<List<double>>>>) { img.Image generatedImage = generateImageFromOutput(outputValue); showDialog( context: context, builder: (BuildContext context) { return Dialog( child: SizedBox( width: generatedImage.width.toDouble(), height: generatedImage.height.toDouble(), child: Image.memory( Uint8List.fromList(img.encodePng(generatedImage)), fit: BoxFit.contain, ), ), ); }, ); } else { debugPrint("Output is of unknown type"); } element?.release(); }); } } img.Image generateImageFromOutput( List<List<List<List<double>>>> outputValue) { int width = outputValue[0][0].length; int height = outputValue[0][0][0].length; int channel = outputValue[0].length; print("Width: $width, Height: $height, Channel: $channel"); // Create the image img.Image generatedImage = img.Image(width: width, height: height); // Set pixel values for (int y = 0; y < height; y++) { for (int x = 0; x < width; x++) { // Extract RGB values from the output tensor data int r = (outputValue[0][0][x][y] * 255).toInt().clamp(0, 255); int g = (outputValue[0][1][x][y] * 255).toInt().clamp(0, 255); int b = (outputValue[0][2][x][y] * 255).toInt().clamp(0, 255); // Set pixel value in the generated image generatedImage.setPixelRgb(x, y, r, g, b); } } return generatedImage; }
问题排查与修复
1. 输出图像维度顺序颠倒(核心问题)
多数CV模型输出维度为[batch_size, 3, height, width](通道在前,高、宽在后),但代码中错误地将宽高读取、像素坐标映射颠倒:
- 错误的宽高读取:把
height当成width,width当成height - 错误的像素取值:将坐标
(x,y)对应到了张量的[x][y],实际应为[y][x]
修复后的generateImageFromOutput函数:
img.Image generateImageFromOutput( List<List<List<List<double>>>> outputValue) { // 模型输出维度:[batch, channel, height, width] int height = outputValue[0][0].length; int width = outputValue[0][0][0].length; int channel = outputValue[0].length; print("Width: $width, Height: $height, Channel: $channel"); img.Image generatedImage = img.Image(width: width, height: height); for (int y = 0; y < height; y++) { for (int x = 0; x < width; x++) { // 按[batch, channel, height, width]顺序读取像素值 int r = (outputValue[0][0][y][x] * 255).toInt().clamp(0, 255); int g = (outputValue[0][1][y][x] * 255).toInt().clamp(0, 255); int b = (outputValue[0][2][y][x] * 255).toInt().clamp(0, 255); generatedImage.setPixelRgb(x, y, r, g, b); } } return generatedImage; }
2. 输入张量维度顺序验证
模型标注输入形状为[batch_size, 3, width, height],但需确认训练时实际采用的是[batch_size, 3, height, width](行业通用格式)。若原图像宽1428、高804,输入shape错误写为[1,3,1428,804],会导致输入维度颠倒,输出异常。
若模型要求[batch,3,height,width],修改输入shape:
final shape = [1, 3, selectedImage!.height, selectedImage!.width];
3. 归一化逻辑匹配模型训练要求
当前代码将像素值转为0-1范围,需和模型训练时的输入归一方式一致:
- 若模型训练用0-1:当前逻辑正确,无需修改
- 若模型训练用-1到1:修改输入归一和输出转换
// 输入归一 floatData = Float32List.fromList( imageData.map((byte) => (byte / 127.5) - 1.0).toList()); // 输出转换 int r = ((outputValue[0][0][y][x] + 1.0) * 127.5).toInt().clamp(0, 255);
4. 通道顺序验证
当前代码用RGB通道顺序,若模型训练时用BGR,需调整通道顺序:
Uint8List imageData = normalizedImage.getBytes(order: img.ChannelOrder.rgb); // RGB转BGR Uint8List bgrData = Uint8List(imageData.length); for (int i = 0; i < imageData.length; i += 3) { bgrData[i] = imageData[i + 2]; bgrData[i + 1] = imageData[i + 1]; bgrData[i + 2] = imageData[i]; } floatData = Float32List.fromList( bgrData.map((byte) => byte / 255.0).toList());
内容的提问来源于stack exchange,提问作者Rifat Khan
相关产品推荐
相关产品推荐

