如何在Python中调用TensorFlow Lite的MobileBERT模型并解决运行错误?
使用Python调用TensorFlow Lite MobileBERT模型实现文本情感分类
我下载了TensorFlow官方的文本分类Android演示应用,该应用包含AverageWordVec和MobileBERT两种文本情感(正负)分类模型,其中MobileBERT的准确率更高。现在希望在Python中使用应用里的mobilebert.tflite文件实现相同的预测效果,找不到合适方案,求示例代码。
更新
使用tokenizer方法运行时,输出如下日志:
2022-11-15 15:22:40.048025: I tensorflow/core/platform/cpu_feature_guard.cc:193] This TensorFlow binary is optimized with oneAPI Deep Neural Network Library (oneDNN) to use the following CPU instructions in performance-critical operations: AVX2 AVX512F FMA To enable them in other operations, rebuild TensorFlow with the appropriate compiler flags. 2022-11-15 15:22:40.177014: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcudart.so.11.0'; dlerror: libcudart.so.11.0: cannot open shared object file: No such file or directory 2022-11-15 15:22:40.177060: I tensorflow/stream_executor/cuda/cudart_stub.cc:29] Ignore above cudart dlerror if you do not have a GPU set up on your machine. 2022-11-15 15:22:40.208395: E tensorflow/stream_executor/cuda/cuda_blas.cc:2981] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered 2022-11-15 15:22:40.831937: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libnvinfer.so.7'; dlerror: libnvinfer.so.7: cannot open shared object file: No such file or directory 2022-11-15 15:22:40.832028: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libnvinfer_plugin.so.7'; dlerror: libnvinfer_plugin.so.7: cannot open shared object file: No such file or directory 2022-11-15 15:22:40.832041: W tensorflow/compiler/tf2tensorrt/utils/py_utils.cc:38] TF-TRT Warning: Cannot dlopen some TensorRT libraries. If you would like to use Nvidia GPU with TensorRT, please make sure the missing libraries mentioned above are installed properly. 2022-11-15 15:22:41.897720: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcuda.so.1'; dlerror: libcuda.so.1: cannot open shared object file: No such file or directory 2022-11-15 15:22:41.897765: W tensorflow/stream_executor/cuda/cuda_driver.cc:263] failed call to cuInit: UNKNOWN ERROR (303) 2022-11-15 15:22:41.897783: I tensorflow/stream_executor/cuda/cuda_diagnostics.cc:156] kernel driver does not appear to be running on this host (Android-CI-CD): /proc/driver/nvidia/version does not exist 2022-11-15 15:22:41.898114: I tensorflow/core/platform/cpu_feature_guard.cc:193] This TensorFlow binary is optimized with oneAPI Deep Neural Network Library (oneDNN) to use the following CPU instructions in performance-critical operations: AVX2 AVX512F FMA To enable them in other operations, rebuild TensorFlow with the appropriate compiler flags. INFO: Created TensorFlow Lite XNNPACK delegate for CPU.
使用tflite_support.task方法运行时,输出如下日志后出现段错误:
2022-11-15 15:25:30.210277: I tensorflow/core/platform/cpu_feature_guard.cc:193] This TensorFlow binary is optimized with oneAPI Deep Neural Network Library (oneDNN) to use the following CPU instructions in performance-critical operations: AVX2 AVX512F FMA To enable them in other operations, rebuild TensorFlow with the appropriate compiler flags. 2022-11-15 15:25:30.342254: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcudart.so.11.0'; dlerror: libcudart.so.11.0: cannot open shared object file: No such file or directory 2022-11-15 15:25:30.342295: I tensorflow/stream_executor/cuda/cudart_stub.cc:29] Ignore above cudart dlerror if you do not have a GPU set up on your machine. 2022-11-15 15:25:30.372550: E tensorflow/stream_executor/cuda/cuda_blas.cc:2981] Unable to register cuBLAS factory: Attempting to register factory for plugin cuBLAS when one has already been registered 2022-11-15 15:25:30.990440: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libnvinfer.so.7'; dlerror: libnvinfer.so.7: cannot open shared object file: No such file or directory 2022-11-15 15:25:30.990522: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libnvinfer_plugin.so.7'; dlerror: libnvinfer_plugin.so.7: cannot open shared object file: No such file or directory 2022-11-15 15:25:30.990532: W tensorflow/compiler/tf2tensorrt/utils/py_utils.cc:38] TF-TRT Warning: Cannot dlopen some TensorRT libraries. If you would like to use Nvidia GPU with TensorRT, please make sure the missing libraries mentioned above are installed properly. INFO: Created TensorFlow Lite XNNPACK delegate for CPU. Segmentation fault (core dumped)
解决方法与示例代码
日志提示说明
上述GPU相关警告均为机器未配置Nvidia GPU的正常提示,不影响CPU推理,可直接忽略。tflite_support出现段错误的核心原因是版本不兼容,以下提供两种可行方案:
方案一:使用TensorFlow Lite Interpreter手动实现推理
此方案无需依赖tflite_support,通过Hugging Face Tokenizer对齐Android端预处理逻辑,保证预测效果一致。
- 安装依赖:
pip install tensorflow transformers
- 示例代码:
import tensorflow as tf from transformers import BertTokenizer # 加载与Android端匹配的MobileBERT分词器 tokenizer = BertTokenizer.from_pretrained("google/mobilebert-uncased") # 加载tflite模型 interpreter = tf.lite.Interpreter(model_path="mobilebert.tflite") interpreter.allocate_tensors() # 获取模型输入输出张量信息 input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() def predict_sentiment(text): # 预处理文本:截断/填充至模型要求的固定长度(Android端为128) inputs = tokenizer( text, padding="max_length", truncation=True, max_length=128, return_tensors="tf" ) # 提取模型所需的三个输入张量 input_ids = inputs["input_ids"].numpy() attention_mask = inputs["attention_mask"].numpy() token_type_ids = inputs["token_type_ids"].numpy() # 设置模型输入 interpreter.set_tensor(input_details[0]['index'], input_ids) interpreter.set_tensor(input_details[1]['index'], attention_mask) interpreter.set_tensor(input_details[2]['index'], token_type_ids) # 执行推理 interpreter.invoke() # 获取输出结果:[负类概率, 正类概率] output = interpreter.get_tensor(output_details[0]['index']) negative_prob = output[0][0] positive_prob = output[0][1] sentiment = "正面" if positive_prob > negative_prob else "负面" return sentiment, negative_prob, positive_prob # 测试示例 test_text = "这部电影太精彩了,强烈推荐!" sentiment, neg_prob, pos_prob = predict_sentiment(test_text) print(f"文本: {test_text}") print(f"情感: {sentiment}, 负面概率: {neg_prob:.4f}, 正面概率: {pos_prob:.4f}")
方案二:修复tflite_support的段错误问题
通过指定兼容版本的TensorFlow和tflite_support解决段错误:
- 安装指定版本依赖:
pip install tensorflow==2.10.0 tflite-support==0.4.4
- 示例代码:
from tflite_support.task import text from tflite_support.task import core # 配置分类器参数 base_options = core.BaseOptions(file_name="mobilebert.tflite") options = text.TextClassifierOptions(base_options=base_options) # 初始化文本分类器 classifier = text.TextClassifier.create_from_options(options) def predict_with_tflite_support(text): # 执行情感分类 result = classifier.classify(text) # 提取最高置信度的结果 top_category = result.classifications[0].categories[0] return top_category.label, top_category.score # 测试示例 test_text = "这个产品质量太差,完全不值得买。" sentiment, confidence = predict_with_tflite_support(test_text) print(f"文本: {test_text}") print(f"情感: {sentiment}, 置信度: {confidence:.4f}")
内容的提问来源于stack exchange,提问作者Hossam Hassan
相关产品推荐
相关产品推荐

