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

PyTorch转TensorFlow Lite:无法获取BABERT中文模型输入形状

问题描述

我要把BABERT中文模型从PyTorch格式转成TensorFlow Lite格式,打算用torch2tflite工具,这个工具要求指定模型输入形状。我写了Python脚本想获取模型形状,但加载pytorch_model.bin时抛出AttributeError,提示'collections.OrderedDict'对象没有'parameters'属性。

我的需求是在Android设备上实现中文分词,本地训练因为CUDA内存不够没法做,现在正在调研预训练模型。相关代码和错误信息如下:

加载模型脚本

import torch;
 
def loadModel():
    model = torch.load("/home/gelassen/Downloads/chinese_babert-base/pytorch_model.bin")
    model_shape = list(model.parameters())[0].shape 
    print(model_shape)
    print("Model shape" + model_shape)

loadModel()

错误信息

/start.py", line 10, in loadModel
    model_shape = list(model.parameters())[0].shape 
AttributeError: 'collections.OrderedDict' object has no attribute 'parameters'

torch2tflite转换命令示例

python3 -m torch2tflite.converter
    --torch-path tests/mobilenetv2_model.pt
    --tflite-path mobilenetv2.tflite
    --target-shape 224 224 3

解决方案

1. 正确加载BABERT模型

pytorch_model.bin只是模型的权重文件,不是完整的模型实例,直接用torch.load加载得到的是权重字典(OrderedDict),自然没有parameters()方法。你需要先加载BABERT的模型结构,再加载权重:

先安装transformers库,然后用对应模型类加载结构,再加载权重:

from transformers import BertTokenizer, BertForTokenClassification
import torch

def loadModel():
    # 加载BABERT的分词任务模型结构
    model = BertForTokenClassification.from_pretrained("uer/chinese_babert-base")
    # 如果用本地权重文件,执行下面一行替换默认权重
    # model.load_state_dict(torch.load("/home/gelassen/Downloads/chinese_babert-base/pytorch_model.bin"))
    
    # 查看模型输入相关参数(BERT类模型输入一般为[batch_size, sequence_length])
    embedding_layer = model.bert.embeddings.word_embeddings
    print("Embedding层形状:", embedding_layer.weight.shape)
    
    # 生成示例输入,确认输入形状
    tokenizer = BertTokenizer.from_pretrained("uer/chinese_babert-base")
    sample_input = tokenizer("测试中文分词", return_tensors="pt")
    print("示例输入形状:")
    print("input_ids:", sample_input["input_ids"].shape)
    print("attention_mask:", sample_input["attention_mask"].shape)

loadModel()

2. 转换模型到TensorFlow Lite

torch2tflite需要完整的PyTorch模型实例和明确的输入形状。注意BERT类模型有多个输入(input_ids、attention_mask等),需要指定所有输入的形状:

步骤1:导出PyTorch模型为TorchScript

先把加载好的模型转成TorchScript格式,方便后续转换:

from transformers import BertForTokenClassification, BertTokenizer
import torch

# 加载模型和分词器
model = BertForTokenClassification.from_pretrained("uer/chinese_babert-base")
tokenizer = BertTokenizer.from_pretrained("uer/chinese_babert-base")

# 设置模型为评估模式
model.eval()

# 生成固定长度的示例输入,确定输入形状(这里指定max_length=128)
inputs = tokenizer("测试中文分词", return_tensors="pt", padding="max_length", max_length=128)

# 导出TorchScript模型
traced_model = torch.jit.trace(model, (inputs["input_ids"], inputs["attention_mask"]))
traced_model.save("babert_token_classification.pt")

步骤2:用torch2tflite转换

转换时需要指定多个输入的形状,每个输入形状用空格分隔,多个输入形状依次排列:

python3 -m torch2tflite.converter
    --torch-path babert_token_classification.pt
    --tflite-path babert_token_classification.tflite
    --target-shape 1 128 1 128

这里1 128对应input_ids的形状(batch_size=1, sequence_length=128),第二个1 128对应attention_mask的形状。

3. 适配Android端分词

  • 转换后的TFLite模型可以在Android中用TensorFlow Lite Interpreter加载
  • 安卓端需要用对应的分词器处理文本,把文字转为模型需要的input_ids和attention_mask格式
  • 如果CUDA内存不足,建议考虑量化后的轻量BERT模型(比如DistilBERT、TinyBERT),体积更小、内存占用更低,更适合移动端

内容的提问来源于stack exchange,提问作者Gleichmut

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 19:59:50