{"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":"markdown","source":"# binary weighted log loss\n+ I've implemented the binary weighted log loss evaluator and set weights as mentioned in: https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/340392\n\n+ ~But I suspect my implementation is different from the true implementation. My evaluation function returns `4.8520` when all of the labels are false and all of the predictions are`0.5`. This result seems to be little bit worse than some published weak baselines.~\n\n+ UPDATE: I found that each weight must be normalized. For example, if given instance has weights of `C1~C7=2` and `overall_patient=7`, their normalized weights should be `C1~C7=2/(2*7+7)` and `overall_patient=7/(2*7+7)` respectively.","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch import nn\nimport torch.nn.functional as F\nimport pandas as pd\nimport numpy as np","metadata":{"execution":{"iopub.status.busy":"2022-10-14T17:52:36.662364Z","iopub.execute_input":"2022-10-14T17:52:36.663082Z","iopub.status.idle":"2022-10-14T17:52:36.668188Z","shell.execute_reply.started":"2022-10-14T17:52:36.663044Z","shell.execute_reply":"2022-10-14T17:52:36.667149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def binary_weighted_log_loss(y_hat_c, y_c, y_overall, y_hat_overall):\n    y_c = y_c.long()\n    y_overall = y_overall.long()\n    y_hat_c = torch.clamp(y_hat_c, min=1e-15, max=1.0-1e-7)\n    y_hat_overall = torch.clamp(y_hat_overall, 1e-15, 1.0-1e-7)\n    vertebrae_rhs = y_c * torch.log(y_hat_c) + (1 - y_c) * torch.log(1.0 - y_hat_c)\n    vertebrae_lhs = 2 * y_c + 1 * (1 - y_c)\n    patient_rhs = y_overall * torch.log(y_hat_overall) + (1 - y_overall) * torch.log(1.0 - y_hat_overall)\n    patient_lhs = 14 * y_overall + 7 * (1 - y_overall)\n    z = torch.sum(torch.cat([torch.sum(vertebrae_lhs, 1).view(-1,1), patient_lhs.view(-1,1)], 1), axis=1)\n    a = torch.cat([vertebrae_lhs * vertebrae_rhs, (patient_lhs * patient_rhs).view(-1, 1)], 1)\n    return -torch.mean(torch.sum(a, 1) * 1/z)\nprint(binary_weighted_log_loss(torch.tensor([[0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0]]).cuda(), torch.tensor([[0,0,0,0,0,0,0]]).cuda(), torch.tensor([0]).cuda(), torch.tensor([0.0]).cuda()))\nprint(binary_weighted_log_loss(torch.tensor([[0.5, 0.5, 0.5, 0.5, 0.5, 0.5, 0.5]]).cuda(), torch.tensor([[0,0,0,0,0,0,0]]).cuda(), torch.tensor([0]).cuda(), torch.tensor([0.5]).cuda()))\nprint(binary_weighted_log_loss(torch.tensor([[1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0]]).cuda(), torch.tensor([[0,0,0,0,0,0,0]]).cuda(), torch.tensor([0]).cuda(), torch.tensor([1.0]).cuda()))","metadata":{"execution":{"iopub.status.busy":"2022-10-14T17:52:36.676819Z","iopub.execute_input":"2022-10-14T17:52:36.677081Z","iopub.status.idle":"2022-10-14T17:52:36.696234Z","shell.execute_reply.started":"2022-10-14T17:52:36.677055Z","shell.execute_reply":"2022-10-14T17:52:36.695088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/rsna-2022-cervical-spine-fracture-detection/train.csv')\ntargets = ['patient_overall', 'C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']\nlabels = torch.tensor(train[targets].values)\nmean_values = train.mean(axis=0).values\nmean_values = torch.tensor(np.vstack([mean_values]*len(train)))","metadata":{"execution":{"iopub.status.busy":"2022-10-14T17:52:36.698632Z","iopub.execute_input":"2022-10-14T17:52:36.699265Z","iopub.status.idle":"2022-10-14T17:52:36.718425Z","shell.execute_reply.started":"2022-10-14T17:52:36.699228Z","shell.execute_reply":"2022-10-14T17:52:36.717145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(mean_values[:,1:])\n# print(labels[:,1:])\n# print(labels[:,0])\n# print(mean_values[:,0])\nbinary_weighted_log_loss(mean_values[:,1:], labels[:, 1:], labels[:,0], mean_values[:,0])","metadata":{"execution":{"iopub.status.busy":"2022-10-14T17:52:36.720473Z","iopub.execute_input":"2022-10-14T17:52:36.721131Z","iopub.status.idle":"2022-10-14T17:52:36.730226Z","shell.execute_reply.started":"2022-10-14T17:52:36.721089Z","shell.execute_reply":"2022-10-14T17:52:36.729241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}