{"cells":[{"metadata":{},"cell_type":"markdown","source":"Hi I am Manoj Akondi from National Institute of technology, Calicut.\n## Intro to notebook:\nIn this notebook, I am trying for a contrastive learning approach to learn cool representations of our data. This notebook uses a new state of the art method [SIMSAIM](https://arxiv.org/abs/2011.10566).<br/>\nSIMSAIM is a simple method where same network is used twice to compute the representations of different augmentations of same image and the similarity between the computed representations is maximized. While backpropagating through whole the network, the gradient flow is detached for one of the branch as shown below.\n![d.PNG](attachment:d.PNG)<br/>\nSince we have lot of noise in our data, I thought of visualizing how noisy the data is!\nSo, I used t-sne on the representations obtained from this model and visualized in [this notebook](https://www.kaggle.com/saimanojakondi/explore-the-rep).","attachments":{"d.PNG":{"image/png":"iVBORw0KGgoAAAANSUhEUgAAAasAAAEbCAYAAABk26sYAAAAAXNSR0IArs4c6QAAAARnQU1BAACxjwv8YQUAAAAJcEhZcwAAEnQAABJ0Ad5mH3gAADOASURBVHhe7Z19VFT3nf/n366cnHDO7snPbmzZUnrISVqTtW5XQ7KSllTLrpikBnUTgnkCNxrUVIymgG3AVMdUsBqikYrIRjCOJhKVGJGIFpJAI0aJwQguGIlgfCBIfIDM+/f9zAOMzAUuM8BcZt6vc97nyL137lxn5vt9fZ/mjgmEEEKIwaGsCCGEGB7KihBCiOGhrAghhBgeyooQQojhoawIIYQYHsqKEEKI4aGsCCGEGB7KihBCiOGhrAghhBgeyooQQojhoawIIYQYHsqKEEKI4aGsCCGEGB7KihBCiOGhrAghhBgeyooQQojhoawIIYQYHsqKjGA6cbFiNWIiF2J7Xbtjm2dYL5bDHBONOdu/wA113rbGT7C/YB1SExKRXX3FcZSHWM+jwjwLkXO2o+6G1bGREDIQKCviMZWVlbhyxcuK3CtuoKloAcKCpiGrutWxzTOsTbswJywEUVl/Rzuu4XztxyhKj4bJNAnmqjbHUR5ibUTRnHEIilqD6nYXWXW04tKVTscfvuXGjRs4ceIELl265NhCiLGgrIjHiKzWr1+PPXv24MyZM46t/kIHmiwJgyMrLayXUL0maWjOPQCksXH48GHk5ubaQohRoayIx0hFJ7JypqCgAEePHrW10kc+Qymra2ja/zIig4ZIhDqoq6uzNTJc3z957wgxKpQV8YqeFZ4z0lofnCElKzqa9mN54iJk5uXjtfmxiPr9e2iR7edrUFKQifnRMfZK/2o9DuavRWrCgxi3ch+O78xA7LjRSjjhmLq8FGfPliMnOQahJhNMoY9gedlX6iyKjhbUlBQga/5UjDNXKU3ZNmrIqhNtNW8hNWkpsjZvQmbiFETOL0BNWyesbf+HD4tykC7P/XIh9qZOQZApBNPyPkVzzQEUZC1A9DgzqjqsuFG/F8vjJqhzhyI6KR3mv2zD3ndXYHqwuq7RcTDvqsZ5dWHW85/gf5MfRPD0Nfj4vP2qvEEaESIkaVRovWe+HdIlpG8oK+IVMvynVfE5s3PnTlsr3nO+RlnKZMRZGu1/3jiG7LgNqOm8gtNVB7Gr57zSlTKkjFGV/ri5yKk6p5RzFXV58TZxjE/cjONKLOg4hYL4cJgistV5rLh6+u8o25WOSCWx0L5k1XoQKeFBXcdYGwsQawpDvKVByaoR1UX2c5jCnkb2+8XIS12MVXvL8XGZBemRSpqhIis5kZYI5f8ZAVPQkyhouO7YdgkV6dORXHLe8bdnSKNBGg9a748zsp8QI0NZEa/praXuGpkPkVb9wFvv51CcNBbBU5djb12r6gldQ+Phj9H4nezTqPQ7qmAOVcKIs6DJvkVtMqveVGi38NCGKvMkJYZ5KGpxLHBosiCuP1l1fImDr6/Chopzth6ZXVYuj3GcIyS9Ql2lK42wxIX2IyvV46p5HVGmYEzMOgKbrtrLkT7pFVS4LsoYANJIkMaC1vvRM83NzY5HEWJMKCuDI0M30nuRil5avzLsZrRs27ZNswLsLQMbIuxE2/FcxIcFKblMRHzme6iT3pENL2VlSoClya4mXbKyoa6nvhQ55leQaZ6HSRqy6j6HEz2yUlgbYJmtenzhy1DW2on2ilcwKb0cA12UL58VaRxovfZaycnJ0XxfjRD5v3CVIhEoK4PiXKUllYmz0EqkBWy0lJaWulWAvUV6YdLiH9giDCWIuveQGT8RQaYghMXn2ofztCr9oZSV9TKO5yRgYtJbOCXP3/Mx3soK36G1bBnCTeGYbTmC8vQEZFW77teHvCf9Dfu5ZtOmTW7vqRHi2kgT+Uov0f9WnRK9UFYGRFqSUolIITX6pLdIR08rvqSkxFYBDZyzeHfNbjTZxt1aUbfzRYw3jUVS8Tm1YThlZVU9neW4q+u5FYMuK8X1I8iaGIygyY/gkSlrUH3dsyFAQT47entZ3s0rDg9yjdLYkcabf6w4JQOBsjIYzlbkSBn2kApEq/KTyP9DKkvvhKsq+tkLYDnrqJyuVSA9ZBLSK+T1GU5ZdeJC8QsINgVjQup+NF88jYq8hUqcMkdVisYTp/D1Wb2yUj2okiUYbboH84ubcKPpE3x42jnYdwNnLXPU84zGgzkn1JGDg7xPUslrvU8S6bWMBERSUkbkeimswIKyMhDSo5IKfiQVQq3FFd6vAHRF9awWPYj7pyYju+BNbEydhwV5R9FmvYL6su3IShivhBKC6JRN2F1ejoN5KYgOUrIan4Tsd4+g4VQZ8lJibMOH4QlZKKqux6myN5E+PUw97n4kZRehvOIAtmclIFyJJig6BZt3H0LFwQKY4+5Rx4Rhenoedle3wHr1BLYmRtjOFRq9BPkfHcbWJ9UxQRF4Nv115KyMR5icIzIZGy2HUH9V9Yqu1qNseyYSwmXOLQYpm/ei+nyH/fZOU8PV+cMxNWM/mjpcelCyojHkMeTV3bxMYzDoa2XgSJoXorACD8rKIEhFMdIqDBnWc63spAIZSddvVKwNBZg1qwANno8A9otU8tI4cm1syPs3kpCeotxFhQQGlJVBcC6iGEnIPJRUdnLdbOEOFu2ozUnAnKIv7V9YHgZk0YJziHAkvY8yvCzXbPR5XTI4UFYGQArbSBv+E7gya7C4gaa9L2LC+Fgs+v0ziHywxw1vhwn5HI60il96g+xdBQaUlQGQ4ZiRNgRDBpNOXCxficggE4LGJ2FrbY8VgqRXpME0UhaHEO+grAyADKeJsAghA0eGAon/Q1kZAJkv8Ow7SISQkfRVD+I5lJUBoKwI8RyWn8CAsjIALGyEeA7LT2BAWRkAFjZCPIflJzCgrAwACxshnsPyExhQVgaAhY0Qz2H5CQwoKwPAwkaI57D8BAaUlQFgYSPEc1h+AgPKygCwsBHiOUNVfvhlY2NBWRkAyooQz6GsAgPKygBQVoR4DstPYEBZGQAWNkI8h+UnMKCsDAALGyGew/ITGFBWBoCFjRDPGarywzkrY0FZGQDKihDPoawCA8rKAFBWhHgOZRUYUFYGgLIixHMoq8CAsjIAlBUhnkNZBQaUlQGgrAjxHJafwICyMgAsbIR4DstPYEBZGQAWNkI8Z6jKD4cBjQVlZQAoK0I8h7IKDCgrA0BZEeI5lFVgQFkZAMqKEM+hrAIDysoAUFaEeA5lFRhQVgaAsiLEc1h+AgPKygCwsBHiOSw/gQFlZQBY2AjxnKEqPxwGNBaUlQGgrAjxHMoqMKCsDABlRYjnUFaBAWVlACgrQjyHsgoMKCsDQFkR4jmUVWBAWRkAyooQz2H5CQwoKwPAwkaI5wxG+bly5YotfXHp0iXHv4gvMKasOq8D354HWk8DF2uBlmpmpOf8MVXaTwFtZ4HrrY43mhDvGQxZnTlzBrm5uTcJyXUYsK6uzraf+A5jyeqGatlIhaZV2TH+FZFX+znA2ul48wnxjMEamXAKySksp6x6bie+wRiykgpLelFalRrj37nwmb2RQoiHDOYwuquYRFYUlXHwvaw6rnKojwGuXnR8IAgZGIM95+sUlMiKojIOvpWVzE3JcJBW5cUEXigs4gGDLSuhvLzcJquGhgbHFuJrfCcrGfpjj4rpGS6+IANkqHpWn376KXtWBsJ3spJVYVqVFRPYkZ42F12QATBUc1ZafxPf4RtZyfCfVkXFMJJvGh0fFEL6Z7Bk1VNMXA1oLHwjK6mMtCophnGGvSuik8GQFb9nZXx8IysuqmD6i3wpnBAdDIastO5g4SorgT0r3zL8spLv1GhVTv6Qpg9QW5SFnCW/RdSzf0GLbfvfcXHPIsREPI7tlR+6P8YXOfcuLHN+jUkh/4CgJ9Y4rtNgkS+HE6KDwV5g4aSnrIhvGX5ZyV0LtConf4jIavcyJNxmgunRV9Fk216FpvzHEDbqAWQdOOz+mN7S9Ddc+r9PtPcNRs69jZxJP0DMuiJYtfYbIYTogLIKDIZfVv5+p4qmN2EOcZWVB2k+hOrls2B+v1x7/yDEeuxVxI+KRPYhg/T2tEKIDoZKVsRYDL+s/P3ef17L6mM0WZ5D5KifD6GsPkHrjqcw+s6FqPpSa79BwtswER1QVoFBAMjqE7QdyUdR5lwkREzAotV/ROqvQ2AymRB090zkHToEa3M5Gvausc01RY5NxPYCkYXa/5t0nDq1EzsXPYTp06cgeuztGD9jMQ4e+7j7/M2HUZs7D/GPPob01KeRMD0Sk27tllXH5ztQsmER5kdF3iQf6+kiFKfPQnzC/2DlvCiE3/0osvfsQ3vlOix/9Gfq+m5H9LNzYf7TBjSek+O3o/D532D6s+r4Fx5W1/kAFr+xHW3Nsm8vPsxfhvQnJmDciyux93cRCDJ9H9PW7cJ3zuu8KYdQseTnCJ7+HN5YvgjZ69KRFR+DVMt76NA83kehrIgOKCsjcA3naw+jKCcdCVHJKGoZ/NW8ASKrbShO+6USwG2IfP4VVBz5QNXXryIl4jaYfvQYij4tQ2NpJtIj/kkdE474lWuxf91zmL8kGStiJyNjz0HbvI619jXMv82E4IdXoEEJROajmgtnI/znSaios88vdXyyEvFdsqrA6ff/il1L7lPndekpNb+P4sR7Eb/hXbscbPNH6rnvmIeKhio0bfrtzcc3WpAz7Q7EZO7ADflb/k8HlmLqqB8gNtOCb08Xozp/LiKVgE0/egjZ219D3u9mY9WuEu35qK+2I/sX/4SwGX/E8dOO635/AUJ/mozqrzSO91UoK6KDkTlnZUXH5cu4YnX8OeIRWZUhJ2GsqrsSYGnqcGwfPAJkGPDv7gJQ25rzZyHY9APEbypWlXoxLI/eDtOYuag4I/uPoH3PPNx1631ISlkAc5rkacSNvUWdJwYFR6qALwuR9W+3IWLldnQ6n8ttGLDnczvPG4fik84FFEo+h7KwPteieko9j3cM2Zl+g7zKSsfxKs37UBT/Y5huewolIspjryJOySpkSR6uOY/pJdbKdMSYfonsQxWObZ/gQmEcgn/xEmooKzLCGImysrZVYU1sJqoGv073IW2oMk+irLyLlqyclbYJoWlvqh6OQ1YhC1DVJPs/Rm1mlMvf7rE//nbEKdl1be9XVh+iZmVkH+fteXwFqjPuVX//FpZjf3c5rhL1637TfZxDVvb/i+v5ekZLTGUoe+FnGPPCX3HF7XgfhrIiOhhxsuo4g/2pUxAUaqasBkBAy8pewd/mWL7dU1YOSdw6C0WfuUqiO7ahM9P3EbthT/dwW7+yKkdV2s/VeV17Vq7pTVY9V+45j3Ns1y0rDTE15CH9jjuRvOODHsf6OJQV0YExZWVFR9N+LE9chMy8fLw2PxZRv38PLWhHfbEZceOCYQqORtIKM9aWNOA7dKKtZhsWT5+FpPTlSEmYjHHRS1FYc1mdSe1r+Mg+HxT5IBZt3ITUqeGq7JsQNH4e8o7LMX1hP/eLCUvw2tYcpMeOV+eehbi4OMTNzUdVvfPc0Xh5+1tIjRwDU9BjyKtrV49TfyctRdbmTchMnILI+QWoaXPOR6nz1r6N9PjZmK+u+fcJ/42pk2Q9AGXlRbRl1XnoJUSY7kNWqWzrKStHD8T0A8xc945jrsi+vXXfMvx5xwHH8m8TRie8hgvNjv39ysp53pDuOStbPsbZN/+EfZ997HZ86645CFc9uJkbdrvMQTl6fj9fiKqGI/plZRu6vANJhSWObY5hSdtw4sdo2LAQ+a7Djb4MZUV0YMwFFl+jLGUy4iyO+1zeOIbsuA2osdXzjbDEhcLU1bOy4urnmxA7ehZyattlg9p0CdVrpiMoOA45n19GW2M1is0Pq3phDCIX/y8qGr5Gy3ELUkQsYQtQ1HTD/jgtWg8iJTys61qsZy2YHRyEseZKXBPhNB5BUXq0OncQwuJfw/v7NyN1/hqUfbFfPS4IoeYqyGVaGwsQawpDvKXBJkdr814sCJ+MjIoLdll2nEJBvEiUsvIiTmHcicU7Su0VfvNBVCy9F+GJ69QHXY7pKatqWE/mIPXuW2Aa9e9IWr0WFe/nYf+6BMycvRJ1Z+UcB1Dy/N1qf4RjJd0naCt/GXGywCJqAYo/KNaYg5LzbkKGLO6Q867MRNm+v6JoZTxm/2ELWpudc1ThmF+4X33GC/BhRSEKZqoPgctCDqgPTlbEOMzP32uXk05ZfVexDA/eNP9l77mNnrcBrWd3YNPiLJx1itfXoayIDowpq3MoThqL4KnLsbeuVVXm19B4+GM0fif7esrqPEqSx8EUk4d6ly6StbkIicGqMZxcglaliyZLgqpHJsFc1eY4ogPNRUmq4dstEC06qswINYV2i9P6BfJilOTiLGiybXCeOwLpFZdtW2x0fImDr6/ChopzdjnZZKXqGJu82lCdpQQXke0QsMBhwEGIUxj/hHvuvhexcxOREv+fSEz7C2plNVxzKapzFyDuR/+gBHIvklf/GWWVsvhAyadidddSd9OoOxG7JNP+GMe57UvKI9WHQe1XvaWpi57GM3eOR9zzySjYtQ3V76xC1hN3qn3fR/QLf8TuUrssO47lIGvG3QiSx4XciyRzLpq6JLkZZttzqvMt3WDbbj29EzuXRGPSfzyMlNT/wfwZs5BZWGRfuv55IXamTUWYOldQxGxs3LQJ9aq11P3/d8beqxv9wDLUds1XqZ5V6e/xyBNJyDUvR5nrsnxfh7IiOjCmrFSP5Xgu4sOCYAqaiPjM91DXNXzWQ1ad1ci6Rx3XJQ8HTqnYjtOSlTqkPs8+7656SVfPn0RVVZVLTuJ8h9VxTDAmZh3Bdduj5PnDcFd6Oez9OO1z21H/j/pS5JhfQaZ5HibZnkvJyk14gr/Jyid3sNAeBmQMHkJ0YNwFFqqir3sPmfETVaNUhthycdwmrF5kdVMvRXAcZ9vei1CaLIgzjUFM3gmcte2XRrMzDmlYz6M8IxpBYc8g5/gFtH2+GXF3zUFBnWPIsTdZWS/jeE4CJia9hVNy3bbncsiqowrmUPUcsQVo7OrS+ZusfHJvQMpqRIYQHRhTVmfx7prdaLKNn7WibueLGG8ai6RiVf+5DQPK/FaE6oE9iYIGe9/HxncnkPNgCCaYP1Y9IG2hdNZkI8IUjazqnj0iV6zoOLsbaQvM2Ji9Glk5u1Hdcs2xT9A6txXtFctxV9c1K1xlZW2AJT4MptGLUHzBaVh/k5X8bLlWxTSkqULjhhjKaiRFhosJ0YExZaWENHsBLGcdCx+uVSA9ZBLSK+RnRhxzVLaK/gqaPqxE7YmtmB02GhMyytFq66lYcbV6DSLDF6HorIjFKZTxWFzSYp+fsl5ARcZkhC/Yi+au3o0G7X9HVtQjSkQX7Y9zQ0tWnbhQ/AKCTcGYkLofzRdPoyJvoRKuCSHppWg8UYva/b9HuCz4SN2HJhlubDuGnDhZYPEwzMXVaOwa9hwchl9WwrD+npXje022bvE/IDRqJpZv2oF2zWMZw4S/Z0V0Ytie1aIHcf/UZGQXvImNqfOwIO8o2my26MTFitWYGhqkelePIKPkjNKFYxl4bBSiEn6PlekLEJ+YhWLb4gzBKZTRuGf8ZMQuSkNKYhwSze+ith8pWFv2Yr7MndnqQGdGY3z8GhxsasP56p0wx92jtinxJK+HpaweV+WBV09ga6Lcui0IodFLkP/RYWx9Uh0XFIFE2//lMmoKlyJa/h9yjPq/pjwzCZPiXkRWwSHU+4Ws+EvBTH/hLwUTnXgjK/nBRXm8/FJwf9y4cQNHjx5FpaykHXa0ej/6kB5P/osrsHX/Xuy0WGCxbEdBzhqsTI7FxJSDaHUcZ3R8I6vO69oVFMNIpDFDiE687VnJ46UXVVBQgBMnTtik5Ir8QvDhw4dtx0jkJ+6HHw9lZW1BWeoUxOR94T4EeKUMy5bKsviRgW9kJbSd1a6omMCODBGzV0UGgLeykl6VU0TOiJxESk6ROZObm+t41HBzA40Fswcuq85jyI4IRvB0M/ZUNdiHITvOo/bQO3g9fTWK6p0rAo2P72QlFdLFWu0KiwncyAIcQgaAt7ISpFflKqXeIsOAw087arJlgZjMNcn80Twst9Q4viPVH522WyalTB9v/06nStC4aUjSMddlNHwnK0GGA4d1sQVj6HBRBfGAwZCVDP9pyalnZI6L+AbfykrouMoeFgNcvej4QBAyMAZDVjJPpSUn18jQIPEdvpeVIEOCPrmzBePzXPiMt1UiXjEYshJcF1FoZTCeg3iOMWTlRCotn9w7kBn2yPCv3M2EiymIlwyWrGTVn5akJDt37nQcRXyFsWTlROayZP5CxCUtb63Kjhl5kfdTVoFyEQUZRAZLVoKcS0tWvlmuTlwxpqwCDCkg0qojhAycwZSVSKmnqGS5es/vXpHhh7LyMbK6SAoEJ28J8YzBlJUgcnKVlW/uWEF6Qln5GOekLltvhHjGYMtKvkvlKisuVzcGlJUPETm5tuLkux6EkIEx2LJyjnZISkpKHFuJr6GsfEjP8XH5Fj0hZGAMtqwE54iHnhvckuGBsvIhWrd4YeEgZGAMhazkfGw8GgvKykdIYegpKgmHHQgZGEMhK4ErdI0FZeUjREpaspJwQpcQ/QyVrIixoKx8gOsErla4VJYQ/VBWgQFl5QN6Lo3tGS5jJ0Q/lFVgQFn5gJ5fOtQKb+9CiD4oq8CAshpmtG7nohXeOJMQfVBWgQFlNcyIhFyl5OxlaS1jZwEkpH8oq8CAshpGXH+CQOQkd6xw/uibIAXOdZUg7xdISP9QVoEBZTWMiHykYPX84q9TVk5ktaAswpBeF5exE9I3lFVgQFkNI72t8OspK1e4KpCQvqGsAgPKygD0JStCSN9QVoEBZWUAKCtCPIeyCgwoKwNAWRHiOZRVYEBZGQDKihDPoawCA8rKAFBWhHgOZRUYUFYGgLIixHMoq8CAsjIAlBUhnkNZBQaUlQGgrAjxHMoqMKCsDABlRYjnUFaBAWVlACgrQjyHsgoMKCsDQFkR4jmUVWBAWRkAyooQz6GsAgPKygBQVoR4DmUVGFBWBoCyIsRzKKvAgLIyAJQVIZ5DWQUGlJUBoKwI8RzKKjCgrAwAZUWI51BWgQFlZQAoK0I8h7IKDAwrq+udVnx9tQNn22+goe06Tly65rcp+bxBc7u/5GTrNdv7KPnmxneOd5gMJZ1Wq+21dr7u8h5ovTf+kMP1X+Foyzea+/wlUgfK+3jpeqetbgxEDCUrKWAiqCMXvsWHLe2Mn6bq/Leo/+Z6wBa6oUQEJWLSet0Z/8mxi1dtdWUgYRhZSYuBkgq8iLSkkUK8Q8QvLXCt15jx30idGSijFYaQlVRYWm8EExiRAtfeweFBT5HKSnqrWq8tExg5963/97J8LiuKipFIZUthDRwZCtJ6PZnAi9Sl/oxPZUVRMa4RYXEeSz8id63XkQnc+HMPy2eykjkqrRebCezIxDHpH5nn4xwvoxV/HaHwiaykoHGMnektgbbKyRNkKbPWa8cw/trg84ms5PsCWi8yw0ikx0B6R4ZKtV43hnHGHxt8PpEVhy+Y/iLDxEQbmZfQes0Yxhn5rp2/Meyy4qQwoyf+vrLJG2SYR+s1YxjX+Nv3F4ddVmwVMnrCocDe0Xq9GKZn/O3LwsMuK04MM3pD3JEKSOu1YpiekbUB/sSwy4q3hGH0hl8SdoeyYvSGsvISyorRm0C559lAoKwYvaGsvISyYvSGsnKHsmL0hrLyEsqK0RvKyh3KitEbyspLKCtGbygrdygrRm8oKy+hrBi9oazcoawYvaGsvISyYvSGsnKHsmL0hrLyEsqK0RvKyh3KitEbyspLKCtGbygrdygrRm8oKy+hrBi9oazcoawYvaGsvISyYvSGsnKHsmL0hrLyEsqK0RvKyh3KitEbyspLKCtGbygrdygrRm8oKy+hrBi9oazcoawYvaGsvISyYvSGsnKHsmL0hrLyEsqqr1xCccV+rMpMw0P/sRCrPvtG45jBzDc4cOivmJnwRyxIeBSP/PVTHNY8zjehrNyhrPrKFRz+/AjW57+GxCemITb/lMYxg5zTHyEtcQHmpc/D/bP/F5azV7SP80EoKy+hrPqKyGofXnripzCZnsKfjrVqHDOIqduHJ8P+C8kHK5D2mzsxLu0DHNA6zkehrNyhrPqKyOoo3sh8Gv9o+hdM2VSrccxg5gxef+E+jF15CG+tm4lbIlZg0+k2jeN8E8rKSyir/tKC3LT7hkFWbSjZsQj/eMcfsLHBOK1B11BW7lBW/edv72fgn4dDVnV78Pht9yFxT5P2fh+HsvISyqq/DJesvlStwn/DLc++g32a+30fysodyqr/DI+sruCDXUtx261JyDw51MP1noWy8hLDyer0J1ixKB4PTH8U9479F9wx41W8fuyi2l6NrMwMPDZ1HMa8uAVrl83AGJMJppBpmLOrHhXOxzefg6XQjIdnzENi6gI8EBaBh1bux7tNjt7K6UpkPD8DDzybgrkvzMa4sf+FuDcqcaDZ8XiZN6ooROKMxzFjyR/w5BMzcN/EH94sq96usbkZ7+zdjpeWPKXOuxSvFKRi3CgTvvebjbCcc56/Z9pwoDQf81Kew723BiP80WTMSzMjvdR4rUPKyh3DyarXz+ZX2LZtA+bN/S3Cf7QEr1hW4Fc/GqU+1z/BfS8VY6/L57+0chcWxMerz/8yPK7K2x1P/AVvyDlkf3MTtr6xGA9MfU6Vr8V4KOLfce/zW7D1tIsg5BpUGZiSoMrYwmcw+dcTcYurrHq7RnnuIwftc8QRk5Hw13w8FXE7TKNmIq3S8fxaOf0h0pelITbqX2AaOwtPpmVg3oYPDTWELqGsvMRQsjp7FOkzp6tu/Bn737XvYMZtJtzy8Ga8fe4KDldk4z4lqO9FLEL6gdM4fPowFv5aieTfVmHLlyKjVuwtfAFhM7bgbZucLmBb5kOqQDqGBhqr8NK0e3B/5t8dCxeUmA6sxn2jfoxfZVbhoNpW8dnbmBkW1X0NTceRPuMn3bLq6xq/UrIq3YbEiP+njh+L6JU7sHZdKmb8YZ9LZdBLTr6D2FsnqvN+pb3fAKGs3DGUrPosP+rvpk/w0q+l8o9C3IZy7D17FluWP4TvmSZjYWmL7TG2z/+dTyP9kwu2v8srXsMvTKMQuuQAPmhpQWHmLPzjlGxscy5ccJTBW6atR2Gj2tZ8CpmJE3HX0gMosX3mL+DtDU91y6rPaxRZVWDVksmq/IzCmBmZ+IvlDTyVoGT2WX+jGqeR+exPMXpJKco09/s+lJWXGEdWqhu/5w8IvXUyYlNUy0haR2nJmDI2WH1wH0f6kcv48NibmKJk9c9ph/E322N6DNE1HEDiHWMRW3i6+7ynq7Di1bdUy++yfU7INEO10i5172/+Aqvi74TptkVYW/cVtmSogvKLLGz9yrH/pue43P81ttTiT4+qFt6YZdh4xlGgdaT8UBbGdp1D+xhfh7Jyxziy0lF+nJ/NkAzkNtkfd/MQXRM2Lrnv5qFo6UnlbsCKQ6qxZ5sTuh33r/u0eyRDGoj5zykZ3YPHdzTgUOkq/Mz0n0g+9LVjv+tzfK7jGlvx7qan1N8DbLh9VYnkX6hG54bPtPcbIJSVlxhHVo5ekEtBcks/srJX+Pdh3vv2VuLNOa9EFNV1bPf2S9ixbobarh63729Im6Jano++iXe79rs+R0v/16hRIfSfb1SBT8T3Bii44Q5l5Y5xZKWj/PQnK1uFH+xSvm7O30rN+InrcJ4jFZUbcb+tXO7HNltZurmMdT/HUR3X6JRVb+W4l3y2HQ+NMvbIBGXlJcaRlUMmtz6HVb11+fuRlb1Q/PTmnlVXnLK6udXXXTjU9tJizAsxwTRtM97pGrZzfQ7pefVzjR7JynFtN0nSeKGs3DGOrHSUn/5k1XTY9vnvbZGPXVbBGLuyEuWu+2zlUrYfQI6trNw8QtD9HJ/ouEbPZGW/tp4NUWOFsvIS48jqG+wrTMItph8jal21y5dhv0HJvmwk7VAC6m8Y0Dbvowpb15yVHHMFhz8tQsq2kyjZ9RJ+qApN1IYal2EMR4v0568gt+EE/jTjxzDdtgCZtc4JY9fnuNT/NXokq1NY9cSd+ElGhWaL1iihrNwxjqx0lJ/+ZOWY9zHd+lTXnJUtTZ/jz2v2Yq/te4Cj8L2HXBtzznmtKCWXRiWap/E901jMKKzv2t/9HJ/puEZPZOUYmfipGVu6hu+NF8rKSwy1wOJkMZ66OximUQ8gdvVObHz/ANauW4Ko2Vvs30TvT1bNjXhj6WRVWG7Hvz6biVW7SvF6fiZmzHwVm+qUfGyTu2OVmF7GRvlbztFYjoUREZiR/7kqPKrA71ishHY7xv2uyL6C8HQlXnpUFljEYF7hh3jnyO6+r9ETWdlatL31CI0Tysod48hKpb/y06+svsH7e9LtK1gnzkVy/j5s2rUNyQkJeH7fl2r/17BseAZjTL/sXiDRck71lv4LP0zYbi8vJ3cjToQWkYq1jlWIhZlPKEHdjnvTdmDLgUI82ec1eiIrez1g5K99SCgrLzGUrFRhkWXjT/1a5GBSH+hx+NWSQmyTZbGNNXjd/BzuUNtvmZqGl3dV450Dm/C4LG01jcNDy/dgh6xGavoCr2fE4w5V4EyqBTfu2XXdy24ltmW1M/Gv/zEbT6amYMaM57Cg8FOUOluKtqW5ybg3RJb1jsI//3ohnnz8Pvzro7/Dwg3vY8fpy71fY3MDtuRmYIosCR4VhcdXF+D1Sh0F7shm/Mr0EF6qcGnNGjCUlTuGklVf5UdV6DveWYtYEcWtv0Xi+vdhOXIYy+ZFqcbdKPzwidXItn1WL+Ldd15F9N2yolV9/iPmYvGeU929IFkCn7sMv5o4GQ+98AckJsTj4fRdsHQtXZdbhm1BXJTjGkKm4ekXHkfYxFmIy9iiysPZPq6xFcWlhZj3qGpQSoNx3l/wp101tlW6N/8/e6T5M6RP+wl+kXnk5uFJg4Wy8hJjySoQ4/gyY9fye61jjBHKyh1jySpA83/78OT3u5ffGzWUlZdQVj5K02ls2bYHWz47gy0Z0/Gfmz53mUczZigrdygrX+US9pbuwZrS0/igdBXumfkm3u3vu4w+DmXlJZSVb1JxbAdm3v0TTHj6d4j7wx77eL/GcUYKZeUOZeWjNJ/E6sQI3PLLZ/Hk86/a58e0jjNQKCsvoawYvaGs3KGsGL2hrLyEsmL0hrJyh7Ji9Iay8hLKitEbysodyorRG8rKSygrRm8oK3coK0ZvKCsvoawYvaGs3KGsGL2hrLyEsmL0hrJyh7Ji9Iay8hLKitEbysodyorRG8rKSygrRm8oK3coK0ZvKCsvoawYvaGs3KGsGL2hrLyEsmL0hrJyh7Ji9Iay8hLKitEbysodyorRG8rKSygrRm8oK3coK0ZvKCsvoawYvaGs3KGsGL2hrLyEsmL0hrJyh7Ji9Iay8pL6b65rvrAM0zPEHcqK0RvKykvkBdR6YRnGNVXnv3V8YogrnVar5uvFMD1z6Xqn41PjHwy7rNgyZPREhouJNkcufKv5mjGMa653Wh2fGP9g2GUlSKtZ68VlGGfOfdvh+LSQnnAonekv0qDxN3wiKxY2pr/IcBfRpr2DoxNM3/HHxp5PZCXdU60XmGEkDW3XHZ8U0htcVcv0Fhm58sfGnk9kJUiFpPVCM4Edfy1ogw17V0xv8dchdJ/JSiqkYxevar7YTODG31YwDSVSKWm9hkzg5mSr/y5M8pmsBBEWF1swznx9lYsqBgrnfxlnpPHvz6MSPpWVIPNX7GExFJXnUFiMv4tK8LmsBHmROYcVmJEltrytkvfI8ClHKQIz/nanit4whKycSKXFVU6BEZGUFDJ/bw0OJ/JaymtKaQVGpEftb1/87QtDycqJvAEyeSyThZSXf0TkJO+l9KC5iGLokddYXmt5zSkv/4i8lyIoGTIPxEaeIWUVaOzZsweXLl1y/EUIGQgFBQW4cuWK4y/ir1BWPkYK2fr163H48GHHFkKIXpqbm23lp7Ky0rGF+CuUlY+RQiaFTXLjRmBMlBIyWJSUlNjKTm5urmML8VcoKx8icpJC5pTV0aNHHXsIIf3hHJVwpq6uzrGH+COUlQ+RwuVa2GTsnRCiD9dRCYnM/RL/hbLyITt37rypsEnYOiSkf3qOSjjDhUr+C2XlI5wTwz3D1iEh/dNzVMIZLlTyXygrHyGFSquwSdg6JKRvZMhcq+xIb4sLlfwTygrX0FSeB7N5Iwq2r8fimQuRd+QoStYsR2buKiSE34Wpyz9AyyB+B08Kk1ZBc4atQ0J6p7dRCWdOnDjhOFIvw18HkIET4LKy4urxzUgrPAX7bVTbUZMdA5NpMsxVl4GL+5EcFoLIjEO4OIgfVFn1p1XIXMPWISHaOJer95aBLVTyTR1ABk5gy8raAMsL61DV7vwUXkBZygSYIrJR08cdgaxt1Xg9dj4sTZ7dKby3IQzXcBk7Ie70XK7eW6T3pYuB1gHWy6gpXI7klOVIT3oYkfFrcLDJf39DykhwGNCVzmPIjgjGmJQyaN68RUmqIGMhYifdhRBTgkeyOnPmjGbh6hkuYyfEnZ7L1XuL9L48os86wIr2qrWYm1cL27iH9QIqMh5EUNQaVHfJjgwVlJUL1sYCxJrCEG9pUB/LPmiyIM5DWfUcwpAJ4fz8fC5jJ6QftJarO0cptEYrPLlfYN91wHmUJE9FYkEtrjq2dNZkI8I0ASllFxxbyFBBWXXxHVpLlmC0aTpyap0fxV7wUFauQxhSuGSoTwqgLFeXYQvZL4srnAWSy9gJ6cZ1ubqUDWdjTv4WZNRCtjuPGfhQen91wNcoS4lA0LQ81Dl/gs1WF4QiztLo2ECGCsqqi1ZUmaNgCslAxTVnm+oyqlatQFFzDyl5KCsZwnAtZE6csnIiApMVTSI0LmMnxI6MPkhjrmePySkrJ85Gnxw/MAZQB9iwor1iOe4Kegx5dZy3GmoCWFaduFi+EpFBEaoL/zWsreXImBAMU0we6m2f0060HcvFi+uPdnX5u/CiZ6VFT1kRQvTTU1ZO+l9R60UdoLC2VWHN1AjE5xxFm/uYIRlkAlhWHWgpM2Nq6D2YnpyK5PlZ2P/5R8ib8xgSX16JlS8vRdqmj9DSofEp9GLOSgvKihDP6U1W/eNFHdBxBvvTnsScPIpquOAwoCcMgaxkvJ0QMnA8l5WH2EQ1F2n7z9i+m/VdYxneruZw/VBDWXnCIMtK5rL4vSpCPGNYZSWiSv1vzM5+D5VVVaiq+hgfZL+IVVVtjgPIUEFZDYTvGlCy1gxzShzGmSYgLmUFzGtL0OhcGeQhIirKihDPGD5ZXUSleSqCTCaYbsrgNVxJ71BWBoCyIsRzhn0YkPgEysoAUFaEeIYsTBr4EnUyEqGsDICISuatCCEDQ2QlC5SI/0NZGQDnN+8JIQODsgocKCsDwAJHiGfI3WA8vmktGVFQVgZAbqnEu6wTMnA43xs4UFYGgSuaCBk4cg/Agf8yMBmJUFYGgTetJWTgyPA5b1UWGFBWBkEKHX+/ipCBISMS/d+wlvgDlJVBkHF3GdIghOhDelSc6w0cKCuDIEOA8qOLhBB9yHcT+f3EwIGyMhDSSuRQICH94/yJe87zBg6UlYGQLwdLAeQYPCF9I0Pm/H5VYEFZGQwphCIs6WFRWoTcjMxTiaTkfoAsH4EFZWVApIclBVJWOjHDGxmKlXmQK1euON4Nd5x3TdB6PDO0EUnJYiSKKvCgrAhxQeZAnL3bnr/eLBWkfMVAKkzZxwqTkOGDsiJEAxluEmE5v3AqYhJJ8esFhPgGyoqQXnBd8CJDg5zQJ8R3UFaE9IEM+8m952S+pK95LELI0EJZEdIHsphCRMVeFSG+hbIipA9kCFBkxZ+hIMS3UFaE9IPrQgtCiG+grAjpB/4MBSG+h7IipB9k3oqLKwjxLZQVIYQQw0NZEUIIMTyUFSGEEMNDWRFCCDE8lBUhhBDDQ1kRQggxPJQVIYQQw0NZEUIIMTyUFSGEEMNDWZERSCcuVqxGTORCbK9rd2wjhPgzlBUZgdxAU9EChAVNQ1Z1q2MbIcSfoawIIYQYHsqKkGHlGprK82A2b0TB9vVYPHMh8o4cRcma5cjMXYWE8LswdfkHaLE6DieE2KCsyAjDio7zNSgpyMT86BiYq9qAq/U4mL8WqQkPYtzKfTi+MwOx40bDZApXFX8pzp4tR05yDEJNJphCH8Hysq/UWRxnazuGwtSFWJy1CZszn0Nk5CIU1Fx27O9EW802vJiwBK9tzUF67HiMi56FuLg4xM0tQO3lE9iZ9jSmPzYL0eNCMT5+DQ42XbM9Uhsrrh7fjLTCU+iw/d2OmuwYdZ2T1f/jMnBxP5LDQhCZcQgXKStCboKyIiOMKzhddRC70qNVJT/JLivb5jKkjFEyGjcXOVXnlAyuoi4vHkGmEIxP3IzjbZ1AxykUxIfDFJGNGvUn8DXKUiKUwMyoEntYT6MgNgRB8RY0iSxaDyIlPAxxlkY5GNazFswODsJYcyWu3TiJgtkzkVFxwSY264VizB9tQvCsAjR8ZzvcHWsDLC+sQ1W700QX1PNPcLkeQkhvUFZkBNKBJkvCzbLqqII5VMkqTonGvkVtMqveVGiXbIA2VJknwRQ0D0UtYodraDqYg+Ubyu3Dbg5ZOeXl9njrF8iLGaOeYzu+qFiOu4KjkbTCDLNZshRx44LVNc1GQeMN+/H90XkM2RHBGJNSphRMCOkLyoqMQLyUlSkBlib7QJxgbTuFAzmrkJ75CpImje6SlbU+DzGmYEzMOoLrtiMbYYkLw13p+/FJzvTuHpmHWBsLEGsKQ7yloWtYkhCiDWVFRiCDJatOtB3PRfzEhSg81aqEITIKdRkWPI/yjGgEhT2DnOMX0Pb5ZsTdNQcFdS2ozpoMU3ASipo9tdV3aC1ZgtGm6cipverYRgjpDcqKjEAGSVbt5Ui/KwjBScW4YNvfQ1ZKXx1ndyNtgRkbs1cjK2c3qltkAUUnLhS/gGDVK5qZV4vuQb9OtFbm4M8HmnX0lFrVtUTBFJKBimvOoy+jatUKLwRIiP9CWZERyCDJ6kIxkoLVYyb8ESXNX6OhYjPmjw9WAknHwcaTONnwIbKiHlHPcdFNPtaLHyBVjg2KQtLGPaioKsf+vFTMnFOIuhtaqurExfKViAyKQErZ17C2liNjgnp8TB7qbYerXt6xXLy4/ijYzyLEHcqKjDCuoL5sO7ISxivphCA6ZRN2l5fjYF4KooOUeMYnIfvdI2g4VYa8lBgEmYIQnpCFoup6nCp7E+nTw9Tj7kdStuolnb+I2q1JGC+PC43B4vxD+Hjrcwgzjcb4xHwcO1WE+WFB6ni1vytqn22JejvaancidWq4fXvQeMSmv41aWXWoSQdaysyYGnoPpienInl+FvZ//hHy5jyGxJdXYuXLS5G26SO0dPTfJyMkEKGsCOkF+Q5W/osrsHX/Xuy0WGCxbEdBzhqsTI7FxJSD4I2eCBk+KCtCtLC2oCx1CmLyvnCff7pShmVLSygrQoYRyooQLRzfgQqebsaeqga0ibE6zqP20Dt4PX01iup5t3dChhPKihBN5FZLbyFl+ngEOeargsZNQ5L53T7mpQghQwVlRQghxPBQVoQQQgwPZUUIIcTwUFaEEEIMD2VFCCHE8FBWhBBCDA7w/wE+Isivx5HABwAAAABJRU5ErkJggg=="}}},{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install efficientnet_pytorch\nfrom prettytable import PrettyTable\n\ndef count_parameters(model):\n    table = PrettyTable([\"Modules\", \"Parameters\"])\n    total_params = 0\n    for name, parameter in model.named_parameters():\n        if not parameter.requires_grad: continue\n        param = parameter.numel()\n        table.add_row([name, param])\n        total_params+=param\n    print(table)\n    print(f\"Total Trainable Params: {total_params}\")\n    return total_params","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import torch\nfrom torch import nn\nimport cv2\nimport numpy as np\nimport time\nimport random\nfrom tqdm import tqdm_notebook as tqdm\nimport pandas as pd\nimport os\n\nfrom efficientnet_pytorch import EfficientNet\nfrom torchvision import models\nfrom torch.nn import functional as F\n\nfrom torch.utils.data import DataLoader,Dataset\nfrom torchvision import transforms\nfrom albumentations.pytorch import ToTensorV2\n\nfrom albumentations import (\n    HorizontalFlip, VerticalFlip, IAAPerspective, ShiftScaleRotate, CLAHE, RandomRotate90,\n    Transpose, ShiftScaleRotate, Blur, OpticalDistortion, GridDistortion, HueSaturationValue,\n    IAAAdditiveGaussianNoise, GaussNoise, MotionBlur, MedianBlur, IAAPiecewiseAffine, RandomResizedCrop,\n    IAASharpen, IAAEmboss, RandomBrightnessContrast, Flip, OneOf, Compose, Normalize, Cutout, CoarseDropout, ShiftScaleRotate, CenterCrop, Resize\n)\nimport albumentations as A","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"#paramters\ntrain_dir= '../input/cassava-leaf-disease-classification/train_images'\ntest_dir = '../input/cassava-leaf-disease-classification/test_images'\ncfg = {\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n    'epochs':25,\n    'batch_size':68,\n    'lr':0.0001,\n    'input_size':256,\n    \n}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"imagenames = [name for name in os.listdir(train_dir)]\ncsv = pd.read_csv('../input/cassava-leaf-disease-classification/train.csv')\nprint(csv.head(10))\nprint(csv['label'].unique())","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def add_guassian_noise(image): \n    return IAAAdditiveGaussianNoise(p=1.0,loc=1.3,scale=(0,255),per_channel=True)(image=image)['image']\ndef cutout(image):\n    return A.augmentations.transforms.Cutout(num_holes=10, max_h_size=10, max_w_size=10, fill_value=0, always_apply=False, p=1.0)(image=image)['image']\ndef get_train_transforms():\n    return Compose([\n            RandomResizedCrop(cfg['input_size'], cfg['input_size'],scale=(0.3,1)),\n            Transpose(p=0.5),\n            HorizontalFlip(p=0.5),\n            VerticalFlip(p=0.5),\n            ShiftScaleRotate(p=0.5),\n            HueSaturationValue(hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2, p=0.2),\n            RandomBrightnessContrast(brightness_limit=(-0.1,0.1), contrast_limit=(-0.1, 0.1), p=0.5),\n        ], p=1.)\n\ndef get_valid_transforms():\n    return Compose([\n            CenterCrop(cfg['input_size'], cfg['input_size'], p=1.),\n            Resize(cfg['input_size'], cfg['input_size']),\n        ], p=1.)\n\n\ndef normalize_and_to_tensor(img):\n    transform = Compose([Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0),\n                       ToTensorV2(p=1.0)],p=1.0)\n    return transform(image=img)['image']\n\n\nclass CASSAVA(Dataset):\n    def __init__(self,\n                 imagenames,\n                 csv,\n                 root_dir,\n                 input_size=cfg['input_size'],\n                 transforms=None,\n                 train=True,\n                contrastive = True):\n        self.imagenames = imagenames\n        self.csv = csv\n        self.root_dir = root_dir\n        self.input_size = input_size\n        self.transforms = transforms\n        self.train = train\n        self.contrastive = contrastive\n    def __len__(self):\n        return len(imagenames)\n    def get_onehot(self,label):\n        onehot = np.zeros(5)\n        onehot[label] = 1\n        return onehot\n    def __getitem__(self,idx):\n        imagename = self.imagenames[idx]\n        label = self.csv[self.csv['image_id']==imagename]['label']\n        label = self.get_onehot(label)\n        image = cv2.imread(self.root_dir+'/'+imagename)\n        image = cv2.resize(image,(self.input_size,self.input_size))\n        image_aug1 = self.transforms(image=image)['image']\n        image_aug2 = self.transforms(image=image)['image']\n        if random.choice([1,2])==1:\n            image_aug1 = add_guassian_noise(image_aug1)\n            image_aug2 = cutout(image_aug2)\n        else:\n            image_aug2 = add_guassian_noise(image_aug2)\n            image_aug1 = cutout(image_aug1)\n        label = torch.from_numpy(label)\n        image = normalize_and_to_tensor(image)\n        image_aug1 = normalize_and_to_tensor(image_aug1)\n        image_aug2 = normalize_and_to_tensor(image_aug2)\n        #image_aug3 = normalize_and_to_tensor(image_aug3)\n        return image,image_aug1,image_aug2,label","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_transforms = get_train_transforms()\nt_dataset = CASSAVA(imagenames,csv,train_dir,transforms = train_transforms)\ntrain_loader = DataLoader(dataset=t_dataset, batch_size=cfg['batch_size'], shuffle=True, num_workers=2)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"backbone = models.resnet50(pretrained=False)\nmodules = list(backbone.children())[:-2]\nbackbone = nn.Sequential(*modules)\nprint(backbone)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CLASSIFIER(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.effn = backbone\n        self.average = nn.AvgPool2d((8,8))\n        self.flatten = nn.Flatten()\n        \n    def forward(self,x):\n        x = self.effn(x)\n        x = self.average(x)\n        x = self.flatten(x)\n        #x = F.relu(self.projection(x))\n        return x\n    \nclass HEAD(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.dense1 = nn.Linear(2048,512)\n        self.dense2 = nn.Linear(512,2048)\n    def forward(self,x):\n        x = F.relu(self.dense1(x))\n        x = F.relu(self.dense2(x))\n        return x\n\nHead = HEAD()\nx = torch.randn((1,2048))\ny = Head(x)\nprint(y.size())\n\nmodel = CLASSIFIER()\nx = torch.randn((1,3,256,256))\ny = model(x)\nprint(y.size())","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"count_parameters(model)\nprint(\"##\"*12,'head')\ncount_parameters(Head)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def cosine_similarity(y_true,y_pred):\n    y_pred = y_pred.detach()\n    y_true = y_true/(y_true.norm(dim=-1)[:,None]+1e-6)\n    y_pred = y_pred/(y_pred.norm(dim=-1)[:,None]+1e-6)\n    loss = y_true*y_pred\n    loss = loss.sum(dim=-1)\n    return -loss","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"optimizer = torch.optim.Adam(model.parameters(), lr=cfg['lr'])\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=10, verbose=True)\nmodel.to(cfg['device'])\nHead.to(cfg['device'])\nfor epoch in range(cfg['epochs']):\n    \n    #epoch parameters\n    epoch_loss = 0\n    model.train()\n    start = time.time()\n    \n    for i,(image2,image3,image4,label) in enumerate(tqdm(train_loader,total=train_loader.__len__(),ncols = 500)):\n        \n        # image1 = image1.to(cfg['device'])\n        image2 = image2.to(cfg['device'])\n        image3 = image3.to(cfg['device'])\n        image4=  image4.to(cfg['device'])\n        #represent\n        rep3,rep4 =model(image3),model(image4)\n        h3,h4 = Head(rep3),Head(rep4)\n        \n        # negativity similarity\n        loss = (cosine_similarity(h3,rep4)+cosine_similarity(h4,rep3))/2\n\n        #backprop\n        optimizer.zero_grad()\n        loss.mean().backward()\n        optimizer.step()\n        \n        epoch_loss+=loss.mean().item()\n\n    scheduler.step(epoch_loss/(i+1))\n    print('Epoch {:03}: | Loss: {:.3f} | Training time: {}'.format(\n            epoch + 1,  \n            epoch_loss/(i+1), \n            str(time.time() - start)[:7]))\n    torch.save(model.state_dict(), 'I_am_trained_{}_{}.pt'.format(epoch,epoch_loss/(i+1)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"del model, optimizer, train_loader, scheduler,Head\ntorch.cuda.empty_cache()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"csv.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","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}