// ======= 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 create() async { final detector = YOLODetector._create(); await detector._loadModel(); return detector; } Future _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> 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> 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 = []; 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>>> reshapeFloat(List flat, List 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; } }