{"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":"gpu","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":4893,"sourceType":"modelInstanceVersion","modelInstanceId":2739}],"dockerImageVersionId":30446,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport numpy as np\nimport os\n\nimport tensorflow_hub as hub\nimport tensorflow as tf\n\nimport torchaudio\nimport torch\nfrom torch.utils.data import DataLoader, Dataset\n\n\ndf = pd.read_csv('/kaggle/input/birdclef-2024/train_metadata.csv')\nAUDIO_PATH = Path('/kaggle/input/birdclef-2024/train_audio')\nmodel_path = 'https://kaggle.com/models/google/bird-vocalization-classifier/frameworks/TensorFlow2/variations/bird-vocalization-classifier/versions/4'\nmodel = hub.load(model_path)\nmodel_labels_df = pd.read_csv(hub.resolve(model_path) + \"/assets/label.csv\")\n\nSAMPLE_RATE = 32000\nWINDOW = 5*SAMPLE_RATE","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:28:36.059643Z","iopub.execute_input":"2024-04-15T22:28:36.060419Z","iopub.status.idle":"2024-04-15T22:28:42.823291Z","shell.execute_reply.started":"2024-04-15T22:28:36.060364Z","shell.execute_reply":"2024-04-15T22:28:42.822250Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"index_to_label = sorted(df.primary_label.unique())\nlabel_to_index = {v: k for k, v in enumerate(index_to_label)}\nmodel_labels = {v: k for k, v in enumerate(model_labels_df.ebird2021)}\nmodel_bc_indexes = [model_labels[label] if label in model_labels else -1 for label in index_to_label]\n\n# filter out birds that the model doesn't predict\nmissing_birds = set(np.array(index_to_label)[np.array(model_bc_indexes) == -1])\nmissing_birds","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:46.540158Z","iopub.execute_input":"2024-04-15T22:30:46.541230Z","iopub.status.idle":"2024-04-15T22:30:46.559412Z","shell.execute_reply.started":"2024-04-15T22:30:46.541175Z","shell.execute_reply":"2024-04-15T22:30:46.558066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save embeddings and predictions for every 5 sec non-overlapping audio","metadata":{}},{"cell_type":"code","source":"# use a torch dataloader to decode audio in parallel on CPU while GPU is running\nclass AudioDataset(Dataset):\n    def __len__(self):\n        return len(df)\n    def __getitem__(self, i):\n        filename = df.filename[i]\n        audio = torchaudio.load(AUDIO_PATH / filename)[0].numpy()[0]\n        return audio, filename\ndataloader = DataLoader(AudioDataset(), batch_size=1, num_workers=os.cpu_count())\n\n\n# embeddings are formated like {\"filename\": np.array(nx1280)} \n# (where n = the number of non overlapping 5 sec chunks in the audio)\nall_embeddings = {}\n\n# predictiones formated like {\"filename\": np.array(nx264)} \nall_predictions = {}\n\nwith tf.device('/gpu:0'):\n    for audio, filename in tqdm(dataloader):\n        audio = audio[0]\n        filename = filename[0]\n        file_embeddings = []\n        file_predictions = []\n        for i in range(0, len(audio), WINDOW):\n            clip = audio[i:i+WINDOW]\n            if len(clip) < WINDOW:\n                clip = np.concatenate([clip, np.zeros(WINDOW - len(clip))])\n            result = model.infer_tf(clip[None, :])\n            file_embeddings.append(result[1][0].numpy())\n            prediction = np.concatenate([result[0].numpy(), -100], axis=None) # add -100 logit for unpredicted birds\n            file_predictions.append(prediction[model_bc_indexes])\n        all_embeddings[filename] = np.stack(file_embeddings)\n        all_predictions[filename] = np.stack(file_predictions)\n\ntorch.save(all_embeddings, 'embeddings.pt')\ntorch.save(all_predictions, 'predictions.pt')","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:28:48.578676Z","iopub.execute_input":"2024-04-15T22:28:48.579119Z","iopub.status.idle":"2024-04-15T22:30:08.611150Z","shell.execute_reply.started":"2024-04-15T22:28:48.579074Z","shell.execute_reply":"2024-04-15T22:30:08.610008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ROC-AUC (score for BirdClef 2024)","metadata":{}},{"cell_type":"code","source":"import sklearn.metrics\n\n'''\nThis script exists to reduce code duplication across metrics.\n'''\n\nimport numpy as np\nimport pandas as pd\nimport pandas.api.types\n\nfrom typing import Union\n\n\nclass ParticipantVisibleError(Exception):\n    pass\n\n\nclass HostVisibleError(Exception):\n    pass\n\n\ndef treat_as_participant_error(error_message: str, solution: Union[pd.DataFrame, np.ndarray]) -> bool:\n    ''' Many metrics can raise more errors than can be handled manually. This function attempts\n    to identify errors that can be treated as ParticipantVisibleError without leaking any competition data.\n\n    If the solution is purely numeric, and there are no numbers in the error message,\n    then the error message is sufficiently unlikely to leak usable data and can be shown to participants.\n\n    We expect this filter to reject many safe messages. It's intended only to reduce the number of errors we need to manage manually.\n    '''\n    # This check treats bools as numeric\n    if isinstance(solution, pd.DataFrame):\n        solution_is_all_numeric = all([pandas.api.types.is_numeric_dtype(x) for x in solution.dtypes.values])\n        solution_has_bools = any([pandas.api.types.is_bool_dtype(x) for x in solution.dtypes.values])\n    elif isinstance(solution, np.ndarray):\n        solution_is_all_numeric = pandas.api.types.is_numeric_dtype(solution)\n        solution_has_bools = pandas.api.types.is_bool_dtype(solution)\n\n    if not solution_is_all_numeric:\n        return False\n\n    for char in error_message:\n        if char.isnumeric():\n            return False\n    if solution_has_bools:\n        if 'true' in error_message.lower() or 'false' in error_message.lower():\n            return False\n    return True\n\n\ndef safe_call_score(metric_function, solution, submission, **metric_func_kwargs):\n    '''\n    Call score. If that raises an error and that already been specifically handled, just raise it.\n    Otherwise make a conservative attempt to identify potential participant visible errors.\n    '''\n    try:\n        score_result = metric_function(solution, submission, **metric_func_kwargs)\n    except Exception as err:\n        error_message = str(err)\n        if err.__class__.__name__ == 'ParticipantVisibleError':\n            raise ParticipantVisibleError(error_message)\n        elif err.__class__.__name__ == 'HostVisibleError':\n            raise HostVisibleError(error_message)\n        else:\n            if treat_as_participant_error(error_message, solution):\n                raise ParticipantVisibleError(error_message)\n            else:\n                raise err\n    return score_result\n\n\ndef verify_valid_probabilities(df: pd.DataFrame, df_name: str):\n    \"\"\" Verify that the dataframe contains valid probabilities.\n\n    The dataframe must be limited to the target columns; do not pass in any ID columns.\n    \"\"\"\n    if not pandas.api.types.is_numeric_dtype(df.values):\n        raise ParticipantVisibleError(f'All target values in {df_name} must be numeric')\n\n    if df.min().min() < 0:\n        raise ParticipantVisibleError(f'All target values in {df_name} must be at least zero')\n\n    if df.max().max() > 1:\n        raise ParticipantVisibleError(f'All target values in {df_name} must be no greater than one')\n\n    if not np.allclose(df.sum(axis=1), 1):\n        raise ParticipantVisibleError(f'Target values in {df_name} do not add to one within all rows')\n\n\ndef roc_auc(solution: pd.DataFrame, submission: pd.DataFrame) -> float: #row_id_column_name: str\n    '''\n    Version of macro-averaged ROC-AUC score that ignores all classes that have no true positive labels.\n    '''\n    #del solution[row_id_column_name]\n    #del submission[row_id_column_name]\n\n    if not pandas.api.types.is_numeric_dtype(submission.values):\n        bad_dtypes = {x: submission[x].dtype  for x in submission.columns if not pandas.api.types.is_numeric_dtype(submission[x])}\n        raise ParticipantVisibleError(f'Invalid submission data types found: {bad_dtypes}')\n\n    solution_sums = solution.sum(axis=0)\n    scored_columns = list(solution_sums[solution_sums > 0].index.values)\n    assert len(scored_columns) > 0\n\n    return safe_call_score(sklearn.metrics.roc_auc_score, solution[scored_columns].values, submission[scored_columns].values, average='macro')\n\n\n# test the score function\nx = pd.get_dummies(df.primary_label)\ny = x.copy()\ny[index_to_label] = 0\nroc_auc(x, x), roc_auc(x, y)","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.613596Z","iopub.execute_input":"2024-04-15T22:30:08.613972Z","iopub.status.idle":"2024-04-15T22:30:08.670191Z","shell.execute_reply.started":"2024-04-15T22:30:08.613928Z","shell.execute_reply":"2024-04-15T22:30:08.669138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Scores of predictions on the first 5 seconds of each recording","metadata":{}},{"cell_type":"code","source":"predicted_classes = torch.tensor([row[0].argmax() for row in all_predictions.values()])\nactual_classes = torch.tensor([label_to_index[label] for label in df.primary_label])\ncorrect = predicted_classes == actual_classes\naccuracy = correct.float().mean()\naccuracy","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.671378Z","iopub.execute_input":"2024-04-15T22:30:08.671733Z","iopub.status.idle":"2024-04-15T22:30:08.684292Z","shell.execute_reply.started":"2024-04-15T22:30:08.671702Z","shell.execute_reply":"2024-04-15T22:30:08.683242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logits = torch.stack([torch.tensor(row[0]) for row in all_predictions.values()])\nce_loss = torch.nn.CrossEntropyLoss()(logits, actual_classes)\nce_loss","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.686773Z","iopub.execute_input":"2024-04-15T22:30:08.687643Z","iopub.status.idle":"2024-04-15T22:30:08.707308Z","shell.execute_reply.started":"2024-04-15T22:30:08.687611Z","shell.execute_reply":"2024-04-15T22:30:08.706135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"actual_probs = torch.eye(len(index_to_label))[actual_classes]\nbce_loss = torch.nn.BCEWithLogitsLoss()(logits, actual_probs)\nbce_loss","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.708667Z","iopub.execute_input":"2024-04-15T22:30:08.708983Z","iopub.status.idle":"2024-04-15T22:30:08.717447Z","shell.execute_reply.started":"2024-04-15T22:30:08.708953Z","shell.execute_reply":"2024-04-15T22:30:08.716378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution = pd.DataFrame(actual_probs.numpy(), columns=index_to_label)\nroc_auc(\n    solution=solution,\n    submission=pd.DataFrame(torch.softmax(logits, 1).numpy(), columns=index_to_label),\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.718563Z","iopub.execute_input":"2024-04-15T22:30:08.718847Z","iopub.status.idle":"2024-04-15T22:30:08.741412Z","shell.execute_reply.started":"2024-04-15T22:30:08.718819Z","shell.execute_reply":"2024-04-15T22:30:08.740441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"roc_auc(\n    solution=solution,\n    submission=pd.DataFrame(torch.sigmoid(logits).numpy(), columns=index_to_label),\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.742610Z","iopub.execute_input":"2024-04-15T22:30:08.742995Z","iopub.status.idle":"2024-04-15T22:30:08.764087Z","shell.execute_reply.started":"2024-04-15T22:30:08.742955Z","shell.execute_reply":"2024-04-15T22:30:08.763152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Analysis","metadata":{}},{"cell_type":"code","source":"df['correct'] = correct.bool().numpy()","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.765309Z","iopub.execute_input":"2024-04-15T22:30:08.765636Z","iopub.status.idle":"2024-04-15T22:30:08.773465Z","shell.execute_reply.started":"2024-04-15T22:30:08.765607Z","shell.execute_reply":"2024-04-15T22:30:08.772582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import plotly.express as px\n\npx.histogram(\n    df,\n    title=\"Distribution of accuracy\",\n    x='primary_label',\n    color='correct',\n).update_xaxes(categoryorder=\"total descending\").show()","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.774599Z","iopub.execute_input":"2024-04-15T22:30:08.774873Z","iopub.status.idle":"2024-04-15T22:30:08.857442Z","shell.execute_reply.started":"2024-04-15T22:30:08.774846Z","shell.execute_reply":"2024-04-15T22:30:08.856232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Failure examples","metadata":{}},{"cell_type":"code","source":"from IPython.display import Audio\nimport torchaudio\nimport matplotlib.pyplot as plt\n\ncompute_melspec = torchaudio.transforms.MelSpectrogram(\n    sample_rate=SAMPLE_RATE,\n    n_mels=128,\n    n_fft=2048, \n    hop_length=512,\n    f_min=0,\n    f_max=SAMPLE_RATE // 2,\n)\n\npower_to_db = torchaudio.transforms.AmplitudeToDB(\n    stype=\"power\",\n    top_db=80.0,\n)\n\ndef show_bird(index, start=0):\n    audio = torchaudio.load(AUDIO_PATH / df.filename[index], start, start+WINDOW)[0][0]\n    display(df.iloc[index])\n    display(Audio(audio, rate=SAMPLE_RATE))\n    plt.figure(figsize=(12, 2.5))\n    plt.subplot(121)\n    plt.plot(audio)\n    plt.gca().get_xaxis().set_visible(False)\n    plt.subplot(122)\n    plt.imshow(power_to_db(compute_melspec(audio)))\n    plt.show()\n    return \n\n\nprobs = torch.sigmoid(logits)\n\nfor count, i in enumerate(df[~df.correct & ~df.primary_label.isin(missing_birds)].sample(10).index):\n    print('### EXAMPLE', count+1,'###')\n    correct_class = df.primary_label[i]\n    predicted_class_index = label_to_index[correct_class]\n    predicted_prob_for_correct_label = probs[i, predicted_class_index]\n    rank = 1 + sorted(probs[i], reverse=True).index(predicted_prob_for_correct_label)\n    \n    predicted_class_index = probs[i].argmax().item()\n    predicted_clas = index_to_label[predicted_class_index]\n    max_predicted_prob = probs[i][predicted_class_index]\n    \n    print(f'correct was {correct_class}, given prob {predicted_prob_for_correct_label} and ranked #{rank}')\n    print(f'predicted {predicted_clas}, with prob {max_predicted_prob}')\n    print()\n    show_bird(i)","metadata":{"execution":{"iopub.status.busy":"2024-04-15T22:30:08.860297Z","iopub.execute_input":"2024-04-15T22:30:08.860635Z","iopub.status.idle":"2024-04-15T22:30:34.496979Z","shell.execute_reply.started":"2024-04-15T22:30:08.860603Z","shell.execute_reply":"2024-04-15T22:30:34.495859Z"},"trusted":true},"execution_count":null,"outputs":[]}]}