{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load in \n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the \"../input/\" directory.\n# For example, running this (by clicking run or pressing Shift+Enter) will list the files in the input directory\n\nimport os\nprint(os.listdir(\"../input\"))\n\n# Any results you write to the current directory are saved as output.","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"4689837fe255ce2294eb7be25586c786e9b8ede1"},"cell_type":"markdown","source":"## Abstract\nI implemented customized loss functions (focal loss and class imbalance loss)  to handle class imbalanced situations inspired by https://arxiv.org/abs/1708.02002.\n\nRight now I only have numpy/Pytorch implementation, but I will add Keras/Tensorflow implementation later."},{"metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","trusted":true},"cell_type":"code","source":"class ImbalancedClassesNumpy:\n    def __init__(self,y):\n        #y has to be binary vectors\n        num_negative = y[y == 0].shape\n        num_positive = y[y == 1].shape\n        self.alpha = num_negative / (num_negative + num_positive)\n        self.gamma = 2\n    def _return_pt(self,target,pred):\n        return np.add(np.multiply(np.subtract(1, target) , np.subtract(1 , pred)) ,np.multiply( target , pred))\n    \n    def class_imbalanced_loss(self,target,pred,alpha = None):\n        if alpha is None:\n            alpha = self.alpha\n        pt = self._return_pt(target,pred)\n        return - alpha * np.sum(np.log(pt))\n    \n    def forcal_loss(self,target,pred,alpha = None):\n        if alpha is None:\n            alpha = self.alpha\n        pt = self._return_pt(target,pred)\n        return - alpha * np.sum(np.multiply(np.subtract(1 , pt) ** gamma,np.log(pt)))\n    ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"3b94151481f93d1c4bb2532dedfc61eb2abaea25"},"cell_type":"code","source":"import torch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"73d7bb5c4e25036825430bf59b72c827b8660bcb"},"cell_type":"code","source":"ten_a = torch.zeros([2,3],dtype=torch.int32)\nten_b = torch.tensor([[1,0,3],[2,0,1]],dtype=torch.int32)\ntorch.masked_select(ten_b,torch.eq(ten_a,ten_b)).shape[0]\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"6bb5fafd96cf2eb6b2fcf8a05ad9e59b58a5b140"},"cell_type":"code","source":"torch.zeros(ten_a.size())","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"aa5072beab4ae993bfcc1ae035ec81f9dc0f037f"},"cell_type":"code","source":"class ImbalancedClassPytorch:\n    def __init__(self,y):\n        num_negative = torch.masked_select(y,torch.eq(y,torch.zeros(y.size(),dtype = torch.uint8))).shape[0]\n        #num_negative = y[y == 0].shape\n        num_positive = torch.masked_select(y,torch.eq(y,torch.ones(y.size(),dtype = torch.uint8))).shape[0]\n        self.alpha = num_negative / (num_negative + num_positive)\n    def _return_pt(self,target,pred):\n        pt = torch.add(torch.mul(torch.sub(1, target) , torch.sub(1 , pred)) ,torch.mul( target , pred))\n        return pt\n    def class_imbalanced_loss(self,target,pred,alpha = None):\n        if alpha is None:\n            alpha = self.alpha\n        pt = self._return_pt(target,pred)\n        return - alpha * torch.sum(torch.mul(torch.sub(1 , pt),torch.log(pt)))\n    def focal_loss(self,target,pred,alpha = None,gamma = 2):\n        if alpha is None:\n            alpha = self.alpha\n        pt = self._return_pt(target,pred)\n        return - alpha * torch.sum(torch.mul(torch.sub(1 , pt).pow(gamma),torch.log(pt)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_uuid":"e599d4de238caed2b1f35e45dd9534cc8befa2d6"},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.6.4","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat":4,"nbformat_minor":1}