不同TensorFlow版本下LSTM模型Int8训练后量化报错求助
LSTM模型INT8训练后量化可行方案
1. 环境选择
直接使用TensorFlow 2.10及以上正式版本,该版本已原生支持LSTM算子INT8量化、多子图模型转换,修复了旧版本MLIR/TOCO转换的各类兼容性bug,无需再测试2.2/2.5/2.6/2.7 nightly等旧版本。
2. 编写符合要求的代表数据集生成器
代表数据集生成器每次返回的内容需要是列表包裹的输入张量,和原有numpy输入格式无冲突,仅需在外层加一层列表即可,同时注意将输入从float64转为float32(TFLite校准默认支持float32输入),示例代码如下:
import pandas as pd import numpy as np def reshape_for_Lstm(data): timesteps=1 samples=int(np.floor(data.shape[0]/timesteps)) data=data.reshape((samples,timesteps,data.shape[1])) return data # 加载校准用数据,可使用全部测试数据或抽取至少100条样本保证校准精度 data = pd.read_csv('./test_x_data_OOP3.csv', index_col=[0]) data = np.array(data, dtype=np.float32) # 转float32适配量化校准 data = reshape_for_Lstm(data) def batch_generator(): # 每次yield一个batch的输入,batch size设为1即可,也可使用更大的batch for i in range(data.shape[0]): # 外层加[]返回列表格式,符合校准器要求 yield [data[i:i+1]]
3. 调整量化转换参数
删除旧版本兼容参数,使用适配新版本的转换配置,代码如下:
import tensorflow as tf converter = tf.lite.TFLiteConverter.from_saved_model('./model/singnature_model_tf_2.7.0-dev20210914') # 开启默认量化优化 converter.optimizations = [tf.lite.Optimize.DEFAULT] # 传入代表数据集生成器 converter.representative_dataset = batch_generator # 配置支持的算子,同时加入INT8内置算子和SELECT_TF_OPS兼容剩余算子 converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8, tf.lite.OpsSet.SELECT_TF_OPS] # 指定输入输出量化类型,若需要输入输出为float32可注释下面两行 converter.inference_input_type = tf.int8 converter.inference_output_type = tf.int8 # 关闭张量列表降级的实验参数避免报错 converter._experimental_lower_tensor_list_ops = False # 执行转换 quantized_tflite_model = converter.convert() # 保存量化后模型 with open('lstm_quantized_int8.tflite', 'wb') as f: f.write(quantized_tflite_model)
4. 验证量化模型精度
转换完成后可使用如下代码测试推理结果和原模型误差是否在可接受范围内:
# 加载量化模型 interpreter = tf.lite.Interpreter(model_path='lstm_quantized_int8.tflite') interpreter.allocate_tensors() input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() # 单条推理测试 test_sample = data[0:1] interpreter.set_tensor(input_details[0]['index'], test_sample) interpreter.invoke() yhat_quant = interpreter.get_tensor(output_details[0]['index']) yclass_quant = interpreter.get_tensor(output_details[1]['index'])
内容的提问来源于stack exchange,提问作者Florida Man
相关产品推荐
相关产品推荐

