如何在Dart/Flutter中使用自定义TFLite模型实现数组预测
在Flutter中加载TFLite模型并获取布尔预测结果
1. 添加依赖
在项目的pubspec.yaml文件中添加TFLite Flutter依赖:
dependencies: flutter: sdk: flutter tflite_flutter: ^0.10.1 # 可根据最新版本调整
执行flutter pub get完成依赖安装。
2. 配置模型资源
将model.tflite文件放到项目的assets目录下,然后在pubspec.yaml中声明该资源:
flutter: assets: - assets/model.tflite
3. 实现模型加载与预测逻辑
以下是完整的Dart代码示例,涵盖模型加载、输入处理、推理执行和结果转换:
import 'package:flutter/services.dart'; import 'package:tflite_flutter/tflite_flutter.dart'; class ModelPredictor { late Interpreter _interpreter; bool _isModelLoaded = false; // 异步加载模型 Future<void> _loadModel() async { if (_isModelLoaded) return; try { final modelData = await rootBundle.load('assets/model.tflite'); _interpreter = await Interpreter.fromBuffer(modelData.buffer); _isModelLoaded = true; } catch (e) { print('模型加载失败: $e'); rethrow; } } // 执行预测:输入数组返回布尔结果 Future<bool> predict(List<double> inputArray) async { await _loadModel(); // 匹配模型输入形状,示例为[1, N]的二维输入 final input = [inputArray]; // 初始化输出张量(根据模型实际输出形状调整) final outputTensorShape = _interpreter.getOutputTensor(0).shape; final outputLength = outputTensorShape.reduce((a, b) => a * b); var output = List<double>.filled(outputLength, 0.0).reshape(outputTensorShape); // 运行推理 _interpreter.run(input, output); // 转换输出为布尔值 final rawOutput = output is List ? (output.first as List).first : output; // 可根据模型输出类型调整规则:若输出为0/1整数,直接判断rawOutput == 1.0即可 return (rawOutput as double) > 0.5; } // 释放模型资源 void dispose() { if (_isModelLoaded) { _interpreter.close(); _isModelLoaded = false; } } }
关键注意事项
- 输入形状匹配:必须保证输入数组的维度和模型训练时的输入维度完全一致,比如模型输入为
[1, 64],则输入列表需包含64个元素,并包装为二维列表[[...]]。 - 输出转换逻辑:根据模型输出类型调整布尔值转换规则:
- 若模型直接输出0/1整数,可直接判断
rawOutput == 1.0 - 若输出是分类概率(0-1区间),选择合适的阈值(如0.5)判定true/false
- 若模型直接输出0/1整数,可直接判断
- 资源释放:在页面销毁或不再使用模型时,调用
dispose()释放Interpreter资源,避免内存泄漏。
内容的提问来源于stack exchange,提问作者Rusiru Gunaratne
相关产品推荐
相关产品推荐

