{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nfrom torch.utils.data.sampler import BatchSampler, Sampler\nfrom torch.utils.data import DataLoader, Dataset\nfrom typing import Iterator, List, Optional, Union, Dict","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"class WeightedClassSampler(Sampler):\n    \"\"\"Abstraction over data sampler.\n\n    Allows you to create stratified sample on unbalanced classes.\n    \"\"\"\n\n    def __init__(\n        self, labels: List[int], class_weights: Dict \n    ):\n        \"\"\"\n        Args:\n            labels (List[int]): list of class label\n                for each elem in the dataset\n            class_weights (dict): give dict of class\n                how may time its repete\n        \"\"\"\n        super().__init__(labels)\n\n        labels = np.array(labels)\n        samples_per_class = {\n            label: int((labels == label).sum() * class_weights[label]) for label in set(labels)\n        }\n        \n\n\n        self.lbl2idx = {\n            label: np.arange(len(labels))[labels == label].tolist()\n            for label in set(labels)\n        }\n        \n        self.length = sum(samples_per_class.values())\n        self.labels = labels\n        self.samples_per_class = samples_per_class\n        self.class_weights = class_weights\n\n\n    def __iter__(self) -> Iterator[int]:\n        \"\"\"\n        Yields:\n            indices of stratified sample\n        \"\"\"\n        indices = []\n        for key in sorted(self.lbl2idx):\n            replace_flag = self.class_weights[key] > 1\n            indices += np.random.choice(\n                self.lbl2idx[key], self.samples_per_class[key], replace=replace_flag\n            ).tolist()\n        assert len(indices) == self.length\n        np.random.shuffle(indices)\n\n        return iter(indices)\n\n\n    def __len__(self) -> int:\n        \"\"\"\n        Returns:\n             length of result sample\n        \"\"\"\n        return self.length","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"### Example","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"train = pd.read_csv(\"../input/siim-isic-melanoma-classification/train.csv\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class BalanceDataset:\n    def __init__(self, df):\n        self.target = df.target.values\n    \n    def __len__(self):\n        return len(self.target)\n    \n    def __getitem__(self, item):\n        return self.target[item]\n    \n    def __get_labels__(self):\n        return self.target","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"- `class_weights= {0:1, 1:10}` this means class 0 only 1time, class 1 repete 10 times","execution_count":null},{"metadata":{"trusted":true},"cell_type":"code","source":"dataset = BalanceDataset(df=train)\n\ndata_loader = DataLoader(\n    dataset,\n    sampler= WeightedClassSampler(labels=dataset.__get_labels__(), class_weights= {0:1, 1:10}),\n    batch_size=5\n)\n\nlen(data_loader)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for i, d in enumerate(data_loader):\n    print(d)\n    if i == 10: break","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}