{"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":"# Ignite Metric Implementation\nImplement the metric according to @barteksadlej123 's the code of the metric. https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/341854","metadata":{}},{"cell_type":"markdown","source":"# Install pytorch-ignite","metadata":{}},{"cell_type":"code","source":"!pip install -q pytorch-ignite","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:53:03.572135Z","iopub.execute_input":"2022-10-05T08:53:03.572512Z","iopub.status.idle":"2022-10-05T08:53:12.387611Z","shell.execute_reply.started":"2022-10-05T08:53:03.572488Z","shell.execute_reply":"2022-10-05T08:53:12.385691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import libraries","metadata":{}},{"cell_type":"code","source":"from typing import Callable, Sequence, Union\n\nimport torch\nfrom ignite.exceptions import NotComputableError\nfrom ignite.engine import Engine\nfrom ignite.metrics.metric import Metric, reinit__is_reduced, sync_all_reduce","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:54:55.550341Z","iopub.execute_input":"2022-10-05T08:54:55.551698Z","iopub.status.idle":"2022-10-05T08:54:57.162228Z","shell.execute_reply.started":"2022-10-05T08:54:55.551633Z","shell.execute_reply":"2022-10-05T08:54:57.160491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create competition metric using Ignite\nTo create a custom metric, I create a new class inheriting from Metric and override three methods :\n\n- reset() : resets internal variables and accumulators\n\n- update() : updates internal variables and accumulators with provided batch output (y_pred, y)\n\n- compute() : computes custom metric and return the result","metadata":{}},{"cell_type":"code","source":"class RSNA_Metric(Metric):\n\n    def __init__(self, output_transform=lambda x: x, device=\"cpu\"):\n        \n        super(RSNA_Metric, self).__init__(output_transform=output_transform,\n                                          device=device)\n\n        self._sum_of_loss = None\n        self._num_examples = None\n        self.competition_weights = {\n            'negative':\n            torch.tensor([7, 1, 1, 1, 1, 1, 1, 1],\n                         dtype=torch.float,\n                         device=self._device),\n            'positive':\n            torch.tensor([14, 2, 2, 2, 2, 2, 2, 2],\n                         dtype=torch.float,\n                         device=self._device),\n        }\n\n    @reinit__is_reduced\n    def reset(self) -> None:\n        self._sum_of_loss = torch.tensor(0.0, device=self._device)\n        self._num_examples = 0\n\n    @reinit__is_reduced\n    def update(self, output: Sequence[torch.Tensor]) -> None:\n        y_pred, y = output[0].detach(), output[1].detach()\n\n        loss    = torch.nn.BCEWithLogitsLoss(reduction='none')(y_pred, y)\n        weights = self.competition_weights['positive'] * y + \\\n                  self.competition_weights['negative'] * (1 - y)\n        loss    = (loss * weights).sum(axis=1)\n        loss    = loss / weights.sum(axis=1)\n\n        self._sum_of_loss  += torch.sum(loss).to(self._device)\n        self._num_examples += y.shape[0]\n\n    @sync_all_reduce(\"_sum_of_loss\", \"_num_examples\")\n    def compute(self) -> Union[float, torch.Tensor]:\n        if self._num_examples == 0:\n            raise NotComputableError\n        return self._sum_of_loss.item() / self._num_examples","metadata":{"execution":{"iopub.status.busy":"2022-10-05T09:00:15.511934Z","iopub.execute_input":"2022-10-05T09:00:15.512306Z","iopub.status.idle":"2022-10-05T09:00:15.527601Z","shell.execute_reply.started":"2022-10-05T09:00:15.512263Z","shell.execute_reply":"2022-10-05T09:00:15.526299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Test code","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ny_preds = torch.tensor([\n    [0, 0, 0, 0, 0, 0, 0, 0],\n    [1, 1, 0, 0, 0, 0, 0, 0],\n    [1, 1, 1, 0, 0, 0, 0, 0],\n], dtype=torch.float, device=device)\ny = torch.tensor([\n    [0, 0, 0, 0, 0, 0, 0, 0],\n    [0, 0, 0, 0, 0, 0, 0, 0],\n    [1, 1, 1, 0, 0, 0, 0, 0],\n], dtype=torch.float, device=device)\n\ndef predict_on_batch(engine, batch):\n    return y_preds, y\n    \nevaluator = Engine(predict_on_batch)\nmetric = RSNA_Metric(device=device)\nmetric.attach(evaluator, 'rsna')\nstate = evaluator.run([[y_preds, y]])\nprint(state.metrics['rsna'])","metadata":{"execution":{"iopub.status.busy":"2022-10-05T08:58:24.756111Z","iopub.execute_input":"2022-10-05T08:58:24.756804Z","iopub.status.idle":"2022-10-05T08:58:24.768582Z","shell.execute_reply.started":"2022-10-05T08:58:24.756775Z","shell.execute_reply":"2022-10-05T08:58:24.767228Z"},"trusted":true},"execution_count":null,"outputs":[]}]}