You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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项目

  1. 在Flutter项目根目录创建assets文件夹,放入转换好的kobert_custom.onnx和vocab.txt。
  2. 修改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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.06 00:43:20