{"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":"# Threshold optimization\n\nI've seen a few of the public notebooks trying out thresholds out of e.g. np.arange(0, 1.01, 0.01), which seems sort of arbitrary to me. What if some of the steps change no predictions? What if your net tends to predict either close to 0 or close to 1? What if there is a lot of change between a threshold of e.g. 0.7 and 0.7005?   \n\nMy answer to these questions is to take the actual outputs of the net into account and split the predictions for positive pixels into quantiles. Now we calculate the dice scores depending on which percentage of the positive labels we want to put over the threshold i.e. turn the corresponding net outputs for these pixels into true positives.   \n\nI also tried out numba to speed up the process a bit, even though this is probably not part of the code which needs speed up.","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"import numba\ndef threshold_optimization_and_get_validation_loss(net, val_dataloader, device, steps=100):\n    results_where_pos = []\n    results_where_neg = []\n    net.eval()\n    with torch.no_grad():\n        for img, label in tqdm(val_dataloader):\n            img = img.to(device)\n            label = label.to(device)\n            result = torch.sigmoid(net(img))\n            results_where_pos.append(result[label == 1].cpu())\n            results_where_neg.append(result[label == 0].cpu())\n    rwp = torch.cat(results_where_pos)\n    rwn = torch.cat(results_where_neg)\n    q = torch.arange(1 / steps, 1 + 1 / steps, 1 / steps)\n    rwp_q = torch.quantile(rwp, q).numpy()\n    val_score_list = dice_score_for_all_thresholds(rwn.numpy(), rwp.numpy(), rwp_q)\n    best_index = val_score_list.argmax()\n    return rwp_q[best_index], val_score_list[best_index]\n\n\n@numba.jit(nopython=True)\ndef dice_score_for_all_thresholds(rwn, rwp, rwp_q):\n    val_score_list = np.ones_like(rwp_q)\n    for i in range(len(rwp_q)):\n        threshold = rwp_q[i]\n        true_pos = (rwp >= threshold).sum()\n        false_pos = (rwn >= threshold).sum()\n        val_score_list[i] = (2 * true_pos) / (true_pos + false_pos + len(rwp))\n    return val_score_list\n\n\n# threshold, max_val_score = threshold_optimization_and_get_validation_loss(net, val_dataloader, device)","metadata":{},"execution_count":null,"outputs":[]}]}