{"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":"## Ensemble Model","metadata":{}},{"cell_type":"markdown","source":"In this note book we will build an ensemble model with 6 models. Model ensembling combines the predictions from multiple models together. Traditionally this is done by running each model on some inputs separately and then combining the predictions.\n\n - resnet101\n - vgg19_bn\n - densenet161\n - squeezenet1_1\n - efficientnet_b6\n - convnext_large\n \n ![ensemble](data:image/jpeg;base64,/9j/4AAQSkZJRgABAQAAAQABAAD/2wCEAAoHCBQUEhcUEhUXFxcXGxgbFxIXFxUaGhcaFx0aGxsaGxgbIC4lGx0pHxgYJTYmKS4wQDMzGyQ5PjkyPSwyMzMBCwsLEA4QHhISHjQqIikyNDIyMjsyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMjIyMv/AABEIAKUBMQMBIgACEQEDEQH/xAAbAAEAAgMBAQAAAAAAAAAAAAAABAUCAwYBB//EAEUQAAIBAgMDBgsGBAUEAwAAAAECEQADBBIhBRMxIkFRUnGRBhQVIzJhcoGSsdJCYoKhosEzU7LRFmOTs/AkRHPhQ6PD/8QAGAEBAQEBAQAAAAAAAAAAAAAAAAMCAQT/xAAhEQACAgIDAQADAQAAAAAAAAAAAQIREiEDMVETIkFhMv/aAAwDAQACEQMRAD8A+zUpSgFKUoBSlKAUpSgFKUoBSlKAUpSgFKUoBSlKAUpSgFKUoBSlKAUpSgFKUoBSlKAUpSgPKVW4vFgnKrLAgk5yJ1MqCOz860777w/1blbUGzDmi4mk1T777w/1XpvvvD/Veu/NjMuKVULcMiCeNuCLjsCGeCNew99W9Yao0nZ7SlK4dFKUoBSlKAUpSgFKUoDylR8VcIAgxJAnT19PZWrM3XPcn9q6o2ZcqJtKhZn657k/tTM/XPcn9q7gxZNpULM3XPcn9qzw1xizKTMBSDA583R2CuOLSCkS68pUS/dObKpiACTpOswBOnN+YolZ1uiXSoMt12/R9NJbrt+j6a7gzmROpUGW67fo+mkt12/R9NMGMidSoMt12/R9NZ2LpzZSZkEg6TpEgxpzj/go4tBSJleUqHcusWIBgAxIiSYk8ebWuJWdbomUqDLddv0fTSW67fo+mu4M5kTqVBluu36PppLddv0fTTBjInUqDLddv0fTW3D3SSVOpEGekGePr0P5UcWgpWSqUpWTRUYduQOytI2laNxrYuDeKVVk1kF1LrPaqsZ9Rr2y3JFU+O2Lau3WY3GVyXzKjAHLctC1lI9UBweYzzEz6WvCKZfJiFIkMpB4EEHp/se41kH9dctd8FbbKo3jqVmWSZLHRWJdmMhWurxjzjequiUgCBoBoAOYdFdURZhbbX8Vv/car6udsnlD2rf+41dDUuVbNw6PaUpUjYpSlAKUpQClKUApSlAQ9ocF9ofI1pzVRDbd69i7lpLavhrbAeMjOnLEhrazIukMCCwygajUisMXtw2rlxbickCbSqJa6BkzEGYkF4yROgOoOluP/JOXZntLGXrVy66G5cRUsFbWW3kBu3Ht3GkKGIRVVyMw55IB00rtq+zKCiJrZLHKzrluQGlgeSZOmnDWTrly/wARrnCi1cIMQ82ssE2QWjPmgb+3zczRwE6v8UrCtu7gUqrZTkZ23ptrby5GI1NyCDqI79VRyzps1Z4I+cf2U+b1THaDlbLqmVbrAOtzMtxMwMDLHGRznvmpGfE528XW0RlSTcZxBl+AUaj31ya/FiL2X9Vt4+cf8PyqsxeyMZftut3FqgZWUWrNlAhzCPOG5mdtZ9Flkd9YbJ2c2Ft7l71y8UC+cunMx04T0DmqfGrZufRNx2LNtMwAPLtLr0XLiIT7g0+6q3E+EttDd83dItFhmCHLcZPSVD9puOg5lJ4Ctu1MdbthEuJnW64QrlVgASBncH7AYoCdYzCoQxuCjK9tFJEFNzmBCHKiyqQT6MLx1UAcKs1smmSn8IACfNXCoLcubcZUCF3jPMDOBHEmdI1rW/hKqAl7bwJ5YyBSwV3CgF5ByodTpw92L7Yw0ckZoDFgLT6CHzLqsZ5tMCnHk8NK0vtPCXBct3EZVGXTdXAWzJbYZSqyH88q5QZ48xpQL7A4reW1uZSuYSFaJidDp0jX31vw586vsv8ANKp9j4mwbe7w8hU1ystxTyyxnzgkywfXpBHNUx8MbjqouXLcBjmtkAn0dDIOlcktMJ7L6qvNyn9s/tVbtfwZa9Ye0MXe5cCXIYCCDMAA5hEgzoQDzVswOGNpN2bj3Chy7y6czvAGrNzmp8atm59E8vTNXN4jZ1xrlxt3bLNdsut8ty8lt7L7sjLIjdtoDGgPEmtNvCY3KA1wg+dBy3Ccoe2oUqWnMwuBiM2YQeA0i1EzqTcHORrw9de5q5JNj3muJcuFcy66O5H/AGg1BnMYs3fSJOo1PNL2VhMUrIcRdZgocuAxhnIthTHEppdOXSJGg4Vyv4Do89Z4I+cb2V+bVS7OstbuXyUCrcuB1IaS3IRCSOYyk9hFZYw31Zr2HYk2kDNhoUreWWLJJEq8A5SCNYmRNcmvxZqL2dRXlacJiUu20uWyGR1VkYc6sAQe41urzlTnkfSqHF7LutduOm6WXNxLme4HabC2cjZVBRZXNIZvRXSukubLuAnIyEcwYMD3ia8TZd0+kyKPVmY/tXrzg12efGXhyl3ZGICNluLna2iNcN++CCjXDowXodOVH2TprI6G1iRlAnMQBJUM0mNdQKkbAwq3LK3bgzuS4JPAZHZRC8Booq8VQOFTfKl0jfzb7KnZeBnzlwAnigggqMzGTPPqKuKUqMpOTtlEklSPaUpXDp5SsLjhQSeb/gFaS7noX1RJ7+APuNdSs42SaVFyHndj8I/pApux0v8AG/8Aeu4iyVSouToZh75/qmnLHBp9oCe9YjuNcoWS6qPCPEumHYWiRcuslq2wiVe6wTPB6gJfsU1YWrsyCIIiRx48CD0aHuqq2qobF4ND9lr133pbNsT/AK9cOm2/hEs2LVq2MqWyiqOgKIFV1zB2mLFrVti8ZyUQl8vDMSOVECJ6KvcfhzcTKpAIIIJmNOyql8HeH/xg+tXHyMVfilGqZKad6MBaTqL8I+79K/COgVpsbPspb3aWkCRBXKpkEAHNPpEhRJPGBNSUwV5vsBfWzD5LM1MtbIXjcYuer6KfCOPvmty5IoyoSZU3EsEKrJaItxkQqhyRwyrHJ4Dh0VbbIfMzkTEIJII4ZuntFWFuyqiFUAdAAFbKlPkyVUbjCndiqTHmLzesKfdqP2q7qNi8GlwDMDI4EEgjsIrEJYu2alHJUc5jcBZvGbttXIVlBaeSHjNl6p5I5Qg6DWoJ2DaNxrktDAAoGYcII5YMiGAcEahtZ5q6fyNb61z4zTyNb61z4zV/rDwlhIpLezrKjKE0mdWckkhwSSTJJ3jyTxLSda8TZtkADKTBVpa5cY5kFsKSzNJ0tW+PV9Zm98jW+tc+M08jW+tc+M0+sPBhIqbFhLZJQQSADqx0Uuw4npuP39lT9nNN0epGn3lY+R7q3+RrfWufGalYXCJbBCjjxJJJPaTWZ8qcaSNRg07ZIqhvtFxx975gRV9UPF4BLhBaQRpmUkGOjTiKnxyUXbNTjktFTnr3PU/yNb61z4zTyNb61z4zV/tD+k/nIgZ6Z6n+RrfWufGaeRrfWufGafaH9HzkV+epmyjLv6lUe+WMf86az8jW+tc+M1Nw+HVFyqIHeSeknnNY5OWMlSNQg07ZU+DXIF+xr5i/cVZ6lwLeQD7qreCj2avKpcA0Y/Fr/l4Z/e2+Se60O6pz3SXKhssAE8JM9E6R38ebnilZRuiZWnEXCqsVALAEqpOUEgaDNBiTzxWo2lPHle0S35Gi2UHBFH4RXcTllL4C45r2FzNba3Fy6FVyM5G8YklR6MMSvEzlnnrpaim0vVXuFYiwnMoHrXknvFMRZMr2oN27u1LZjA+y2s+oHjPf2VNrjVHU7PaV5XtcOkTaIm2Qelf6lqOCw9FvcwzfnIPeTUnaH8M9q/1LUXNVYLROXZnvX6FP4iP2Ne75uqPi/wDVa81M1bxRyzPfP1V+I/TXhZz9oD2Rr3mR+VY5qZqYizZglh29lJJ1J1fnqHiZO0rEcEw+JLeovcwwX+h+6puB9N/ZT5vUHZJ3uKxN8eipTD2zMg7nM1xh0ecuMh/8dRn2bj0Wt+7lAgSSQBJjj7q179+qvxH6aY/gvtD5GtOatRimjkns379+qvxH6aeMP1V+I/TVDjdsNauXBcVTbRbJBUsXZr9x7aLlIgctRrP2vVXi7fBKqtm7LG0Mr5EIFz7UM0sFOhjSdJruKOWy/wB+/VX4j9NZWL5LFSACADoZ4z6h0VGzVswf8R/ZT5vXJRSQTdk2tF6/BygSYnjAA9Z9x7q31XXT5x/w/KsxVs1J0jf4w/VX4j9NPGH6q/EfpqHiMSttczTGZF06XZUX82Farm07Klg11AU9MFhye3vHZW8UZyZY+MP1V+I/TTxh+qvxH6aqm2zhw2U3rYYMFK51kMQDl7YIPYQaLtixMG6isAWKM65gonUwTpCk+49BruCGTLXxh+qvxH6azs3sxgiCNYmQR0g1BsYhbih0YMpmGGoMGD+YNbbB86vsv80rkoKgpOywqNdxEGFEkcSTAHq4HWpNVgbV/bP7ViKtmpOiT4w/VX4j9NPGH6q/EfprRmpmqmCM5M3+MP1V+I/TTxh+qvxH6a0ZqZqYIZM3b9+qvxH6a22b2adII4jt4EdIqJmrPBnlt7K/NqzKCSsKWyHZXLtK6f5mGswP/Dcv5v8AeWpF5QXcET6PHsqLtoi3iMLiD6IZrFxp9FcRlyn/AFbdpfxmpV4+cf8AD8q5Ds1Po8CkcGYe+f6pr2W6/eB+0Vjmpmq1E7MpfrD3L/c0153b9I+QmsA4MQeOo9Yr3NXKFmF9AEc8+VtSSTwPOdatl4VT4h5R46rfI1cLwqfJ+jcDKlKVM2RNo/wz2r/UKgk1N2l/DbtX+papNqYprdougBbNbUAif4lxE4SJjNMSKvx9MlPsgb3EW1JXO8eMsc9trjEWmi0iAMmrAyOObLp01oTaOMZC624aGJDWrsMLe/yhULAqz5LXGfT4aitt3wkS3vFdHL2lQuFCjMXXOAi5iSYDnLqRl10INbV28M+XdXMpYLvJtxrcFuYzZvTYDh0mtV+jJExm1cVaS7cyA27Yutla3cLGDiWU58wGWLdkRHC5x1FW+xcY12yLjlSxa4DlVkjI7KFKsTDAAA6kTMaRVR/ilTkO6fK6TkhXclzh93G7ZhEXzI1OnfZbM2kt0si23tZFttkuLkMXFzCF5gDKn1qw5q6uwyJ4QeENzCtu7dq4Wui2oxAVTbsli4BbMwBefRUwCYBPTZbN2kli0lpcLi1VRALWw5J5yxViSxMkk8STU3D4dLou27ih0dFVlYSGBzggisPBu6wR8PcYs+Gfd52mXSA1pyT6RKMoLc7K1eef+mVh0VieFCX8S2FFq6pQowuFGywVnK8gZH6FMyIM6wNzbZti5dtsGG6Cl2gH0gCCEUlyNfSyxKsJ0NW+01AUEASWWfXxqhx2yLV5i1zOxiBLtCglSQo4CSi91V41cdGJvZle2lg3c23a2zOFRgUzBxKlVJywyg3VPGAbg4ZhMdcXgdCqWipCkvulA82U3ehWW1ZMkA6gRUkbMtCNDoIHKPTbP/42+49JqNa2BZFsI+a4QqrnZmJ5GQrEnQA21IHNr0mtuLM5IsMRtVES3cAZ0uMih1ywpuMqLIYg6swEAE6HStp2i1q4wXD3rsqmtoWoGr6HPcXWoj7PtndekN0Zt5WIAJESQNCYkerMemrfZbS79ifN6zyKos7B7RX4vaOOe25sYTd8loN27bF3NBjLbVXQnhGZxrxEVr2RcxRtzjVtrehcwtlsvDTjz9NdPVLjWi6/YvyqfErkb5OiNtE22QW7lzJndAhDKGNxSHQLmBBaUmIPA1XeQcPcVzmZt7OZ4tZmMw8nd8oEiSpkAgEAQIy2ts43yp3jJkBKAKrDPKMjnNxylOAg6nWqq5sC415iWTIVhXKAxqHIKzLKWBVkmCpbUE1dw/hNS/pfXcBaCwxIBFxZLc1wLm1jotjXtqLiNk2rtt7du4QQ2rTmysbcagRINu7zH7QqLb8HkAOZlJMgkWwBkYXZtqCxy2wbui6wEA141ing4nFjbzMbefLayqyIthTbjMeQdyTBJjPzxqxfgTXpfbOsG3aW2zZysjMFC6EkgZR0AgeuJOprO497Om4FstDzvCwEcnhl55iq3ZWzVsFshEMFGVVyiVa4c2h1MOq9iDsFxs9puj2X+aVycfxYi9kHa9zae4fcrZ3mmTKSTMjmcRl6eeJjWK24B7uT/qAguz5wWs2TNAnLm1jtroaorzcu57R+QqXCrZub0U2Jxt3O8PcDLdtKtkW+QbZuWczh8hLSrPJDQNRErWm3t++ygi0JO9jMrrJW2txMwBORZJUtytRwE6Xu8pvK9GDJ5HPLtHFXLiMudUmWQ2yC2mEGnHKPO3TEtwOumkzZW18RdZA9pbYYOzyGlQot+b4wHm4RJ45SYHCrXeU3lc+b9GaIWysU73b4Lu1tWCoLiBWDAvnKwom36IWZJysZIINT8+I3h3C2jyRm3jOOdojKD66x3lS9lmXf2V+bVnkjUWdg7kUPhViMcMFiM9nDFd28vvrsroYITdcpgYgBhrFavAu5jDhEOPEXYXUznZMoyG5P244/nrNXW2xvb+Gwx9Fna/cUiQyYfKVHq87cst+A1uxR86/4flUeLspPo5xtmXlDZEthjcus91LrW7l1HNxkVnCZlylrekkebgaaVut7OxIZXa8xYMpYbxwjRcSeREAG3vBHSw6ARpuYi7vH1v5hcWLa2m3e6DJ6LhIYkSTy5ksDoBGm1tnFFQTaIJFyJw+IGqiULrqVHNAzEngIM1ZJExs/Y+KRLaM+UIiIQLrswAOGzlHKgqGW1dAUREjp03eT8ZJ86SMqhRvWEgZcykxoxhjvBryuOlR/GcU9y24FxElSbZt3FJzeKAzDQsC5e0OYDK/GCRlhdqY10zNbS22S4zK1rEaMq28qdJOZ35ShswTQcYYnbLfAYdreGKP6XnSYYv6bOw5ZALGGGpGtdQvCuas3y9gOysrNbkqy5GBK6grJyn1SYrpV4VPlXRqH7MqUpUShD2n/AAj2r/UKqLqK6lXUMp4qwBB7QeNXG0bZa0wAk6GOmCDH5VRpdB0B15xwI7QdRXp4KaaI8nZ4cHaK5Tbt5R9nIkc3NEcw7hW3dJ1V7h0z89e2sZpNXxJWarOz7KWxbW3bCARlyqRGkzI1nKszxgVutWUQkoiqTElVUTlECY4wNK8mvHuACSQPWTFMRZZbKPLudifN6jSE2lHPfw5JGv8A21wCe2MT+XqqRsYE52ggHKASImJkj1a1qxojH4Vum3iUn2jZeP8A6z3V4uT/AEz0Q6RL2pbZk5Ikgg5REkDomqZ7hX0kde1G+YEV0N+6FEkEyQABEye0itfjX3G70+qtQm4rSOSgmyjR2bREdj6lI/NoFTLOzbjauwQdC8o950HdVj419xu9Pqp419xu9Pqrr5ZPo4oRRpXZNocVLet2Zvmak2MMiAhFCzxgRNYeNfcbvT6qys3wxIggiDBy8DMcCeg1N3+zar9G+q/H4AuQyMFaIMiQw5p7JPfVhWi9eA0gkngBE9uug99ci2no60mtlV5Mvda33PTyZe61vuerPxr7jd9v6qeNfcbvt/VVfpP0nhErPJl7rW+56eTL3Wt9z1Z+Nfcbvt/VTxr7jd9v6qfSfowiVnky91rfc1Tdn4EpLO2ZjpoIAHQBW7xr7jd6fVWdm8DPEEcQeInhw0PurMpya2ajGKejdVXjdnszZkYAn0lYEgxpOnAxVpWi7fCmILHjAjQeskgCsRbT0akk1sqvJl7rW+56eTL3Wt9z1Z+Nfcbvt/VTxr7jd9v6qr9J+k8IlZ5Mvda33PTyZe61vuerPxr7jd9v6qeNfcbvt/VT6T9GESs8mXutb7mqwwGD3YMnMzek0Rw4ADmHGs/GvuN3p9VbbV0MJHNoQeIPQazKcmtmoxSeipsvm2ldH8vDWI7b1y/m/wBlKYw+db8PypgFnH4thzJhkPaouvHddXvrLaWHfMXVSykCQPSBHqPEU4mlLY5E2tEfNTNUY4hR6Ry+pgV+dei8p4MO8V7Vs89kjNTNWg3R0jvFYeMp1h2Aye4Uo5ZvvtyG7D8q6JeFc5bsvc0VCAdC7DKADxgHUn3V0leXmadUX4092e0pSoFRUe/hbb+mit2gfOpFKArTsa1zBl9l3H71j5Gt9a58ZqzpWs5emcV4Vo2Nb5zcPa7fsa32dn2kMqiz0kSe81LpXHJvtnVFI9qj8JeQLGI1ixeRmjqXA1lyfuqt3OfYq8rRisOl1Ht3FDI4Kup1DKwgg+41w6asedF9ofI1qzVzez9sut7yffl7lnJlxKnOtxADG8IHm7uWJB9LiIkCrXyjbzumeDbjOSGCKSFIBuEZc0MpiZ5Qq0FonJ7Mb+2Ft3HS4jIttbbG6ShU71mRAFBLliyFYjiR014u3rBjIzPLW1lLdxgDdGZZIXTQgnokTWGLw+HuMy3CpZxaVlzwfNu1y3AB0IZiQRx91R/FcLnBzHNlEOLr8lbJVvSDQsEqSOeT0mu4nLL3NWeDPnH9lPm9V2I2hbt5C7aOVVWCsyyxCrLKCFBLKATAkih2tasXGF1mBKqQFt3H0BfqKaTWmdi9l/Vdcbzj/hHugn9zVVivCVmR2wuGxFwBWIutaNtVIBIm3dZLj9irrzGa82Pj7t+3vMRYOHuMFzWWYMRpoZHMeg6ipwVs1Pos72IVBmYwJVZ14uwVRp0lgPfWD4u2ubNcQZYzSyjLm4TJ0nmmo2Os7y3kzFTmRgwAMFHVxoeOqiqm54Po4uHevN2czr970wIOikgGBHAAkjSqtE7RftjLYOU3EBkDKWWZPARPE14mOtn7ajQmCQpCrxJU6geuql9iWyCCdTvBmyrM3AgnhxAtiKj43YKvbZbb8uZzaAg5GTKWAJXS4SOgwaYizo7d5WAZWDA8GUgg9hFZWT51fZf5pVbsqy1uyltyCVnVZjiSNTqxgiTpJkwJitl7E3FdDbtG6YcFQyrA5Ost/wA1rklo6nsv6rM3Kf2z+QA/aqna+3MXbsPct4J86xlBZHDHMBGVDm1mJHCZMgGpGBxL3Ez3LZtOxlrTMrFCQJUsuhqfGrZqfRY568zVzmI2vcDXIa0Ft3bVrdmS/Le0GcnMAOTcJAy8MpnWKxt+E6MoK23M7wCD6TJbW4qrIGrq2kwNDBOk2pGTps9M9cl/iK49xRbCbsnlEgyBGFECWBOuIOpA4ARoZmbN8IBfa2EtXALgZgzcmEUIc0MASfOqIE6zqeNc0c2dBnrPCNy29lfm1UWz9ou9+9afLFvIyFRBKuXGvKYHW3x5JmeSIBO3GpfuE2bMKtxQt2+Wg20ls2RRxuMCQDIC8dYg8ktM1F7JXgxy0u4gcMRed1MyCiBbVth6mt2kb8VXdasPZVEVFAVVAVVHABRAA9wrXcxEMVCzABOoHGY+VQSso2b2UHiAa0tg7Z4oh/CKx8ZPV/UK14hi6MmVhmBGZXysJESrDUH1itUzlorvBfB2vFgcifxMRrlH8+7V4tpRwUDsAFc94K4G9hcPu7ztecu7FywAGZjAVRoukExxYseerrxk9X9QpTFolUqI2KI1K6SNZHOY/epdZaaCdntK8r2h0UpSgFKUoBSlKAUpSgK/aYAUGB6S/vXNY/You3HubxkLKFIQBcwlTyypBeMpgyCMzCYMV0u153cgEwykwCTHYKphil6w7CYPca9PEk0R5HTIQ2BazBvtD7WVJmcOQZiZ/wCmTvPQKjWPBpd0qPcOYLbWUAUA2zbZYywTraEnQkHm0q3OJXhmBPQNT3DWpNnC3W4KEHWfj7lH7kVuWK7Mxcn0Vt7ZKtbtW80LaZGEKCZRlYFWYkoZBE66Mavdm6u/YnzevE2QPtXHPZCj8hP51LwmDW3OWdYksxJ04ce01Gc01SKRg07ZJqmxjRdfsX5Vc1WbRwbM2e3BMAMpMTEwQenU1jjklLZ2abWjnduYK7eNs22Rd0TcXOHM3FK7vVSMojOCeVo/A1WeLYsXSilggHJyvcVNWDtLDkiRmQEDMGaeAFdP4pf/AJY+NaeKX/5Y+Na9DcPSX5eFENl3mXzlwkmQBvbhCowuwJgS4z2xn4nIT242Nk3lLEXCmdredReunkqmHRoJHp+aucrjDDXov/FL/wDLHxrTxS//ACx8a0uHo/Lwg7Lw1y2W3lxnUhYzO7kMGuTq3AZTaHap55Jt8A03R7L/ADSo3il/+WPjWp+zcGykvcjNEBQZgc8nnJgd1Z5JRxpM7GMrLKqO83Lue0fkKvKp8dgrmctbAYNqVJggxEg840FS4pJPZSabWiKUQtmKrm0GaBMAyBPbrWIw9vhu059Mq/a0PNzjjWfit/8Al/rWnit/+X+ta9OUPSNS8PFRBwVR2Aer6V7h0UtoimVVQTJJAAkmJOnTA7q98Vvfy/1rTxW//L/WtMoeipeBFVZyqBJkwAJPSY4mpuzDLv7K/NqheK3/AOX+tastm4RkBZ4zNGg4KBMCefidanySjjSZuEZXssKrr5863sp83qwqrxZ86fZT5vUYK2Un0bM1M1ctYfGq2aQVe4wFt1ZyqIcQ06BMmcCygksBE65qjrtTGOLQZCma5azFLGJkqLlnOrFo3YCtdlmkMBpwNWowdhmpmrk7O0MUQgOa2Q1oMDhr7nK6LLF80Ny21jVQusQTTD7SxhCru4bLaBz2bxy5hZzO1zMA5Be6CmhGSSdDIydRfbk+9f6hVrXO4W67WkNwAPIDQCokPEhSSQDExJ410QqXIuikGK9pSpmxSlKAUpSgFKUoBSlKAVqeyrekqntANbaUBqSyq+ioHYAK20pQClKUApSlAKUpQClKUApSlAKUpQClKUApSlAKUpQHlU+0Gi6fZT5vVxVfjNnZ2z52UwBACmYk849ZrfG1GVsxNNrRX56Z6k+Rv81/hT+1PI3+a/wp/avR9oE8JEbPTPUryN/mt3J/ankb/NbuT+1PtAfOREZ+HtL/AFCuhFVKbHggm45ggxCawZ5hVtUeWak1RuEWuxXteUqRQ9pSlAKUpQClKUApSlAKUpQClKUApSlAKUpQClKUApSlAKUpQClKUApSlAKUpQClKUApSlAKUpQClKUApSlAKUpQH//Z)","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm import tqdm\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import classification_report\n\nimport cv2\nfrom PIL import Image, ImageOps, ImageFilter\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms.functional as F\nfrom torchvision.utils import make_grid\nfrom torchvision import models\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import GradScaler, autocast","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-08-04T18:27:44.041385Z","iopub.execute_input":"2022-08-04T18:27:44.041843Z","iopub.status.idle":"2022-08-04T18:27:44.050869Z","shell.execute_reply.started":"2022-08-04T18:27:44.041809Z","shell.execute_reply":"2022-08-04T18:27:44.049739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class cfg():\n    num_classes = 2\n    epochs = 12\n    batch_size = 80\n    learning_rate = 5e-3\n    device = 'cuda'\n    input_size = 380\n    \nmodel_dict = {'resnet101': models.resnet101, \n              'vgg19bn': models.vgg19_bn, \n              'densenet161': models.densenet161,\n              'squeezenet11': models.squeezenet1_1,\n              'efficientnetb4': models.efficientnet_b4,\n              'convnextlarge': models.convnext_large}","metadata":{"execution":{"iopub.status.busy":"2022-08-04T18:27:44.258171Z","iopub.execute_input":"2022-08-04T18:27:44.258919Z","iopub.status.idle":"2022-08-04T18:27:44.264846Z","shell.execute_reply.started":"2022-08-04T18:27:44.258883Z","shell.execute_reply":"2022-08-04T18:27:44.26389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The dataset was preprocessed in another notebook and loaded here","metadata":{}},{"cell_type":"code","source":"path = '/kaggle/input/'\ndf = pd.DataFrame({'path': glob(path + 'mayo-clinic-eda-and-preprocessing/train_patches/*.png')}).assign(\n    image_id=lambda x: x['path'].str.split('/').str.get(-1).str[:-4]\n)\nprint(df.shape)\ntarget_mapping = {'CE': 0, 'LAA': 1}\ntrain_df = pd.read_csv(path + 'mayo-clinic-strip-ai/train.csv')\ndf = df.merge(train_df, on='image_id').assign(\n        target=lambda x: x['label'].map(target_mapping)\n)\nprint(df['label'].value_counts())","metadata":{"execution":{"iopub.status.busy":"2022-08-04T18:27:44.665693Z","iopub.execute_input":"2022-08-04T18:27:44.66653Z","iopub.status.idle":"2022-08-04T18:27:45.009063Z","shell.execute_reply.started":"2022-08-04T18:27:44.66648Z","shell.execute_reply":"2022-08-04T18:27:45.007022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#upsample = int(df['label'].value_counts().diff().abs().dropna().iloc[0])\n#cutmix_df = pd.concat([pd.DataFrame({'path': 'cutmix', 'image_id': -1, 'center_id': -1, 'patient_id': 'a', 'image_num': -1, 'label': 'LAA', 'target': 1}, index=[0])]*upsample).reset_index(drop=True)\n\ndef cut_mix(df, lamb=0.3):\n    smp1 = df.sample(1).iloc[0]['path']\n    smp2 = df.sample(1).iloc[0]['path']\n\n    W = H = cfg.input_size\n    cut_rat = np.sqrt(1. - lamb)\n    cut_w = int(W * cut_rat)\n    cut_h = int(H * cut_rat)\n\n    # uniform\n    cx = np.random.randint(W)\n    cy = np.random.randint(H)\n\n    x1 = np.clip(cx - cut_w // 2, 0, W)\n    y1 = np.clip(cy - cut_h // 2, 0, H)\n    x2 = np.clip(cx + cut_w // 2, 0, W)\n    y2 = np.clip(cy + cut_h // 2, 0, H)\n\n    img1 =  cv2.resize(np.array(Image.open(smp1)), (W, H))\n    img2 =  cv2.resize(np.array(Image.open(smp2)), (W,H))\n    temp = cv2.resize(img1[y1:y2, x1:x2], (W//2, H//2))\n    height, width, _ = temp.shape\n    choice = np.random.choice(2)\n    if choice == 0: \n        img2[:height, :width] = temp #top left\n    else:\n        img2[height:, width:] = temp = temp #bottom right\n    return img2\n","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-04T18:38:53.566977Z","iopub.execute_input":"2022-08-04T18:38:53.567543Z","iopub.status.idle":"2022-08-04T18:38:53.5815Z","shell.execute_reply.started":"2022-08-04T18:38:53.567487Z","shell.execute_reply":"2022-08-04T18:38:53.58038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"This class will read the image and create the batches","metadata":{}},{"cell_type":"code","source":"class FlowFromDataFrame(Dataset):\n    def __init__(self, df, img_transforms, x='path', y='target', test=False):\n        super(FlowFromDataFrame, self).__init__()\n        self.df = df\n        self.img_transforms = img_transforms\n        self.x, self.y = x, y\n        self.test = False\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        image = np.array(Image.open(row[self.x]))    \n        aug = self.img_transforms(image=image)\n        image = aug['image']\n\n        if self.test:\n            return image\n        else:\n            return (image, row[self.y])","metadata":{"execution":{"iopub.status.busy":"2022-08-04T18:27:46.884015Z","iopub.execute_input":"2022-08-04T18:27:46.884475Z","iopub.status.idle":"2022-08-04T18:27:46.89332Z","shell.execute_reply.started":"2022-08-04T18:27:46.884431Z","shell.execute_reply":"2022-08-04T18:27:46.892003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here we use albumentations,for data augmentation","metadata":{}},{"cell_type":"code","source":"image_net = dict(mean = [0.485, 0.456, 0.406], std = [0.229, 0.224, 0.225])\nnormalize = image_net\ninput_size = cfg.input_size\n\ndata_transforms = {\n        'train': A.Compose([\n            A.Resize(input_size, input_size),\n            A.RandomRotate90(p=0.5),\n            A.Flip(p=0.5),\n            A.Transpose(p=0.5),\n            A.GaussNoise(p=0.5),\n            A.Normalize(**normalize),\n            ToTensorV2()\n        ]),\n\n        'valid': A.Compose([\n            A.Resize(input_size, input_size),\n            A.Normalize(**normalize),\n            ToTensorV2()\n        ])\n    }","metadata":{"execution":{"iopub.status.busy":"2022-08-04T18:28:42.192386Z","iopub.execute_input":"2022-08-04T18:28:42.193224Z","iopub.status.idle":"2022-08-04T18:28:42.201777Z","shell.execute_reply.started":"2022-08-04T18:28:42.193186Z","shell.execute_reply":"2022-08-04T18:28:42.200298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here is where all the magic happens, the method ensemble reads the model_dict remove the last layer and add the classifier.\nThe backbone for all models is frozen","metadata":{}},{"cell_type":"code","source":"def get_classifier(in_features):\n    classifier = nn.Sequential(\n                    nn.Linear(in_features=in_features, out_features=1024),\n                    nn.ReLU(inplace=True),\n                    #nn.BatchNorm1d(1024),\n                    nn.Dropout(0.3),\n                    nn.Linear(1024, 512),\n                    nn.ReLU(inplace=True),\n                    #nn.BatchNorm1d(512),\n                    nn.Dropout(0.3),\n                    nn.Linear(512, cfg.num_classes)\n    )\n    return classifier\n\ndef freeze_backbone(model):\n    for idx, (name, param) in enumerate(model.named_parameters()):\n        param.requires_grad = False\n        \ndef ensemble(model_name):\n    model = model_dict[model_name](pretrained=True)\n    freeze_backbone(model)\n\n    if model_name == 'resnet101':\n        in_features = model.fc.in_features\n        model.fc = get_classifier(in_features)\n    elif model_name == 'vgg19bn':\n        model.classifier[6] = nn.Linear(4096, cfg.num_classes)\n    elif model_name == 'densenet161':\n        in_features = model.classifier.in_features\n        model.classifier = get_classifier(in_features)\n    elif model_name == 'squeezenet11':\n        model.classifier[1] = nn.Conv2d(512, cfg.num_classes, kernel_size=(1,1), stride=(1,1))\n    elif model_name == 'efficientnetb4':\n        in_features = model.classifier[-1].in_features\n        model.classifier = get_classifier(in_features)\n    elif model_name == 'convnextlarge':\n        model.classifier[2] = nn.Linear(1536, cfg.num_classes)\n\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-08-03T01:03:44.201638Z","iopub.execute_input":"2022-08-03T01:03:44.203074Z","iopub.status.idle":"2022-08-03T01:03:44.217222Z","shell.execute_reply.started":"2022-08-03T01:03:44.203024Z","shell.execute_reply":"2022-08-03T01:03:44.215519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def loss_optimizer(model, dataloaders, batch_size, num_classes):\n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=cfg.learning_rate)\n    #scheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, 1e-1, steps_per_epoch=len(dataloaders['train']), epochs=cfg.epochs)\n    scheduler = None\n    return criterion, optimizer, scheduler","metadata":{"execution":{"iopub.status.busy":"2022-08-03T01:03:45.581131Z","iopub.execute_input":"2022-08-03T01:03:45.582683Z","iopub.status.idle":"2022-08-03T01:03:45.590959Z","shell.execute_reply.started":"2022-08-03T01:03:45.582627Z","shell.execute_reply":"2022-08-03T01:03:45.58942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, criterion, optimizer, device, dataloaders, num_epochs=10, scheduler=None, model_name=None):\n    print('Training: {}'.format(model_name))\n    history = []\n    best_acc = 0\n    scaler = GradScaler()\n\n    for epoch in range(num_epochs):\n        print('Epoch {}/{}'.format(epoch+1, num_epochs))\n        print('-' * 75)\n        lr_rate = []\n\n        for phase in ['train', 'valid']:\n            if phase == 'train':\n\n                model.train()\n            else:\n                model.eval()\n\n            batch_acc = 0\n            batch_loss = 0\n\n            for inputs, labels in tqdm(dataloaders[phase], leave=False):\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n                optimizer.zero_grad()\n\n                with autocast():\n                    outputs = model(inputs)\n                    loss = criterion(outputs, labels)\n                \n                if phase == 'train':\n                    scaler.scale(loss).backward()\n                    scaler.step(optimizer)\n                    scaler.update()\n                    #loss.backward()\n                    #optimizer.step()\n                    if scheduler is not None:\n                        scheduler.step()\n                    lr_rate.append(optimizer.param_groups[0]['lr'])\n\n                _, preds = torch.max(outputs, 1)\n                batch_loss += loss.item() * inputs.size(0)\n                batch_acc += torch.sum(preds == labels.data)\n\n            epoch_loss = batch_loss / len(dataloaders[phase].dataset)\n            epoch_acc = batch_acc.double() / len(dataloaders[phase].dataset)\n            epoch_lr = np.mean(lr_rate)\n\n            #if phase == 'valid':\n                #scheduler.step(epoch_acc)\n\n            if phase == 'valid' and epoch_acc > best_acc:\n                print('Saving model: Last best: \\t{:.4f} <  Val acc: \\t{:.4f}' .format(best_acc, epoch_acc))\n                best_acc = epoch_acc\n                torch.save(model, '/kaggle/working/{}_best.pt'.format(model_name))\n\n            history.append({phase + '_loss': epoch_loss, phase + '_acc': epoch_acc, 'lr': epoch_lr, 'model': model_name})\n            print('{}-loss: {:.4f}, acc: {:.4f} lr: {:.6f}'.format(phase, epoch_loss, epoch_acc, epoch_lr))\n\n    return parse_history(history)","metadata":{"execution":{"iopub.status.busy":"2022-08-03T01:03:46.961101Z","iopub.execute_input":"2022-08-03T01:03:46.962663Z","iopub.status.idle":"2022-08-03T01:03:46.978662Z","shell.execute_reply.started":"2022-08-03T01:03:46.962607Z","shell.execute_reply":"2022-08-03T01:03:46.977047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def parse_history(history):\n    df = pd.DataFrame(history)\n    \n    df[['valid_loss', 'valid_acc']] = df[['valid_loss', 'valid_acc']].bfill()\n    df = df.dropna().reset_index(drop=True)\n    df['train_acc'] = df['train_acc'].map(lambda x: x.item())\n    df['valid_acc'] = df['valid_acc'].map(lambda x: x.item())\n    return df","metadata":{"execution":{"iopub.status.busy":"2022-08-03T01:02:10.083412Z","iopub.execute_input":"2022-08-03T01:02:10.083861Z","iopub.status.idle":"2022-08-03T01:02:10.096201Z","shell.execute_reply.started":"2022-08-03T01:02:10.083824Z","shell.execute_reply":"2022-08-03T01:02:10.0945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train, valid = train_test_split(df.index)\ntrain_df = df.loc[train].reset_index(drop=True)\nvalid_df = df.loc[valid].reset_index(drop=True)\n\ndf_dataset = FlowFromDataFrame(train_df, img_transforms=data_transforms['train'])\ntrain_dataloader = DataLoader(df_dataset, batch_size=cfg.batch_size, shuffle=True, num_workers=2, pin_memory=True)\n\ndf_dataset = FlowFromDataFrame(valid_df, img_transforms=data_transforms['valid'])\nvalid_dataloader = DataLoader(df_dataset, batch_size=cfg.batch_size, shuffle=False, num_workers=2, pin_memory=True)\n\ndataloaders = {'train': train_dataloader, 'valid': valid_dataloader}","metadata":{"execution":{"iopub.status.busy":"2022-08-04T18:30:01.75021Z","iopub.execute_input":"2022-08-04T18:30:01.750708Z","iopub.status.idle":"2022-08-04T18:30:01.765042Z","shell.execute_reply.started":"2022-08-04T18:30:01.750664Z","shell.execute_reply":"2022-08-04T18:30:01.763696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_batch(dataloader):\n    rows = cols = int(cfg.batch_size  ** 0.5)-1\n    img, lbl = next(iter(dataloader))\n    fig, ax = plt.subplots(4, 4, figsize=(16, 10), squeeze=True)\n    for idx, i in enumerate(ax.flatten()):\n        i.imshow(img[idx].permute(1, 2, 0).numpy())\n        i.set(xticklabels=[], yticklabels=[], xticks=[], yticks=[])\n        i.set_title(lbl[idx].item())\n    plt.tight_layout()\nplot_batch(dataloaders['train'])","metadata":{"execution":{"iopub.status.busy":"2022-08-03T01:03:57.516144Z","iopub.execute_input":"2022-08-03T01:03:57.516599Z","iopub.status.idle":"2022-08-03T01:04:00.593499Z","shell.execute_reply.started":"2022-08-03T01:03:57.516562Z","shell.execute_reply":"2022-08-03T01:04:00.592212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Train the set of models and save the best epochs by validation accuracy","metadata":{}},{"cell_type":"code","source":"ensemble_train_df = pd.DataFrame()\n\nfor k, _ in model_dict.items():\n    model = ensemble(k)\n    model.to(cfg.device)\n    \n    criterion, optimizer, scheduler = loss_optimizer(model, dataloaders, cfg.batch_size, cfg.num_classes)\n    history_df = train_model(model, criterion, optimizer, cfg.device, dataloaders, \n                             num_epochs=cfg.epochs, scheduler=scheduler, model_name=k)\n    ensemble_train_df = pd.concat([ensemble_train_df, history_df])\n    \nensemble_train_df = ensemble_train_df.assign(epoch=list(range(cfg.epochs))*len(model_dict)).reset_index(drop=True)\nensemble_train_df['epoch'] += 1\nensemble_train_df","metadata":{"execution":{"iopub.status.busy":"2022-08-03T00:28:32.815693Z","iopub.execute_input":"2022-08-03T00:28:32.81607Z","iopub.status.idle":"2022-08-03T00:47:54.994056Z","shell.execute_reply.started":"2022-08-03T00:28:32.816035Z","shell.execute_reply":"2022-08-03T00:47:54.991876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here we can see the performance of each model","metadata":{}},{"cell_type":"code","source":"fig, ax = plt.subplots(figsize=(18, 8))\nsns.lineplot(data=ensemble_train_df.reset_index(drop=True), x='epoch', y='valid_acc', hue='model', ax=ax)\nax.grid(True)","metadata":{"execution":{"iopub.status.busy":"2022-08-03T01:01:08.904058Z","iopub.execute_input":"2022-08-03T01:01:08.905275Z","iopub.status.idle":"2022-08-03T01:01:09.322342Z","shell.execute_reply.started":"2022-08-03T01:01:08.905224Z","shell.execute_reply":"2022-08-03T01:01:09.320859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Load the best models and make predictions","metadata":{}},{"cell_type":"code","source":"def evaluate_ensemble():\n    y_pred = pd.DataFrame()\n    softmax = nn.Softmax(1)\n    for i in glob('/kaggle/working/*.pt'):\n        model = torch.load(i)\n        model_name = i.split('/')[-1].split('_')[0]\n        print('Evaluating: {}'.format(model_name))\n        with torch.no_grad():\n            model.eval()\n            for inputs, labels in tqdm(dataloaders['valid'], leave=False):\n                inputs = inputs.to(cfg.device)\n                labels = labels.to(cfg.device)\n\n                with autocast():\n                    outputs = model(inputs)\n                    outputs = softmax(outputs)\n        y_pred[model_name + '_CE'] = outputs.cpu().numpy()[:, 0]\n        y_pred[model_name + '_LAA'] = outputs.cpu().numpy()[:, 1]\n            \n    return y_pred\n\nensemble_eval_df = evaluate_ensemble()","metadata":{"execution":{"iopub.status.busy":"2022-08-03T00:53:42.491003Z","iopub.execute_input":"2022-08-03T00:53:42.491952Z","iopub.status.idle":"2022-08-03T00:56:33.298337Z","shell.execute_reply.started":"2022-08-03T00:53:42.491907Z","shell.execute_reply":"2022-08-03T00:56:33.296588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ensemble_eval_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-03T00:56:33.301698Z","iopub.execute_input":"2022-08-03T00:56:33.302074Z","iopub.status.idle":"2022-08-03T00:56:33.327102Z","shell.execute_reply.started":"2022-08-03T00:56:33.30204Z","shell.execute_reply":"2022-08-03T00:56:33.325894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Here we aggregate all of the model outputs to make predictions","metadata":{}},{"cell_type":"code","source":"ensemble_eval_df['CE'] = ensemble_eval_df.filter(like='CE').mean(1)\nensemble_eval_df['LAA'] = ensemble_eval_df.filter(like='LAA').mean(1)\nensemble_eval_df['y_pred'] = ensemble_eval_df[['CE', 'LAA']].idxmax(1).map(target_mapping)\nensemble_eval_df['y_true'] = valid_df['target']","metadata":{"execution":{"iopub.status.busy":"2022-08-03T00:59:15.666679Z","iopub.execute_input":"2022-08-03T00:59:15.667779Z","iopub.status.idle":"2022-08-03T00:59:15.680347Z","shell.execute_reply.started":"2022-08-03T00:59:15.667733Z","shell.execute_reply":"2022-08-03T00:59:15.679158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ensemble_eval_df[['CE', 'LAA']].head()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Finally we can see how the ensemble performed","metadata":{}},{"cell_type":"code","source":"print(classification_report(y_true=ensemble_eval_df['y_true'], y_pred = ensemble_eval_df['y_pred'], target_names=list(target_mapping.keys())))","metadata":{"execution":{"iopub.status.busy":"2022-08-03T01:00:45.815504Z","iopub.execute_input":"2022-08-03T01:00:45.817075Z","iopub.status.idle":"2022-08-03T01:00:45.834992Z","shell.execute_reply.started":"2022-08-03T01:00:45.817013Z","shell.execute_reply":"2022-08-03T01:00:45.832787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Conclusion\n\nAs you can see a lot of improvments can be made\n   - Improve preprocessing of the images\n   - Better data augmentation\n   - Hyperparameter tuning","metadata":{}}]}