{"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":"# About this notebook  \n- PyTorch resnext50_32x4d starter code  \n- StratifiedKFold 5 folds  \n\nIf this notebook is helpful, feel free to upvote :)","metadata":{}},{"cell_type":"markdown","source":"# EXPERIMENT 1 : Baseline model    \n-Backbone : resnet50                                                  \n-No TTA                 \n-5 folds training                 \n-Resolution :256               \n-Final accuracy : **0.87372**                                      \n-Data : only use competition dataset.                                                                      -scheduler='CosineAnnealingWarmRestarts'                                   \n-Loss,score,data----> compute average parameter\n![image.png](attachment:b515a416-1565-4ef1-9f64-4b9803102c8d.png)                 \n ","metadata":{},"attachments":{"b515a416-1565-4ef1-9f64-4b9803102c8d.png":{"image/png":"iVBORw0KGgoAAAANSUhEUgAAAsEAAACBCAYAAAAhdmcxAAAAAXNSR0IArs4c6QAAAARnQU1BAACxjwv8YQUAAAAJcEhZcwAAEnQAABJ0Ad5mH3gAADDVSURBVHhe7Z1faxxH1v9Pfi/BMLAiFmx04QXbBDL4QSissCABoYT1hSML3zgXEoLYs0sUdCsZE+lWRPskci6MdGHfGFsxxGEjBDHIaIkQT3YCQTasL4YF2SjsgN+Cf3XqT3dV9f/pGWmk+X5Cx5qpnuru6lNV3zp1uuutNwJqNAgAAAAAAIBe4f/pfwEAAAAAAOgZIIIBAAAAAEDPAREMAAAAAAB6DohgAAAAAADQc2SI4Ge0MjJCA952/clrnd5p1PFXnuuPLWOu4x7V9Tchr+nxLev6bv1ETZ0C2sjze7J8y99Lj+ZPdN3cu07kfxh411CsfkXraLQM8th4Wh1hvDzuPtPfg+aT+bBcEssvhZNgw7p+t3YNeWxYkFVOZdM7jq5DXdnHePfAq9+ujYst5hrqd630VuoBAEdALk/w7O0tamyF27cfnNIpxwDZONepf35Yf+FSv3uZZgZW9LU9omVapMEe6+BlA3dcxX/lQ/pW3rsVmtVfHS9E53Nlkc6bOvZwjmjhcoEO+hzVdL2U2+0pWrrhdkCZNp5RR5R4uEwbw4/C43x6TieCygcLQdm3xHG3YbaPGy9o+aG2DWmD8/Q4d4OSbcOyjbrylMbMMcRWO6sTBWXT24EUgce076jfrREF/byww7WaMxgPbFxuqg25aaVz+Y435mhX77M+uUrjcCiBY8AJD4cQAuNBn6iY12hIf+MgGu87a8O0PGY69FN0aUJ0ZGv3CzTgIBdnr8nGsd0dz3Gn+eQ+LY3M0YQpFyGIpieJlh602IFU+miUXtBL8+NMG8+oI4L6xiLR/KPjNfgFhwbbx+bkVbpU0V+c/ZiWR7ZpZqNFQejbsLDRBwskRPZCeAyHsumHhah7t4RIvPUhHelpxFD91G6bz9GEGBBvbtcT2qBTdHqAaPPVgf7M5btNsxPhdVXH5mh0a5EeHLq3HYBilBPB7EHi0Z4zFeZNg3jTZHEjZXcaJW462J6qKehhSGlwmr89pc2RizQU7PCaHj9YFf9u037uY5THv/7ACyjLzr9ePaUWlKM3TW22HB4J6R0R+w6KBoxEgzUY/D48ptyH87Lvo5932j3OmII0+Ztzify+DTh5iy1iX945Rmw4K53RZVAslOE17Wxv0+hwNbRRKVrFv1sHtK++KYRv09k2nl5HuO7tCBE99m4nBbBvw579RcpbtQdhWbdeB/LVsZQ6ekhk2rBfB30vXFa6wByj2LWxfRDNXrBmBp7/g2a2xL+Ng5YGchGbfV4XA0Xbhj3KppfE2MY419u1mlXOod069uOVvSz3uz8FNrzy3PR3tl16Nh7rZdW/i01rJ9495/KlKRqyHBxyYCT+3TsI7TS7nwfg8MklgpduuMbrNJIsnqQnSU+TjKzSuOk8uOG1p8k4vSEaCa9zsadRePM9Tks37lO/zmN9cptmVtpYyQf6lACQQucy7U/wNbiVt6OIMtq5EF77rhiBB1OB2qOy8Zt1Ls06bWyFnj051U32NJT4cnIl13S1meLiY9KIfQ88jwk37DeI1mWamioLbUA0vL9U9e/E9nCORu30PFO9Yv/BV1fjf18S7mAGF87oc1fncX7hstUAi/O3wxHkdo2qOjU7vTzn+5S9y87wygFNcxk4nrAMLJEur9UXtWVsvHlAe3SGTjdtEeWLxnI0n/yDqBaWr13HKx9cFXazSju2PehOd1q3E2XqQJ46llpHD4FMG+b76rSzYrNtICu9NMPUrzOTQof7Aw4NKTKQS7Hh5sELacP7CUKybHpZ2IvKZRrYnSljq50w+8i2No61RVEvVR6qv+M6auySBbAd0iT2GxD9bsfErvLsOoNzQShia7Q3/8id1Rvpo375hxLid95+JK/FeIvZhrP6eQCOgpZigt0p7SmrwdJTrdoDUP9lVTQK1jRZMBVbVx2IaPjkVG0tvUGevR2KsuqFgo1rDgLxYV2bESYd5+w1pzwr7160BNApGhp2p6WklyQo09f0siHKx56G4vJp0QOTjLjHQYMuGjtHQJ2jmi02KlUaKzqIYAFu8mjl94loT+v8x5ZojZ/qW/ol3WuYlW7CPVpr2JWXRw0ETDkL4ZlWKWyCgYbYHvbRHdEJ+YOIcjYuBrbWQGd3nmjmSvtEYOWDa86gy63j52iIhYFV/tyuhPe0bB3IqmOC1DraafLasCfkI2Sl831Qg2K3fc9LKH4CgR0Ioxxk2bAYGIcDESEQ/bj2sulHjR0SJW1PhRxI5KBMtMFWO6vCDZ7SjmODOra65ACH44PtQabBCHnepl9djg4k5GyDcliZdnD07T75ryRyvgAcPZ2JCZYdmOqcnErA2PFezQPaLNLZd4LAC2nEhzh3nsrLi+XBUFtRL5nqPILfX1HTSAbpCQsaD9UhhlOPqqEM40eFmOKpbuP5axdOZ6bi2myx505zXVZToV2Bupe+2Kv0nXFEVk17t801uNN0WenlWbrB3lnRuZhOTtaLFjExxbZoL2vjPAiyOuBY72wZ/Dp0g8M1QmSHbwbOor644Rnl60B6HWPS62hnyWHDLCD1A5XmHB0BmZVemm0xKHLFj/S+tkqcDU+uWOLcc6YwZdO7GdkeiIGovndy65ANyplZdkw9DL3YcUREeDAjbBxWqv838ABLDZ7NNRzeTAoAabRdBMvGT4omP3heYwvfyAMQh4vy6NgPDQnk9K8b35SK7cGQW5GHL9gDKEbd9hSanAq3YU+Y9uKwR4Asj4EhiOfV08K2Z7bDyGmuNfYUm+vnaTydeOT4XmtFaKMG7UHR5R99O0NWeqsoL6TjBRLI85uspnZCWZjBZ2kbF3X0fGwdDafAy6HCTfjBu6AO+G9ZkLMDWnRzKIQzu6QpVQfS6lieOtpJctqw3Q7FvZ0hK71llKfenfEj2n8VnU4vSmDDLPjjPPv6+sumdz2yn7TbWLO190E/1ZazAM6br+7Hz1ZlqJs9G6MGb+6zBPYbJvD2CNAttFkEu7FEclrSedOC6PBuWFOZsnNrc4xvESLHFx3eCj/pXE6A5Ed5eUJvuT6+/mTgUTdt16n+21OhPeyORXnFysb6Ka9Sa1NV3NnZnUnzyVdd5AnWU90L/wi9Ds2f6Kb3JLODFH0pJKXrB4+KeomlSBUCLnjdkDk/xxOpOijpQcmawhXnwR1Z0PmUtnElEO0n/dUbLdr0oJEU5LanU7URLspzx55BDoXwvbTtqAPJdSxfHe0crdlwqkhPSDc2VnSAp9p5K47ft0FNyzbMIkvUkfBNA+IePBD9iLlPZdPbRP/bYkDbCe+yHgQGz9okomcsWhGXosz5AWk79DAZvw3RIUtWnLwMqUhpI2RZAdAFvPVGQA1r3sKBKxXHB7mMmtclccfvdVgcP2zHlHHDJ98+oAl+GyAqFAf9W8Ip3Ecdn99fGOQpj8kPaaVP1yiieSvs0a63D3t8DtGT6pfh7O0VIvlghN0YmXO0Y3MVfvkqovtloabB9AerfGT+2xdpN1FYeDYiym+dajI2kO9h/PmF9ziav7pWfietayfxZOXPuNfm2ShPxXtTi46NZqUb9H2MTcvCO4Zfh5jgOn37jNTBuHufZuMt1BGO4S4Zd2jj3kNx3NsXaeMGxy/b16HtLObY8TZQtA4k17H0Opqn/NIpbcMRGyiYrjHnEZeWiXOM+GsvZcNeHYnUs7LpbcG3BXMd8f2oSe/nctFtID+8N06qfPiem3Y0tp5G6kJyHUkn6fyMnUSPHVd+jo165+Dbb/H6CUBnyBDBGcjGK68gBW1HNuz8Ani7w9ENFj9JfJhiHoCjAHUAAABAi3TmwThwOMj4ah9/+haAEwzqAAAAgBaBJ/iYEzeVqqaqDhKnuAwtTXseItEpNI/DDl0BXUn31oHkaWZDt9dBAAA4yZQTwQAAAAAAABxDEA4BAAAAAAB6DohgAAAAAADQc0AEAwAAAACAniNDBPODHbwKk7u1e9nYZNTxy6/OZa4jbqlGfp2SdX1HsoqNOYd2reLUY/ADmub+ia2wvXi/t22AH85z0mL2kZg84l5on3p+0Trmnz8/+BWmJ9gIvyosMQ/vGIkv3Y+vJ3FlYLcBWekhSfXQLwP/GrPqaPr1ueUXvw8AAIDeI5cnmJ9gDlZjElv7XzLeQaQAqVP/fPwKNfW7+n2i8toe0TIt0uBhd5B6qdZls3QryA+LvxsvKFgxrOiSsGwf8g0nxr5dG6h+ar4Pt3VeJnagz3kR/MAvffHLRcv8rfN7OEd7zvlZSzLzJs/fEoni+m7yi/R1ulp/3xWRUuTJd+WG+YRvHGCBWKO9ebMssbi+Ri1WpKpVnoZjVxPjty2YvHnz24CsdCY+f3V+vCCO+a1/jVl1lPMNf79Cs2sx18dvEtH5yw1vFQEAgJ7nhIdDiA72QZ8QENdoSH/jIATGHV5ydcx0iGp5Vnep587T1Eu1XrowRZvb9SPwRB9f6huLtDl5NVwo4ezHQoy6y/ym0Tx44Sz7zDZwekD/GYdvM0Lk8qpOjU/jXxLIy/ySfX6VD2naW4bYQS5p+4JeGiMQ+39rrbxU+eAqzdIq7QSeXl6qnJJXJ3tepyWaoulAlCobd5bhZeR1TdF67aL+os0k5S+XTR6mfuvc5TLehhx1lAcqoeg/RxNiwIt6BAAAIItyIpi9XDw1Kb2tZqrRm+p00sQW42X1p1OjXip7urNIyMA5qqUsH8nic9NZ3/y1XFOeaJv2D60HfU0729t0vk+IFBZAW09pxz+2EAL2VHekjBPT1TSyU55y37AMpReR74l9n7x75N+faLiBN10d7KPum38/5TG9KW0zZR3NO41ntLNGNHvB8uo9/4da3rNxkEsESVG5tUiD5nxEOYw7ostFim7bZs5ey5wZ8RdtkOvmJ5xf1CYzYJGbtb8j8gW+0Ob7t7JINP9xh973nZK/GRRcMTYpbOYGLz+t9u2OOgoAAOAkkksEL92IEzgaFhDS28rTjI9oeWSVxo2IYmFlTwXrqVhbZLHAGm/MBdO9vPmiYkmu06/S1rnDXPFjAktgprWlOLxM+xN8DUR7B9Hp4o7AoRBbUzTEnqxKlcZGvJAIPq8ri3TeCUmxFifJSs/DmrgnQUiAmk4O7rG4hzsXwrx354fd6XopgO3parUpz1ycV449l9s0O1Fkbfs0Qi+iFOtsi7enhF0e0L76OgMVjrA+IOyY7VuWQ4JX1YjuAufOgtf1uqrrd7AGMYMLZ2g9beD25L707Ep74c/syRY2vG8PVOwBhhxYLdIDq85KIa//lvDAgeboyxQxv7lwOT5/TWp6Rv4y5OT2GSGE+fcqtMFpAwrVUVW+o8NVtwzZxs35RWKSAQAA9CItxQS7KxxNWZ22nqrUXq7IVHAwlVlXnZDo1ORUZy1dVMzeDkVJ9UIRgZMPFVN5QNPWtUnP7CEgPV2TVS1aT9HQsCsa1XT/ilfmIVnp+RD3MBDOfdRvC4yz15y8K+9edL2I0vMqfp8QYyn3t73benreiDhD5YOFGNvKi/I4y7AEY4u+9zMR5cU2A7H1STGIE3nFeaSlAB2Zo4kC58jXZfJUAiwmPp1DHnTdajzsozsJx+f6cpMFnu9RFQIvHKgIgWjHzHLeMs7YHH+Edt6es+Jylec1Tdi7cdE6f0vopqdn5y/rnxksm3P1ZiPy1lEZd+yEf4S2ZTZ1PyCEAQCg1+lMTLAUqa/pZSM6FexMxcp1/8/Q6TQF3GmEgBh8dVV0jkYEinPn6fS8WF48tRUJ11ChEHYZyXjIQOQnlGFAVnpOvJjYS7dsT5wSmMH1XXG9iNGYWg893W282zwwioi4UvBUupopMOcszyknzSdfKRGvxTMLOn7wbemB7+1s3YPtisRrdPrVtvNgnYMsL3H8X1wRKO2My14MeCLhF84gyBtoMmIgY4vAWl9Y76SwLzSI0vknDkTd9Mz8tbAPBrp8ruzJt+Pyc9ZROavEg+qH6TMh1TEeBNjhIAAAAHqRtovgUBSpB4w2Xx3oFI0tfCOxiYeL8mp68Z/yQZ2opzIR24snt6Sp9BhkKIQ3lXyD4x3Ng08JZRiQlV4W9pLWpIgJru+h7UXkW2iL9njYe6+82xxO4HrpynGOhvhNDc5sA9G+EJmR6fAEeF9fxMuYXQ8/DKF1YuKYY3AHNmIgogWw/1YDWf5x8cUpAxM1Q8OzD2oQ5oQKyEGO9lx73lhD1iAjTM+Rv2wP3AfjVLug/8xZR9lTrARwjvonjwkAAKDXabMIduPxZOiC86YFNTUaeAJ1DGxbY3yLEDm+eoAnDE/oLOqhHzcemjfpidSeQFWGVoyuR3q6L5K1mNKfslEet1CQ6fLRnyRnq/JtBUEceBz8xgZapJu37tNeghe4tQfjYq5fP9g29q4rtM3Dff5DevL3zsOIOmbX8dRqj31pDzZ71dWgItEzGjl/PRBhO4kLOeHyd2J+xf4PRB1LGAQob6kJXzklvf627alBDofHiL/jjpcUkmFw0nPkL+3HjYNXMctmoJyjjooyG7S9yamoNsgfOAEAAOg93nojoEZDf/TRnbb+ZOB3gsopWfngG3suQzh+2O7gWdxwB2UIfhvAnfxl9US/JtxHHZ8flAnylMfkh5fyPPwVzVthe4y8fWK8bZ2CBQnHsUamt/1rjJQziwjr+tPSzTS6/J6v+yrty/ABdf3y/vB7aJMexvLynr29QiQfVLQFR9RO4u0g+VVexk783+XCOcd4b6ASf3H2F7VR3wZU+pl4m3PK14JFK5eplx65vox7Gzk3g8mf//aO4V5jQfuWeXHsrTkHvw755ZuV7hHJXxApQ8++U68hvo1iTFmbe29oycYAAACcODJEcAayA88rSEEvI8VcmtgGAAAAADhEOvNgHAA2TZ4ip8y3gAAAAAAAHBYQwaBz8EyBfhjqfK54TQAAAACAw6FcOAQAAAAAAADHECWCAQAAAAAA6CEQDgEAAAAAAHqOnhDB9a8HaGDAbCs9tlxqkx7/1br+vz4+mncyg3z8ukIDXx83C63TSs/VK9Ayx9LGTzrcT1ynx7/rj+B4cyLrWGf6mQwRzAe1BJTerv8gZBQX8kBcpVGiS+5jkPuK30ZuSsy+neKLdWo0GmKrxbzOzVxntIBdAR29hqz0CLIs/OMYoRotz+YP1938E0VsUh4VuvQNX3eDdv9ur/WWg98f03X72Lx1Y8Uy9qW3lV/193nxrjNij2XT88LX8ckeLV/xLLQt+SfbuMLYj96s+5xt41Wa+Psejbc4wIrkz5uVl5N+HAZxsXU8g7I2nEakHieLHdPedKxNTrLxbkCXU2eu3atfbbfjsvmLfmLqPM2836LIaMXmu42SddBtxxLKwjlGcn+faoMmj6S+OKmOlepHolowtny8Y0T3Se5nFGnp5fqZRDgmOJl/vfnmnXfefFPXHx1U2meP/6s/aw6+f/PZO5+9+f5Afxb8639FHo/5+2/Er2z+++b7WkwebYaP/87/ukcOqH/z5h1xXt8//kz+m7CXJq08mKx0dS7O9Zryeszn4ZZblOTy+i+ff+2zSNnbqH2+F7nkJHIv1fEL5cHE2ETb8POW97PIsbx7JvOz72HZ9Lwk2U4b8s+ycZ1nvnqYdJ7aNpLqWQqp9dOisP0eFbq8c5eELP8yNlyMxHLU5/FZx9rkJNvpEgrVg2K4Nt56XUmiXfnLfFqpY0VtvtsoWwd5f6u848pRfpdYRqZv/17WkWQbVHXoM9HXx9/fDvYjNjH3W7YraWWWVb9y1b/2150S4RBVGvqCaHNzx1Hlzf/boM2PxmjoD/oLMYLY+WqWhv4yRGMfLdFOOz0cpRGjm9V+2m3UaEh/k04/9X9EtPcyaRySkS5GSXe+GqWx/zHvChOjnsV9mm58S5dO669SqdDpP4ky/8++/myo04PPiZbnpum8/qYzsFd5nWZ/nKGbnfIUFaT+cIY2v5imS8be3pug5Y82aeZhPp9E84c7tPTRMk28p7/4wyWaFna9tKpGm2XT8xLJR1M+/2wb5zKkv+/St3/J8w67JBsXtjG3TKNf3Un0MoJ4ytpwUSqnRSvx4z5FWhFpB1/StGhjOkGSjXcNom5922jkrAcF0O1+6Jljr+ssUbvqShvzr/5Nte8Puqqf7jyl6+B7NWr8LfS8Vq+ItvDHDdox5S/v0Sytx85Ec924SftTbHvpQkDWoS/W6cuESd3O9SMeff00Snv0MrAvrUF+Flom0H4uWf1MVrqi/f1MqZjg6p9FRfMa0/3/bNLo6JA4Vc2vO6LwRccpvmEBt/TPVhp27Ypv+xRSlWrfXArPNYvfd2jjR1vEemSkywGCXdH4hn4TXyni4QEF0eyf3V/Uvx4XFcPOt5PowY8txL1ppGAKw0yNvC8aGPHfzPvhPs5UTNLvM4kpj18f0MyP4t9/v8xhK03a2fTsVTZW4l9p12XT86LymZ3ybbEd+WfZOJdhik37pNn4H3igu0kb/9feWpoHP6QinIZT02srv7rTbP5UoP/79PToVKeT/kncIs5JZNkwt33X6fEPuo6INrBuQqRy1xOX+j/F+X0x5LY7og6Oi056OqkD8uto4bY4ycYZfY2/2/fImyo2bYnewvuj7+8PJl3cG3Ou9jl6v3fPX/cveotOE5vzs/eLTmUnEXUMiXNeZRvZpP0D8Ynv59ePg2tf+dUcJ98xsvJXuPavjqOTHFT73lo/nYZ//Oi1FauDcfdJl6VIKxbKULYfyUb1/V6ds6j85VuqZQ0OhQ3fZKGZGErUyX7EJWJzUufZNuiT1c8U6Ifa3M/kEsFLn7jGFxjYe0M0S7Z3VxnT+dPhhXCDawpfiuavdsRexwmr8goxd/47f6STlW7gkZIwUE/A5iGs/OO0J0ZKTmWRxixGmNYo9FAIGgfRYP9zSMdbi+1nHqWNKxvRnhX5nfhv+We9j+NtSfl9Lkapv0/9JcuJvZ7fRQdnaRh7lQ3o+/s0Lc83HOWWTc9EC0tzHT6l80/j95cip/N0+sAWOX4HldfGKzQ0OhozU5EDcc/D44utiMATomfnz6Ftcez70ieuUF36ZFB6WuQ+wj42P78ZXCOX6/i/l2nX2KDYbG8E25WdvstxaZYQ9tM5/2Jk2bAYQG6K7/i+/zhD4/+ZVsco0pZaIlaeq9NeiPsrRNPsdwkDchaQHGNo1d9GEecBk2Hj8hrfH6SN0V1dxkQzi1qo8vGl3Znjr9P5zwcdEbT0Oc+orcv+aJzLT5aV8cSJ63tI9KU5d97Pmc0SA0X5/S4tf6S/isDnd4f6dRmsfyE+m/PLw59Oq/LiaxlgW1THCmZUvpqR9rnOAvQTPg6nF+joM/Kvfz1IM38yz8WoLUl0yX66TeJPwe2Hd/zvOP44bGey6mBWennK9yM20rNsiULpHPwjuQOBgoNY5Sn9khKdXZ3uR6RtqXMf/Pw8rVttQPPlnrTBfS47c33OIDSjn8nVDxlK9DMx5BLBs0Hjo7aw8nijRh4NiGZoKEj31H1ENOdFN1JFG962wN5ac+271L8qbo5jvFnpGjlSam0qsPo3k3+Dpv8zaBmXaFwWuWJMJI4wO4+4N3aHKkdpaSEjPmV/zyjPyZ0/ig7U2IicfciLEnmDLC6C6SpRIYPGpmx6Bgf7ootN279k/pkI4WANRKQAcR6QyWnjhlY60ODBVb0VGdS9V3M69Mr/jEUa91F78KinOkMvmcCeunRQbdjyXNj2VP4ybbVj0fTWSLfh0Lujp73ldGQBeLrWlO3UPg1aHQxPxc5QVttU0vOSaePqHhlhY4dsKC/autPvTIiBjh2KN2q1gbKs/nBaHM0g7Pdv9v2Jmc3Kwaw1+IubBc0iEB/iHphrMcKE7L5BzuqpmdMipObPFBk0tSj+YtHizPFg6jro2FRiHdRkpQvYo8o2brcH+Snbj2jkrIpdZxWbn98hmjNtnBiIiYF/nDc7FplnykwN0+l+xDi1ePu5n+6IsnKcVeJ6QmeEGISJVmXQ6Sey+pmsdI82DdRKhUMwtnc3Ms0mRbFdyEnxhMcF0ZjKWKukxiQpXRgfe1pipwKL4cQa8ZSNMLUv2zoizqJJL/8t/jGeB4EcOcuRG29ixM/TSAUo9/vQQ2M6UDkqLUDgJTTCSzYmIWXTy9Lp/EVz7cwkuCLPJ6sOHAWq8wpsSIbfZGPaIe44VYNr8rAaXumhYBuz8h8YF6WjkellKW/DhXAEiHmeIKVt4s7v52Wiz8UAXJdBsenmfDiCTYp21VGzF82fKRj8vFgNkALR+j2LlENFnL8rPoTILNhOppKRPztS1r8QIsOUQeFwlhLEijP3+ZbUOijISi9Pm+ogz7h8suQOujWjjhc3OpBLRuuHpJmanLS1HzExxXbYjDNQjesnsvqZIv1Q+ygtgpVHgk9UhULY0/1SFIu0oOJpgZPvxncxWaNDP12OhG0PeVlUgyLL90cx2grKlztn3WF3qpGT1xLeZzlNJQP+ReWSW9qUYpRyv1ceHeU5Ud8wkbj0RNS0iuOFEcjGTw7myqbnJPKQgaFN+achPWZxx06buhak1QFrgNR52LvBMfGWJ1lO86WhBIItuowHiTcpFkz9keXjhvKYrTVvk09ZG24def3SUeGKfBaIm1LwWtORthfoO9E1fZI0VZlAoo1n0/9HcTf9mQLe8s4MCmEyKB/aCX/LYQeHhZqZ8DyhcvDUnj4hb/7hjGKcl86jFQ9oErH3XjlTRv8YHiWxDmqy0lunTXVQhg0pAeyHarANx8485GkrdZ9rh6XKQaAeGAYD0iPoR8z9kzM3cZ5ZY0dZ/Uwr/VCb+pnyIlhPXy99wgLMD4UQYskLpZAdVI5pDRft6TnM0Wss4jzYyBMrRny6iuVpR8iC6PAX+SlWZbh2mITaOCZOd9gdCR0R18deNmvEJ700VoMpp1YtD4REGnj8dGqu36egZiKsGGI5bRQNsDfeZn/6SXYgdnygaMhuWrHbZdNzkRLonzf/wNOV1rHFwh2AEEHWU9DqSeKkhxzS6oB+AMPq2DqPErThMXUd0Z/ikNeXIkCk6ArQ5ZMU/+nfO+0JKkJeG24Xqo7p67fDJPTGApE78ga/tSbOBmRnW5AUG89C1gG7fAoiO3vRAgWeSFm++u/DQF97aENuO16awvmnh1pI50o7B7Lm/Jw2Rtlg0vS+WwejJKWbdrCorZTtR+y+0RfAjLJh+40GBZ4RsgegepPv/NcDw2AwnlLH2t6P+OXDoa4i//CtIsIGV+1+IqufKdIPtbmf0a9KS0C9W+4db/Pf46beD+e9u02+8y3unXjmfXich37nm5d/9B1w+jxafEeo+w5Fm4TjB++6i16/+169rHSG90l+d556d2B0U/lEzy/9HXrpx2rtPcFZx/fKQJQzX1NkP/lewbh8cv4+DSfv+Os35Rybr3edkXtYNj0HqfcmR/6xdVCSZeOMt49zHnlsXCPPM9n+koitA8E5RI+vNqtt8Wzrm7pdD2Ku3yvn6PGj7VbqPs79Ed/Lz3FtXwqJNqyuX5a5Xb4FjhHYhtky2oBI/fPKl7f22rh1jUk4Zaw2dY7q/qq/7XxSbECcw7/4XExdibk+uQV1Keb85G+K3GPvHKx6apeLtDOdFrkPqSTnH0njLdEGctyLOGLLMGrHYZpbdll1MCvdYGy9FftMroMh5jz8+xI9P73Z98Ero6g9Wb8zm3MfQ+R1xqQl1zGBV4fiysiUXyTvyPnFlL+Xf9R2o/XQ3SMrXWO3g23gLf6f1sMnFh69jZMYNVnxJocFj6xkrNYRHNtHnsvmGO12xEsMylGnFf32jzhPQvcjRv5/HZRP93fX+XfrefUix93GTz5H2VeCdnDS61j72/Py4RAghdZfiwZ6jSrVvnNf3XWcMG8YONyHNMHx4njb+IlHTnG7DyeB48bJrmOd6Gd6RwTrIPL2P1WahjBIO2bnSOCRk4pjKvpENThkOD6T35+Zd5WiroEHe+57IwGI5dja+EmHYzj5XdDl3kAAuoATW8c608/0RDgEAAAAAAAANkoENxr6IwAAAAAAACcfxAQDAAAAAICeAyIYAAAAAAD0HBDBAAAAAACg54AIBgAAAAAAPQdEMAAAAAAA6DkgggEAAAAAQM8BEQwAAAAAAHoOiGAAAAAAANBzQAR3Dc9oZWSEBsx295n+Pp3mk/nwNyP3speEfn4v334tYM7l+pPX+pvjgFvux+vcAQAAANAqEMFdwWt6fKtGS5Mr1NjaUtun53RaOpUPFtT+t6f0N0dE8ye6uUA0OqI/HxvOUU2W+SNaTjx3JZRXnuuPAAAAADj2QAR3BQe0v0U0eyGf8O1G6huLRPNf0PSA/gIAAAAAoIuBCO562EscTtcPjMzT46ZOykn9rvX7G6v6WxtzjBbDJJ7fo/G1KZr+4JT+ogDNn+i6PK4dluCdhwzhMGkxHlkvfeDWT2SKSIZoWJ/NteYLezDnVKMl8WnphnWMnOEqAAAAAOhO3nojoEZDf3RhATG4sK0/WYzM0e6tD4k6nF5hgXRlkTb11yHDtPxwgS5Rl6dX9MckWLzFilLzexZsl2lmYCUMj5C/eRHNX35PtL51jar6K4YF8HhDlyd/EbufPs7WVOT32ajf7k9sUe2sOt6dtx/Rt3kFsXWPZ2+HeYyTvmb/evX+5/W+6vNTGksob2nD2xfD69fnuzHsn2PS9wyL4RqROWYBUIdKpsfcUwAAAKAdpIpgcFgkiKxYgZcg1mLFLed7n/rt3yeI5VbxRWarIjgQtfyVled+TH6OSNa/p/n4Yx61CAYAAABAdwJPcCfTc3uxEkRWmnfXiEBD3L5xIrqtIjgqslsTwUmeXCVMZ7b0Rxt+iNBcv2cnxqPMHLUIRh0qmZ67DgEAAADFgCe4K0gQWbECMUGsHYUIlnnFhXMwOUVMqghuQVTLcwrDJ+AJBgAAAEAceDCum6lUaWxkm2Y2woewmk++krG7uR5C07/f+E0/BJYoWlkA8gNfBR+MO3stfKWb3tYniUbnH4m/cwjgHFQvTNHmwlf5Hwas9NGo/pOp9J0h2jqgff25fjfBs5xKH/WPEC39gofhAAAAgJMCRHBXc4ou3Vqh2bVa8FaCwYUzlhfXiFexSXG7SuNyP/MGCfH72hzRwmW9D9H6wzlHJHY9LLRvn6GZK/o69Ra8IYKFvfX9wBWOL7YE+NmPaXnElIvyKrNQDwh+r8Txpikr5+0PqhxHrfuAt0MAAAAAxxuEQwAAAAAAgJ4DnmAAAAAAANBzQAQDAAAAAICeAyIYAAAAAAD0HBDBAAAAAACg54AIBgAAAAAAPQdEMAAAAAAA6DmyX5EWWXb1pC1nqlYKCxZQMMvN6o/ZqNXElvQnp3wSl6wNl/aNLKvrHb9sOgAAAAAAiJIugvUKY0awnUR4Wd5xWqHGp+fEJy2IB8znLKLL6SpRai9o4ZG6THDa0r1M2XQAAAAAAMCkhkM0D15Iz+LECRXALEjvrA3T8pgRvKfo0sQU0dr9fMv0Ng9oj4ap3xKzcpneFOobi7Q5eTXBk36KTg8Qbb460J99yqYDAAAAAAAmOyZ4a5EemCVqE2BvarCcrNiuP3mtUwT+srbOcrPsSeUlftmDafYxS/4a7DSx3fqJovqU80lKS6b521PaHLlIQ4EgFcd6wMsPb9N+nowqH9L05DbNXDHnLM7jxiqNzn+c6AW+s0Y0eyHJy/yMdjqaDgAAAAAAmFQRXPlggdYniZZuKAG6EiOGZThBY452t7aoobdgKl6GU7yg5Ycm7REtN2qeEGYRqabweZ/deaKZFSNm1fS+DE/Qea8PLNJgQbGbykCfip/lMIWRy7Q/Ic5xhGjvwBLyKVQ/Fed1+4y4Bi4jFRqRFIqgRHfUsx4OImq0N/8oEnpSNh0AAAAAALhkeoKlyGPxGYjhe1TXaUE4QS3+Qaz6L6tEztS/CTeoh3kIRoVwM8JRhhNsHdA+f2jWaWNritat+Nzq2ByNbj2lHUcFn6Mai+QWHwjjON6BKwc0LfIwAvJ8X76YWvlbI/RvT6kyckS+4Rk9WNim2YnoOZoy5m361eWIR7tsOgAAAAAAcMn9ijQltNhLukrjRuQ1D2iTztDpWOX5ml42hMB9u09/1lT6aJRe0EtLpTmC8+w1cRz9UJnMXxxPejn1lvC2hZZZq9Hgq6vhMUkIcPOmiCzEIOAmC9vb+iE3PnchhONiiptP7tNSjvjqeJEfUjYdAAAAAAAUEMGKUzQ0PEzUOFCexhhBG5LwkFaqcPaQ+U/RuvZyhlt7XtFWefeiyN9+ME4gH3aboqE8IQXyWtwH49Q5+ygv8OhwNaenOqt8yqYDAAAAAPQ2BUWwJ+YqVRob2bZieF2qF3yvaMaDYz4yf8vznEhrD8ZFz/81PV7htzdU3fOT8cLsibZCQZizVZqlbdr4LYwflm9/8EXo8zotCWE9nfnasoTjB5RNBwAAAAAATOp7guVDb2v6g8aO31Xoh9esEAJ7H38xB/f3LF7d9+xGieYfXRBC5cPhBsUXivDyn4x5RzCLYBmGwV5p7/2/QZrB30edGz+wFn1gLr3syqcDAAAAAIA4sleMAwAAAAAA4IRRMBwCAAAAAACA4w9EMAAAAAAA6DkgggEAAAAAQM8BEQwAAAAAAHoOiGAAAAAAANBzQAQDAAAAAICeAyK4GwgW49CbXByE3wFs/vZ4fi+6cAcAAAAAAMgNRPBRoxfboPlH4bLQF+q08lwvUb1Wj4jd+i+rRFgVDgAAAACgZSCCj5TXepnjFXeVt7PX5Ap6lXcv0iit0s5z/b3kGe2sEc1e8Fa1AwAAAAAAuYEIPkqaddrYGqblsQRBW6nS2AjR3sFr/YWgeUB7NEVDictMAwAAAACALFKXTW4+mafBhW39yWJkjnZvfUjU4fSKDhXY1F+HCOH4cIEuUZenV/THJDi298aL1H3lPdi+qMoj5nM6z2hlpEZL+pPN7O0tqp0tm05UvztC42v6S5vJFWp8eq7r07NsPLuMAQAAAHAcSRXBoMPkEMEqZvgpjcl9+GG5y7Qx/MgNnwAAAAAAAIWAJ7iT6VluREfg6u8iWML33XqO/W3gCc5KhycYAAAA6E3gCT5SlMjcm0/37Eoh15ij9eGnNJ47FAIAAAAAACSBB+OOlHM0MT9MmwuXacV+A8Tze87n6oUpoq1FGl/YptHhKgQwAAAAAEBJ4AnuBmRs8Kr+INBT9SEmLCFnmAUAAAAAAEgFIhgAAAAAAPQcCIcAAAAAAAA9B0QwAAAAAADoOSCCAQAAAABAzwERDAAAAAAAeg6IYAAAAAAA0HNABAMAAAAAgJ4j+xVpkaWLT9q7atWyxDNb+mPh5XL9pYWt8klc9jlcdjiybK99/By/jxw/8o5hAAAAAADgky6C9SIOoeA6ecglickIRy2IB/IKSSVAySofJWrP0PrWNaqqr1yksH1KY7EDCXX8jeGUZZS93/P571xwBXHWMswAAAAAAL1OajhE8+CF9ExOnFABzILyztowLY8ZwXuKLk1MEa3dp8dN/VUazQPao2Hqt8Rspe+M/iue+sYibU5eTfCkn6LTA0Sbrw705yj+76uf2gMUvQzzdp3ynD4AAAAAQK+SHRO8tUgPnuu/E2Bv5MBIuF1/8lqnCNibbKUN3H2mExj2XM4LwckeULMPf9bJEjtNbLd+ihF4nE9SWjLN357S5shFGgoEqTjWA16+eJv282RU+ZCmJ7dp5oo5Z3EeN1ZpdP7jRC/wnTWi2QtJXuZntJOWnvl7AAAAAACQh1QRXPlggdYniZZuKAG6EiOGZThBY452t7aoobdgKl6GU7yg5Ycm7REtN2qeEGYRqUIAeJ/deaKZFSNmrfAEnff6wCINFhS7qQz0hfG3I5dpf0Kc4wjR3oEl5FNgT2zj9hlxDVxGKjQiKRRBie6oZz0cRKhQhqTQk6TfhzyjBwvbNDpcLRDTDAAAAADQe2R6gqXIY/EZiOF7VNdpQThBLf5Bsvovq0TO1L8JN6iHeQhGrRhWGU6wdUD7/KFZp42tKVq34nOrY3M0uvWUdhwVfI5qLJILPdAWwnG8A1cOaFrkYQTo+b58MbXyt0bo355SZeSIfIMSqLMT0XM0Zczb9KvLid7upN8b6nf5AbkpmkY8MAAAAABAKrlfkaaEGntJV2nciLzmAW3SGTodq8pe08uGELhv9+nPmkofjdILemmpPEdwnr0mjqMfKpP5i+OZUAjeEt6W0DJrNRp8dTU8JgkBbt4UkYUYBNxkYXpbP+TG5y6EcFxMcfPJfVrKEV8dL/Kzfy898jwgeZjwQB4AAAAAAAjILYIVp2hoeJiocaA8lTGCNiThIa9U4ewh85+ide0lDbf2vKKt8u5Fkb/9YJxAPuw2RUMZYlUir8V9ME6ds0/RMAW/fNJ/z95oJYDbUy4AAAAAACedgiLYE2OVKo2NbFsxvC7VC75XNOPBMR+Zv+V5TqS1B+Oi5/+aHq/w2xeq7vnJeGH2RFuhIMzZKs3SNm38FsYPy7c3+CL2eT1nmELC8dN+//yefM9w4I0GAAAAAACZpL4nWE2x6w8aO35XoR9es0II7H38xSDc37N4dd+zGyWaf3RBC5UPhwsUW+iC8fKPW2yCRbAMw2CvtBduEKQZ/H3UucW/uze97BRpv9fXrT/ZnOR3OwMAAAAAlCV7xTgAAAAAAABOGAXDIQAAAAAAADj+QAQDAAAAAICeAyIYAAAAAAD0HComGAAAAAAAgB4CnmAAAAAAANBjEP1/e+AvEh15lngAAAAASUVORK5CYII="}}},{"cell_type":"markdown","source":"# Data Loading","metadata":{}},{"cell_type":"code","source":"import os\n\nimport pandas as pd\n\nfrom matplotlib import pyplot as plt\nimport seaborn as sns","metadata":{"execution":{"iopub.status.busy":"2022-11-29T10:02:02.246264Z","iopub.execute_input":"2022-11-29T10:02:02.246598Z","iopub.status.idle":"2022-11-29T10:02:03.623735Z","shell.execute_reply.started":"2022-11-29T10:02:02.246567Z","shell.execute_reply":"2022-11-29T10:02:03.622809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"os.listdir('../input/cassava-leaf-disease-classification')","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:05.513533Z","iopub.execute_input":"2022-11-29T09:58:05.513907Z","iopub.status.idle":"2022-11-29T09:58:05.521875Z","shell.execute_reply.started":"2022-11-29T09:58:05.513867Z","shell.execute_reply":"2022-11-29T09:58:05.521069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\ntest = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')\nlabel_map = pd.read_json('../input/cassava-leaf-disease-classification/label_num_to_disease_map.json', \n                         orient='index')\ndisplay(train.head())\ndisplay(test.head())\ndisplay(label_map)","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:05.523116Z","iopub.execute_input":"2022-11-29T09:58:05.523719Z","iopub.status.idle":"2022-11-29T09:58:05.942519Z","shell.execute_reply.started":"2022-11-29T09:58:05.523585Z","shell.execute_reply":"2022-11-29T09:58:05.941673Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.distplot(train['label'], kde=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:05.944213Z","iopub.execute_input":"2022-11-29T09:58:05.944601Z","iopub.status.idle":"2022-11-29T09:58:06.238866Z","shell.execute_reply.started":"2022-11-29T09:58:05.944555Z","shell.execute_reply":"2022-11-29T09:58:06.238041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Directory settings","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = './'\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nTRAIN_PATH = '../input/cassava-leaf-disease-classification/train_images'\nTEST_PATH = '../input/cassava-leaf-disease-classification/test_images'","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:06.241918Z","iopub.execute_input":"2022-11-29T09:58:06.242629Z","iopub.status.idle":"2022-11-29T09:58:06.249046Z","shell.execute_reply.started":"2022-11-29T09:58:06.242580Z","shell.execute_reply":"2022-11-29T09:58:06.248158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CFG","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    debug=False\n    apex=False\n    print_freq=100\n    num_workers=4\n    model_name='resnext50_32x4d'\n    size=512\n    scheduler='CosineAnnealingWarmRestarts' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts']\n    epochs=10\n    #factor=0.2 # ReduceLROnPlateau\n    #patience=4 # ReduceLROnPlateau\n    #eps=1e-6 # ReduceLROnPlateau\n    #T_max=10 # CosineAnnealingLR\n    T_0=10 # CosineAnnealingWarmRestarts\n    lr=1e-4\n    min_lr=1e-6\n    batch_size=32\n    weight_decay=1e-6\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    seed=42\n    target_size=5\n    target_col='label'\n    n_fold=5\n    trn_fold=[0, 1, 2, 3, 4]\n    train=True\n    inference=False\n    \nif CFG.debug:\n    CFG.epochs = 1\n    train = train.sample(n=1000, random_state=CFG.seed).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:06.251091Z","iopub.execute_input":"2022-11-29T09:58:06.251513Z","iopub.status.idle":"2022-11-29T09:58:06.262791Z","shell.execute_reply.started":"2022-11-29T09:58:06.251424Z","shell.execute_reply":"2022-11-29T09:58:06.261921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Library","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\nsys.path.append('../input/pytorch-image-models/pytorch-image-models-master')\n\nimport os\nimport math\nimport time\nimport random\nimport shutil\nfrom pathlib import Path\nfrom contextlib import contextmanager\nfrom collections import defaultdict, Counter\n\nimport scipy as sp\nimport numpy as np\nimport pandas as pd\n\nfrom sklearn import preprocessing\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.model_selection import StratifiedKFold\n\nfrom tqdm.auto import tqdm\nfrom functools import partial\n\nimport cv2\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nimport torchvision.models as models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\n\nfrom albumentations import (\n    Compose, OneOf, Normalize, Resize, RandomResizedCrop, RandomCrop, HorizontalFlip, VerticalFlip, \n    RandomBrightness, RandomContrast, RandomBrightnessContrast, Rotate, ShiftScaleRotate, Cutout, \n    IAAAdditiveGaussianNoise, Transpose\n    )\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import ImageOnlyTransform\n\nimport timm\n\nimport warnings \nwarnings.filterwarnings('ignore')\n\nif CFG.apex:\n    from apex import amp\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:06.265683Z","iopub.execute_input":"2022-11-29T09:58:06.265925Z","iopub.status.idle":"2022-11-29T09:58:10.081159Z","shell.execute_reply.started":"2022-11-29T09:58:06.265901Z","shell.execute_reply":"2022-11-29T09:58:10.079757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    return accuracy_score(y_true, y_pred)\n\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    LOGGER.info(f'[{name}] start')\n    yield\n    LOGGER.info(f'[{name}] done in {time.time() - t0:.0f} s.')\n\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_torch(seed=CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:10.094140Z","iopub.execute_input":"2022-11-29T09:58:10.096254Z","iopub.status.idle":"2022-11-29T09:58:10.131738Z","shell.execute_reply.started":"2022-11-29T09:58:10.096211Z","shell.execute_reply":"2022-11-29T09:58:10.130252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CV split","metadata":{}},{"cell_type":"code","source":"folds = train.copy()\nFold = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed)\nfor n, (train_index, val_index) in enumerate(Fold.split(folds, folds[CFG.target_col])):\n    folds.loc[val_index, 'fold'] = int(n)\nfolds['fold'] = folds['fold'].astype(int)\nprint(folds.groupby(['fold', CFG.target_col]).size())","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:10.133547Z","iopub.execute_input":"2022-11-29T09:58:10.134197Z","iopub.status.idle":"2022-11-29T09:58:10.178967Z","shell.execute_reply.started":"2022-11-29T09:58:10.134151Z","shell.execute_reply":"2022-11-29T09:58:10.177876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass TrainDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['image_id'].values\n        self.labels = df['label'].values\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.file_names[idx]\n        file_path = f'{TRAIN_PATH}/{file_name}'\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n        label = torch.tensor(self.labels[idx]).long()\n        return image, label\n    \n\nclass TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['image_id'].values\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.file_names[idx]\n        file_path = f'{TEST_PATH}/{file_name}'\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n        return image","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:10.180106Z","iopub.execute_input":"2022-11-29T09:58:10.180567Z","iopub.status.idle":"2022-11-29T09:58:10.194925Z","shell.execute_reply.started":"2022-11-29T09:58:10.180533Z","shell.execute_reply":"2022-11-29T09:58:10.194149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TrainDataset(train, transform=None)\n\nfor i in range(1):\n    image, label = train_dataset[i]\n    plt.imshow(image)\n    plt.title(f'label: {label}')\n    plt.show() ","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:10.198045Z","iopub.execute_input":"2022-11-29T09:58:10.198463Z","iopub.status.idle":"2022-11-29T09:58:10.573811Z","shell.execute_reply.started":"2022-11-29T09:58:10.198433Z","shell.execute_reply":"2022-11-29T09:58:10.572440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Transforms","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data):\n    \n    if data == 'train':\n        return Compose([\n            #Resize(CFG.size, CFG.size),\n            RandomResizedCrop(CFG.size, CFG.size),\n            Transpose(p=0.5),\n            HorizontalFlip(p=0.5),\n            VerticalFlip(p=0.5),\n            ShiftScaleRotate(p=0.5),\n            Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n\n    elif data == 'valid':\n        return Compose([\n            Resize(CFG.size, CFG.size),\n            Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:10.576760Z","iopub.execute_input":"2022-11-29T09:58:10.579614Z","iopub.status.idle":"2022-11-29T09:58:10.592968Z","shell.execute_reply.started":"2022-11-29T09:58:10.579571Z","shell.execute_reply":"2022-11-29T09:58:10.589260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TrainDataset(train, transform=get_transforms(data='train'))\n\nfor i in range(1):\n    image, label = train_dataset[i]\n    plt.imshow(image[0])\n    plt.title(f'label: {label}')\n    plt.show() ","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:10.594295Z","iopub.execute_input":"2022-11-29T09:58:10.594777Z","iopub.status.idle":"2022-11-29T09:58:10.864037Z","shell.execute_reply.started":"2022-11-29T09:58:10.594741Z","shell.execute_reply":"2022-11-29T09:58:10.863261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MODEL","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# MODEL\n# ====================================================\nclass CustomResNext(nn.Module):\n    def __init__(self, model_name='resnext50_32x4d', pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        n_features = self.model.fc.in_features\n        self.model.fc = nn.Linear(n_features, CFG.target_size)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:10.868090Z","iopub.execute_input":"2022-11-29T09:58:10.868669Z","iopub.status.idle":"2022-11-29T09:58:10.878942Z","shell.execute_reply.started":"2022-11-29T09:58:10.868630Z","shell.execute_reply":"2022-11-29T09:58:10.877890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CustomResNext(model_name=CFG.model_name, pretrained=False)\ntrain_dataset = TrainDataset(train, transform=get_transforms(data='train'))\ntrain_loader = DataLoader(train_dataset, batch_size=4, shuffle=True,\n                          num_workers=4, pin_memory=True, drop_last=True)\n\nfor image, label in train_loader:\n    output = model(image)\n    print(output)\n    break","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:10.880175Z","iopub.execute_input":"2022-11-29T09:58:10.880666Z","iopub.status.idle":"2022-11-29T09:58:22.169956Z","shell.execute_reply.started":"2022-11-29T09:58:10.880632Z","shell.execute_reply":"2022-11-29T09:58:22.168876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n\ndef asMinutes(s):\n    m = math.floor(s / 60)\n    s -= m * 60\n    return '%dm %ds' % (m, s)\n\n\ndef timeSince(since, percent):\n    now = time.time()\n    s = now - since\n    es = s / (percent)\n    rs = es - s\n    return '%s (remain %s)' % (asMinutes(s), asMinutes(rs))\n\n\ndef train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device):\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to train mode\n    model.train()\n    start = end = time.time()\n    global_step = 0\n    for step, (images, labels) in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        y_preds = model(images)\n        loss = criterion(y_preds, labels)\n        # record loss\n        losses.update(loss.item(), batch_size)\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        if CFG.apex:\n            with amp.scale_loss(loss, optimizer) as scaled_loss:\n                scaled_loss.backward()\n        else:\n            loss.backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            optimizer.step()\n            optimizer.zero_grad()\n            global_step += 1\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Data {data_time.val:.3f} ({data_time.avg:.3f}) '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f}  '\n                  #'LR: {lr:.6f}  '\n                  .format(\n                   epoch+1, step, len(train_loader), batch_time=batch_time,\n                   data_time=data_time, loss=losses,\n                   remain=timeSince(start, float(step+1)/len(train_loader)),\n                   grad_norm=grad_norm,\n                   #lr=scheduler.get_lr()[0],\n                   ))\n    return losses.avg\n\n\ndef valid_fn(valid_loader, model, criterion, device):\n    batch_time = AverageMeter()\n    data_time = AverageMeter()\n    losses = AverageMeter()\n    scores = AverageMeter()\n    # switch to evaluation mode\n    model.eval()\n    preds = []\n    start = end = time.time()\n    for step, (images, labels) in enumerate(valid_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        # compute loss\n        with torch.no_grad():\n            y_preds = model(images)\n        loss = criterion(y_preds, labels)\n        losses.update(loss.item(), batch_size)\n        # record accuracy\n        preds.append(y_preds.softmax(1).to('cpu').numpy())\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Data {data_time.val:.3f} ({data_time.avg:.3f}) '\n                  'Elapsed {remain:s} '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(\n                   step, len(valid_loader), batch_time=batch_time,\n                   data_time=data_time, loss=losses,\n                   remain=timeSince(start, float(step+1)/len(valid_loader)),\n                   ))\n    predictions = np.concatenate(preds)\n    return losses.avg, predictions\n\n\ndef inference(model, states, test_loader, device):\n    model.to(device)\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    for i, (images) in tk0:\n        images = images.to(device)\n        avg_preds = []\n        for state in states:\n            model.load_state_dict(state['model'])\n            model.eval()\n            with torch.no_grad():\n                y_preds = model(images)\n            avg_preds.append(y_preds.softmax(1).to('cpu').numpy())\n        avg_preds = np.mean(avg_preds, axis=0)\n        probs.append(avg_preds)\n    probs = np.concatenate(probs)\n    return probs","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:22.171831Z","iopub.execute_input":"2022-11-29T09:58:22.172210Z","iopub.status.idle":"2022-11-29T09:58:22.200623Z","shell.execute_reply.started":"2022-11-29T09:58:22.172171Z","shell.execute_reply":"2022-11-29T09:58:22.199241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train loop","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Train loop\n# ====================================================\ndef train_loop(folds, fold):\n\n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    # ====================================================\n    # loader\n    # ====================================================\n    trn_idx = folds[folds['fold'] != fold].index\n    val_idx = folds[folds['fold'] == fold].index\n\n    train_folds = folds.loc[trn_idx].reset_index(drop=True)\n    valid_folds = folds.loc[val_idx].reset_index(drop=True)\n\n    train_dataset = TrainDataset(train_folds, \n                                 transform=get_transforms(data='train'))\n    valid_dataset = TrainDataset(valid_folds, \n                                 transform=get_transforms(data='valid'))\n\n    train_loader = DataLoader(train_dataset, \n                              batch_size=CFG.batch_size, \n                              shuffle=True, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset, \n                              batch_size=CFG.batch_size, \n                              shuffle=False, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n    \n    # ====================================================\n    # scheduler \n    # ====================================================\n    def get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=CFG.factor, patience=CFG.patience, verbose=True, eps=CFG.eps)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, T_max=CFG.T_max, eta_min=CFG.min_lr, last_epoch=-1)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, T_mult=1, eta_min=CFG.min_lr, last_epoch=-1)\n        return scheduler\n\n    # ====================================================\n    # model & optimizer\n    # ====================================================\n    model = CustomResNext(CFG.model_name, pretrained=True)\n    model.to(device)\n\n    optimizer = Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay, amsgrad=False)\n    scheduler = get_scheduler(optimizer)\n\n    # ====================================================\n    # apex\n    # ====================================================\n    if CFG.apex:\n        model, optimizer = amp.initialize(model, optimizer, opt_level='O1', verbosity=0)\n\n    # ====================================================\n    # loop\n    # ====================================================\n    criterion = nn.CrossEntropyLoss()\n\n    best_score = 0.\n    best_loss = np.inf\n    \n    for epoch in range(CFG.epochs):\n        \n        start_time = time.time()\n        \n        # train\n        avg_loss = train_fn(train_loader, model, criterion, optimizer, epoch, scheduler, device)\n\n        # eval\n        avg_val_loss, preds = valid_fn(valid_loader, model, criterion, device)\n        valid_labels = valid_folds[CFG.target_col].values\n        \n        if isinstance(scheduler, ReduceLROnPlateau):\n            scheduler.step(avg_val_loss)\n        elif isinstance(scheduler, CosineAnnealingLR):\n            scheduler.step()\n        elif isinstance(scheduler, CosineAnnealingWarmRestarts):\n            scheduler.step()\n\n        # scoring\n        score = get_score(valid_labels, preds.argmax(1))\n\n        elapsed = time.time() - start_time\n\n        LOGGER.info(f'Epoch {epoch+1} - avg_train_loss: {avg_loss:.4f}  avg_val_loss: {avg_val_loss:.4f}  time: {elapsed:.0f}s')\n        LOGGER.info(f'Epoch {epoch+1} - Accuracy: {score}')\n\n        if score > best_score:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best.pth')\n    \n    check_point = torch.load(OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best.pth')\n    valid_folds[[str(c) for c in range(5)]] = check_point['preds']\n    valid_folds['preds'] = check_point['preds'].argmax(1)\n\n    return valid_folds","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-11-29T09:58:22.201904Z","iopub.execute_input":"2022-11-29T09:58:22.202573Z","iopub.status.idle":"2022-11-29T09:58:22.224117Z","shell.execute_reply.started":"2022-11-29T09:58:22.202535Z","shell.execute_reply":"2022-11-29T09:58:22.223378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# main\n# ====================================================\ndef main():\n\n    \"\"\"\n    Prepare: 1.train  2.test  3.submission  4.folds\n    \"\"\"\n\n    def get_result(result_df):\n        preds = result_df['preds'].values\n        labels = result_df[CFG.target_col].values\n        score = get_score(labels, preds)\n        LOGGER.info(f'Score: {score:<.5f}')\n    \n    if CFG.train:\n        # train \n        oof_df = pd.DataFrame()\n        for fold in range(CFG.n_fold):\n            if fold in CFG.trn_fold:\n                _oof_df = train_loop(folds, fold)\n                oof_df = pd.concat([oof_df, _oof_df])\n                LOGGER.info(f\"========== fold: {fold} result ==========\")\n                get_result(_oof_df)\n        # CV result\n        LOGGER.info(f\"========== CV ==========\")\n        get_result(oof_df)\n        # save result\n        oof_df.to_csv(OUTPUT_DIR+'oof_df.csv', index=False)\n    \n    if CFG.inference:\n        # inference\n        model = CustomResNext(CFG.model_name, pretrained=False)\n        states = [torch.load(OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best.pth') for fold in CFG.trn_fold]\n        test_dataset = TestDataset(test, transform=get_transforms(data='valid'))\n        test_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, \n                                 num_workers=CFG.num_workers, pin_memory=True)\n        predictions = inference(model, states, test_loader, device)\n        # submission\n        test['label'] = predictions.argmax(1)\n        test[['image_id', 'label']].to_csv(OUTPUT_DIR+'submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-29T09:58:22.226817Z","iopub.execute_input":"2022-11-29T09:58:22.227064Z","iopub.status.idle":"2022-11-29T09:58:22.240555Z","shell.execute_reply.started":"2022-11-29T09:58:22.227040Z","shell.execute_reply":"2022-11-29T09:58:22.239904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    main()","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2022-11-29T09:58:22.242312Z","iopub.execute_input":"2022-11-29T09:58:22.242760Z","iopub.status.idle":"2022-11-29T09:58:41.688093Z","shell.execute_reply.started":"2022-11-29T09:58:22.242724Z","shell.execute_reply":"2022-11-29T09:58:41.685560Z"},"trusted":true},"execution_count":null,"outputs":[]}]}