import 'dart:typed_data'; import 'dart:ui' as ui; import 'package:flutter/services.dart' show rootBundle; import 'package:flutter_onnxruntime/flutter_onnxruntime.dart'; import 'ctc_decoder.dart'; import 'ocr_backend.dart'; /// ONNX-based text recognition backend. /// /// NOTE: assumes PP-OCRv4 mobile rec (input 3x48xW, CTC blank=0). Verify /// on-device; the dictionary/blank convention may need adjustment per the exact /// exported model. /// /// The model and dictionary are bundled as assets and are optional: if either /// is missing this backend reports unavailable and recognition is a clean /// no-op (returns null). It never throws out of [recognize]. class OnnxRecognitionBackend implements OcrBackend { static const _modelAsset = 'assets/models/ocr/rec.onnx'; static const _dictAsset = 'assets/models/ocr/ppocr_keys_v1.txt'; // Rec model input geometry. static const _targetHeight = 48; static const _minWidth = 16; static const _maxWidth = 320; OrtSession? _session; CtcDecoder? _decoder; bool _initAttempted = false; bool _available = false; @override String get name => 'onnx'; @override Future isAvailable() async { await _ensureInit(); return _available; } @override Future recognize(Uint8List pngBytes) async { await _ensureInit(); final session = _session; final decoder = _decoder; if (!_available || session == null || decoder == null) return null; OrtValue? input; Map? outputs; try { final pre = await _preprocess(pngBytes); if (pre == null) return null; final inputName = session.inputNames[0]; input = await OrtValue.fromList(pre.data, [ 1, 3, _targetHeight, pre.width, ]); outputs = await session.run({inputName: input}); final out = outputs[session.outputNames[0]]; if (out == null) return null; // Expected output shape: [1, T, C]. final shape = out.shape; if (shape.length != 3) return null; final timeSteps = shape[1]; final numClasses = shape[2]; // asFlattenedList() returns the data flat (row-major); asList() would // return a list nested per the output shape. final flat = (await out.asFlattenedList()) .map((v) => (v as num).toDouble()) .toList(); final text = decoder.decodeFlat(flat, timeSteps, numClasses).trim(); if (text.isEmpty) return null; return text; } catch (_) { return null; } finally { if (input != null) { await input.dispose(); } if (outputs != null) { for (final t in outputs.values) { await t.dispose(); } } } } /// Lazily load the dictionary and create the inference session. On any /// failure the backend is marked unavailable. Future _ensureInit() async { if (_initAttempted) return; _initAttempted = true; try { final dictRaw = await rootBundle.loadString(_dictAsset); final charset = dictRaw .split('\n') .map((line) => line.replaceAll('\r', '')) .toList(); // Drop a single trailing empty entry from a final newline, then append a // space character as PP-OCR does. if (charset.isNotEmpty && charset.last.isEmpty) { charset.removeLast(); } charset.add(' '); final ort = OnnxRuntime(); final session = await ort.createSessionFromAsset( _modelAsset, options: OrtSessionOptions( intraOpNumThreads: 2, providers: [OrtProvider.CPU], ), ); _session = session; _decoder = CtcDecoder(charset, blankIndex: 0); _available = true; } catch (_) { _session = null; _decoder = null; _available = false; } } /// Decode and preprocess the PNG into the CHW Float32 tensor the rec model /// expects. Returns null on any decode failure. Future<_PreprocessResult?> _preprocess(Uint8List pngBytes) async { final codec = await ui.instantiateImageCodec(pngBytes); final frame = await codec.getNextFrame(); final src = frame.image; try { final origW = src.width; final origH = src.height; if (origW <= 0 || origH <= 0) return null; // Width that preserves aspect ratio at the target height, clamped. final scaledW = (_targetHeight * origW / origH).round(); final targetW = scaledW.clamp(_minWidth, _maxWidth); // Render the resized image onto a white canvas. If the scaled width is // narrower than the target, the right side stays white (padding). final recorder = ui.PictureRecorder(); final canvas = ui.Canvas(recorder); final paintWidth = scaledW < targetW ? scaledW : targetW; canvas.drawRect( ui.Rect.fromLTWH(0, 0, targetW.toDouble(), _targetHeight.toDouble()), ui.Paint()..color = const ui.Color(0xFFFFFFFF), ); canvas.drawImageRect( src, ui.Rect.fromLTWH(0, 0, origW.toDouble(), origH.toDouble()), ui.Rect.fromLTWH(0, 0, paintWidth.toDouble(), _targetHeight.toDouble()), ui.Paint(), ); final picture = recorder.endRecording(); final resized = await picture.toImage(targetW, _targetHeight); picture.dispose(); try { final byteData = await resized.toByteData( format: ui.ImageByteFormat.rawRgba, ); if (byteData == null) return null; final rgba = byteData.buffer.asUint8List(); // Layout CHW (3 x H x W), normalize (v/255 - 0.5) / 0.5, RGB only. final hw = _targetHeight * targetW; final data = Float32List(3 * hw); for (var y = 0; y < _targetHeight; y++) { for (var x = 0; x < targetW; x++) { final pixel = (y * targetW + x) * 4; final r = rgba[pixel] / 255.0; final g = rgba[pixel + 1] / 255.0; final b = rgba[pixel + 2] / 255.0; final idx = y * targetW + x; data[idx] = (r - 0.5) / 0.5; data[hw + idx] = (g - 0.5) / 0.5; data[2 * hw + idx] = (b - 0.5) / 0.5; } } return _PreprocessResult(data, targetW); } finally { resized.dispose(); } } finally { src.dispose(); } } } class _PreprocessResult { _PreprocessResult(this.data, this.width); final Float32List data; final int width; }