{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"slr_model = SPOTER(num_classes=250)    #自己创建的模型，前面需要加入自己的模型\nmodel_pth = 'train5_best.pth'  #自己模型的路径\nslr_model.load_state_dict(torch.load(model_pth,map_location=torch.device('cpu')))   #导入模型权重\nslr_model.train(False)  #关闭训练模式\n ","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:49:41.388834Z","iopub.execute_input":"2023-04-11T08:49:41.389702Z","iopub.status.idle":"2023-04-11T08:49:41.429817Z","shell.execute_reply.started":"2023-04-11T08:49:41.389643Z","shell.execute_reply":"2023-04-11T08:49:41.42894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **将pytorch模型转换为onnx模型**","metadata":{}},{"cell_type":"code","source":"input_onnx_file = \"/kaggle/working/model.onnx\"\ndef run_convert_onnx(slr_model): \n    torch.onnx.export(\n                slr_model,\n                torch.randn((60,543,3)),   # 输入的大小\n                input_onnx_file,             # 保存输出onnx的路径\n                export_params = True,        \n                opset_version = 12,          # tONNX 版本\n                do_constant_folding=True,    # 是否执行常数折叠进行优化\n                input_names =  ['inputs'],    # the model's input names\n                output_names = ['outputs'],   # the model's output names\n                dynamic_axes={\n                    'inputs': {0: 'length'},\n                    #'output': {0: 'length'},\n                },\n                #verbose = True,\n            )\n\nrun_convert_onnx(slr_model)\nprint('model.onnx saved !!')","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:49:43.169374Z","iopub.execute_input":"2023-04-11T08:49:43.170099Z","iopub.status.idle":"2023-04-11T08:49:44.006072Z","shell.execute_reply.started":"2023-04-11T08:49:43.170059Z","shell.execute_reply":"2023-04-11T08:49:44.003965Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **检查模型输出是否相同**\n导入tf模型","metadata":{}},{"cell_type":"code","source":"!pip install onnxruntime","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:49:44.007625Z","iopub.execute_input":"2023-04-11T08:49:44.008108Z","iopub.status.idle":"2023-04-11T08:49:53.404564Z","shell.execute_reply.started":"2023-04-11T08:49:44.008066Z","shell.execute_reply":"2023-04-11T08:49:53.403345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import onnxruntime as ort\nmodel_path = \"/kaggle/working/model.onnx\"\nort_session = ort.InferenceSession(model_path)\n\n# 定义测试输入\nimport numpy as np\ninput_tensor = torch.randn(60,543,3)\ninput_data = input_tensor.numpy()\n\n# # 检查onnx输出\n# print(ort_session.get_inputs())\n \n    \nort_inputs = {ort_session.get_inputs()[0].name: input_data}\nort_outputs = ort_session.run(None,ort_inputs )\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:49:53.406652Z","iopub.execute_input":"2023-04-11T08:49:53.40731Z","iopub.status.idle":"2023-04-11T08:49:53.516248Z","shell.execute_reply.started":"2023-04-11T08:49:53.407266Z","shell.execute_reply":"2023-04-11T08:49:53.515276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **将onnx模型转换为tf_model模型**","metadata":{}},{"cell_type":"code","source":"!pip install onnx_tf    # 安装依赖包","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:49:53.517574Z","iopub.execute_input":"2023-04-11T08:49:53.518426Z","iopub.status.idle":"2023-04-11T08:50:03.056238Z","shell.execute_reply.started":"2023-04-11T08:49:53.518388Z","shell.execute_reply":"2023-04-11T08:50:03.054944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from onnx_tf.backend import prepare\nimport onnx\nTF_PATH = \"tf_model\" # 保存tf模型的位置\nONNX_PATH = \"/kaggle/working/model.onnx\"   # onnx模型的path\nonnx_model = onnx.load(ONNX_PATH)  # 导入onnx模型权重\ntf_rep = prepare(onnx_model)  # 创建TensorflowRep对象\ntf_rep.export_graph(TF_PATH)\nprint('tf.saved_model() passed !!')","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:03.059313Z","iopub.execute_input":"2023-04-11T08:50:03.059785Z","iopub.status.idle":"2023-04-11T08:50:23.964434Z","shell.execute_reply.started":"2023-04-11T08:50:03.059736Z","shell.execute_reply":"2023-04-11T08:50:23.963258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **测试tf_model模型是否准确**","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nmobilenet_save_path = '/kaggle/working/tf_model'\nloaded = tf.saved_model.load(mobilenet_save_path)\n# print(list(loaded.signatures.keys())) \ninfer = loaded.signatures[\"serving_default\"]\n# print(infer)\n# 检查tf输出\nx1=tf.constant(input_tensor)\nlabel = infer(x1)\n# print(label['outputs'])","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:23.96633Z","iopub.execute_input":"2023-04-11T08:50:23.966754Z","iopub.status.idle":"2023-04-11T08:50:30.440723Z","shell.execute_reply.started":"2023-04-11T08:50:23.966715Z","shell.execute_reply":"2023-04-11T08:50:30.439596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" # **将转换好的.pb模型打包**","metadata":{}},{"cell_type":"code","source":"packagePath = '/kaggle/working/tf_model'\nzipPath = '/kaggle/working/'\n!zip tf_model.zip $model_path","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:30.442622Z","iopub.execute_input":"2023-04-11T08:50:30.443035Z","iopub.status.idle":"2023-04-11T08:50:32.189382Z","shell.execute_reply.started":"2023-04-11T08:50:30.442994Z","shell.execute_reply":"2023-04-11T08:50:32.188116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **将tf_model模型转为tflite模型**","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\n# print(tf.__version__)\nsaved_model_dir = '/kaggle/working/tf_model'   # 保存tensorflow模型的文件夹名称\ntflite_path = '/kaggle/working/model.tflite'   #保存 tflite 模型的路径名称\n# Convert the model\nconverter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir) \n\nconverter.target_spec.supported_ops = [\n  tf.lite.OpsSet.TFLITE_BUILTINS,  # enable TensorFlow Lite ops.\n  tf.lite.OpsSet.SELECT_TF_OPS   # enable TensorFlow ops.\n]\ntflite_model = converter.convert()\nopen(tflite_path, \"wb\").write(tflite_model)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:32.191476Z","iopub.execute_input":"2023-04-11T08:50:32.191834Z","iopub.status.idle":"2023-04-11T08:50:38.17068Z","shell.execute_reply.started":"2023-04-11T08:50:32.191803Z","shell.execute_reply":"2023-04-11T08:50:38.169646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **测试转换的tflite模型是否准确**","metadata":{}},{"cell_type":"code","source":"!pip install tflite_runtime==2.9.1","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:38.172141Z","iopub.execute_input":"2023-04-11T08:50:38.172605Z","iopub.status.idle":"2023-04-11T08:50:48.631596Z","shell.execute_reply.started":"2023-04-11T08:50:38.172567Z","shell.execute_reply":"2023-04-11T08:50:48.630371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tqdm\nimport tflite_runtime.interpreter as tflite\n\nmodel_path = '/kaggle/working/model.tflite'\ninterpreter = tflite.Interpreter(model_path)\nfound_signatures = list(interpreter.get_signature_list().keys())\nprediction_fn = interpreter.get_signature_runner(\"serving_default\")\n# input_tensor = torch.randn(50,543,3)\n# input_data = input_tensor.numpy()\n# print(input_data)\noutput = prediction_fn(inputs=input_data)\nsign = np.argmax(output[\"outputs\"])\nprint(output['outputs'])\n","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:48.634533Z","iopub.execute_input":"2023-04-11T08:50:48.635074Z","iopub.status.idle":"2023-04-11T08:50:48.738344Z","shell.execute_reply.started":"2023-04-11T08:50:48.635025Z","shell.execute_reply":"2023-04-11T08:50:48.73706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **submission**","metadata":{}},{"cell_type":"code","source":"ROWS_PER_FRAME = 543  # number of landmarks per frame\nimport pandas as pd\ndef load_relevant_data_subset(pq_path):\n    data_columns = ['x', 'y', 'z']\n    data = pd.read_parquet(pq_path, columns=data_columns)\n    n_frames = int(len(data) / ROWS_PER_FRAME)\n    data = data.values.reshape(n_frames, ROWS_PER_FRAME, len(data_columns))\n    return data.astype(np.float32)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:48.740039Z","iopub.execute_input":"2023-04-11T08:50:48.740576Z","iopub.status.idle":"2023-04-11T08:50:48.74774Z","shell.execute_reply.started":"2023-04-11T08:50:48.740536Z","shell.execute_reply":"2023-04-11T08:50:48.74655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tqdm\nimport tflite_runtime.interpreter as tflite\nimport time\nt1 = time.time()\nmodel_path = '/kaggle/working/model.tflite'\ninterpreter = tflite.Interpreter(model_path)\nfound_signatures = list(interpreter.get_signature_list().keys())\npq_path = '/kaggle/input/asl-signs/train_landmark_files/18796/1020380433.parquet'\nprediction_fn = interpreter.get_signature_runner(\"serving_default\")\nframes = load_relevant_data_subset(pq_path)\n# print(frames)\noutput = prediction_fn(inputs=frames)\n# print(output)\nsign = np.argmax(output[\"outputs\"])\nprint(sign, output['outputs'].shape)\nt2 =time.time()\nprint('{:.2f}ms'.format((t2-t1)*1000))","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:48.749348Z","iopub.execute_input":"2023-04-11T08:50:48.750034Z","iopub.status.idle":"2023-04-11T08:50:48.816216Z","shell.execute_reply.started":"2023-04-11T08:50:48.749996Z","shell.execute_reply":"2023-04-11T08:50:48.81502Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **测试Tflite模型准确率**","metadata":{}},{"cell_type":"code","source":"import json\nimport time\ndef load_dataset(file_location):\n    sign_dic = {}\n    label_dic_path = '/kaggle/input/asl-signs/sign_to_prediction_index_map.json'\n    with open(label_dic_path, 'r') as sign_f:\n        sign_dic = json.load(sign_f)\n    # 读入数据集中所有的data数据与label，返回np.array类型(读进内存中)\n    # Load the datset csv file\n    df = pd.read_csv(file_location, encoding=\"utf-8\")\n    df_signs = df['sign'].to_list()\n    df_labels = [sign_dic.get(i) for i in df_signs]   # 从字典返回labels\n    df_paths = df['path'].to_list()\n    return df_paths, df_labels\n\ndata_file = '/kaggle/input/asl-signs/train.csv'\ndf_paths, df_labels = load_dataset(data_file)\n# print(df_paths[:50])\nsign_list = [0]*250\nsign_all = [0]*250\n\ndef run_tflite(df_path,de_label):\n    df_path = '/kaggle/input/asl-signs/'+df_path\n    model_path = '/kaggle/working/model.tflite'\n    interpreter = tflite.Interpreter(model_path)\n    found_signatures = list(interpreter.get_signature_list().keys())\n    prediction_fn = interpreter.get_signature_runner(\"serving_default\")\n    frames = load_relevant_data_subset(df_path)\n#     print(frames.shape)\n    output = prediction_fn(inputs=frames)\n    sign = np.argmax(output[\"outputs\"])\n#     print(output)\n    print('pre:{}\\t true:{}'.format(sign,de_label))\n    if sign==de_label:\n        sign_list[de_label] += 1\n    sign_all[de_label] +=1\nt1 = time.time()\nvideo_sum = 100\nfor i in range(100):\n    run_tflite(df_paths[i], df_labels[i])  \nt2 = time.time()\nprint(sign_list)\nprint(sum(sign_list))\nprint(sum(sign_all))\nprint('cost {:.4f}'.format(t2-t1))\nprint('pre video {:.4f}'.format((t2-t1)/video_sum))","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:48.826852Z","iopub.execute_input":"2023-04-11T08:50:48.827313Z","iopub.status.idle":"2023-04-11T08:50:55.012751Z","shell.execute_reply.started":"2023-04-11T08:50:48.827269Z","shell.execute_reply":"2023-04-11T08:50:55.010636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\n!zip submission.zip  'model.tflite'\n!ls\n\nprint(f'submit ok')","metadata":{"execution":{"iopub.status.busy":"2023-04-11T08:50:55.014463Z","iopub.execute_input":"2023-04-11T08:50:55.014859Z","iopub.status.idle":"2023-04-11T08:50:57.735124Z","shell.execute_reply.started":"2023-04-11T08:50:55.014819Z","shell.execute_reply":"2023-04-11T08:50:57.73376Z"},"trusted":true},"execution_count":null,"outputs":[]}]}