{"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":"# Weighted Multi-Class Logarithmic Loss Function Implementation\n\nThe mathematical formula of loss function:\n\n$$ Log\\:Loss = - \\left( \\frac{\\sum_{i=1}^{M} w_{i} . \\sum_{j=1}^{N_{i}} \\frac{y_{ij}}{N_{i}}.\\ln p_{ij}}{\\sum_{i=1}^{M} w_{i}} \\right) $$\n\n**Note:** In order to avoid the extremes of the log function, each predicted probability 𝑝 is replaced with $max(min(\\rho,1−10^{-15}),10^{−15})$.\n\n**Where:** \n- ***Wi* = Class weights** \n- ***N* = Number of images in the batch** \n- ***M* = Number of classes**\n- ***Ln* = Natural logarithm** \n- ***Yij* = 1 if observation belongs to class *j* and 0 otherwise** \n- ***Pij* = Predicted probability that image *i* belongs to class *j***","metadata":{}},{"cell_type":"code","source":"import torch\n\ndef weighted_multiclass_log_loss(preds, labels, n_classes, weights):\n    preds = preds.to(device)\n    labels = labels.to(device)\n\n    labels_class_counts = torch.bincount(labels)\n    labels_class_counts[torch.where(labels_class_counts == 0)] = 1\n    labels_onehot = torch.eye(n_classes)[labels].to(device)\n    labels_scaled = torch.divide(labels_onehot, labels_class_counts)\n\n    preds_clamped = torch.clamp(preds.type(torch.DoubleTensor), min=10**-15, max=1-10**-15).to(device)\n    preds_sum = torch.sum(preds_clamped, dim=1).reshape(len(labels), 1)\n    preds_scaled = torch.divide(preds_clamped, preds_sum)\n    preds_log = torch.log(preds_scaled)\n\n    log_loss = -(((weights[0] * torch.sum(preds_log[:,0] * labels_scaled[:,0])) + (weights[1] * torch.sum(preds_log[:,1] * labels_scaled[:,1]))) / torch.sum(weights))\n\n    return log_loss","metadata":{"execution":{"iopub.status.busy":"2022-08-10T11:50:24.238347Z","iopub.execute_input":"2022-08-10T11:50:24.238791Z","iopub.status.idle":"2022-08-10T11:50:24.248886Z","shell.execute_reply.started":"2022-08-10T11:50:24.238754Z","shell.execute_reply":"2022-08-10T11:50:24.247889Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nn_classes = 2\nweights = torch.tensor([0.5, 0.5]).to(device)\ny_preds = torch.tensor([\n    [0.2, 0.8],\n    [0.9, 0.1],\n    [0.7, 0.3],\n    [0.1, 0.9],\n    [0.8, 0.2],\n    [0.5, 0.5],\n    [0.6, 0.4],\n    [0.4, 0.6]\n])\ny_true = torch.tensor([1, 0, 0, 1, 0, 0, 0, 0])\n\nloss = weighted_multiclass_log_loss(y_preds, y_true, n_classes, weights)\nprint('Loss: ', loss.item())","metadata":{"execution":{"iopub.status.busy":"2022-08-10T11:52:07.19572Z","iopub.execute_input":"2022-08-10T11:52:07.196657Z","iopub.status.idle":"2022-08-10T11:52:07.208101Z","shell.execute_reply.started":"2022-08-10T11:52:07.196607Z","shell.execute_reply":"2022-08-10T11:52:07.206678Z"},"trusted":true},"execution_count":null,"outputs":[]}]}