1
This commit is contained in:
@@ -8,6 +8,7 @@ import 'camera/camera_screen.dart';
|
||||
import 'container.dart';
|
||||
import 'home/home_screen.dart';
|
||||
import 'legal/terms_screen.dart';
|
||||
import 'models/model_manager.dart';
|
||||
import 'payment/paywall_screen.dart';
|
||||
import 'update/update_checker.dart';
|
||||
import 'update/update_screen.dart';
|
||||
@@ -93,6 +94,9 @@ class _StartupGateState extends State<StartupGate> {
|
||||
}
|
||||
final token = await session.readToken();
|
||||
if (!mounted) return;
|
||||
// 模型热更新:后台拉取模型目录并下载/更新各数据集模型(不阻塞启动,
|
||||
// 相机页打开前若未就绪会兜底等待;下载失败回退内置资产模型)
|
||||
ModelManager.instance.refresh();
|
||||
Navigator.of(context)
|
||||
.pushReplacementNamed(token == null ? '/login' : '/home');
|
||||
}
|
||||
|
||||
@@ -6,6 +6,7 @@ import 'package:permission_handler/permission_handler.dart';
|
||||
import 'package:wakelock_plus/wakelock_plus.dart';
|
||||
|
||||
import '../detection/detector_worker.dart';
|
||||
import '../models/model_manager.dart';
|
||||
import '../reminder/reminder.dart';
|
||||
import 'app_camera_controller.dart';
|
||||
import 'camera_view_model.dart';
|
||||
@@ -119,8 +120,18 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
setState(() => _permissionGranted = granted);
|
||||
if (!granted) return;
|
||||
|
||||
// 模型热更新:优先使用已下载的数据集模型(启动时后台拉取;此处兜底等待,
|
||||
// 下载慢/失败不阻塞相机启动——无下载模型时 worker 回退内置资产)
|
||||
if (!ModelManager.instance.ready) {
|
||||
try {
|
||||
await ModelManager.instance
|
||||
.refresh()
|
||||
.timeout(const Duration(seconds: 15));
|
||||
} catch (_) {}
|
||||
}
|
||||
// 模型加载/推理在后台 isolate,不阻塞 UI;worker 为 null 时仅预览并提示
|
||||
final worker = await DetectorWorker.create();
|
||||
final worker = await DetectorWorker.create(
|
||||
models: ModelManager.instance.models);
|
||||
final viewModel = CameraViewModel(reminder: Reminder());
|
||||
viewModel.setModelReady(worker != null);
|
||||
final analyzer = FrameAnalyzer(worker: worker, viewModel: viewModel);
|
||||
@@ -276,6 +287,27 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
return Stack(
|
||||
fit: StackFit.expand,
|
||||
children: [
|
||||
// 模型热更新下载失败提示(已加载模型仍可用,仅提示补更新)
|
||||
if (vm.state.modelReady && ModelManager.instance.error != null)
|
||||
Positioned(
|
||||
left: 16,
|
||||
right: 16,
|
||||
top: MediaQuery.of(context).padding.top + 56,
|
||||
child: Container(
|
||||
padding:
|
||||
const EdgeInsets.symmetric(horizontal: 12, vertical: 6),
|
||||
decoration: BoxDecoration(
|
||||
color: Colors.black54,
|
||||
borderRadius: BorderRadius.circular(8),
|
||||
),
|
||||
child: Text(
|
||||
ModelManager.instance.error!,
|
||||
textAlign: TextAlign.center,
|
||||
style: const TextStyle(color: Colors.orange, fontSize: 12),
|
||||
),
|
||||
),
|
||||
),
|
||||
|
||||
// 模型未加载时仅显示相机预览,不做检测标注(横幅置于顶栏下方,避免与底部诊断行重叠)
|
||||
if (!vm.state.modelReady)
|
||||
Positioned(
|
||||
@@ -305,7 +337,7 @@ class _CameraScreenState extends State<CameraScreen> {
|
||||
crossAxisAlignment: CrossAxisAlignment.center,
|
||||
children: [
|
||||
Text(
|
||||
'阈值:${(_minScore * 100).toStringAsFixed(0)}% 模型:${vm.state.modelReady ? '已加载' : '未加载'} 帧:${vm.state.framesReceived} 流:${camera?.streamCallbacks ?? 0} 推理:${vm.state.debugDetectCalls}次 异常:${vm.state.debugDetectErrors}次 处理:${vm.state.debugLastMs}ms 最高分:${(vm.state.debugHighestScore * 100).toStringAsFixed(1)}% 图:${vm.state.imageWidthPx}x${vm.state.imageHeightPx} 传感:${camera?.sensorOrientation ?? '-'} 屏转:${camera?.displayDegrees ?? '-'} 旋:${camera?.rotationDegrees ?? 0} turn:${camera?.quarterTurns ?? '-'}',
|
||||
'阈值:${(_minScore * 100).toStringAsFixed(0)}% 模型:${vm.state.modelReady ? ModelManager.instance.modelsLabel : '未加载'} 帧:${vm.state.framesReceived} 流:${camera?.streamCallbacks ?? 0} 推理:${vm.state.debugDetectCalls}次 异常:${vm.state.debugDetectErrors}次 处理:${vm.state.debugLastMs}ms 最高分:${(vm.state.debugHighestScore * 100).toStringAsFixed(1)}% 图:${vm.state.imageWidthPx}x${vm.state.imageHeightPx} 传感:${camera?.sensorOrientation ?? '-'} 屏转:${camera?.displayDegrees ?? '-'} 旋:${camera?.rotationDegrees ?? 0} turn:${camera?.quarterTurns ?? '-'}',
|
||||
style: const TextStyle(color: Colors.white70, fontSize: 12),
|
||||
),
|
||||
if (_nativeStats.isNotEmpty)
|
||||
|
||||
@@ -87,10 +87,13 @@ class _OverlayPainter extends CustomPainter {
|
||||
_drawDashedRect(canvas, box, paint);
|
||||
}
|
||||
|
||||
// 标签:框上方,含距离
|
||||
// 标签:框上方,含距离;多模型时标注来源模型名(内置资产不标)
|
||||
final dist = _distanceLabel(r);
|
||||
final modelTag = r.modelName.isNotEmpty && r.modelName != '内置'
|
||||
? '[${r.modelName}]'
|
||||
: '';
|
||||
final text =
|
||||
'${_labels[r.label] ?? r.label} ${(r.score * 100).toInt()}%$dist';
|
||||
'${_labels[r.label] ?? r.label}$modelTag ${(r.score * 100).toInt()}%$dist';
|
||||
final textPainter = TextPainter(
|
||||
text: TextSpan(
|
||||
text: text,
|
||||
|
||||
@@ -9,6 +9,10 @@ class DetectionResult {
|
||||
/// 轨迹已确认(多帧稳定/高分/活动确认),false = 候选,渲染为虚线
|
||||
final bool confirmed;
|
||||
|
||||
/// 产出该框的模型(数据集 id 与名称;内置资产模型为 -1/空)
|
||||
final int modelId;
|
||||
final String modelName;
|
||||
|
||||
const DetectionResult({
|
||||
required this.label,
|
||||
required this.score,
|
||||
@@ -17,6 +21,8 @@ class DetectionResult {
|
||||
required this.right,
|
||||
required this.bottom,
|
||||
this.confirmed = true,
|
||||
this.modelId = -1,
|
||||
this.modelName = '',
|
||||
});
|
||||
|
||||
double get width => right - left;
|
||||
@@ -40,6 +46,8 @@ class DetectionResult {
|
||||
right: right ?? this.right,
|
||||
bottom: bottom ?? this.bottom,
|
||||
confirmed: confirmed ?? this.confirmed,
|
||||
modelId: modelId,
|
||||
modelName: modelName,
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -7,17 +7,27 @@ import 'package:flutter/foundation.dart' show debugPrint;
|
||||
import 'package:flutter/services.dart' show rootBundle;
|
||||
|
||||
import '../camera/motion_detector.dart';
|
||||
import '../models/model_manager.dart';
|
||||
import 'background_model.dart';
|
||||
import 'detection_result.dart';
|
||||
import 'nms.dart';
|
||||
import 'tflite_detector.dart';
|
||||
import 'visual_prior.dart';
|
||||
|
||||
/// 推理工作单元:模型加载与检测全部在后台 isolate 执行,
|
||||
/// 主 isolate 只投递帧数据、接收结果,UI 不被推理阻塞(iOS 真机卡顿根因)。
|
||||
///
|
||||
/// 多模型并行推理:传入 [models](各数据集下载模型)后,每帧逐模型推理,
|
||||
/// 结果按类别分组跨模型 NMS 合并(同标签重复框取高分,不同标签互不压制);
|
||||
/// 无下载模型时回退内置资产模型。
|
||||
class DetectorWorker {
|
||||
static const String modelAsset = 'assets/model.tflite';
|
||||
static const String labelsAsset = 'assets/labels.txt';
|
||||
|
||||
/// 内置资产回退模型的标识
|
||||
static const int builtinModelId = -1;
|
||||
static const String builtinModelName = '内置';
|
||||
|
||||
final Isolate _isolate;
|
||||
final ReceivePort _responses;
|
||||
|
||||
@@ -54,16 +64,26 @@ class DetectorWorker {
|
||||
});
|
||||
}
|
||||
|
||||
/// 读取模型资产并启动后台推理 isolate;加载失败返回 null(App 降级为仅预览)。
|
||||
static Future<DetectorWorker?> create() async {
|
||||
/// 加载模型并启动后台推理 isolate;加载失败返回 null(App 降级为仅预览)。
|
||||
/// [models] 为空时回退内置资产模型(模型缺失同样返回 null)。
|
||||
static Future<DetectorWorker?> create({List<ModelBundle>? models}) async {
|
||||
try {
|
||||
final data = await rootBundle.load(modelAsset);
|
||||
final modelBytes =
|
||||
data.buffer.asUint8List(data.offsetInBytes, data.lengthInBytes);
|
||||
final labels = (await rootBundle.loadString(labelsAsset))
|
||||
.split('\n')
|
||||
.where((l) => l.trim().isNotEmpty)
|
||||
.toList();
|
||||
final payload = <List<Object?>>[];
|
||||
if (models != null && models.isNotEmpty) {
|
||||
for (final m in models) {
|
||||
payload.add(
|
||||
[m.bytes, m.labels, m.datasetId, m.datasetName]);
|
||||
}
|
||||
} else {
|
||||
final data = await rootBundle.load(modelAsset);
|
||||
final modelBytes =
|
||||
data.buffer.asUint8List(data.offsetInBytes, data.lengthInBytes);
|
||||
final labels = (await rootBundle.loadString(labelsAsset))
|
||||
.split('\n')
|
||||
.where((l) => l.trim().isNotEmpty)
|
||||
.toList();
|
||||
payload.add([modelBytes, labels, builtinModelId, builtinModelName]);
|
||||
}
|
||||
|
||||
final responses = ReceivePort();
|
||||
final isolate = await Isolate.spawn(_workerMain, responses.sendPort);
|
||||
@@ -73,7 +93,7 @@ class DetectorWorker {
|
||||
.timeout(const Duration(seconds: 10),
|
||||
onTimeout: () => throw TimeoutException('worker port timeout'));
|
||||
worker._port = port;
|
||||
port.send(['load', modelBytes, labels]);
|
||||
port.send(['load', payload]);
|
||||
await worker._ready.future
|
||||
.timeout(const Duration(seconds: 20), onTimeout: () {
|
||||
throw TimeoutException('model load timeout');
|
||||
@@ -147,6 +167,8 @@ class DetectorWorker {
|
||||
top: v[3] as double,
|
||||
right: v[4] as double,
|
||||
bottom: v[5] as double,
|
||||
modelId: v.length > 6 ? (v[6] as num).toInt() : -1,
|
||||
modelName: v.length > 7 ? v[7] as String : '',
|
||||
);
|
||||
}).toList();
|
||||
final motion = (list[5] as List)
|
||||
@@ -202,7 +224,7 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
mainPort.send(['port', control.sendPort]);
|
||||
mainPort.send(['log', 'worker-start']);
|
||||
|
||||
TfliteDetector? detector;
|
||||
List<TfliteDetector> detectors = const [];
|
||||
MotionDetector? motion;
|
||||
BackgroundModel? background;
|
||||
var lastDualMs = 0; // 双字节序推理诊断节流
|
||||
@@ -214,12 +236,35 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
case 'load':
|
||||
mainPort.send(['log', 'load-received']);
|
||||
try {
|
||||
detector = await TfliteDetector.fromBuffer(
|
||||
list[1] as Uint8List, (list[2] as List).cast<String>());
|
||||
if (detector == null) {
|
||||
mainPort.send(['load-error', 'fromBuffer 返回 null']);
|
||||
// 多模型:逐模型加载,单个失败不阻塞其余;全部失败才报错
|
||||
final loaded = <TfliteDetector>[];
|
||||
final failures = <String>[];
|
||||
for (final entry in list[1] as List) {
|
||||
final e = entry as List;
|
||||
final name = e.length > 3 ? e[3] as String : '';
|
||||
final d = await TfliteDetector.fromBuffer(
|
||||
e[0] as Uint8List,
|
||||
(e[1] as List).cast<String>(),
|
||||
modelId: (e[2] as num).toInt(),
|
||||
modelName: name,
|
||||
);
|
||||
if (d == null) {
|
||||
failures.add(name.isEmpty ? 'unknown' : name);
|
||||
} else {
|
||||
loaded.add(d);
|
||||
}
|
||||
}
|
||||
if (loaded.isEmpty) {
|
||||
mainPort.send([
|
||||
'load-error',
|
||||
'模型加载失败:${failures.join(',')} '
|
||||
'(fromBuffer 返回 null)'
|
||||
]);
|
||||
} else {
|
||||
mainPort.send(['log', 'fromBuffer-ok']);
|
||||
detectors = loaded;
|
||||
mainPort.send(['log',
|
||||
'loaded=${loaded.map((d) => d.modelName).join(',')} '
|
||||
'failed=${failures.isEmpty ? '-' : failures.join(',')}']);
|
||||
motion = MotionDetector();
|
||||
background = BackgroundModel();
|
||||
mainPort.send(['ready']);
|
||||
@@ -229,10 +274,9 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
}
|
||||
break;
|
||||
case 'frame':
|
||||
final d = detector;
|
||||
final m = motion;
|
||||
final b = background;
|
||||
if (d == null || m == null || b == null) break;
|
||||
if (detectors.isEmpty || m == null || b == null) break;
|
||||
final frame = list[1] as List;
|
||||
final planes = (frame[0] as List).cast<Uint8List>();
|
||||
final strides = (frame[1] as List).cast<int>();
|
||||
@@ -302,42 +346,49 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
uv1 += ' v:min=$vMin max=$vMax mean=${(vSum / n2).toStringAsFixed(0)}';
|
||||
}
|
||||
}
|
||||
// 首帧(或相机重启后)自适应判定 YUV 值域与色序,再跑正式推理
|
||||
if (!isBgra && !d.yuvModeKnown) {
|
||||
d.decideYuvChroma(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
);
|
||||
}
|
||||
var results = d.detectRaw(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
rgbaOrder: rgbaOrder,
|
||||
);
|
||||
// 自愈:判定后 1.5s 内无检测且帧可用 → 用实时帧重跑完整判定
|
||||
// (首帧模糊/暗帧导致启发式猜错时,画面稳定后 oracle 可分胜负)
|
||||
if (!isBgra && d.yuvModeKnown && !d.yuvRetried &&
|
||||
results.length <= 1 &&
|
||||
DateTime.now().millisecondsSinceEpoch - d.yuvDecisionMs > 1500 &&
|
||||
d.retryDecision(
|
||||
// 多模型并行推理:每模型先首帧自适应判定 YUV 值域/色序,再逐模型推理;
|
||||
// 汇总后按类别分组跨模型 NMS 合并(同标签重复框取高分,异标签互不压制)
|
||||
var results = <DetectionResult>[];
|
||||
for (final d in detectors) {
|
||||
if (!isBgra && !d.yuvModeKnown) {
|
||||
d.decideYuvChroma(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
)) {
|
||||
results = d.detectRaw(
|
||||
);
|
||||
}
|
||||
var dets = d.detectRaw(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
rgbaOrder: rgbaOrder,
|
||||
);
|
||||
// 自愈:判定后 1.5s 内无检测且帧可用 → 用实时帧重跑完整判定
|
||||
// (首帧模糊/暗帧导致启发式猜错时,画面稳定后 oracle 可分胜负)
|
||||
if (!isBgra && d.yuvModeKnown && !d.yuvRetried &&
|
||||
dets.length <= 1 &&
|
||||
DateTime.now().millisecondsSinceEpoch - d.yuvDecisionMs >
|
||||
1500 &&
|
||||
d.retryDecision(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
)) {
|
||||
dets = d.detectRaw(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: isBgra,
|
||||
);
|
||||
}
|
||||
results.addAll(dets);
|
||||
}
|
||||
results = mergeAcrossModels(results, TfliteDetector.iouThreshold);
|
||||
// 低分野鸡框过视觉先验(颜色/位置),减少户外误报
|
||||
results = VisualPrior.filter(
|
||||
results,
|
||||
@@ -358,14 +409,14 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
if (isBgra &&
|
||||
DateTime.now().millisecondsSinceEpoch - lastDualMs > 3000) {
|
||||
lastDualMs = DateTime.now().millisecondsSinceEpoch;
|
||||
final a = d.diagnoseOrder(
|
||||
final a = detectors.first.diagnoseOrder(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
height: height,
|
||||
isBgra: true,
|
||||
rgbaOrder: false);
|
||||
final b = d.diagnoseOrder(
|
||||
final b = detectors.first.diagnoseOrder(
|
||||
planes: planes,
|
||||
strides: strides,
|
||||
width: width,
|
||||
@@ -382,8 +433,16 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
width,
|
||||
height,
|
||||
results
|
||||
.map((r) =>
|
||||
[r.label, r.score, r.left, r.top, r.right, r.bottom])
|
||||
.map((r) => [
|
||||
r.label,
|
||||
r.score,
|
||||
r.left,
|
||||
r.top,
|
||||
r.right,
|
||||
r.bottom,
|
||||
r.modelId,
|
||||
r.modelName,
|
||||
])
|
||||
.toList(),
|
||||
motionRegions
|
||||
.map((mr) => [mr.left, mr.top, mr.right, mr.bottom])
|
||||
@@ -395,17 +454,21 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
'planes=${planes.length} yLen=${yPlane.length} stride=${strides[0]} '
|
||||
'y:min=$yMin max=$yMax mean=${yMean.toStringAsFixed(1)} '
|
||||
'diff=${yDiff < 0 ? '-' : yDiff.toStringAsFixed(3)} '
|
||||
'uv1:[$uv1] | ${d.yuvDiag}$dualDiag',
|
||||
'uv1:[$uv1] | ${detectors.first.yuvDiag}$dualDiag',
|
||||
]);
|
||||
break;
|
||||
case 'reset':
|
||||
motion?.reset();
|
||||
background?.reset();
|
||||
detector?.resetYuvMode();
|
||||
for (final d in detectors) {
|
||||
d.resetYuvMode();
|
||||
}
|
||||
break;
|
||||
case 'set-min-score':
|
||||
detector?.minScore = (list[1] as num).toDouble();
|
||||
mainPort.send(['log', 'min-score=${detector?.minScore}']);
|
||||
for (final d in detectors) {
|
||||
d.minScore = (list[1] as num).toDouble();
|
||||
}
|
||||
mainPort.send(['log', 'min-score=${detectors.isEmpty ? '-' : detectors.first.minScore}']);
|
||||
}
|
||||
} catch (e, st) {
|
||||
mainPort.send([
|
||||
@@ -415,3 +478,20 @@ Future<void> _workerMain(SendPort mainPort) async {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// 多模型结果合并:按类别分组,组内 NMS(不同模型检出同一目标时取高分)。
|
||||
/// 各模型类别体系独立(如野鸡/疑似 vs 野兔/疑似),不同类别互不压制。
|
||||
List<DetectionResult> mergeAcrossModels(
|
||||
List<DetectionResult> all, double iouThreshold) {
|
||||
if (all.length <= 1) return all;
|
||||
final byLabel = <String, List<DetectionResult>>{};
|
||||
for (final r in all) {
|
||||
byLabel.putIfAbsent(r.label, () => []).add(r);
|
||||
}
|
||||
final merged = <DetectionResult>[];
|
||||
for (final group in byLabel.values) {
|
||||
merged.addAll(nms(group, iouThreshold));
|
||||
}
|
||||
merged.sort((a, b) => b.score.compareTo(a.score));
|
||||
return merged;
|
||||
}
|
||||
|
||||
@@ -11,7 +11,9 @@ import 'nms.dart';
|
||||
/// cx/cy/w/h 已归一化,类别得分已过 sigmoid;按 out[dim][anchor] 索引。
|
||||
/// 输入为 NCHW [1, 3, 704, 704](litert 导出保留 torch 布局)。
|
||||
class TfliteDetector {
|
||||
static const int inputSize = 704;
|
||||
// 输入尺寸取自模型本身(ultralytics litert 导出 NCHW [1,3,H,W],各数据集
|
||||
// 训练 imgsz 可不同),默认 704 兜底
|
||||
static const int defaultInputSize = 704;
|
||||
// 野鸡数据置信度普遍偏低(0.1~0.2 量级),保留低分池供运动检测提升;
|
||||
// 可运行时调整(设置页滑块),默认 0.10
|
||||
double minScore = 0.10;
|
||||
@@ -24,39 +26,55 @@ class TfliteDetector {
|
||||
final List<String> _labels;
|
||||
final int _numClasses;
|
||||
final int _numAnchors;
|
||||
final int inputSize;
|
||||
|
||||
final Float32List _input =
|
||||
Float32List(1 * inputSize * inputSize * 3);
|
||||
/// 模型身份(多模型并行推理区分来源;内置资产模型为 -1/空)
|
||||
final int modelId;
|
||||
final String modelName;
|
||||
|
||||
late final Float32List _input;
|
||||
|
||||
/// 输出按模型形状 [1, 4+nc, anchors] 的嵌套 List 组织,
|
||||
/// run() 要求输出对象形状与模型完全一致(扁平 List 会被拒)。
|
||||
final List<List<List<double>>> _output;
|
||||
|
||||
TfliteDetector._(this._interpreter, this._labels, this._numClasses,
|
||||
this._numAnchors, this._output);
|
||||
this._numAnchors, this._output, this.inputSize, this.modelId,
|
||||
this.modelName) {
|
||||
_input = Float32List(1 * inputSize * inputSize * 3);
|
||||
}
|
||||
|
||||
/// 模型缺失或加载失败返回 null(App 降级为仅预览)。
|
||||
/// 在后台 isolate 内调用(模型字节由主 isolate 读取后传入)。
|
||||
static Future<TfliteDetector?> fromBuffer(
|
||||
Uint8List bytes, List<String> labels) async {
|
||||
Uint8List bytes,
|
||||
List<String> labels, {
|
||||
int modelId = -1,
|
||||
String modelName = '',
|
||||
}) async {
|
||||
try {
|
||||
final interpreter = Interpreter.fromBuffer(
|
||||
bytes,
|
||||
options: InterpreterOptions()..threads = 4,
|
||||
);
|
||||
return TfliteDetector._fromModel(interpreter, labels);
|
||||
return TfliteDetector._fromModel(
|
||||
interpreter, labels, modelId, modelName);
|
||||
} catch (_) {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
/// 输出布局 [1, 4+nc, anchors] 取自模型本身,类别数不与 labels 文件长度耦合。
|
||||
factory TfliteDetector._fromModel(
|
||||
Interpreter interpreter, List<String> labels) {
|
||||
factory TfliteDetector._fromModel(Interpreter interpreter,
|
||||
List<String> labels, int modelId, String modelName) {
|
||||
final shape = interpreter.getOutputTensor(0).shape;
|
||||
final numClasses =
|
||||
shape.length >= 3 && shape[1] > 4 ? shape[1] - 4 : labels.length;
|
||||
final numAnchors = shape.length >= 3 && shape[2] > 0 ? shape[2] : 2100;
|
||||
final inputShape = interpreter.getInputTensor(0).shape;
|
||||
final inputSize = inputShape.length >= 4
|
||||
? inputShape[3]
|
||||
: defaultInputSize;
|
||||
final output = List.generate(
|
||||
1,
|
||||
(_) => List.generate(
|
||||
@@ -64,8 +82,8 @@ class TfliteDetector {
|
||||
(_) => List<double>.filled(numAnchors, 0),
|
||||
),
|
||||
);
|
||||
return TfliteDetector._(
|
||||
interpreter, labels, numClasses, numAnchors, output);
|
||||
return TfliteDetector._(interpreter, labels, numClasses, numAnchors,
|
||||
output, inputSize, modelId, modelName);
|
||||
}
|
||||
|
||||
/// 原始数据接口(后台 isolate 用,不依赖 CameraImage)。
|
||||
@@ -463,6 +481,8 @@ class TfliteDetector {
|
||||
top: (cy - h / 2).clamp(0.0, 1.0),
|
||||
right: (cx + w / 2).clamp(0.0, 1.0),
|
||||
bottom: (cy + h / 2).clamp(0.0, 1.0),
|
||||
modelId: modelId,
|
||||
modelName: modelName,
|
||||
));
|
||||
}
|
||||
final kept = nms(boxes, iouThreshold);
|
||||
|
||||
@@ -0,0 +1,285 @@
|
||||
import 'dart:convert';
|
||||
import 'dart:io';
|
||||
|
||||
import 'package:crypto/crypto.dart' show sha256;
|
||||
import 'package:flutter/foundation.dart';
|
||||
import 'package:http/http.dart' as http;
|
||||
import 'package:path_provider/path_provider.dart';
|
||||
|
||||
import '../config/app_config.dart';
|
||||
|
||||
/// 模型目录条目(GET /api/v1/app/update 响应 data.models[])。
|
||||
/// 服务器发布模型后随版本检查一同下发,App 按目录逐数据集下载/更新。
|
||||
class ModelCatalogItem {
|
||||
final int datasetId;
|
||||
final String datasetName;
|
||||
final String version;
|
||||
final List<String> labels;
|
||||
final int sizeBytes;
|
||||
final String sha256;
|
||||
final String downloadUrl;
|
||||
|
||||
const ModelCatalogItem({
|
||||
required this.datasetId,
|
||||
required this.datasetName,
|
||||
required this.version,
|
||||
required this.labels,
|
||||
required this.sizeBytes,
|
||||
required this.sha256,
|
||||
required this.downloadUrl,
|
||||
});
|
||||
|
||||
factory ModelCatalogItem.fromJson(Map<String, dynamic> j) =>
|
||||
ModelCatalogItem(
|
||||
datasetId: (j['datasetId'] as num?)?.toInt() ?? 0,
|
||||
datasetName: j['datasetName'] as String? ?? '',
|
||||
version: j['version'] as String? ?? '',
|
||||
labels: (j['labels'] as List? ?? const [])
|
||||
.map((e) => e.toString())
|
||||
.toList(),
|
||||
sizeBytes: (j['sizeBytes'] as num?)?.toInt() ?? 0,
|
||||
sha256: j['sha256'] as String? ?? '',
|
||||
downloadUrl: j['downloadUrl'] as String? ?? '',
|
||||
);
|
||||
}
|
||||
|
||||
/// 已就绪模型(字节 + 标签,供推理 worker 加载;含内置资产回退模型)
|
||||
class ModelBundle {
|
||||
final int datasetId;
|
||||
final String datasetName;
|
||||
final String version;
|
||||
final List<String> labels;
|
||||
final Uint8List bytes;
|
||||
|
||||
const ModelBundle({
|
||||
required this.datasetId,
|
||||
required this.datasetName,
|
||||
required this.version,
|
||||
required this.labels,
|
||||
required this.bytes,
|
||||
});
|
||||
}
|
||||
|
||||
/// 模型热更新管理:启动时拉取模型目录(随 /app/update 公开接口下发,无需登录态),
|
||||
/// 按需下载/校验/持久化各数据集模型,供相机页多模型并行推理。
|
||||
///
|
||||
/// 存储:应用私有目录 `models/<datasetId>/`(model.tflite + labels.json + meta.json),
|
||||
/// meta 记录 {version, sha256},服务器发布新版本时按版本+摘要重下,不重复下载旧模型。
|
||||
class ModelManager extends ChangeNotifier {
|
||||
static final ModelManager instance = ModelManager._();
|
||||
|
||||
final String baseUrl;
|
||||
final http.Client _client;
|
||||
final Future<Directory> Function()? _rootDirOverride;
|
||||
|
||||
List<ModelBundle> _models = const [];
|
||||
bool _ready = false;
|
||||
bool _refreshing = false;
|
||||
String? _error;
|
||||
Future<void>? _inFlight;
|
||||
|
||||
ModelManager._({String? baseUrl, http.Client? client})
|
||||
: this(baseUrl: baseUrl, client: client);
|
||||
|
||||
/// 可注入 baseUrl / client / 存储根目录(单测用)
|
||||
@visibleForTesting
|
||||
ModelManager({
|
||||
String? baseUrl,
|
||||
http.Client? client,
|
||||
Future<Directory> Function()? rootDir,
|
||||
}) : baseUrl = baseUrl ?? AppConfig.apiBaseUrl,
|
||||
_client = client ?? http.Client(),
|
||||
_rootDirOverride = rootDir;
|
||||
|
||||
/// 已就绪模型列表(空 = 无服务器模型,回退内置资产)
|
||||
List<ModelBundle> get models => _models;
|
||||
|
||||
/// 是否成功拉取过目录(即使下载失败也为 true,用于区分"从未联网"与"目录为空")
|
||||
bool get ready => _ready;
|
||||
|
||||
/// 最近一次同步的错误信息(下载失败/校验失败等;目录为空不算错误)
|
||||
String? get error => _error;
|
||||
|
||||
bool get refreshing => _refreshing;
|
||||
|
||||
/// 模型名摘要(诊断行展示):内置 / 数据集名×n
|
||||
String get modelsLabel {
|
||||
if (_models.isEmpty) return '内置';
|
||||
return _models.map((m) => m.datasetName).join(',');
|
||||
}
|
||||
|
||||
/// 拉取目录并同步本地模型;并发调用共享同一进行中的刷新。
|
||||
Future<void> refresh() {
|
||||
if (_refreshing) return _inFlight ?? Future.value();
|
||||
_refreshing = true;
|
||||
_inFlight = _doRefresh().whenComplete(() {
|
||||
_refreshing = false;
|
||||
_inFlight = null;
|
||||
notifyListeners();
|
||||
});
|
||||
return _inFlight!;
|
||||
}
|
||||
|
||||
Future<void> _doRefresh() async {
|
||||
try {
|
||||
final res = await _client
|
||||
.get(Uri.parse('$baseUrl/api/v1/app/update'))
|
||||
.timeout(const Duration(seconds: 8));
|
||||
// 服务器 Content-Type 无 charset,http 包默认按 latin1 解码会乱码 → 显式 utf8
|
||||
final body =
|
||||
jsonDecode(utf8.decode(res.bodyBytes)) as Map<String, dynamic>;
|
||||
final data = body['data'] as Map<String, dynamic>? ?? const {};
|
||||
final list = data['models'] as List? ?? const [];
|
||||
final catalog = list
|
||||
.map((e) => ModelCatalogItem.fromJson(e as Map<String, dynamic>))
|
||||
.toList();
|
||||
|
||||
final failed = <String>[];
|
||||
for (final item in catalog) {
|
||||
if (!await _ensureLocal(item)) failed.add(item.datasetName);
|
||||
}
|
||||
await _prune(catalog);
|
||||
_models = await _loadBundles(catalog);
|
||||
_ready = true;
|
||||
_error = failed.isEmpty
|
||||
? null
|
||||
: '模型下载失败:${failed.join(',')}(重试旧模型或稍后再试)';
|
||||
} catch (e) {
|
||||
if (!_ready) _error = '模型目录拉取失败:$e';
|
||||
// 已就绪过则保留旧模型,不覆盖 error(下载级错误优先展示)
|
||||
}
|
||||
}
|
||||
|
||||
/// 保证目录条目在本地可用:meta 匹配且文件在 → 跳过;否则下载并校验 sha256。
|
||||
Future<bool> _ensureLocal(ModelCatalogItem item) async {
|
||||
final dir = await _modelDir(item.datasetId);
|
||||
try {
|
||||
final meta = await _readMeta(dir);
|
||||
final file = File('${dir.path}/model.tflite');
|
||||
if (meta != null &&
|
||||
meta['version'] == item.version &&
|
||||
meta['sha256'] == item.sha256 &&
|
||||
await file.exists()) {
|
||||
return true;
|
||||
}
|
||||
// 版本更新或文件缺失:下载校验(失败重试一次)
|
||||
for (var attempt = 0; attempt < 2; attempt++) {
|
||||
final ok = await _downloadAndVerify(item, dir, file);
|
||||
if (ok) return true;
|
||||
await file.delete().catchError((_) => file);
|
||||
await File('${dir.path}/model.tflite.part')
|
||||
.delete()
|
||||
.catchError((_) => file);
|
||||
}
|
||||
debugPrint('[ModelManager] 下载失败: ${item.datasetName} ${item.version}');
|
||||
return false;
|
||||
} catch (e) {
|
||||
debugPrint('[ModelManager] _ensureLocal ${item.datasetName}: $e');
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
Future<bool> _downloadAndVerify(
|
||||
ModelCatalogItem item, Directory dir, File file) async {
|
||||
final part = File('${file.path}.part');
|
||||
final sink = part.openWrite();
|
||||
try {
|
||||
final res = await _client
|
||||
.send(http.Request('GET', Uri.parse('$baseUrl${item.downloadUrl}')))
|
||||
.timeout(const Duration(minutes: 3));
|
||||
if (res.statusCode != 200) return false;
|
||||
await res.stream.pipe(sink);
|
||||
await sink.close();
|
||||
final bytes = await part.readAsBytes();
|
||||
final hex = sha256.convert(bytes).toString();
|
||||
if (item.sha256.isNotEmpty && hex != item.sha256) {
|
||||
debugPrint('[ModelManager] sha256 不匹配: ${item.datasetName} '
|
||||
'want=${item.sha256} got=$hex');
|
||||
return false;
|
||||
}
|
||||
await part.rename(file.path);
|
||||
await dir.create(recursive: true);
|
||||
await File('${dir.path}/labels.json')
|
||||
.writeAsString(jsonEncode(item.labels));
|
||||
await File('${dir.path}/meta.json').writeAsString(jsonEncode({
|
||||
'version': item.version,
|
||||
'sha256': item.sha256,
|
||||
}));
|
||||
debugPrint('[ModelManager] 已下载 ${item.datasetName} '
|
||||
'${bytes.length}B -> ${file.path}');
|
||||
return true;
|
||||
} catch (e) {
|
||||
await sink.close().catchError((_) {});
|
||||
debugPrint('[ModelManager] 下载异常 ${item.datasetName}: $e');
|
||||
return false;
|
||||
}
|
||||
}
|
||||
|
||||
/// 清理服务器目录中已下线的数据集模型(不再发布则删本地)
|
||||
Future<void> _prune(List<ModelCatalogItem> catalog) async {
|
||||
final root = await _rootDir();
|
||||
if (!await root.exists()) return;
|
||||
final keep = catalog.map((c) => '${c.datasetId}').toSet();
|
||||
await for (final e in root.list()) {
|
||||
if (e is Directory) {
|
||||
// 目录 URI 末尾带 '/',pathSegments 末位为空串 → 过滤后取目录名
|
||||
final name = e.uri.pathSegments.where((s) => s.isNotEmpty).last;
|
||||
if (!keep.contains(name)) {
|
||||
await e.delete(recursive: true).catchError((_) => e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Future<List<ModelBundle>> _loadBundles(
|
||||
List<ModelCatalogItem> catalog) async {
|
||||
final bundles = <ModelBundle>[];
|
||||
for (final item in catalog) {
|
||||
try {
|
||||
final dir = await _modelDir(item.datasetId);
|
||||
final file = File('${dir.path}/model.tflite');
|
||||
if (!await file.exists()) continue;
|
||||
final labels = await File('${dir.path}/labels.json').exists()
|
||||
? (jsonDecode(
|
||||
await File('${dir.path}/labels.json').readAsString())
|
||||
as List)
|
||||
.map((e) => e.toString())
|
||||
.toList()
|
||||
: item.labels;
|
||||
bundles.add(ModelBundle(
|
||||
datasetId: item.datasetId,
|
||||
datasetName: item.datasetName,
|
||||
version: item.version,
|
||||
labels: labels,
|
||||
bytes: await file.readAsBytes(),
|
||||
));
|
||||
} catch (e) {
|
||||
debugPrint('[ModelManager] 读取 ${item.datasetName} 失败: $e');
|
||||
}
|
||||
}
|
||||
return bundles;
|
||||
}
|
||||
|
||||
Future<Map<String, dynamic>?> _readMeta(Directory dir) async {
|
||||
final f = File('${dir.path}/meta.json');
|
||||
if (!await f.exists()) return null;
|
||||
try {
|
||||
return jsonDecode(await f.readAsString()) as Map<String, dynamic>;
|
||||
} catch (_) {
|
||||
return null;
|
||||
}
|
||||
}
|
||||
|
||||
Future<Directory> _rootDir() async {
|
||||
if (_rootDirOverride != null) return _rootDirOverride();
|
||||
final support = await getApplicationSupportDirectory();
|
||||
return Directory('${support.path}/models');
|
||||
}
|
||||
|
||||
Future<Directory> _modelDir(int datasetId) async {
|
||||
final root = await _rootDir();
|
||||
final dir = Directory('${root.path}/$datasetId');
|
||||
await dir.create(recursive: true);
|
||||
return dir;
|
||||
}
|
||||
}
|
||||
@@ -39,7 +39,9 @@ class UpdateChecker {
|
||||
final res = await _client
|
||||
.get(Uri.parse('$baseUrl/api/v1/app/update'))
|
||||
.timeout(const Duration(seconds: 8));
|
||||
final body = jsonDecode(res.body) as Map<String, dynamic>;
|
||||
// 服务器 Content-Type 无 charset,http 包默认按 latin1 解码会乱码 → 显式 utf8
|
||||
final body =
|
||||
jsonDecode(utf8.decode(res.bodyBytes)) as Map<String, dynamic>;
|
||||
final data = body['data'] as Map<String, dynamic>? ?? const {};
|
||||
return AppUpdateInfo(
|
||||
version: data['version'] as String? ?? '',
|
||||
|
||||
Reference in New Issue
Block a user