如何在Flutter中加载自定义训练的KoBERT模型?
在Flutter中加载自定义微调的KoBERT模型
核心前提
Flutter本身无法直接加载PyTorch格式的.safetensors模型,需要先将模型转换为移动端支持的格式(如ONNX或TensorFlow Lite),再通过对应的Flutter ML插件完成加载与推理。
步骤1:将PyTorch模型转换为ONNX格式
ONNX是跨平台的模型格式,Flutter有成熟的插件支持。用以下Python代码完成转换:
import torch from transformers import BertForSequenceClassification # 加载本地微调后的模型 model = BertForSequenceClassification.from_pretrained('.') model.eval() # 切换到推理模式 # 创建与训练时匹配的示例输入(需与实际输入维度一致,示例为batch=1、序列长度64) dummy_input_ids = torch.randint(0, 1000, (1, 64)) dummy_attention_mask = torch.ones((1, 64)) # 导出为ONNX模型 torch.onnx.export( model, (dummy_input_ids, dummy_attention_mask), "kobert_custom.onnx", export_params=True, opset_version=15, do_constant_folding=True, input_names=['input_ids', 'attention_mask'], output_names=['logits'], # 支持动态batch和序列长度 dynamic_axes={ 'input_ids': {0: 'batch_size', 1: 'sequence_length'}, 'attention_mask': {0: 'batch_size', 1: 'sequence_length'}, 'logits': {0: 'batch_size'} } )
同时需要导出KoBERT的词汇表文件vocab.txt(可从原skt/kobert-base-v1模型中获取),用于后续文本分词。
步骤2:将模型资源集成到Flutter项目
- 在Flutter项目根目录创建
assets文件夹,放入转换好的kobert_custom.onnx和vocab.txt。 - 修改
pubspec.yaml声明资源:
flutter: assets: - assets/kobert_custom.onnx - assets/vocab.txt
步骤3:Flutter中加载模型并实现推理
以flutter_onnxruntime插件为例(支持ONNX模型推理),同时需要实现KoBERT的韩文分词逻辑:
1. 添加依赖
修改pubspec.yaml:
dependencies: flutter: sdk: flutter flutter_onnxruntime: ^0.3.0 flutter_mecab: ^0.2.0 # 用于韩文分词,匹配KoBERT的分词规则
2. 模型加载与推理代码
import 'package:flutter_onnxruntime/flutter_onnxruntime.dart'; import 'package:flutter/services.dart'; import 'package:flutter_mecab/flutter_mecab.dart'; class CustomKoBERT { late OrtSession _modelSession; late Map<String, int> _vocabMap; late Mecab _mecab; // 初始化模型与分词器 Future<void> init() async { // 加载词汇表 final vocabText = await rootBundle.loadString('assets/vocab.txt'); _vocabMap = {}; final vocabLines = vocabText.split('\n'); for (int i = 0; i < vocabLines.length; i++) { final token = vocabLines[i].trim(); if (token.isNotEmpty) _vocabMap[token] = i; } // 加载ONNX模型 final modelBytes = await rootBundle.load('assets/kobert_custom.onnx'); _modelSession = await OrtSession.fromBuffer(modelBytes.buffer.asUint8List()); // 初始化Mecab分词器(匹配KoBERT分词规则) _mecab = await Mecab.init(); } // 文本分词,转换为模型所需的input_ids和attention_mask Map<String, List<int>> _tokenize(String text, int maxSeqLen) { // 用Mecab做韩文分词 final tokens = _mecab.parse(text).map((item) => item.surface).toList(); // 拼接CLS、SEP标记,处理padding final inputIds = <int>[_vocabMap['[CLS]']!]; for (final token in tokens) { if (inputIds.length >= maxSeqLen - 1) break; inputIds.add(_vocabMap[token] ?? _vocabMap['[UNK]']!); } inputIds.add(_vocabMap['[SEP]']!); // 填充到指定长度 while (inputIds.length < maxSeqLen) { inputIds.add(_vocabMap['[PAD]']!); } // 生成attention_mask final attentionMask = inputIds.map((id) => id == _vocabMap['[PAD]']! ? 0 : 1).toList(); return {'input_ids': inputIds, 'attention_mask': attentionMask}; } // 推理预测 Future<List<double>> predict(String text, {int maxSeqLen = 64}) async { final tokenizedData = _tokenize(text, maxSeqLen); // 创建模型输入张量 final inputIdsTensor = OrtTensor.createTensorInt64(tokenizedData['input_ids']!, [1, maxSeqLen]); final attentionMaskTensor = OrtTensor.createTensorInt64(tokenizedData['attention_mask']!, [1, maxSeqLen]); // 运行模型 final outputs = await _modelSession.run({ 'input_ids': inputIdsTensor, 'attention_mask': attentionMaskTensor }); // 提取输出logits并释放资源 final logits = outputs['logits']!.value as List<List<double>>; inputIdsTensor.dispose(); attentionMaskTensor.dispose(); outputs.forEach((_, tensor) => tensor.dispose()); return logits[0]; } }
注意事项
- 移动端性能优化:在Android上开启NNAPI加速、iOS上开启Metal加速,
flutter_onnxruntime可通过配置OrtSessionOptions实现。 - 模型大小:KoBERT模型体积较大,建议对模型进行量化(如INT8量化)后再转换,减少移动端内存占用。
内容的提问来源于stack exchange,提问作者EASY
相关产品推荐
相关产品推荐

