{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"#pip install apex","execution_count":null,"outputs":[]},{"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\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 read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 5GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"#dataset.py\nimport cv2\nimport logging\nimport math\nimport numpy as np\nimport pandas as pd\nimport random\nfrom collections import defaultdict\nfrom itertools import chain\nfrom operator import itemgetter\nfrom pathlib import Path\n\nimport torch\nimport torchvision.transforms.functional as F\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\n\n\n\n\ndef tta(args, images):\n    \"\"\"Augment all images in a batch and return list of augmented batches\"\"\"\n\n    ret = []\n    n1 = math.ceil(args.tta ** 0.5)\n    n2 = math.ceil(args.tta / n1)\n    k = 0\n    for i in range(n1):\n        for j in range(n2):\n            if k >= args.tta:\n                break\n\n            dw = round(args.tta_size * images.size(2))\n            dh = round(args.tta_size * images.size(3))\n            w = i * (images.size(2) - dw) // max(n1 - 1, 1)\n            h = j * (images.size(3) - dh) // max(n2 - 1, 1)\n\n            imgs = images[:, :, w:w + dw, h:h + dh]\n            if k & 1:\n                imgs = imgs.flip(3)\n            if k & 2:\n                imgs = imgs.flip(2)\n            if k & 4:\n                imgs = imgs.transpose(2, 3)\n\n            ret.append(nn.functional.interpolate(imgs, images.size()[2:], mode='nearest'))\n            k += 1\n\n    return ret\n\ndef worker_init_fn(worker_id):\n    np.random.seed(random.randint(0, 10 ** 9) + worker_id)\n\ndef get_train_val_loader(args, predict=False):\n    def train_transform1(image):\n        if random.random() < 0.5:\n            image = image[:, ::-1, :]\n        if random.random() < 0.5:\n            image = image[::-1, :, :]\n        if random.random() < 0.5:\n            image = image.transpose([1, 0, 2])\n        image = np.ascontiguousarray(image)\n\n        if args.scale_aug != 1:\n            size = random.randint(round(512 * args.scale_aug), 512)\n            x = random.randint(0, 512 - size)\n            y = random.randint(0, 512 - size)\n            image = image[x:x + size, y:y + size]\n            image = cv2.resize(image, (512, 512), interpolation=cv2.INTER_NEAREST)\n\n        return image\n\n    def train_transform2(image):\n        a, b = np.random.normal(1, args.pw_aug[0], (6, 1, 1)), np.random.normal(0, args.pw_aug[1], (6, 1, 1))\n        a, b = torch.tensor(a, dtype=torch.float32), torch.tensor(b, dtype=torch.float32)\n        return image * a + b\n\n    if not predict:\n        train_dataset = CellularDataset(args.data, 'train_all_controls' if args.all_controls_train else 'train_controls',\n                transform=(train_transform1, train_transform2), cv_number=args.cv_number,\n                split_seed=args.data_split_seed, normalization=args.data_normalization)\n        train = DataLoader(train_dataset, args.batch_size, shuffle=True, drop_last=True,\n                num_workers=args.num_data_workers, worker_init_fn=worker_init_fn)\n\n    for i in range(1 if not predict else 2):\n        dataset = CellularDataset(args.data, 'val' if i == 0 else 'train', cv_number=args.cv_number,\n                split_seed=args.data_split_seed, normalization=args.data_normalization)\n        loader = DataLoader(dataset, args.batch_size, shuffle=False, num_workers=args.num_data_workers,\n                worker_init_fn=worker_init_fn)\n        if i == 0:\n            val = loader\n        else:\n            train = loader\n\n    assert len(set(train.dataset.data).intersection(set(val.dataset.data))) == 0\n    return train, val\n\ndef get_test_loader(args, exclude_leak=False):\n    test_dataset = CellularDataset(args.data, 'test' if not exclude_leak else 'test_noleak',\n            normalization=args.data_normalization)\n    return DataLoader(test_dataset, args.batch_size, shuffle=False, num_workers=args.num_data_workers,\n            worker_init_fn=worker_init_fn)\n\n\nclass CellularDataset(Dataset):\n    treatment_classes = 1108\n\n    def __init__(self, root_dir, mode, split_seed=0, cv_number=0, transform=None, normalization='global'):\n        \"\"\"\n        :param split_seed: seed for train/val split of labeled experiments and HUVEC-18\n        :param mode: possible choices:\n                        train -- dataset containing only non-control images from training set\n                        train_controls -- dataset containing non-control and control images from training set\n                        train_all_controls -- dataset containing non-control and control images from training set and\n                                              control images from validation and test set\n                        val -- dataset containing only non-control images from validation set\n                        test -- dataset containing only non-control images from test set\n                        test_noleak -- dataset containing only non-control images from test set excluding HUVEC-18\n        :param transform: tuple of 2 functions for image transformation. First is called right after loading with image\n                          in numpy format. Second is called after normalization and converting to tensor\n        \"\"\"\n\n        super().__init__()\n\n        self.root = Path(root_dir)\n        self.transform = transform\n\n        assert normalization in ['global', 'experiment', 'sample']\n        self.normalization = normalization\n\n        if mode == 'train_controls':\n            mode = 'train'\n            move_controls = True\n            all_controls = False\n        elif mode == 'train_all_controls':\n            mode = 'train'\n            move_controls = True\n            all_controls = True\n        else:\n            move_controls = False\n            all_controls = False\n\n        if mode == 'test_noleak':\n            mode = 'test'\n            exclude_leak = True\n        else:\n            exclude_leak = False\n        assert mode in ['train', 'val', 'test']\n        self.mode = mode\n\n        csv = pd.read_csv(self.root / ('train.csv' if mode in ['train', 'val'] else 'test.csv'))\n        csv_controls = pd.read_csv(self.root / ('train_controls.csv' if mode in ['train', 'val'] else 'test_controls.csv'))\n        if all_controls:\n            csv_controls_test = pd.read_csv(self.root / 'test_controls.csv')\n        self.data = []  # (experiment, plate, well, site, cell_type, sirna or None)\n        experiments = {}\n        for row in chain(csv.iterrows(), csv_controls.iterrows(), *([csv_controls_test.iterrows()] if all_controls else [])):\n            r = row[1]\n            typ = r.experiment[:r.experiment.find('-')]\n            self.data.append((r.experiment, r.plate, r.well, 1, typ, r.sirna if hasattr(r, 'sirna') else None))\n            self.data.append((r.experiment, r.plate, r.well, 2, typ, r.sirna if hasattr(r, 'sirna') else None))\n            if not hasattr(r, 'sirna') or r.sirna < self.treatment_classes:\n                if typ not in experiments:\n                    experiments[typ] = set()\n                experiments[typ].add(r.experiment)\n        if mode in ['train', 'val']:\n            data_dict = {(e, p, w): sir for e, p, w, s, typ, sir in self.data}\n            for row in pd.read_csv(self.root / 'test.csv').iterrows():\n                r = row[1]\n                typ = r.experiment[:r.experiment.find('-')]\n                if r.experiment == 'HUVEC-18':\n                    sirna = data_dict[('RPE-03', (r.plate - 2) % 4 + 1, r.well)]\n                    assert sirna < self.treatment_classes\n                    self.data.append((r.experiment, r.plate, r.well, 1, typ, sirna))\n                    self.data.append((r.experiment, r.plate, r.well, 2, typ, sirna))\n                    if typ not in experiments:\n                        experiments[typ] = set()\n                    experiments[typ].add(r.experiment)\n            if not all_controls:\n                for row in pd.read_csv(self.root / 'test_controls.csv').iterrows():\n                    r = row[1]\n                    typ = r.experiment[:r.experiment.find('-')]\n                    if r.experiment == 'HUVEC-18':\n                        sirna = data_dict[('RPE-03', (r.plate - 2) % 4 + 1, r.well)]\n                        assert sirna == r.sirna or sirna == 1138 or r.sirna == 1138\n                        self.data.append((r.experiment, r.plate, r.well, 1, typ, r.sirna))\n                        self.data.append((r.experiment, r.plate, r.well, 2, typ, r.sirna))\n        if exclude_leak:\n            self.data = list(filter(lambda x: x[0] != 'HUVEC-18', self.data))\n\n        self.cell_types = sorted(experiments.keys())\n        all_data = self.data.copy()\n\n        if mode != 'test':\n            state = random.Random(split_seed)\n            cells = list(map(itemgetter(1), sorted(experiments.items())))\n            for i in range(len(cells)):\n                cells[i] = sorted(cells[i])\n                if i == 3:\n                    cells[i] = cells[i] + cells[i]  # duplicate U2OS experiments for validation\n                state.shuffle(cells[i])\n\n            # cell[i] is a list of experiments for i-th cell type\n            assert list(map(len, cells)) == [7, 17, 7, 6]\n\n            # counts of experiments from given cell type for given fold\n            counts = [\n                [2, 2, 1, 1],\n                [1, 3, 2, 1],\n                [1, 3, 1, 1],\n                [1, 3, 1, 1],\n                [1, 3, 1, 1],\n                [1, 3, 1, 1],\n            ]\n\n            splits = []\n            start = [0, 0, 0, 0]\n            for count in counts:\n                splits.append(sorted(cells[0][start[0]:start[0] + count[0]]) +\n                              sorted(cells[1][start[1]:start[1] + count[1]]) +\n                              sorted(cells[2][start[2]:start[2] + count[2]]) +\n                              sorted(cells[3][start[3]:start[3] + count[3]]))\n                for i in range(4):\n                    start[i] += count[i]\n            assert start == [7, 17, 7, 6]\n            logging.info('Splits: {}'.format(splits))\n\n            if cv_number != -1:\n                val = sorted(splits[cv_number])\n            else:\n                val = []\n            all = []\n            for k, v in sorted(experiments.items()):\n                v = sorted(v)\n                all.extend(v)\n            tr = sorted(set(all) - set(val))\n\n            if mode == 'train':\n                logging.info('Train dataset: {}'.format(sorted(tr)))\n                self.data = list(filter(lambda d: d[0] in tr, self.data))\n            elif mode == 'val':\n                logging.info('Val dataset: {}'.format(val))\n                self.data = list(filter(lambda d: d[0] in val, self.data))\n            else:\n                assert 0\n\n        assert len(set(self.data)) == len(self.data)\n        assert len(set(all_data)) == len(all_data)\n\n        controls = list(filter(lambda d: d[-1] is not None and d[-1] >= self.treatment_classes,\n            (all_data if all_controls else self.data)))\n        self.data = list(filter(lambda d: not (d[-1] is not None and d[-1] >= self.treatment_classes),\n            self.data))\n        if move_controls:\n            self.data += controls\n\n        self.filter()\n\n        logging.info('{} dataset size: data: {}'.format(mode, len(self.data)))\n\n    def filter(self, func=None):\n        \"\"\"\n        Filter dataset by given function. If function is not specified, it will clear current filter\n        :param func: func((index, (experiment, plate, well, site, cell_type, sirna or None))) -> bool\n        \"\"\"\n        if func is None:\n            self.data_indices = None\n        else:\n            self.data_indices = list(filter(lambda i: func(i, self.data[i]), range(len(self.data))))\n\n    def __len__(self):\n        return len(self.data_indices if self.data_indices is not None else self.data)\n\n    def __getitem__(self, i):\n        i = self.data_indices[i] if self.data_indices is not None else i\n        d = self.data[i]\n\n        images = []\n        for channel in range(1, 7):\n            for dir in ['train', 'test']:\n                path = self.root / dir / d[0] / 'Plate{}'.format(d[1]) / '{}_s{}_w{}.png'.format(d[2], d[3], channel)\n                if path.exists():\n                    break\n            else:\n                assert 0\n            images.append(cv2.imread(str(path), cv2.IMREAD_GRAYSCALE))\n            assert images[-1] is not None\n        image = np.stack(images, axis=-1)\n\n        if self.transform is not None:\n            image = self.transform[0](image)\n\n        image = F.to_tensor(image)\n\n        if self.normalization == 'experiment':\n            pixel_mean = torch.tensor(P.pixel_stats[d[0]][0]) / 255\n            pixel_std = torch.tensor(P.pixel_stats[d[0]][1]) / 255\n        elif self.normalization == 'global':\n            pixel_mean = torch.tensor(list(map(lambda x: x[0], P.pixel_stats.values()))).mean(0) / 255\n            pixel_std = torch.tensor(list(map(lambda x: x[1], P.pixel_stats.values()))).mean(0) / 255\n        elif self.normalization == 'sample':\n            pixel_mean = image.mean([1, 2])\n            pixel_std = image.std([1, 2]) + 1e-8\n        else:\n            assert 0\n\n        image = (image - pixel_mean.reshape(-1, 1, 1)) / pixel_std.reshape(-1, 1, 1)\n\n        if self.transform is not None:\n            image = self.transform[1](image)\n\n        cell_type = nn.functional.one_hot(torch.tensor(self.cell_types.index(d[-2]), dtype=torch.long),\n                len(self.cell_types)).float()\n\n        r = [image, cell_type, torch.tensor(i, dtype=torch.long)]\n        if self.mode != 'test':\n            r.append(torch.tensor(d[-1], dtype=torch.long))\n        return tuple(r)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#model\nimport math\n\nimport torch\nimport torchvision\nfrom torch import nn\nfrom torch.nn import functional as F\n\n\nclass Model(nn.Module):\n    def __init__(self, args):\n        super().__init__()\n\n        kwargs = {}\n        backbone = args.backbone\n        if args.backbone.startswith('mem-'):\n            kwargs['memory_efficient'] = True\n            backbone = args.backbone[4:]\n\n        if backbone.startswith('densenet'):\n            channels = 96 if backbone == 'densenet161' else 64\n            first_conv = nn.Conv2d(6, channels, 7, 2, 3, bias=False)\n            pretrained_backbone = getattr(torchvision.models, backbone)(pretrained=True, **kwargs)\n            self.features = pretrained_backbone.features\n            self.features.conv0 = first_conv\n            features_num = pretrained_backbone.classifier.in_features\n        elif backbone.startswith('resnet') or backbone.startswith('resnext'):\n            first_conv = nn.Conv2d(6, 64, 7, 2, 3, bias=False)\n            pretrained_backbone = getattr(torchvision.models, backbone)(pretrained=True, **kwargs)\n            self.features = nn.Sequential(\n                first_conv,\n                pretrained_backbone.bn1,\n                pretrained_backbone.relu,\n                pretrained_backbone.maxpool,\n                pretrained_backbone.layer1,\n                pretrained_backbone.layer2,\n                pretrained_backbone.layer3,\n                pretrained_backbone.layer4,\n            )\n            features_num = pretrained_backbone.fc.in_features\n        elif backbone.startswith('efficientnet'):\n            from efficientnet_pytorch import EfficientNet\n            self.efficientnet = EfficientNet.from_pretrained(backbone)\n            first_conv = nn.Conv2d(6, self.efficientnet._conv_stem.out_channels, kernel_size=3, stride=2, padding=1, bias=False)\n            self.efficientnet._conv_stem = first_conv\n            self.features = self.efficientnet.extract_features\n            features_num = self.efficientnet._conv_head.out_channels\n        else:\n            raise ValueError('wrong backbone')\n\n        self.concat_cell_type = args.concat_cell_type\n        self.classes = args.classes\n\n        features_num = features_num + (4 if self.concat_cell_type else 0)\n\n        self.neck = nn.Sequential(\n            nn.BatchNorm1d(features_num),\n            nn.Linear(features_num, args.embedding_size, bias=False),\n            nn.ReLU(inplace=True),\n            nn.BatchNorm1d(args.embedding_size),\n            nn.Linear(args.embedding_size, args.embedding_size, bias=False),\n            nn.BatchNorm1d(args.embedding_size),\n        )\n        self.arc_margin_product = ArcMarginProduct(args.embedding_size, args.classes)\n\n        if args.head_hidden is None:\n            self.head = nn.Linear(args.embedding_size, args.classes)\n        else:\n            self.head = []\n            for input_size, output_size in zip([args.embedding_size] + args.head_hidden, args.head_hidden):\n                self.head.extend([\n                    nn.Linear(input_size, output_size, bias=False),\n                    nn.BatchNorm1d(output_size),\n                    nn.ReLU(),\n                ])\n            self.head.append(nn.Linear(args.head_hidden[-1], args.classes))\n            self.head = nn.Sequential(*self.head)\n\n        for m in self.modules():\n            if isinstance(m, nn.BatchNorm1d) or isinstance(m, nn.BatchNorm2d):\n                m.momentum = args.bn_mom\n\n    def embed(self, x, s):\n        x = self.features(x)\n\n        x = F.adaptive_avg_pool2d(x, (1, 1))\n        x = x.view(x.size(0), -1)\n        if self.concat_cell_type:\n            x = torch.cat([x, s], dim=1)\n\n        embedding = self.neck(x)\n        return embedding\n\n    def metric_classify(self, embedding):\n        return self.arc_margin_product(embedding)\n\n    def classify(self, embedding):\n        return self.head(embedding)\n\n\nclass ModelAndLoss(nn.Module):\n    def __init__(self, args):\n        super().__init__()\n\n        self.args = args\n        self.model = Model(args)\n        self.metric_crit = ArcFaceLoss()\n        self.crit = DenseCrossEntropy()\n\n    def train_forward(self, x, s, y):\n        embedding = self.model.embed(x, s)\n\n        metric_output = self.model.metric_classify(embedding)\n        metric_loss = self.metric_crit(metric_output, y)\n\n        output = self.model.classify(embedding)\n        loss = self.crit(output, y)\n\n        acc = (output.max(1)[1] == y.max(1)[1]).float().mean().item()\n\n        coeff = self.args.metric_loss_coeff\n        return loss * (1 - coeff) + metric_loss * coeff, acc\n\n    def eval_forward(self, x, s):\n        embedding = self.model.embed(x, s)\n        output = self.model.classify(embedding)\n        return output\n\n    def embed(self, x, s):\n        return self.model.embed(x, s)\n\n\nclass DenseCrossEntropy(nn.Module):\n    def forward(self, x, target):\n        x = x.float()\n        target = target.float()\n        logprobs = torch.nn.functional.log_softmax(x, dim=-1)\n\n        loss = -logprobs * target\n        loss = loss.sum(-1)\n        return loss.mean()\n\n\nclass ArcFaceLoss(nn.modules.Module):\n    def __init__(self, s=30.0, m=0.5):\n        super().__init__()\n        self.crit = DenseCrossEntropy()\n        self.s = s\n        self.cos_m = math.cos(m)\n        self.sin_m = math.sin(m)\n        self.th = math.cos(math.pi - m)\n        self.mm = math.sin(math.pi - m) * m\n\n    def forward(self, logits, labels):\n        logits = logits.float()\n        cosine = logits\n        sine = torch.sqrt(1.0 - torch.pow(cosine, 2))\n        phi = cosine * self.cos_m - sine * self.sin_m\n        phi = torch.where(cosine > self.th, phi, cosine - self.mm)\n\n        output = (labels * phi) + ((1.0 - labels) * cosine)\n        output *= self.s\n        loss = self.crit(output, labels)\n        return loss / 2\n\n\nclass ArcMarginProduct(nn.Module):\n    def __init__(self, in_features, out_features):\n        super().__init__()\n        self.weight = nn.Parameter(torch.FloatTensor(out_features, in_features))\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        stdv = 1. / math.sqrt(self.weight.size(1))\n        self.weight.data.uniform_(-stdv, stdv)\n\n    def forward(self, features):\n        cosine = F.linear(F.normalize(features), F.normalize(self.weight))\n        return cosine","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#main.py\nimport itertools\nimport logging\nimport math\nimport pickle\nimport random\nimport sys\nimport time\nfrom argparse import ArgumentParser\nfrom collections import defaultdict\nfrom pathlib import Path\n\nimport numpy as np\nimport torch\n#from apex import amp\nfrom torch import nn\n\n\n\n\ndef parse_args():\n    def lr_type(x):\n        x = x.split(',')\n        return x[0], list(map(float, x[1:]))\n\n    def bool_type(x):\n        if x.lower() in ['1', 'true']:\n            return True\n        if x.lower() in ['0', 'false']:\n            return False\n        raise ValueError()\n\n    parser = ArgumentParser()\n    parser.add_argument('-m', '--mode', default='train', choices=('train', 'val', 'predict'))\n    parser.add_argument('--backbone', default='mem-densenet161',\n            help='backbone for the architecture. '\n                 'Supported backbones: ResNets, ResNeXts, DenseNets (from torchvision), EfficientNets. '\n                 'For DenseNets, add prefix \"mem-\" for memory efficient version')\n    parser.add_argument('--head-hidden', type=lambda x: None if not x else list(map(int, x.split(','))),\n            help='hidden layers sizes in the head. Defaults to absence of hidden layers')\n    parser.add_argument('--concat-cell-type', type=bool_type, default=True)\n    parser.add_argument('--metric-loss-coeff', type=float, default=0.2)\n    parser.add_argument('--embedding-size', type=int, default=1024)\n    parser.add_argument('--bn-mom', type=float, default=0.05)\n    parser.add_argument('--wd', '--weight-decay', type=float, default=1e-5)\n    parser.add_argument('--label-smoothing', '--ls', type=float, default=0)\n    parser.add_argument('--mixup', type=float, default=0,\n            help='alpha parameter for mixup. 0 means no mixup')\n    parser.add_argument('--cutmix', type=float, default=1,\n            help='parameter for beta distribution. 0 means no cutmix')\n\n    parser.add_argument('--classes', type=int, default=1139,\n            help='number of classes predicting by the network')\n    parser.add_argument('--fp16', type=bool_type, default=True,\n            help='mixed precision training/inference')\n    parser.add_argument('--disp-batches', type=int, default=50,\n            help='frequency (in iterations) of printing statistics of training / inference '\n                 '(e.g. accuracy, loss, speed)')\n\n    parser.add_argument('--tta', type=int,\n            help='number of TTAs. Flips, 90 degrees rotations and resized crops (for --tta-size != 1) are applied')\n    parser.add_argument('--tta-size', type=float, default=1,\n            help='crop percentage for TTA')\n\n    parser.add_argument('--save',\n            help='path for the checkpoint with best accuracy. '\n                 'Checkpoint for each epoch will be saved with suffix .<number of epoch>')\n    parser.add_argument('--load',\n            help='path to the checkpoint which will be loaded for inference or fine-tuning')\n    parser.add_argument('--start-epoch', type=int, default=0)\n    parser.add_argument('--pred-suffix', default='',\n            help='suffix for prediction output. '\n                 'Predictions output will be stored in <loaded checkpoint path>.output<pred suffix>')\n\n    parser.add_argument('--pw-aug', type=lambda x: tuple(map(float, x.split(','))), default=(0.1, 0.1),\n            help='pixel-wise augmentation in format (scale std, bias std). scale will be sampled from N(1, scale_std) '\n                 'and bias from N(0, bias_std) for each channel independently')\n    parser.add_argument('--scale-aug', type=float, default=0.5,\n            help='zoom augmentation. Scale will be sampled from uniform(scale, 1). '\n                 'Scale is a scale for edge (preserving aspect)')\n    parser.add_argument('--all-controls-train', type=bool_type, default=True,\n            help='train using all control images (also these from the test set)')\n    parser.add_argument('--data-normalization', choices=('global', 'experiment', 'sample'), default='sample',\n            help='image normalization type: '\n                 'global -- use statistics from entire dataset, '\n                 'experiment -- use statistics from experiment, '\n                 'sample -- use mean and std calculated on given example (after normalization)')\n    parser.add_argument('--data', type=Path, default=Path('../data'),\n            help='path to the data root. It assumes format like in Kaggle with unpacked archives')\n    parser.add_argument('--cv-number', type=int, default=0, choices=(-1, 0, 1, 2, 3, 4, 5),\n            help='number of fold in 6-fold split. '\n                 'For number of given cell type experiment in certain fold see dataset.py file. '\n                 '-1 means not using validation set (training on all data)')\n    parser.add_argument('--data-split-seed', type=int, default=0,\n            help='seed for splitting experiments for folds')\n    parser.add_argument('--num-data-workers', type=int, default=10,\n            help='number of data loader workers')\n    parser.add_argument('--seed', type=int,\n            help='global seed (for weight initialization, data sampling, etc.). '\n                 'If not specified it will be randomized (and printed on the log)')\n\n    parser.add_argument('--pl-epoch', type=int, default=None,\n            help='first epoch where pseudo-labeling starts')\n    parser.add_argument('--pl-size-func', type=str, default='x',\n            help='function indicating percentage of the test set transferred to the training set. '\n                 'Function is called once an epoch and argument \"x\" is number from 0 to 1 indicating '\n                 'training progress (0 is first epoch of pseudo-labeling, and 1 is last epoch of traning). '\n                 'For example: \"x\" -- constant number of test examples is added each epoch; '\n                 '\"x*0.6+0.4\" -- 40% of test set added at the begining of pseudo-labeling and '\n                 'then constant number each epoch')\n\n    parser.add_argument('-b', '--batch_size', type=int, default=24)\n    parser.add_argument('--gradient-accumulation', type=int, default=2,\n            help='number of iterations for gradient accumulation')\n    parser.add_argument('-e', '--epochs', type=int, default=90)\n    parser.add_argument('-l', '--lr', type=lr_type, default=('cosine', [1.5e-4]),\n            help='learning rate values and schedule given in format: schedule,value1,epoch1,value2,epoch2,...,value{n}. '\n                 'in epoch range [0, epoch1) initial_lr=value1, in [epoch1, epoch2) initial_lr=value2, ..., '\n                 'in [epoch{n-1}, total_epochs) initial_lr=value{n}, '\n                 'in every range the same learning schedule is used. Possible schedules: cosine, const')\n    args = parser.parse_args()\n\n    if args.mode == 'train':\n        assert args.save is not None\n    if args.mode == 'val':\n        assert args.save is None\n    if args.mode == 'predict':\n        assert args.load is not None\n        assert args.save is None\n\n    if args.seed is None:\n        args.seed = random.randint(0, 10 ** 9)\n\n    return args\n\ndef setup_logging(args):\n    head = '{asctime}:{levelname}: {message}'\n    handlers = [logging.StreamHandler(sys.stderr)]\n    if args.mode == 'train':\n        handlers.append(logging.FileHandler(args.save + '.log', mode='w'))\n    if args.mode == 'predict':\n        handlers.append(logging.FileHandler(args.load + '.output.log', mode='w'))\n    logging.basicConfig(level=logging.DEBUG, format=head, style='{', handlers=handlers)\n    logging.info('Start with arguments {}'.format(args))\n\ndef setup_determinism(args):\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    torch.manual_seed(args.seed)\n    np.random.seed(args.seed)\n    random.seed(args.seed)\n\n\n@torch.no_grad()\ndef infer(args, model, loader):\n    \"\"\"Infer and return prediction in dictionary formatted {sample_id: logits}\"\"\"\n\n    if not len(loader):\n        return {}\n    res = {}\n\n    model.eval()\n    tic = time.time()\n    for i, (X, S, I, *_) in enumerate(loader):\n        X = X.cuda()\n        S = S.cuda()\n\n        Xs = dataset.tta(args, X) if args.tta else [X]\n        ys = [model.eval_forward(X, S) for X in Xs]\n        y = torch.stack(ys).mean(0).cpu()\n\n        for j in range(len(I)):\n            assert I[j].item() not in res\n            res[I[j].item()] = y[j].numpy()\n\n        if (i + 1) % args.disp_batches == 0:\n            logging.info('Infer Iter: {:4d}  ->  speed: {:6.1f}'.format(\n                i + 1, args.disp_batches * args.batch_size / (time.time() - tic)))\n            tic = time.time()\n\n    return res\n\n\ndef predict(args, model):\n    \"\"\"Entrypoint for predict mode\"\"\"\n\n    test_loader = dataset.get_test_loader(args)\n    train_loader, val_loader = dataset.get_train_val_loader(args, predict=True)\n\n    if args.fp16:\n        model = amp.initialize(model, opt_level='O1')\n\n    logging.info('Starting prediction')\n\n    output = {}\n    for k, loader in [('test', test_loader),\n                      ('val', val_loader)]:\n        output[k] = {}\n        res = infer(args, model, loader)\n\n        for i, v in res.items():\n            d = loader.dataset.data[i]\n            name = '{}_{}_{}'.format(d[0], d[1], d[2])\n            if name not in output[k]:\n                output[k][name] = []\n            output[k][name].append(v)\n\n    logging.info('Saving predictions to {}'.format(args.load + '.output' + args.pred_suffix))\n    with open(args.load + '.output' + args.pred_suffix, 'wb') as file:\n        pickle.dump(output, file)\n\n\ndef score(args, model, loader):\n    \"\"\"Return accuracy of the model on validation set\"\"\"\n\n    logging.info('Starting validation')\n\n    res = infer(args, model, loader)\n\n    cell_type_c = np.array([0, 0, 0, 0])  # number of examples for given cell type\n    cell_type_s = np.array([0, 0, 0, 0])  # number of correctly classified examples for given cell type\n    for i, v in res.items():\n        d = loader.dataset.data[i]\n        r = v[:loader.dataset.treatment_classes].argmax() == d[-1]\n\n        ser = loader.dataset.cell_types.index(d[4])\n        cell_type_c[ser] += 1\n        cell_type_s[ser] += r\n\n    acc = (cell_type_s.sum() / cell_type_c.sum()).item() if cell_type_c.sum() != 0 else 0\n    logging.info('Eval: acc: {} ({})'.format(cell_type_s / cell_type_c, acc))\n    return acc\n\n\ndef get_learning_rate(args, epoch):\n    assert len(args.lr[1][1::2]) + 1 == len(args.lr[1][::2])\n    for start, end, lr, next_lr in zip([0] + args.lr[1][1::2],\n                                       args.lr[1][1::2] + [args.epochs],\n                                       args.lr[1][::2],\n                                       args.lr[1][2::2] + [0]):\n        if start <= epoch < end:\n            if args.lr[0] == 'cosine':\n                return lr * (math.cos((epoch - start) / (end - start) * math.pi) + 1) / 2\n            elif args.lr[0] == 'const':\n                return lr\n            else:\n                assert 0\n    assert 0\n\n@torch.no_grad()\ndef smooth_label(args, Y):\n    nY = nn.functional.one_hot(Y, args.classes).float()\n    nY += args.label_smoothing / (args.classes - 1)\n    nY[range(Y.size(0)), Y] -= args.label_smoothing / (args.classes - 1) + args.label_smoothing\n    return nY\n\n@torch.no_grad()\ndef transform_input(args, X, S, Y):\n    \"\"\"Apply mixup, cutmix, and label-smoothing\"\"\"\n\n    Y = smooth_label(args, Y)\n\n    if args.mixup != 0 or args.cutmix != 0:\n        perm = torch.randperm(args.batch_size).cuda()\n\n    if args.mixup != 0:\n        coeffs = torch.tensor(np.random.beta(args.mixup, args.mixup, args.batch_size), dtype=torch.float32).cuda()\n        X = coeffs.view(-1, 1, 1, 1) * X + (1 - coeffs.view(-1, 1, 1, 1)) * X[perm,]\n        S = coeffs.view(-1, 1) * S + (1 - coeffs.view(-1, 1)) * S[perm,]\n        Y = coeffs.view(-1, 1) * Y + (1 - coeffs.view(-1, 1)) * Y[perm,]\n\n    if args.cutmix != 0:\n        img_height, img_width = X.size()[2:]\n        lambd = np.random.beta(args.cutmix, args.cutmix)\n        column = np.random.uniform(0, img_width)\n        row = np.random.uniform(0, img_height)\n        height = (1 - lambd) ** 0.5 * img_height\n        width = (1 - lambd) ** 0.5 * img_width\n        r1 = round(max(0, row - height / 2))\n        r2 = round(min(img_height, row + height / 2))\n        c1 = round(max(0, column - width / 2))\n        c2 = round(min(img_width, column + width / 2))\n        if r1 < r2 and c1 < c2:\n            X[:, :, r1:r2, c1:c2] = X[perm, :, r1:r2, c1:c2]\n\n            lambd = 1 - (r2 - r1) * (c2 - c1) / (img_height * img_width)\n            S = S * lambd + S[perm] * (1 - lambd)\n            Y = Y * lambd + Y[perm] * (1 - lambd)\n\n    return X, S, Y\n\ndef pseudo_label(args, epoch, pl_data, model, val_loader, test_loader, train_loader):\n    \"\"\"Pseudo-label some test and validation examples and move them to the training set\"\"\"\n\n    if args.pl_epoch is None or epoch < args.pl_epoch:\n        return\n\n    logging.info('Starting pseudo-labeling')\n\n    test_loader.dataset.filter(lambda i, d: ('test', i) not in pl_data)\n    test_res = infer(args, model, test_loader)\n    test_loader.dataset.filter()\n\n    val_loader.dataset.filter(lambda i, d: ('val', i) not in pl_data)\n    val_res = infer(args, model, val_loader)\n    val_loader.dataset.filter()\n\n    test_res = sorted(test_res.items())\n    val_res = sorted(val_res.items())\n\n\n    set_classes = defaultdict(lambda: [])  # classes that are already in the training set for the plate\n    for j in range(len(train_loader.dataset.data)):\n        experiment_plate = train_loader.dataset.data[j][:2]\n        sirna = train_loader.dataset.data[j][-1]\n        set_classes[experiment_plate].append(sirna)\n\n    confs = []\n    last = None\n    for k, (i, v) in itertools.chain(\n            zip(itertools.repeat('val'), val_res),\n            zip(itertools.repeat('test'), test_res)):\n        loader = val_loader if k == 'val' else test_loader\n\n        # assumes that both sides of an example will be next to each other\n        if i % 2 == 0:\n            assert last is None\n            last = i, v\n            continue\n        else:\n            last_i, last_v = last\n            assert last_i == i - 1\n            last = None\n\n            logits = v + last_v  # ensemble two sites\n            plate = loader.dataset.data[i][1] - 1\n            experiment = loader.dataset.data[i][0]\n            class_group_id = P.group_assignment[experiment][plate]\n            possible_classes = P.groups[class_group_id]\n            remaining_classes = list(set(range(loader.dataset.treatment_classes)) - possible_classes)\n            logits[remaining_classes] = -10e6\n\n            experiment_plate = loader.dataset.data[i][:2]\n            if set_classes[experiment_plate]:\n                logits[set_classes[experiment_plate]] = -10e6\n            logits = logits[:loader.dataset.treatment_classes]\n            r = logits.argmax().item()\n\n            logits.sort()\n            c = logits[-1] - logits[-2]\n            confs.append(((k, i - 1), c, r))\n\n\n    x = (epoch - args.pl_epoch + 1) / (args.epochs - args.pl_epoch + 1)\n    val_test_examples = len(val_loader.dataset.data) // 2 + len(test_loader.dataset.data) // 2\n    added_examples = len(pl_data) // 2\n    n = round(eval('lambda x: ' + args.pl_size_func)(x) * val_test_examples) - added_examples\n    n = max(n, 0)\n\n    confs = list(filter(lambda x: x[0] not in pl_data, confs))\n    confs.sort(key=lambda x: -x[1])\n    confs = confs[:n]\n\n    val_misclass = 0\n    val_count = 0\n    test_count = 0\n    not_added_count = 0\n    added_sirnas = defaultdict(set)\n    for (k, i), c, r in confs:\n        if k == 'val':\n            d1 = val_loader.dataset.data[i]\n            d2 = val_loader.dataset.data[i + 1]\n        elif k == 'test':\n            d1 = test_loader.dataset.data[i]\n            d2 = test_loader.dataset.data[i + 1]\n        else:\n            assert 0\n        assert d1[:3] == d2[:3] and d1[-2:] == d2[-2:]\n\n        if r in added_sirnas[d1[:2]]:\n            not_added_count += 1\n            continue\n\n        if k == 'val':\n            val_count += 1\n            if d1[-1] != r:\n                val_misclass += 1\n        elif k == 'test':\n            test_count += 1\n        else:\n            assert 0\n\n        added_sirnas[d1[:2]].add(r)\n        pl_data.add((k, i))\n        pl_data.add((k, i + 1))\n        train_loader.dataset.data.append((*d1[:-1], r))\n        train_loader.dataset.data.append((*d2[:-1], r))\n\n    logging.info('Pseudo-labeling: Added {} ({} val, {} test), {} ({:.3f}%) val misclassified, '\n                 '{} ({:.3f}%) not added, pl_data size {}, train size {}, threshold {}'.format(\n                     n, val_count, test_count, val_misclass, val_misclass / val_count * 100 if val_count != 0 else 0,\n                     not_added_count, not_added_count / (not_added_count + n) * 100 if not_added_count + n != 0 else 0,\n                     len(pl_data), len(train_loader.dataset.data), confs[-1][1] if len(confs) != 0 else 'None'))\n\n\ndef train(args, model):\n    train_loader, val_loader = dataset.get_train_val_loader(args)\n\n    optimizer = torch.optim.Adam(model.parameters(), lr=0, weight_decay=args.wd)\n\n    if args.fp16:\n        model, optimizer = amp.initialize(model, optimizer, opt_level='O1')\n\n    if args.load is not None:\n        best_acc = score(args, model, val_loader)\n    else:\n        best_acc = float('-inf')\n\n    if args.mode == 'val':\n        return\n\n    if args.pl_epoch is not None:\n        test_loader = dataset.get_test_loader(args, exclude_leak=True)\n        pl_data = set()\n\n    for epoch in range(args.start_epoch, args.epochs):\n        if args.pl_epoch is not None:\n            pseudo_label(args, epoch, pl_data, model, val_loader, test_loader, train_loader)\n\n        with torch.no_grad():\n            avg_norm = np.mean([v.norm().item() for v in model.parameters()])\n\n        logging.info('Train: epoch {}   avg_norm: {}'.format(epoch, avg_norm))\n\n        model.train()\n        optimizer.zero_grad()\n\n        cum_loss = 0\n        cum_acc = 0\n        cum_count = 0\n        tic = time.time()\n        for i, (X, S, _, Y) in enumerate(train_loader):\n            lr = get_learning_rate(args, epoch + i / len(train_loader))\n            for g in optimizer.param_groups:\n                g['lr'] = lr\n\n            X = X.cuda()\n            S = S.cuda()\n            Y = Y.cuda()\n            X, S, Y = transform_input(args, X, S, Y)\n\n            loss, acc = model.train_forward(X, S, Y)\n            if args.fp16:\n                with amp.scale_loss(loss, optimizer) as scaled_loss:\n                    scaled_loss.backward()\n            else:\n                loss.backward()\n            if (i + 1) % args.gradient_accumulation == 0:\n                optimizer.step()\n                optimizer.zero_grad()\n\n            cum_count += 1\n            cum_loss += loss.item()\n            cum_acc += acc\n            if (i + 1) % args.disp_batches == 0:\n                logging.info('Epoch: {:3d} Iter: {:4d}  ->  speed: {:6.1f}   lr: {:.9f}   loss: {:.6f}   acc: {:.6f}'.format(\n                    epoch, i + 1, cum_count * args.batch_size / (time.time() - tic), optimizer.param_groups[0]['lr'],\n                    cum_loss / cum_count, cum_acc / cum_count))\n                cum_loss = 0\n                cum_acc = 0\n                cum_count = 0\n                tic = time.time()\n\n        acc = score(args, model, val_loader)\n        torch.save(model.state_dict(), str(args.save + '.{}'.format(epoch)))\n        if acc >= best_acc:\n            best_acc = acc\n            logging.info('Saving best to {} with score {}'.format(args.save, best_acc))\n            torch.save(model.state_dict(), str(args.save))\n\ndef main(args):\n    model = ModelAndLoss(args).cuda()\n    logging.info('Model:\\n{}'.format(str(model)))\n\n    if args.load is not None:\n        logging.info('Loading model from {}'.format(args.load))\n        model.load_state_dict(torch.load(str(args.load)))\n\n    if args.mode in ['train', 'val']:\n        train(args, model)\n    elif args.mode == 'predict':\n        predict(args, model)\n    else:\n        assert 0\n\n\n\nif __name__ == '__main__':\n    args = parse_args()\n    setup_logging(args)\n    setup_determinism(args)\n    main(args)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import argparse\nimport logging\nimport math\nimport numpy as np\nimport pandas as pd\nimport pickle\nimport sys\nfrom collections import defaultdict\nfrom functools import reduce\nfrom itertools import permutations, groupby, chain\nfrom multiprocessing import Pool\nfrom operator import itemgetter\nfrom pathlib import Path\n\nfrom scipy.optimize import linear_sum_assignment\n\n\nclass Dataset:\n    CLASSES = 1108\n\n    def __init__(self, path):\n        self.data = {}\n        self.controls = {}\n\n        path = Path(path)\n        for is_control, file in [(0, 'train.csv'), (1, 'train_controls.csv'), (1, 'test_controls.csv')]:\n            csv = pd.read_csv(path / file)\n            for row in csv.iterrows():\n                r = row[1]\n                (self.controls if is_control else self.data)[self.split(r.id_code)] = r.sirna\n\n        # HUVEC-18 leak\n        for file in ['test.csv']:\n            csv = pd.read_csv(path / file)\n            for row in csv.iterrows():\n                r = row[1]\n                if self.split(r.id_code)[0:2] == ('HUVEC', '18'):\n                    s = self.split(r.id_code)\n                    s = list(s)\n                    s[0] = 'RPE'\n                    s[1] = '03'\n                    s[2] = (s[2] - 1) % 4\n                    s = tuple(s)\n                    assert self.data[s] < self.CLASSES\n                    self.data[self.split(r.id_code)] = self.data[s]\n\n        self.groups, self.group_assignment = self._get_groups()\n\n    @staticmethod\n    def split(id_code):\n        \"\"\"Return (cell_type, experiment number of given cell type, plate number, well)\"\"\"\n\n        a = id_code.find('-')\n        b = id_code.find('_')\n        c = id_code.rfind('_')\n        return id_code[:a], id_code[a + 1:b], int(id_code[b + 1:c]) - 1, id_code[c + 1:]\n\n    def _get_groups(self):\n        \"\"\"Calculate class groups that are on plates and assignment for the labeled set\"\"\"\n\n        data = defaultdict(lambda: defaultdict(lambda: defaultdict(lambda: defaultdict(lambda: 0))))\n        for (serie, exper, plate, _), sirna in self.data.items():\n            data[serie][exper][plate][sirna] += 1\n\n        groups = set()\n        for serie in data:\n            for exper in data[serie]:\n                for plate in data[serie][exper]:\n                    k = tuple(sorted(list(data[serie][exper][plate].keys())))\n                    if len(k) == self.CLASSES // 4:\n                        groups.add(k)\n        groups = sorted(groups)\n        assert len(groups) == 4\n        for i in range(len(groups)):\n            for j in range(i + 1, len(groups)):\n                assert len(set(groups[i]).intersection(set(groups[j]))) == 0\n\n        assignment = {}\n        for serie in data:\n            for exper in data[serie]:\n                gs = []\n                for plate in data[serie][exper]:\n                    k = tuple(sorted(list(data[serie][exper][plate].keys())))\n                    sc = [len(set(g).intersection(set(k))) for g in groups]\n                    assert sum(sc) == max(sc)\n                    g = sc.index(max(sc))\n                    gs.append(g)\n                assignment[(serie, exper)] = tuple(gs)\n                assert(sorted(gs) == [0, 1, 2, 3])\n\n        return groups, assignment\n\n    def assign_groups(self, data):\n        \"\"\"Find group assignments as dictionary in format {code_id: list_of_classes}\"\"\"\n\n        ret = {}\n        for exper_name, exper in groupby(sorted(data), key=lambda x: self.split(x[0])[:2]):\n            exper = list(exper)\n            ks, vs = [], []\n            for _, v in groupby(sorted(exper), key=lambda x: self.split(x[0])[2]):\n                v = list(v)\n                ks.append(list(map(itemgetter(0), v)))\n                vs.append(list(map(itemgetter(1), v)))\n            # ks[i][j] -- code id of j-th well on i-th plate of experiment 'exper_name'\n            # vs[i][j] -- logits for j-th well on i-th plate of experiment 'exper_name'\n\n            scs = []\n            for v in vs:\n                v = np.array(v)\n                v = v.argmax(1)\n                sc = [len(list(filter(lambda x: x in g, v))) for g in map(set, self.groups)]\n                scs.append(sc)\n            # scs[i][j] -- number of best classes that are on i-th plate and are in j-th class group\n\n            scs = np.array(scs)\n            scs = scs / scs.sum(0, keepdims=True)\n\n            perms = []\n            for perm in permutations(range(len(vs))):\n                score = 0\n                for i, j in enumerate(perm):\n                    score += scs[i, j]\n                perms.append((score, perm))\n            perms.sort(key=lambda x: -x[0])\n\n            best_perm = perms[0][1]\n            conf = perms[0][0] - (perms[1][0] if len(perms) > 1 else perms[0][0])\n            score = perms[0][0]\n\n            if exper_name in self.group_assignment:\n                if self.group_assignment[exper_name] == best_perm:\n                    assignment_type = 'correct_assignment'\n                else:\n                    assignment_type = 'incorrect_assignment'\n            else:\n                assignment_type = 'prediction'\n\n            logging.info('groups: {:8} -> {} ( score: {:.5f}  conf: {:.5f} ) {}  size: {}'.format(\n                '-'.join(exper_name), best_perm, score, conf, assignment_type, sum(map(len, ks))))\n\n            for i, k in enumerate(ks):\n                for n in k:\n                    assert n not in ret\n                    ret[n] = self.groups[best_perm[i]]\n        return ret\n\n    def accuracy(self, data):\n        if isinstance(data, dict):\n            data = data.items()\n\n        correct_hits = 0\n        total = 0\n        correct_hits_exper = defaultdict(lambda: 0)\n        total_exper = defaultdict(lambda: 0)\n        for k, v in data:\n            split = self.split(k)\n            total += 1\n            total_exper[split[:2]] += 1\n            if v == self.data[split]:\n                correct_hits += 1\n                correct_hits_exper[split[:2]] += 1\n\n        if total == 0:\n            return 0, {}\n        return correct_hits / total, dict(map(lambda x: (x[0][0], x[0][1] / x[1][1] if x[1][1] != 0 else 0),\n            zip(correct_hits_exper.items(), total_exper.items())))\n\n\nclass PredictionGroup:\n    def __init__(self, x):\n        if isinstance(x, dict):\n            x = x.items()\n        self.data = []\n        for k, v in x:\n            for pred in (v if isinstance(v, list) else [v]):\n                self.data.append((k, pred[:Dataset.CLASSES]))\n\n    def __len__(self):\n        return len(self.data)\n\n    def __iter__(self):\n        return iter(self.data)\n\n    def combine(self, f=None):\n        if f is None:\n            f = lambda x: x.sum(0)\n\n        r = {}\n        for code_id, iterable in groupby(sorted(self.data, key=lambda x: x[0]), key=lambda x: x[0]):\n            iterable = list(iterable)\n            pred = np.array(list(map(lambda x: x[1], iterable)))\n            r[code_id] = f(pred)\n        return PredictionGroup(r)\n\n    def retain_plate_classes(self, assignment):\n        r = []\n        for code_id, pred in self:\n            new_pred = pred.copy()\n            new_pred[list(set(range(len(new_pred))) - set(assignment[code_id]))] = -np.inf\n            r.append((code_id, new_pred))\n        return PredictionGroup(r)\n\n    def assign_argmax(self):\n        for k, v in self:\n            yield k, v.argmax()\n\n    def _assign_unique_in_plate(self, plate):\n        preds = np.array(list(map(itemgetter(1), plate)))\n        preds = np.vectorize(lambda x: x if x != -np.inf else -1e10)(preds)\n        _, indices = linear_sum_assignment(-preds)\n        return [(k, v.item()) for (k, _), v in zip(plate, indices)]\n\n    def assign_unique(self, pool=__builtins__):\n        plates = (list(plate) for _, plate in groupby(sorted(self, key=itemgetter(0)),\n            key=lambda x: Dataset.split(x[0])[:3]))\n        return chain(*pool.map(self._assign_unique_in_plate, plates))\n\n    def concat(*args):\n        r = []\n        for w in args:\n            for k, v in w:\n                r.append((k, [v]))\n        return PredictionGroup(r)\n\n    def normalize(self):\n        return self.map(lambda x: (x - x.mean()) / max(x.std(), 1e-8))\n\n    def map(self, f=None):\n        if not self.data:\n            return PredictionGroup([])\n\n        preds = np.array(list(map(itemgetter(1), self)))\n        if f is not None:\n            preds = f(preds)\n        return PredictionGroup(((k, preds[i])) for i, (k, _) in enumerate(self))\n\n\nclass Prediction:\n    def __init__(self, data, y=None):\n        if y is not None:\n            self.val, self.test = data, y\n\n        else:\n            if isinstance(data, Path) or isinstance(data, str):\n                with Path(data).open('rb') as f:\n                    data = pickle.load(f)\n\n            self.val = PredictionGroup(data['val'])\n            self.test = PredictionGroup(data['test'])\n\n    def _map(self, f):\n        if isinstance(self, Prediction):\n            return Prediction(f(self.val), f(self.test))\n        else:\n            return Prediction(\n                    f(list(map(lambda x: x.val, self))),\n                    f(list(map(lambda x: x.test, self))),\n            )\n\n    def combine(self, *args, **kwargs):\n        return self._map(lambda x: x.combine(*args, **kwargs))\n\n    def retain_plate_classes(self, dataset):\n        return self._map(lambda x: x.retain_plate_classes(dataset.assign_groups(x)))\n\n    def concat(*args):\n        return Prediction._map(args, lambda x: PredictionGroup.concat(*x))\n\n    def normalize(self, *args, **kwargs):\n        return self._map(lambda x: x.normalize(*args, **kwargs))\n\n    def map(self, *args, **kwargs):\n        return self._map(lambda x: x.map(*args, **kwargs))\n\ndef parse_args():\n    parser = argparse.ArgumentParser()\n    parser.add_argument('--data', type=Path, default=Path('../data/'))\n    parser.add_argument('-t', '--threads', type=int, default=12)\n    parser.add_argument('-w', '--weights', type=lambda x: list(map(float, x.split(','))))\n    parser.add_argument('-o', '--output', type=Path, required=True)\n    parser.add_argument('files', nargs='+', type=Path)\n    args = parser.parse_args()\n\n    if args.weights is None:\n        args.weights = [1] * len(args.files)\n\n    return args\n\nif __name__ == '__main__':\n    args = parse_args()\n    logging.basicConfig(level=logging.DEBUG, format='{asctime}:{levelname}: {message}', style='{',\n            handlers=[logging.StreamHandler(sys.stderr)])\n    logging.info('Args: {}'.format(args))\n\n    pool = Pool(args.threads)\n\n    logging.info('Loading dataset')\n    dataset = Dataset(args.data)\n\n    logging.info('Loading predictions')\n    preds = []\n    for i, file in enumerate(args.files):\n        pred = Prediction(args.files[i])\n        pred = pred.combine()\n        score = dataset.accuracy(pred.val.assign_argmax())\n        preds.append(pred)\n        logging.info('File {} -> score: {}'.format(args.files[i], score))\n\n\n    preds = list(map(lambda p: p[0].map(lambda x: (x * p[1])), zip(preds, args.weights)))\n    pred = Prediction.concat(*preds)\n\n    logging.info('Evaluating...')\n    logging.info('Average score:                       {}'.format(dataset.accuracy(pred.val.assign_argmax())))\n    pred = pred.combine()\n    logging.info('Score after ensemble:                {}'.format(dataset.accuracy(pred.val.assign_argmax())))\n    pred = pred.retain_plate_classes(dataset)\n    logging.info('Score after retaining plate classes: {}'.format(dataset.accuracy(pred.val.assign_argmax())))\n    logging.info('Score after linear sum assignment:   {}'.format(dataset.accuracy(pred.val.assign_unique(pool=pool))))\n\n    logging.info('Saving csv submission into {}'.format(args.output))\n    with args.output.open('w') as f:\n        print('id_code,sirna', file=f)\n        for k, v in sorted(pred.test.assign_unique(pool=pool)):\n            print(','.join([str(k), str(v)]), file=f)","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}