{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":44224,"databundleVersionId":5188730,"sourceType":"competition"},{"sourceId":5212841,"sourceType":"datasetVersion","datasetId":3006078,"isSourceIdPinned":true},{"sourceId":5668986,"sourceType":"datasetVersion","datasetId":3012613},{"sourceId":5763578,"sourceType":"datasetVersion","datasetId":3006070},{"sourceId":8556431,"sourceType":"datasetVersion","datasetId":5070689},{"sourceId":8771305,"sourceType":"datasetVersion","datasetId":5271088},{"sourceId":8815069,"sourceType":"datasetVersion","datasetId":5273973},{"sourceId":8856072,"sourceType":"datasetVersion","datasetId":5331160},{"sourceId":8935987,"sourceType":"datasetVersion","datasetId":5294716,"isSourceIdPinned":false},{"sourceId":69148,"sourceType":"modelInstanceVersion","modelInstanceId":57681}],"dockerImageVersionId":30408,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:47:12.811148Z","iopub.execute_input":"2024-07-17T07:47:12.811545Z","iopub.status.idle":"2024-07-17T07:47:13.917766Z","shell.execute_reply.started":"2024-07-17T07:47:12.811507Z","shell.execute_reply":"2024-07-17T07:47:13.916534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install /kaggle/input/bird-clef-2023-addones/onnxruntime-1.14.0-cp37-cp37m-manylinux_2_27_x86_64.whl --no-deps","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:47:13.920064Z","iopub.execute_input":"2024-07-17T07:47:13.920428Z","iopub.status.idle":"2024-07-17T07:47:17.251881Z","shell.execute_reply.started":"2024-07-17T07:47:13.920386Z","shell.execute_reply":"2024-07-17T07:47:17.250805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip list | grep onnx","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:47:17.253397Z","iopub.execute_input":"2024-07-17T07:47:17.253738Z","iopub.status.idle":"2024-07-17T07:47:17.25855Z","shell.execute_reply.started":"2024-07-17T07:47:17.2537Z","shell.execute_reply":"2024-07-17T07:47:17.257489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install onnxoptimizer\n!pip install onnxruntime\n!pip install onnxsim\n!pip install onnx2pytorch\n!pip install onnx2torch\n!pip install nnAudio","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-07-17T07:47:17.261235Z","iopub.execute_input":"2024-07-17T07:47:17.261692Z","iopub.status.idle":"2024-07-17T07:48:29.085136Z","shell.execute_reply.started":"2024-07-17T07:47:17.261654Z","shell.execute_reply":"2024-07-17T07:48:29.083658Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#!pip install seaborn\n#!pip install timm\n#!pip install librosa","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/bird-clef-2023-code/main_folder/main_folder/')\nsys.path.append('/kaggle/input/onnx2pytroch/onnx2pytorch-master/')\nsys.path.append('/kaggle/input/o2pt-noclamp/onnx2torch-main/')","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:29.086959Z","iopub.execute_input":"2024-07-17T07:48:29.087285Z","iopub.status.idle":"2024-07-17T07:48:29.092872Z","shell.execute_reply.started":"2024-07-17T07:48:29.087252Z","shell.execute_reply":"2024-07-17T07:48:29.091862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport librosa\nimport seaborn as sns\nimport os\nimport json\nimport IPython.display as ipd\nimport soundfile as sf\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport h5py\nimport onnxruntime as ort\nimport onnx\nimport onnxoptimizer\nimport timm\n\nfrom glob import glob\nfrom tqdm import tqdm\nfrom matplotlib import pyplot as plt\nfrom itertools import chain\nfrom os.path import join as pjoin\nfrom copy import deepcopy\nfrom scipy.spatial.distance import cosine\nfrom torchvision import transforms, models\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nfrom matplotlib.backends.backend_agg import FigureCanvasAgg as FigureCanvas\nfrom onnx2pytorch import ConvertModel\n\nfrom code_base.models import WaveCNNClasifier, WaveCNNAttenClasifier\nfrom code_base.datasets import WaveDataset, WaveAllFileDataset\nfrom code_base.utils.inference_utils import apply_avarage_weights_on_swa_path\nfrom code_base.inefernce import BirdsInference\nfrom code_base.utils import load_json, compose_submission_dataframe, groupby_np_array, stack_and_max_by_samples\nfrom code_base.utils.metrics import padded_cmap_numpy\n%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2024-07-17T08:01:46.446786Z","iopub.execute_input":"2024-07-17T08:01:46.44768Z","iopub.status.idle":"2024-07-17T08:01:46.465412Z","shell.execute_reply.started":"2024-07-17T08:01:46.447636Z","shell.execute_reply":"2024-07-17T08:01:46.464339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"EXP_NAME = \"convnext_small_fb_in22k_ft_in1k_384__convnextv2_tiny_fcmae_ft_in22k_in1k_384__eca_nfnet_l0_noval_v32_075Clipwise025TimeMax_GausMean\"\nTRAIN_PERIOD = 5\nprint(\"Possible checkpoints:\\n\\n{}\".format(\"\\n\".join(set([\n    os.path.basename(el) for el in glob(f\"/kaggle/input/bird-clef-2023-models/{EXP_NAME}/{EXP_NAME}/*/checkpoints/*.pt*\") if \"train\" not in os.path.basename(el)\n]))))","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:34.486437Z","iopub.execute_input":"2024-07-17T07:48:34.486782Z","iopub.status.idle":"2024-07-17T07:48:34.524521Z","shell.execute_reply.started":"2024-07-17T07:48:34.486749Z","shell.execute_reply":"2024-07-17T07:48:34.52354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG = {\n    # Main\n    \"run_validation\": False,\n    \"run_test\": True,\n    # Inference Class\n    \"use_sigmoid\": False,\n    # Data config\n#     \"train_df_path\":\"/home/vova/data/exps/BirdCLEF_2023/birdclef_2023/train_metadata_extended.csv\",\n#     \"split_path\":\"/home/vova/data/exps/BirdCLEF_2023/cv_split_2023_v1.npy\",\n    \"folds\":[0],\n#     \"train_data_root\":\"/home/vova/data/exps/BirdCLEF_2023/birdclef_2023/train_audio/\",\n    \"test_data_root\": \"/kaggle/input/birdclef-2023/test_soundscapes/*.ogg\",\n    \"label_map_data_path\":'/kaggle/input/bird-clef-2023-models/bird2int_2023.json',\n\n    \"lookback\":None,\n    \"lookahead\":None,\n    \"segment_len\":5,\n    \"step\": None,\n    \"late_normalize\": True,\n    # Model config\n    \"exp_name\":EXP_NAME,\n     \"model_class\": WaveCNNAttenClasifier,\n     \"model_config\": dict(\n         #backbone=\"convnext_tiny_in22ft1k\",\n         #backbone=\"convnext_small.fb_in22k_ft_in1k_384\",\n         #mel_spec_paramms={\n         #   \"sample_rate\": 32000,\n         #   \"n_mels\": 128,\n         #    \"f_min\": 20,\n         #    \"n_fft\": 2048,\n         #    \"hop_length\": 512,\n         #    \"normalized\": True,\n         #},\n         #head_config={\n         #    \"p\": 0.5,\n         #    \"num_class\": 264,\n         #    \"train_period\": TRAIN_PERIOD,\n         #    \"infer_period\": TRAIN_PERIOD,\n         #},\n         #pretrained=False\n     ),\n     \"chkp_name\":\"model.last.pth\",\n     \"swa_checkpoint\": None,\n     \"distributed_chkp\": False,\n}\n\nif CONFIG.get(\"use_sed_mode\", False):\n    assert CONFIG[\"step\"] is not None\nelse:\n    assert CONFIG[\"step\"] is None\n    \nif \"folds\" not in CONFIG:\n    CONFIG[\"folds\"] = list(range(CONFIG[\"n_folds\"]))","metadata":{"execution":{"iopub.status.busy":"2024-07-17T08:05:28.622525Z","iopub.execute_input":"2024-07-17T08:05:28.623627Z","iopub.status.idle":"2024-07-17T08:05:28.632879Z","shell.execute_reply.started":"2024-07-17T08:05:28.623585Z","shell.execute_reply":"2024-07-17T08:05:28.631716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"code","source":"bird2id = load_json(CONFIG[\"label_map_data_path\"])","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:34.53719Z","iopub.execute_input":"2024-07-17T07:48:34.537489Z","iopub.status.idle":"2024-07-17T07:48:34.5491Z","shell.execute_reply.started":"2024-07-17T07:48:34.53746Z","shell.execute_reply":"2024-07-17T07:48:34.548136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CONFIG[\"run_validation\"]:\n    df = pd.read_csv(CONFIG[\"train_df_path\"])\n    split = np.load(CONFIG[\"split_path\"], allow_pickle=True)\n    val_df = [df.iloc[split[fold_id][1]].reset_index(drop=True) for fold_id in CONFIG[\"folds\"]]","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:34.550372Z","iopub.execute_input":"2024-07-17T07:48:34.550751Z","iopub.status.idle":"2024-07-17T07:48:34.556613Z","shell.execute_reply.started":"2024-07-17T07:48:34.550715Z","shell.execute_reply":"2024-07-17T07:48:34.555491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CONFIG[\"run_test\"]:\n    test_au_pathes = glob(CONFIG[\"test_data_root\"])\n\n    test_df = pd.DataFrame({\n        \"filename\": test_au_pathes,\n        \"duration_s\": [librosa.get_duration(path=el) for el in test_au_pathes]\n    })","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:34.557778Z","iopub.execute_input":"2024-07-17T07:48:34.558122Z","iopub.status.idle":"2024-07-17T07:48:38.378297Z","shell.execute_reply.started":"2024-07-17T07:48:34.558095Z","shell.execute_reply":"2024-07-17T07:48:38.377328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CONFIG[\"run_validation\"]:\n    val_ds_conig = {\n       \"root\": CONFIG[\"train_data_root\"],\n       \"label_str2int_mapping_path\": CONFIG[\"label_map_data_path\"],\n       \"use_audio_cache\": True,\n       \"n_cores\": 64,\n       \"verbose\": False,\n       \"segment_len\": CONFIG[\"segment_len\"],\n       \"lookback\":CONFIG[\"lookback\"],\n       \"lookahead\":CONFIG[\"lookahead\"],\n       \"sample_id\": None,\n       \"late_normalize\": CONFIG[\"late_normalize\"],\n       \"step\": CONFIG[\"step\"],\n       \"validate_sr\": 32_000\n    }\nif CONFIG[\"run_test\"]:\n    ds_config_test = {\n       \"root\": \"\",\n       \"label_str2int_mapping_path\": CONFIG[\"label_map_data_path\"],\n       \"n_cores\": 64,\n       \"use_audio_cache\": True,\n       \"test_mode\": True,\n       \"segment_len\": CONFIG[\"segment_len\"],\n       \"lookback\":CONFIG[\"lookback\"],\n       \"lookahead\":CONFIG[\"lookahead\"],\n        \"sample_id\": None,\n        \"late_normalize\": CONFIG[\"late_normalize\"],\n        \"step\": CONFIG[\"step\"],\n        \"validate_sr\": 32_000\n    }\nloader_config = {\n    \"batch_size\": 4,\n    \"drop_last\": False,\n    \"shuffle\": False,\n    \"num_workers\": 0,\n}","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:38.379554Z","iopub.execute_input":"2024-07-17T07:48:38.380101Z","iopub.status.idle":"2024-07-17T07:48:38.389702Z","shell.execute_reply.started":"2024-07-17T07:48:38.380068Z","shell.execute_reply":"2024-07-17T07:48:38.388546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CONFIG[\"run_test\"]:\n    ds_test = WaveAllFileDataset(df=test_df, **ds_config_test)\nif CONFIG[\"run_validation\"]:\n    ds_val = [WaveAllFileDataset(df=df, **val_ds_conig) for df in val_df]","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:38.391069Z","iopub.execute_input":"2024-07-17T07:48:38.391363Z","iopub.status.idle":"2024-07-17T07:48:38.407801Z","shell.execute_reply.started":"2024-07-17T07:48:38.391334Z","shell.execute_reply":"2024-07-17T07:48:38.406718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CONFIG[\"run_validation\"]:\n    loader_val = [torch.utils.data.DataLoader(\n        ds,\n        **loader_config,\n    )for ds in ds_val]\nif CONFIG[\"run_test\"]:\n    loader_test = torch.utils.data.DataLoader(\n        ds_test,\n        **loader_config,\n    )","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:38.408897Z","iopub.execute_input":"2024-07-17T07:48:38.40917Z","iopub.status.idle":"2024-07-17T07:48:38.415066Z","shell.execute_reply.started":"2024-07-17T07:48:38.409137Z","shell.execute_reply":"2024-07-17T07:48:38.414097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_model_and_upload_chkp(\n    model_class,\n    model_config,\n    model_device,\n    model_chkp,\n    use_distributed=False,\n    swa_checkpoint=None\n):\n    print(model_chkp)\n    if \"swa\" in model_chkp:\n        print(\"swa by {}\".format(os.path.splitext(os.path.basename(model_chkp))[0]))\n        t_chkp = apply_avarage_weights_on_swa_path(model_chkp, use_distributed=use_distributed, take_best=swa_checkpoint)\n    else:\n        print(\"vanilla model\")\n        t_chkp = torch.load(model_chkp, map_location=\"cpu\")\n        \n    t_model = model_class(**model_config, device=model_device)\n    t_model.load_state_dict(t_chkp)\n    t_model.eval()\n    return t_model","metadata":{"execution":{"iopub.status.busy":"2024-07-17T07:48:38.418391Z","iopub.execute_input":"2024-07-17T07:48:38.418712Z","iopub.status.idle":"2024-07-17T07:48:38.42663Z","shell.execute_reply.started":"2024-07-17T07:48:38.418684Z","shell.execute_reply":"2024-07-17T07:48:38.425609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = [create_model_and_upload_chkp(\n#         model_class=CONFIG[\"model_class\"],\n#         model_config=CONFIG['model_config'],\n#         model_device=\"cpu\",\n#         model_chkp=f\"/kaggle/input/bird-clef-2023-models/{CONFIG['exp_name']}/{CONFIG['exp_name']}/fold_{m_i}/checkpoints/{CONFIG['chkp_name']}\",\n#         swa_checkpoint=CONFIG['swa_checkpoint'],\n#         use_distributed=CONFIG['distributed_chkp']\n# ) for m_i in CONFIG[\"folds\"]]","metadata":{"execution":{"iopub.status.busy":"2024-07-13T15:27:32.821574Z","iopub.execute_input":"2024-07-13T15:27:32.821946Z","iopub.status.idle":"2024-07-13T15:27:32.830979Z","shell.execute_reply.started":"2024-07-13T15:27:32.821899Z","shell.execute_reply":"2024-07-13T15:27:32.829845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ort.InferenceSession(f\"/kaggle/input/bird-clef-2023-models/{CONFIG['exp_name']}/{CONFIG['exp_name']}/checkpoints/model_simpl.onnx\")","metadata":{"execution":{"iopub.status.busy":"2024-07-13T15:27:32.832536Z","iopub.execute_input":"2024-07-13T15:27:32.833151Z","iopub.status.idle":"2024-07-13T15:27:38.653947Z","shell.execute_reply.started":"2024-07-13T15:27:32.833098Z","shell.execute_reply":"2024-07-13T15:27:38.652799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **重訓練模型：模型轉換和載入**","metadata":{}},{"cell_type":"markdown","source":"加載預訓練的 PyTorch 模型權重：","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn as nn\nfrom torchvision import models\n\nclass CustomResNet(nn.Module):\n    def __init__(self, pretrained_model):\n        super(CustomResNet, self).__init__()\n        self.features = nn.Sequential(*list(pretrained_model.children())[:-1])  # 移除原始的全連接層\n        self.fc = nn.Linear(2048, 264)  # 假設有264個類別\n\n    def forward(self, x):\n        x = self.features(x)\n        x = torch.flatten(x, 1)\n        x = self.fc(x)\n        return x\n# 設置設備\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n# 加載預訓練模型\npretrained_model = models.resnet50(weights=None)\nmodel = WaveCNNAttenClasifier(pretrained_model,device,CONFIG[\"model_config\"],CONFIG[\"model_config\"])\n\n# 加載預訓練權重\ncheckpoint_path = \"/kaggle/input/test-model/converted_model.pth\"\nstate_dict = torch.load(checkpoint_path, map_location=torch.device('cpu'))\n\n# 加載權重\nmissing_keys, unexpected_keys = model.load_state_dict(state_dict, strict=False)\nprint(\"Missing keys:\", missing_keys)\nprint(\"Unexpected keys:\", unexpected_keys)\n\n# 設置模型為推理模式\nmodel.eval()\n\n# 模擬輸入數據\ndummy_input = torch.randn(1, 160000).to(device)  # 根據模型的輸入尺寸進行修改\n\n# 定義 ONNX 模型的保存路徑\nonnx_dir = \"/kaggle/working/model\"\nonnx_path = os.path.join(onnx_dir, \"model_simpl.onnx\")\n\n# 確保目標目錄存在\nos.makedirs(onnx_dir, exist_ok=True)\n\n# 將模型轉換為 ONNX 格式\ntorch.onnx.export(model, dummy_input, onnx_path, opset_version=11,\n                  input_names=['input'], output_names=['output'])\n\nprint(f\"Model has been converted to ONNX and saved at {onnx_path}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-07-17T08:13:49.786257Z","iopub.execute_input":"2024-07-17T08:13:49.787079Z","iopub.status.idle":"2024-07-17T08:13:50.205986Z","shell.execute_reply.started":"2024-07-17T08:13:49.787038Z","shell.execute_reply":"2024-07-17T08:13:50.204572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"將 PyTorch 模型轉換為 ONNX 模型：","metadata":{}},{"cell_type":"markdown","source":"加載並檢查 ONNX 模型：","metadata":{}},{"cell_type":"code","source":"from pprint import pprint\nimport onnxruntime\nimport os\n\nonnx_path = \"/kaggle/working/model/model_simpl.onnx\"\n\n# 检查文件是否存在\nif not os.path.exists(onnx_path):\n    print(f\"文件路径不存在: {onnx_path}\")\nelse:\n    provider = \"CPUExecutionProvider\"\n    onnx_session = onnxruntime.InferenceSession(onnx_path, providers=[provider])\n\n    print(\"----------------- 输入部分 -----------------\")\n    input_tensors = onnx_session.get_inputs()  # 该 API 会返回列表\n    for input_tensor in input_tensors:         # 因为可能有多个输入，所以为列表\n        input_info = {\n            \"name\" : input_tensor.name,\n            \"type\" : input_tensor.type,\n            \"shape\": input_tensor.shape,\n        }\n        pprint(input_info)\n\n    print(\"----------------- 输出部分 -----------------\")\n    output_tensors = onnx_session.get_outputs()  # 该 API 会返回列表\n    for output_tensor in output_tensors:         # 因为可能有多个输出，所以为列表\n        output_info = {\n            \"name\" : output_tensor.name,\n            \"type\" : output_tensor.type,\n            \"shape\": output_tensor.shape,\n        }\n        pprint(output_info)\n","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:25.88613Z","iopub.execute_input":"2024-07-12T08:54:25.886427Z","iopub.status.idle":"2024-07-12T08:54:26.153761Z","shell.execute_reply.started":"2024-07-12T08:54:25.886397Z","shell.execute_reply":"2024-07-12T08:54:26.152622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n#看模型節點的\nimport onnx\n\n# 加载ONNX模型\nonnx_model_path = \"/kaggle/working/model/model_simpl.onnx\"\nonnx_model = onnx.load(onnx_model_path)\n\n# 打印所有节点的名称和类型\nfor node in onnx_model.graph.node:\n    print(f\"Node name: {node.name}, Node type: {node.op_type}\")\n    for attr in node.attribute:\n        print(f\"  Attribute name: {attr.name}, Attribute value: {attr}\")","metadata":{"_kg_hide-output":false,"_kg_hide-input":false,"execution":{"iopub.status.busy":"2024-07-12T08:54:26.156192Z","iopub.execute_input":"2024-07-12T08:54:26.15656Z","iopub.status.idle":"2024-07-12T08:54:26.259715Z","shell.execute_reply.started":"2024-07-12T08:54:26.156515Z","shell.execute_reply":"2024-07-12T08:54:26.258575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import onnx\nimport onnxruntime as ort\nimport torch\nimport torch.nn as nn\n\n# 加载 ONNX 模型\nonnx_path = f\"/kaggle/input/bird-clef-2023-models/{CONFIG['exp_name']}/{CONFIG['exp_name']}/checkpoints/model_simpl.onnx\"\nonnx_model = onnx.load(onnx_path)\n\n# 使用 onnxruntime 加载模型并查看输入输出信息\nonnx_session = ort.InferenceSession(onnx_path)\ninput_info = onnx_session.get_inputs()\noutput_info = onnx_session.get_outputs()\n\n# 创建一个空的 PyTorch 模型\nclass DynamicModule(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n# 遍历 ONNX 模型的节点来构建 PyTorch 模型\ntorch_model = DynamicModule()\nfor node in onnx_model.graph.node:\n    # 对每个节点进行适当的转换\n    # 示例：如果是 Conv 节点，添加对应的 PyTorch Conv 操作到 torch_model 中\n    # 这里需要根据您模型的具体结构来实现转换\n    pass  # 暂时保留空语句，用于遍历节点后的实现\n\n# 可以根据实际模型结构，逐步实现对应节点的转换\n\n# 打印 PyTorch 模型的结构\nprint(torch_model)\n","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:26.261215Z","iopub.execute_input":"2024-07-12T08:54:26.261683Z","iopub.status.idle":"2024-07-12T08:54:27.908354Z","shell.execute_reply.started":"2024-07-12T08:54:26.261633Z","shell.execute_reply":"2024-07-12T08:54:27.907205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# *數據準備*","metadata":{}},{"cell_type":"code","source":"class BirdAudioDataset(Dataset):\n    def __init__(self, annotations_file, audio_base_dir, transform=None):\n        self.audio_labels = pd.read_csv(annotations_file, encoding='latin1')\n        self.audio_base_dir = audio_base_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.audio_labels)\n\n    def __getitem__(self, idx):\n        label = self.audio_labels.iloc[idx, 0]\n        audio_file = self.audio_labels.iloc[idx, -1]\n        audio_path = os.path.join(self.audio_base_dir, audio_file)\n        y, sr = librosa.load(audio_path, sr=None)\n        melspec = librosa.feature.melspectrogram(y=y, sr=sr, n_mels=128)\n        melspec_db = librosa.power_to_db(melspec, ref=np.max)\n\n        fig = plt.figure(figsize=(2.24, 2.24))\n        canvas = FigureCanvas(fig)\n        plt.imshow(melspec_db, aspect='auto', origin='lower')\n        plt.axis('off')\n        canvas.draw()\n        image = np.frombuffer(canvas.tostring_rgb(), dtype='uint8')\n        image = image.reshape(fig.canvas.get_width_height()[::-1] + (3,))\n        plt.close(fig)\n\n        image = Image.fromarray(image)\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\n\ntrain_dataset = BirdAudioDataset(annotations_file='/kaggle/input/newbirdset/bird3/train_metadata.csv', audio_base_dir='/kaggle/input/newbirdset/bird3/train_data', transform=transform)\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:27.909691Z","iopub.execute_input":"2024-07-12T08:54:27.909986Z","iopub.status.idle":"2024-07-12T08:54:28.044128Z","shell.execute_reply.started":"2024-07-12T08:54:27.909957Z","shell.execute_reply":"2024-07-12T08:54:28.042906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 模型微調","metadata":{}},{"cell_type":"code","source":"class CustomResNet(nn.Module):\n    def __init__(self, pretrained_model):\n        super(CustomResNet, self).__init__()\n        self.features = nn.Sequential(*list(pretrained_model.children())[:-1])  # 移除原始的全連接層\n        self.fc = nn.Linear(2048, 264)  # 假設有264個類別\n\n    def forward(self, x):\n        x = self.features(x)\n        x = torch.flatten(x, 1)\n        x = self.fc(x)\n        return x\n\n# 加載預訓練模型\npretrained_model = models.resnet50(weights=None)\nmodel = CustomResNet(pretrained_model)\n\n# 加載預訓練權重\ncheckpoint_path = \"/kaggle/input/test-model/converted_model.pth\"\nstate_dict = torch.load(checkpoint_path, map_location=torch.device('cpu'))\n\n# 加載權重\nmissing_keys, unexpected_keys = model.load_state_dict(state_dict, strict=False)\nprint(\"Missing keys:\", missing_keys)\nprint(\"Unexpected keys:\", unexpected_keys)\n\n# 設置設備\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n\n# 設置模型為推理模式\nmodel.eval()\n\n# 模擬輸入數據\ndummy_input = torch.randn(1, 3, 224, 224).to(device)  # 根據模型的輸入尺寸進行修改\n\n# 定義 ONNX 模型的保存路徑\nonnx_dir = \"/kaggle/working/model\"\nonnx_path = os.path.join(onnx_dir, \"model_simpl.onnx\")\n\n# 確保目標目錄存在\nos.makedirs(onnx_dir, exist_ok=True)\n\n# 將模型轉換為 ONNX 格式\ntorch.onnx.export(model, dummy_input, onnx_path, opset_version=11,\n                  input_names=['input'], output_names=['output'])\n\nprint(f\"Model has been converted to ONNX and saved at {onnx_path}\")","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:28.045729Z","iopub.execute_input":"2024-07-12T08:54:28.046113Z","iopub.status.idle":"2024-07-12T08:54:30.591536Z","shell.execute_reply.started":"2024-07-12T08:54:28.04607Z","shell.execute_reply":"2024-07-12T08:54:30.590393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 訓練循環","metadata":{}},{"cell_type":"code","source":"\"\"\"num_epochs = 10\n\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    for inputs, labels in train_loader:\n        inputs, labels = inputs.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n    print(f'Epoch [{epoch+1}/{num_epochs}], Loss: {running_loss/len(train_loader):.4f}')\"\"\"\n","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.59274Z","iopub.execute_input":"2024-07-12T08:54:30.593095Z","iopub.status.idle":"2024-07-12T08:54:30.602336Z","shell.execute_reply.started":"2024-07-12T08:54:30.593063Z","shell.execute_reply":"2024-07-12T08:54:30.601047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 評估和保存","metadata":{}},{"cell_type":"code","source":"\"\"\"torch.save(model.state_dict(), 'fine_tuned_model.pth')\"\"\"","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.603621Z","iopub.execute_input":"2024-07-12T08:54:30.603956Z","iopub.status.idle":"2024-07-12T08:54:30.618952Z","shell.execute_reply.started":"2024-07-12T08:54:30.603922Z","shell.execute_reply":"2024-07-12T08:54:30.617706Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference Class","metadata":{}},{"cell_type":"code","source":" #音频预处理函数\ndef preprocess_audio(audio_path, segment_len=CONFIG[\"segment_len\"], sr=32000):\n    y, _ = librosa.load(audio_path, sr=sr)\n    segments = []\n    for i in range(0, len(y), segment_len * sr):\n        segment = y[i:i + segment_len * sr]\n        if len(segment) == segment_len * sr:\n            segments.append(segment)\n    return np.array(segments)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.620262Z","iopub.execute_input":"2024-07-12T08:54:30.62067Z","iopub.status.idle":"2024-07-12T08:54:30.630188Z","shell.execute_reply.started":"2024-07-12T08:54:30.620636Z","shell.execute_reply":"2024-07-12T08:54:30.629131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inference_class = BirdsInference(\n    device=\"gpu\",\n    verbose_tqdm=True,\n    use_sigmoid=CONFIG[\"use_sigmoid\"],\n#     model_output_key=CONFIG[\"model_output_key\"],\n)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.631522Z","iopub.execute_input":"2024-07-12T08:54:30.632202Z","iopub.status.idle":"2024-07-12T08:54:30.640189Z","shell.execute_reply.started":"2024-07-12T08:54:30.632155Z","shell.execute_reply":"2024-07-12T08:54:30.63931Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Val Predict","metadata":{}},{"cell_type":"code","source":"if CONFIG[\"run_validation\"]:\n    val_tgts, val_preds, val_preds_long = inference_class.predict_val_loaders(\n        nn_models=model,\n        data_loaders=loader_val\n    )\n    print(\n        f\"Min prob {val_preds.min()}. Max prob {val_preds.max()}.\\n\"\n        f\"Padded CMAP: {padded_cmap_numpy(val_tgts, val_preds)}\"\n    )","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.641315Z","iopub.execute_input":"2024-07-12T08:54:30.64236Z","iopub.status.idle":"2024-07-12T08:54:30.651913Z","shell.execute_reply.started":"2024-07-12T08:54:30.642312Z","shell.execute_reply":"2024-07-12T08:54:30.65077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test Pred","metadata":{}},{"cell_type":"code","source":"#if CONFIG[\"run_test\"]:\n#    test_preds, test_preds_long, test_dfidx, test_end = inference_class.predict_test_loader(\n#        nn_models=model,\n#        data_loader=loader_test,\n#        is_onnx_model=True\n#    )\n#    test_pred_df = compose_submission_dataframe(\n#        probs=test_preds,\n#        dfidxs=test_dfidx,\n#        end_seconds=test_end,\n #       filenames=loader_test.dataset.df[loader_test.dataset.name_col].copy(),\n #       bird2id=bird2id,\n#    )\n#    sample_submission = pd.read_csv(\"/kaggle/input/birdclef-2023/sample_submission.csv\")\n#    if test_pred_df.shape[1] > sample_submission.shape[1]:\n#        print(\"Shrinking columns\")\n#        test_pred_df = test_pred_df[sample_submission.columns]","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2024-07-12T08:54:30.65362Z","iopub.execute_input":"2024-07-12T08:54:30.654125Z","iopub.status.idle":"2024-07-12T08:54:30.662557Z","shell.execute_reply.started":"2024-07-12T08:54:30.654085Z","shell.execute_reply":"2024-07-12T08:54:30.661431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#test_pred_df.iloc[:,1:].values.max(), test_pred_df.iloc[:,1:].values.min()","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.663863Z","iopub.execute_input":"2024-07-12T08:54:30.664202Z","iopub.status.idle":"2024-07-12T08:54:30.678515Z","shell.execute_reply.started":"2024-07-12T08:54:30.664169Z","shell.execute_reply":"2024-07-12T08:54:30.67755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#plt.hist(test_pred_df.iloc[:,1:].values.max(axis=1), bins=30);","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.679722Z","iopub.execute_input":"2024-07-12T08:54:30.680027Z","iopub.status.idle":"2024-07-12T08:54:30.689412Z","shell.execute_reply.started":"2024-07-12T08:54:30.679998Z","shell.execute_reply":"2024-07-12T08:54:30.68824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"#使用再訓練模型預測\nmodel.load_state_dict(torch.load('fine_tuned_model.pth'))\nmodel.eval()\n\ndef predict(audio_path):\n    y, sr = librosa.load(audio_path, sr=None)\n    melspec = librosa.feature.melspectrogram(y=y, sr=sr, n_mels=128)\n    melspec_db = librosa.power_to_db(melspec, ref=np.max)\n    \n    fig = plt.figure(figsize=(2.24, 2.24))\n    canvas = FigureCanvas(fig)\n    plt.imshow(melspec_db, aspect='auto', origin='lower')\n    plt.axis('off')\n    canvas.draw()\n    image = np.frombuffer(canvas.tostring_rgb(), dtype='uint8')\n    image = image.reshape(fig.canvas.get_width_height()[::-1] + (3,))\n    plt.close(fig)\n    \n    image = Image.fromarray(image)\n    image = transform(image)\n    image = image.unsqueeze(0).to(device)\n    \n    with torch.no_grad():\n        outputs = model(image)\n        _, predicted = torch.max(outputs, 1)\n    return predicted.item()\n\n# 使用测试音频文件进行预测\ntest_audio_path = 'test_audio/test_file.wav'\nprediction = predict(test_audio_path)\nprint(f'Predicted class: {prediction}')\"\"\"\n","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.691016Z","iopub.execute_input":"2024-07-12T08:54:30.691431Z","iopub.status.idle":"2024-07-12T08:54:30.703073Z","shell.execute_reply.started":"2024-07-12T08:54:30.691388Z","shell.execute_reply":"2024-07-12T08:54:30.701906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 原onnx預測","metadata":{}},{"cell_type":"code","source":"# 推理函數\ndef predict_audio_class(audio_path):\n    segments = preprocess_audio(audio_path)\n    predictions = []\n    \n    for segment in segments:\n        input_data = segment[np.newaxis, :]  # 形狀應為 (1, num_samples)\n        ort_inputs = {model.get_inputs()[0].name: input_data}\n        ort_outs = model.run(None, ort_inputs)\n        predictions.append(ort_outs[0])\n    \n    # 平均所有片段的預測結果\n    avg_predictions = np.mean(predictions, axis=0)\n    predicted_class_idx = np.argmax(avg_predictions)\n    \n    # 反轉 bird2id 字典以獲取類別名稱\n    id2bird = {v: k for k, v in bird2id.items()}\n    predicted_class = id2bird[predicted_class_idx]\n    \n    # 計算預測的品種與每個品種的相似度\n    similarity_scores = {}\n    for bird_id, bird_name in id2bird.items():\n        similarity_scores[bird_name] = 1 - cosine(avg_predictions, np.eye(len(id2bird))[bird_id])\n    \n    # 找到最高相似度的品種\n    max_similarity_bird = max(similarity_scores, key=similarity_scores.get)\n    max_similarity_score = similarity_scores[max_similarity_bird]\n    \n    return predicted_class, avg_predictions, max_similarity_bird, max_similarity_score\n\n# 測試推理函數\naudio_path = \"/kaggle/input/birdtest/XC907019 h.ogg\"\npredicted_class, avg_predictions, max_similarity_bird, max_similarity_score = predict_audio_class(audio_path)\nprint(f\"The predicted class is: {predicted_class}\")\nprint(f\"The most similar class is: {max_similarity_bird} with a similarity score of: {max_similarity_score}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:30.7044Z","iopub.execute_input":"2024-07-12T08:54:30.7051Z","iopub.status.idle":"2024-07-12T08:54:37.657224Z","shell.execute_reply.started":"2024-07-12T08:54:30.705067Z","shell.execute_reply":"2024-07-12T08:54:37.655276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# pytorch預測","metadata":{}},{"cell_type":"code","source":"def predict_audio_class_pytorch(audio_path, model):\n    segments = preprocess_audio(audio_path)\n    predictions = []\n    \n    model.eval()\n    with torch.no_grad():\n        for segment in segments:\n            input_data = torch.tensor(segment[np.newaxis, np.newaxis, :], dtype=torch.float32)  # 添加額外維度以匹配模型輸入\n            output = model(input_data)\n            predictions.append(output.numpy())\n    \n    # 平均所有片段的預測結果\n    avg_predictions = np.mean(predictions, axis=0)\n    predicted_class_idx = np.argmax(avg_predictions)\n    \n    # 反轉 bird2id 字典以獲取類別名稱\n    id2bird = {v: k for k, v in bird2id.items()}\n    predicted_class = id2bird[predicted_class_idx]\n    \n    # 計算預測的品種與每個品種的相似度\n    similarity_scores = {}\n    for bird_id, bird_name in id2bird.items():\n        similarity_scores[bird_name] = 1 - cosine(avg_predictions, np.eye(len(id2bird))[bird_id])\n    \n    # 找到最高相似度的品種\n    max_similarity_bird = max(similarity_scores, key=similarity_scores.get)\n    max_similarity_score = similarity_scores[max_similarity_bird]\n    \n    return predicted_class, avg_predictions, max_similarity_bird, max_similarity_score\n\nclass CustomResNet(nn.Module):\n    def __init__(self):\n        super(CustomResNet, self).__init__()\n        self.features = nn.Sequential(\n            nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(kernel_size=3, stride=2, padding=1),\n            # 其他層...\n        )\n        self.fc = nn.Linear(2048, len(bird2id))  # 假設有多個類別\n\n    def forward(self, x):\n        x = self.features(x)\n        x = torch.flatten(x, 1)\n        x = self.fc(x)\n        return x\n\n# 加載PyTorch模型\npytorch_model_path = \"/kaggle/working/converted_model_full.pth\"\nmodelpytorch = torch.load(pytorch_model_path)\n\n# 測試模型預測\naudio_path = \"/kaggle/input/birdtest/XC907019 h.ogg\"\npredicted_class_pytorch, avg_predictions_pytorch, max_similarity_bird_pytorch, max_similarity_score_pytorch = predict_audio_class_pytorch(audio_path, modelpytorch)\nprint(f\"PyTorch model predicted class: {predicted_class_pytorch}\")\nprint(f\"PyTorch model most similar class: {max_similarity_bird_pytorch} with similarity score: {max_similarity_score_pytorch}\")","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.658523Z","iopub.status.idle":"2024-07-12T08:54:37.659055Z","shell.execute_reply.started":"2024-07-12T08:54:37.65878Z","shell.execute_reply":"2024-07-12T08:54:37.65881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Map Predictions","metadata":{}},{"cell_type":"code","source":"if CONFIG.get(\"check_inf_class\", False):\n    val_ds_check = WaveAllFileDataset(df=val_df[0], **{\n           \"root\": CONFIG[\"train_data_root\"],\n           \"label_str2int_mapping_path\": CONFIG[\"label_map_data_path\"],\n           \"n_cores\": 64,\n           \"use_audio_cache\": True,\n           \"test_mode\": True,\n           \"segment_len\": CONFIG[\"segment_len\"],\n           \"lookback\":CONFIG[\"lookback\"],\n           \"lookahead\":CONFIG[\"lookahead\"],\n            \"sample_id\": None,\n            \"late_normalize\": CONFIG[\"late_normalize\"],\n            \"step\": CONFIG[\"step\"],\n        }\n    )\n    val_loader_check = torch.utils.data.DataLoader(\n        val_ds_check,\n        **loader_config\n    )\n    \n    test_preds_check, test_preds_long_check, test_dfidx_check, test_end_check = inference_class.predict_test_loader(\n        nn_models=model,\n        data_loader=val_loader_check\n    )\n    test_preds_check_grouped = groupby_np_array(\n        groupby_f=test_dfidx_check,\n        array_to_group=test_preds_check,\n        apply_f=stack_and_max_by_samples,\n    )\n    print(np.allclose(\n        test_preds_check_grouped,\n        val_preds\n    ))","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.661386Z","iopub.status.idle":"2024-07-12T08:54:37.662055Z","shell.execute_reply.started":"2024-07-12T08:54:37.661727Z","shell.execute_reply":"2024-07-12T08:54:37.661756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Boost Probs","metadata":{}},{"cell_type":"code","source":"# CLASSES = list(set(test_pred_df.columns[1:]))\n# print(len(CLASSES))","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.664019Z","iopub.status.idle":"2024-07-12T08:54:37.664408Z","shell.execute_reply.started":"2024-07-12T08:54:37.664214Z","shell.execute_reply":"2024-07-12T08:54:37.664234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def connected_region_indices(bool_array):\n#     connected_regions = []\n#     region = []\n\n#     for i, value in enumerate(bool_array):\n#         if value:\n#             region.append(i)\n#         elif region:\n#             connected_regions.append(region)\n#             region = []\n\n#     if region:\n#         connected_regions.append(region)\n\n#     return connected_regions\n\n# def modify_probabilities(group, lower_prob=0.5, max_prob=0.75, min_seq_len=2):\n#     class_cols = ['shesta1', 'colsun2', 'amesun2', 'bcbeat1', 'marsun2']\n    \n#     mask_low = (group[CLASSES] > lower_prob).values\n#     mask_high = (group[CLASSES] > max_prob).values\n    \n#     # At least 2 chunks AND exceeds max_prob\n#     classes_to_boost = np.where((mask_low.sum(axis=0) >= min_seq_len) & mask_high.any(axis=0))[0]\n#     if len(classes_to_boost) > 0:\n#         for cls in classes_to_boost:\n#             new_probs = group[CLASSES[cls]].values.copy()\n#             connected_regions = connected_region_indices(mask_low[:,cls])\n#             for region in connected_regions:\n#                 if len(region) >= min_seq_len:\n#                     max_region_prob = group[CLASSES[cls]].iloc[region].max()\n#                     if max_region_prob > max_prob:\n#                         # print(f\"Boosting: {CLASSES[cls]}\")\n#                         new_probs[region] = max_region_prob\n#             group[CLASSES[cls]] = new_probs\n                \n#     return group","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.665794Z","iopub.status.idle":"2024-07-12T08:54:37.666182Z","shell.execute_reply.started":"2024-07-12T08:54:37.66597Z","shell.execute_reply":"2024-07-12T08:54:37.665989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_pred_df[\"id\"] = test_pred_df[\"row_id\"].apply(lambda x: \"_\".join(x.split(\"_\")[:-1]))\n# test_pred_df[\"sec\"] = test_pred_df[\"row_id\"].apply(lambda x: int(x.split(\"_\")[-1]))\n\n# test_pred_df = test_pred_df.sort_values([\"id\", \"sec\"]).reset_index(drop=True)\n\n# test_pred_df = test_pred_df.groupby(\"id\").apply(modify_probabilities).reset_index(drop=True)\n\n# test_pred_df = test_pred_df.drop(columns=[\"id\", \"sec\"])","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.667186Z","iopub.status.idle":"2024-07-12T08:54:37.667583Z","shell.execute_reply.started":"2024-07-12T08:54:37.667361Z","shell.execute_reply":"2024-07-12T08:54:37.667379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# np.where((test_pred_df_new[CLASSES].values != test_pred_df[CLASSES].values).sum(axis=0))\n# test_pred_df_new.loc[(test_pred_df_new[CLASSES[185]] - test_pred_df[CLASSES[185]]) > 0, CLASSES[185]]\n# test_pred_df.loc[(test_pred_df_new[CLASSES[185]] - test_pred_df[CLASSES[185]]) > 0, CLASSES[185]]","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.669697Z","iopub.status.idle":"2024-07-12T08:54:37.670232Z","shell.execute_reply.started":"2024-07-12T08:54:37.669954Z","shell.execute_reply":"2024-07-12T08:54:37.669983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def boost_pred_prob(\n#     input_df,\n#     class_columns,\n#     select_prob=0.8,\n#     boost_prob=0.2\n# ):\n#     result_df = input_df.copy()\n#     for sample_id in tqdm(set(input_df[\"id\"])):\n#         sample_id_pred_df = result_df.loc[result_df[\"id\"] == sample_id, class_columns]\n#         boost_classes = class_columns[sample_id_pred_df.values.max(axis=0) > select_prob]\n#         boost_mask = sample_id_pred_df[boost_classes].values < select_prob - boost_prob\n#         result_df.loc[result_df[\"id\"] == sample_id, boost_classes] = (\n#             (boost_mask.astype(np.float32) * (sample_id_pred_df[boost_classes].values + boost_prob)) +\n#             ((~boost_mask).astype(np.float32) * sample_id_pred_df[boost_classes].values)\n#         )\n#     return result_df\n\n# test_pred_df[\"id\"] = test_pred_df[\"row_id\"].apply(lambda x: \"_\".join(x.split(\"_\")[:-1]))\n# test_pred_df[\"sec\"] = test_pred_df[\"row_id\"].apply(lambda x: int(x.split(\"_\")[-1]))\n\n# test_pred_df = boost_pred_prob(\n#     input_df=test_pred_df,\n#     class_columns=test_pred_df.columns[1:-2]\n# )\n# test_pred_df = test_pred_df.drop(columns=[\"id\", \"sec\"])","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.67204Z","iopub.status.idle":"2024-07-12T08:54:37.672611Z","shell.execute_reply.started":"2024-07-12T08:54:37.672278Z","shell.execute_reply":"2024-07-12T08:54:37.672306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_pred_df.iloc[:,1:].values.max(), test_pred_df.iloc[:,1:].values.min()","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.674149Z","iopub.status.idle":"2024-07-12T08:54:37.674698Z","shell.execute_reply.started":"2024-07-12T08:54:37.67438Z","shell.execute_reply":"2024-07-12T08:54:37.674405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plt.hist(test_pred_df.iloc[:,1:].values.max(axis=1), bins=30);","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.67635Z","iopub.status.idle":"2024-07-12T08:54:37.676766Z","shell.execute_reply.started":"2024-07-12T08:54:37.676576Z","shell.execute_reply":"2024-07-12T08:54:37.676596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save Prediction","metadata":{}},{"cell_type":"code","source":"#sample_submission = pd.read_csv(\"/kaggle/input/birdclef-2023/sample_submission.csv\")\n#assert set(sample_submission.columns) == set(test_pred_df.columns)\n#test_pred_df = test_pred_df[sample_submission.columns]\n#test_pred_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2024-07-12T08:54:37.678583Z","iopub.status.idle":"2024-07-12T08:54:37.679132Z","shell.execute_reply.started":"2024-07-12T08:54:37.678836Z","shell.execute_reply":"2024-07-12T08:54:37.678864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}