{"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":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2023-06-27T22:38:21.257240Z","iopub.execute_input":"2023-06-27T22:38:21.257738Z","iopub.status.idle":"2023-06-27T22:38:22.404312Z","shell.execute_reply.started":"2023-06-27T22:38:21.257696Z","shell.execute_reply":"2023-06-27T22:38:22.402450Z"},"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":"2023-06-27T22:38:22.408215Z","iopub.execute_input":"2023-06-27T22:38:22.408769Z","iopub.status.idle":"2023-06-27T22:38:46.769974Z","shell.execute_reply.started":"2023-06-27T22:38:22.408698Z","shell.execute_reply":"2023-06-27T22:38:46.768322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip list | grep onnx","metadata":{"execution":{"iopub.status.busy":"2023-06-27T22:38:46.778233Z","iopub.execute_input":"2023-06-27T22:38:46.778632Z","iopub.status.idle":"2023-06-27T22:39:11.640468Z","shell.execute_reply.started":"2023-06-27T22:38:46.778583Z","shell.execute_reply":"2023-06-27T22:39:11.638487Z"},"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/')","metadata":{"execution":{"iopub.status.busy":"2023-06-27T22:39:11.644002Z","iopub.execute_input":"2023-06-27T22:39:11.644503Z","iopub.status.idle":"2023-06-27T22:39:11.651438Z","shell.execute_reply.started":"2023-06-27T22:39:11.644461Z","shell.execute_reply":"2023-06-27T22:39:11.650215Z"},"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 h5py\nimport onnxruntime as ort\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\n\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\n","metadata":{"execution":{"iopub.status.busy":"2023-06-27T22:39:11.653519Z","iopub.execute_input":"2023-06-27T22:39:11.654061Z","iopub.status.idle":"2023-06-27T22:39:18.270657Z","shell.execute_reply.started":"2023-06-27T22:39:11.654017Z","shell.execute_reply":"2023-06-27T22:39:18.269579Z"},"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":"2023-06-27T22:39:18.272389Z","iopub.execute_input":"2023-06-27T22:39:18.273505Z","iopub.status.idle":"2023-06-27T22:39:18.289962Z","shell.execute_reply.started":"2023-06-27T22:39:18.273463Z","shell.execute_reply":"2023-06-27T22:39:18.288369Z"},"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/bird-clef-2023-addones/fake_test_20/fake_test_20/*.ogg\",\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#     \"label_map_data_path\":'/kaggle/input/bird-clef-2023-models/bird221int_202x.json',\n#     \"label_map_data_path\": \"/kaggle/input/bird-clef-2023-models/xc_birds_202x_only_scored.json\",\n#     \"label_map_data_path\": '/kaggle/input/bird-clef-2023-models/bird2id_xc_pretrain.json',\n#     \"label_map_data_path\": \"/kaggle/input/bird-clef-2023-models/bird2id.json\",\n#     \"label_map_data_path\":\"/kaggle/input/bird-clef-2023-models/bird2id_v1.json\",\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#         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":"2023-06-27T22:39:18.292352Z","iopub.execute_input":"2023-06-27T22:39:18.292786Z","iopub.status.idle":"2023-06-27T22:39:18.304877Z","shell.execute_reply.started":"2023-06-27T22:39:18.292739Z","shell.execute_reply":"2023-06-27T22:39:18.303488Z"},"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":"2023-06-27T22:39:18.306742Z","iopub.execute_input":"2023-06-27T22:39:18.307546Z","iopub.status.idle":"2023-06-27T22:39:18.327665Z","shell.execute_reply.started":"2023-06-27T22:39:18.307506Z","shell.execute_reply":"2023-06-27T22:39:18.326185Z"},"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":"2023-06-27T22:39:18.334931Z","iopub.execute_input":"2023-06-27T22:39:18.335707Z","iopub.status.idle":"2023-06-27T22:39:18.342488Z","shell.execute_reply.started":"2023-06-27T22:39:18.335663Z","shell.execute_reply":"2023-06-27T22:39:18.341246Z"},"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(filename=el) for el in test_au_pathes]\n    })","metadata":{"execution":{"iopub.status.busy":"2023-06-27T22:39:18.344109Z","iopub.execute_input":"2023-06-27T22:39:18.344827Z","iopub.status.idle":"2023-06-27T22:39:24.493961Z","shell.execute_reply.started":"2023-06-27T22:39:18.344776Z","shell.execute_reply":"2023-06-27T22:39:24.492416Z"},"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":"2023-06-27T22:39:24.496298Z","iopub.execute_input":"2023-06-27T22:39:24.497808Z","iopub.status.idle":"2023-06-27T22:39:24.509655Z","shell.execute_reply.started":"2023-06-27T22:39:24.497764Z","shell.execute_reply":"2023-06-27T22:39:24.507936Z"},"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":"2023-06-27T22:39:24.511242Z","iopub.execute_input":"2023-06-27T22:39:24.511596Z","iopub.status.idle":"2023-06-27T22:39:24.535056Z","shell.execute_reply.started":"2023-06-27T22:39:24.511561Z","shell.execute_reply":"2023-06-27T22:39:24.533566Z"},"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":"2023-06-27T22:39:24.536909Z","iopub.execute_input":"2023-06-27T22:39:24.537970Z","iopub.status.idle":"2023-06-27T22:39:24.545038Z","shell.execute_reply.started":"2023-06-27T22:39:24.537916Z","shell.execute_reply":"2023-06-27T22:39:24.544066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"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":"2023-06-27T22:39:24.546366Z","iopub.execute_input":"2023-06-27T22:39:24.547171Z","iopub.status.idle":"2023-06-27T22:39:24.559242Z","shell.execute_reply.started":"2023-06-27T22:39:24.547133Z","shell.execute_reply":"2023-06-27T22:39:24.558096Z"},"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":"2023-06-27T22:39:24.561014Z","iopub.execute_input":"2023-06-27T22:39:24.562074Z","iopub.status.idle":"2023-06-27T22:39:24.577795Z","shell.execute_reply.started":"2023-06-27T22:39:24.562028Z","shell.execute_reply":"2023-06-27T22:39:24.576763Z"},"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":"2023-06-27T22:39:24.579486Z","iopub.execute_input":"2023-06-27T22:39:24.580128Z","iopub.status.idle":"2023-06-27T22:39:30.978848Z","shell.execute_reply.started":"2023-06-27T22:39:24.580090Z","shell.execute_reply":"2023-06-27T22:39:30.977439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference Class","metadata":{}},{"cell_type":"code","source":"inference_class = BirdsInference(\n    device=\"cpu\",\n    verbose_tqdm=True,\n    use_sigmoid=CONFIG[\"use_sigmoid\"],\n#     model_output_key=CONFIG[\"model_output_key\"],\n)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T22:39:30.980765Z","iopub.execute_input":"2023-06-27T22:39:30.981947Z","iopub.status.idle":"2023-06-27T22:39:30.994455Z","shell.execute_reply.started":"2023-06-27T22:39:30.981902Z","shell.execute_reply":"2023-06-27T22:39:30.993159Z"},"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":"2023-06-27T22:39:30.996670Z","iopub.execute_input":"2023-06-27T22:39:30.997188Z","iopub.status.idle":"2023-06-27T22:39:31.010450Z","shell.execute_reply.started":"2023-06-27T22:39:30.997135Z","shell.execute_reply":"2023-06-27T22:39:31.008831Z"},"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":{"execution":{"iopub.status.busy":"2023-06-27T22:39:31.012149Z","iopub.execute_input":"2023-06-27T22:39:31.012660Z","iopub.status.idle":"2023-06-27T22:40:21.196405Z","shell.execute_reply.started":"2023-06-27T22:39:31.012605Z","shell.execute_reply":"2023-06-27T22:40:21.195106Z"},"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":"2023-06-27T22:40:21.198256Z","iopub.execute_input":"2023-06-27T22:40:21.199043Z","iopub.status.idle":"2023-06-27T22:40:21.213449Z","shell.execute_reply.started":"2023-06-27T22:40:21.199000Z","shell.execute_reply":"2023-06-27T22:40:21.212026Z"},"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":"2023-06-27T22:40:21.215050Z","iopub.execute_input":"2023-06-27T22:40:21.215446Z","iopub.status.idle":"2023-06-27T22:40:21.541924Z","shell.execute_reply.started":"2023-06-27T22:40:21.215408Z","shell.execute_reply":"2023-06-27T22:40:21.540406Z"},"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":"2023-06-27T22:40:21.543539Z","iopub.execute_input":"2023-06-27T22:40:21.544426Z","iopub.status.idle":"2023-06-27T22:40:21.557955Z","shell.execute_reply.started":"2023-06-27T22:40:21.544375Z","shell.execute_reply":"2023-06-27T22:40:21.555768Z"},"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":"2023-06-27T22:40:21.559672Z","iopub.execute_input":"2023-06-27T22:40:21.560105Z","iopub.status.idle":"2023-06-27T22:40:21.570331Z","shell.execute_reply.started":"2023-06-27T22:40:21.560065Z","shell.execute_reply":"2023-06-27T22:40:21.568717Z"},"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":"2023-06-27T22:40:21.572500Z","iopub.execute_input":"2023-06-27T22:40:21.573106Z","iopub.status.idle":"2023-06-27T22:40:21.589876Z","shell.execute_reply.started":"2023-06-27T22:40:21.573049Z","shell.execute_reply":"2023-06-27T22:40:21.588725Z"},"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":"2023-06-27T22:40:21.591281Z","iopub.execute_input":"2023-06-27T22:40:21.591682Z","iopub.status.idle":"2023-06-27T22:40:21.605462Z","shell.execute_reply.started":"2023-06-27T22:40:21.591622Z","shell.execute_reply":"2023-06-27T22:40:21.603948Z"},"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":"2023-06-27T22:40:21.607450Z","iopub.execute_input":"2023-06-27T22:40:21.607972Z","iopub.status.idle":"2023-06-27T22:40:21.621103Z","shell.execute_reply.started":"2023-06-27T22:40:21.607919Z","shell.execute_reply":"2023-06-27T22:40:21.619623Z"},"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":"2023-06-27T22:40:21.629259Z","iopub.execute_input":"2023-06-27T22:40:21.629649Z","iopub.status.idle":"2023-06-27T22:40:21.638748Z","shell.execute_reply.started":"2023-06-27T22:40:21.629614Z","shell.execute_reply":"2023-06-27T22:40:21.637702Z"},"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":"2023-06-27T22:40:21.640217Z","iopub.execute_input":"2023-06-27T22:40:21.640567Z","iopub.status.idle":"2023-06-27T22:40:21.652270Z","shell.execute_reply.started":"2023-06-27T22:40:21.640534Z","shell.execute_reply":"2023-06-27T22:40:21.651093Z"},"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":"2023-06-27T22:40:21.654143Z","iopub.execute_input":"2023-06-27T22:40:21.655402Z","iopub.status.idle":"2023-06-27T22:40:21.665638Z","shell.execute_reply.started":"2023-06-27T22:40:21.655345Z","shell.execute_reply":"2023-06-27T22:40:21.664368Z"},"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\")\nassert set(sample_submission.columns) == set(test_pred_df.columns)\ntest_pred_df = test_pred_df[sample_submission.columns]\ntest_pred_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-06-27T22:40:21.667752Z","iopub.execute_input":"2023-06-27T22:40:21.668663Z","iopub.status.idle":"2023-06-27T22:40:21.769006Z","shell.execute_reply.started":"2023-06-27T22:40:21.668603Z","shell.execute_reply":"2023-06-27T22:40:21.767741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}