Files
SDKmuta/lib/screens/yolo_detector.dart
T
2025-04-05 11:47:40 +02:00

163 lines
5.1 KiB
Dart

// ======= yolo_detector.dart (updated with imageToInputTensor helper) =======
import 'dart:io';
import 'dart:typed_data';
import 'package:tflite_flutter/tflite_flutter.dart';
import 'package:image/image.dart' as img;
class DetectionResult {
final double x1, y1, x2, y2;
final double confidence;
final int classId;
DetectionResult(this.x1, this.y1, this.x2, this.y2, this.confidence, this.classId);
}
class CroppedImageResult {
final img.Image image;
final int offsetX;
final int offsetY;
final int cropWidth;
final int cropHeight;
CroppedImageResult(this.image, this.offsetX, this.offsetY, this.cropWidth, this.cropHeight);
}
class YOLODetector {
late Interpreter _interpreter;
late int inputWidth, inputHeight;
final double confidenceThreshold = 0.4;
YOLODetector._create();
static Future<YOLODetector> create() async {
final detector = YOLODetector._create();
await detector._loadModel();
return detector;
}
Future<void> _loadModel() async {
final options = InterpreterOptions()..threads = 2;
_interpreter = await Interpreter.fromAsset('assets/best_float32_nms.tflite', options: options);
final inputTensor = _interpreter.getInputTensor(0);
final inputShape = inputTensor.shape;
inputHeight = inputShape[1];
inputWidth = inputShape[2];
print("📐 Model input shape: $inputShape");
print("📦 Model input type: ${inputTensor.type}");
}
Future<List<DetectionResult>> detect(File imageFile) async {
final bytes = await imageFile.readAsBytes();
final decodedImage = img.decodeImage(bytes);
if (decodedImage == null) throw Exception("Could not decode image.");
return detectFromImage(decodedImage);
}
Future<List<DetectionResult>> detectFromImage(img.Image originalImage) async {
final croppedResult = cropToModelAspectRatio(originalImage, inputWidth, inputHeight);
final fixedImage = img.bakeOrientation(croppedResult.image);
final inputTensor = imageToInputTensor(fixedImage, inputWidth, inputHeight);
final inputBuffer = reshapeFloat(inputTensor.toList(), [1, inputHeight, inputWidth, 3]);
final outputShape = _interpreter.getOutputTensor(0).shape;
final numDetections = outputShape[1];
final outputBuffer = List.generate(1, (_) =>
List.generate(numDetections, (_) => List.filled(6, 0.0)));
print("🚀 Running inference...");
_interpreter.run(inputBuffer, outputBuffer);
final rawDetections = outputBuffer[0];
final results = <DetectionResult>[];
for (var detection in rawDetections) {
final x1 = detection[0];
final y1 = detection[1];
final x2 = detection[2];
final y2 = detection[3];
final score = detection[4];
final classId = detection[5].toInt();
if (score < confidenceThreshold) continue;
results.add(DetectionResult(
x1 * croppedResult.cropWidth + croppedResult.offsetX,
y1 * croppedResult.cropHeight + croppedResult.offsetY,
x2 * croppedResult.cropWidth + croppedResult.offsetX,
y2 * croppedResult.cropHeight + croppedResult.offsetY,
score,
classId,
));
}
print("🎯 Detections found: ${results.length}");
return results;
}
void close() {
_interpreter.close();
}
List<List<List<List<double>>>> reshapeFloat(List<double> flat, List<int> dims) {
final d0 = dims[0], d1 = dims[1], d2 = dims[2], d3 = dims[3];
final reshaped = List.generate(d0, (_) =>
List.generate(d1, (_) =>
List.generate(d2, (_) => List.filled(d3, 0.0))));
int index = 0;
for (int i = 0; i < d0; i++) {
for (int j = 0; j < d1; j++) {
for (int k = 0; k < d2; k++) {
for (int l = 0; l < d3; l++) {
reshaped[i][j][k][l] = flat[index++];
}
}
}
}
return reshaped;
}
CroppedImageResult cropToModelAspectRatio(img.Image image, int targetWidth, int targetHeight) {
final targetAspectRatio = targetWidth / targetHeight;
final originalAspectRatio = image.width / image.height;
int cropWidth = image.width;
int cropHeight = image.height;
if (originalAspectRatio > targetAspectRatio) {
cropWidth = (image.height * targetAspectRatio).toInt();
} else {
cropHeight = (image.width / targetAspectRatio).toInt();
}
final offsetX = ((image.width - cropWidth) / 2).toInt();
final offsetY = ((image.height - cropHeight) / 2).toInt();
final croppedImage = img.copyCrop(image, x: offsetX, y: offsetY, width: cropWidth, height: cropHeight);
return CroppedImageResult(croppedImage, offsetX, offsetY, cropWidth, cropHeight);
}
Float32List imageToInputTensor(img.Image inputImage, int width, int height) {
final resized = img.copyResize(inputImage, width: width, height: height);
final Float32List input = Float32List(width * height * 3);
int index = 0;
for (int y = 0; y < height; y++) {
for (int x = 0; x < width; x++) {
final pixel = resized.getPixel(x, y);
final r = pixel.r / 255.0;
final g = pixel.g / 255.0;
final b = pixel.b / 255.0;
input[index++] = r;
input[index++] = g;
input[index++] = b;
}
}
return input;
}
}